doc-fraud-ocr-validate / evaluate_ocr.py
Offlin33er's picture
TrOCR OCR evaluation harness (CER + exact match on SROIE test lines)
b61332d verified
Raw History Blame Contribute Delete
4.53 kB
"""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()