aleksasp's picture
Polish metrical HTR — current public release
9c471a9
Raw History Blame Contribute Delete
1.53 kB
#!/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()