#!/usr/bin/env python3 """Run the MelDynamics Polish Metrical HTR model on a single line crop. Example: python inference.py masked_line_crop.png --device cuda """ from __future__ import annotations import argparse import torch from PIL import Image from transformers import TrOCRProcessor, VisionEncoderDecoderModel def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("image", help="one segmented/masked text-line crop") ap.add_argument("--model", default="meldynamics/polish-metrical-htr-experimental") ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") ap.add_argument("--max-new-tokens", type=int, default=96) args = ap.parse_args() # The trained checkpoint includes the compatible processor/tokenizer. Force the # non-square line geometry; the model otherwise falls back to a ViT square size. processor = TrOCRProcessor.from_pretrained(args.model, size={"height": 192, "width": 1024}) model = VisionEncoderDecoderModel.from_pretrained(args.model).to(args.device).eval() image = Image.open(args.image).convert("RGB") pixels = processor(images=image, return_tensors="pt").pixel_values.to(args.device) with torch.inference_mode(): ids = model.generate( pixels, max_new_tokens=args.max_new_tokens, num_beams=1, interpolate_pos_encoding=True, ) print(processor.batch_decode(ids, skip_special_tokens=True)[0]) if __name__ == "__main__": main()