File size: 4,533 Bytes
b61332d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 | """Evaluate printed-text OCR (TrOCR) on a subset of SROIE 2019 test lines.
Reports CER (character error rate, Levenshtein/len(ref)) and exact-match
accuracy, raw and after punctuation/case normalization. Optionally logs
to a trackio dashboard.
Run: python evaluate_ocr.py --n 1000 [--batch-size 16] [--trackio-space <id>]
Note: requires trocr_compat.py alongside (transformers 5.x needs a generated
tokenizer.json for microsoft/trocr-base-printed).
"""
import argparse
import json
import re
import torch
from datasets import load_dataset
from trocr_compat import ensure_trocr_tokenizer
from transformers import TrOCRProcessor, VisionEncoderDecoderModel
def levenshtein(a, b):
if a == b:
return 0
if len(a) < len(b):
a, b = b, a
prev = list(range(len(b) + 1))
for i, ca in enumerate(a, 1):
cur = [i]
for j, cb in enumerate(b, 1):
cur.append(min(prev[j] + 1, cur[j - 1] + 1, prev[j - 1] + (ca != cb)))
prev = cur
return prev[-1]
def normalize(s):
return re.sub(r"[^A-Z0-9]", "", s.upper())
def cer(ref, hyp):
ref = ref.strip()
if not ref:
return None
return levenshtein(ref, hyp) / len(ref)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--model", default="microsoft/trocr-base-printed")
ap.add_argument("--dataset", default="priyank-m/SROIE_2019_text_recognition")
ap.add_argument("--split", default="test")
ap.add_argument("--n", type=int, default=1000)
ap.add_argument("--batch-size", type=int, default=16)
ap.add_argument("--seed", type=int, default=7)
ap.add_argument("--out", default="results_ocr.json")
ap.add_argument("--trackio-project", default="doc-fraud-ocr")
ap.add_argument("--trackio-space", default=None)
args = ap.parse_args()
device = "cuda" if torch.cuda.is_available() else "cpu"
print("device:", device, flush=True)
ensure_trocr_tokenizer(args.model)
processor = TrOCRProcessor.from_pretrained(args.model)
model = VisionEncoderDecoderModel.from_pretrained(args.model).to(device).eval()
ds = load_dataset(args.dataset, split=args.split)
ds = ds.shuffle(seed=args.seed).select(range(min(args.n, len(ds))))
print("evaluating on", len(ds), "lines", flush=True)
if args.trackio_space:
import trackio
trackio.init(project=args.trackio_project, space_id=args.trackio_space)
refs, hyps = [], []
for start in range(0, len(ds), args.batch_size):
batch = ds[start:start + args.batch_size]
images = [im.convert("RGB") for im in batch["image"]]
pixel_values = processor(images=images, return_tensors="pt").pixel_values.to(device)
with torch.no_grad():
ids = model.generate(pixel_values, max_new_tokens=64)
texts = processor.batch_decode(ids, skip_special_tokens=True)
refs.extend(batch["text"])
hyps.extend(texts)
done = start + len(texts)
if done % 100 < args.batch_size:
part = [c for r, h in zip(refs, hyps) if (c := cer(normalize(r), normalize(h))) is not None]
running = sum(part) / max(1, len(part))
print("%d/%d running CER %.4f" % (done, len(ds), running), flush=True)
if args.trackio_space:
trackio.log({"cer": running, "samples": done}, step=done)
raw_cers = [c for c in (cer(r, h) for r, h in zip(refs, hyps)) if c is not None]
norm_cers = [c for c in (cer(normalize(r), normalize(h)) for r, h in zip(refs, hyps)) if c is not None]
raw_exact = sum(1 for r, h in zip(refs, hyps) if r.strip() == h.strip()) / len(refs)
norm_exact = sum(1 for r, h in zip(refs, hyps) if normalize(r) == normalize(h)) / len(refs)
metrics = {
"model": args.model, "dataset": args.dataset, "split": args.split, "n": len(refs),
"cer_raw": round(sum(raw_cers) / len(raw_cers), 4),
"cer_normalized": round(sum(norm_cers) / len(norm_cers), 4),
"exact_match_raw": round(raw_exact, 4),
"exact_match_normalized": round(norm_exact, 4),
}
print(json.dumps(metrics, indent=2), flush=True)
with open(args.out, "w") as f:
json.dump(metrics, f, indent=2)
worst = sorted(zip(refs, hyps), key=lambda rh: cer(normalize(rh[0]), normalize(rh[1])) or 0, reverse=True)[:10]
print("worst 10 (ref | hyp):")
for r, h in worst:
print(" ", repr(r), "|", repr(h))
if args.trackio_space:
trackio.log(metrics)
trackio.finish()
if __name__ == "__main__":
main() |