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()