Download evaluate_ocr.py from Offlin33er/doc-fraud-ocr-validate: direct link, hf CLI and curl.
- Browser
- Download file 4.53 kB
-
https://huggingface.co/Offlin33er/doc-fraud-ocr-validate/resolve/main/evaluate_ocr.py
- Command line
-
hf download hf://Offlin33er/doc-fraud-ocr-validate/evaluate_ocr.py
-
curl -L -o evaluate_ocr.py https://huggingface.co/Offlin33er/doc-fraud-ocr-validate/resolve/main/evaluate_ocr.py
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() |