Files
zulassung-ocr-compare/ocr_compare/benchmark_synthetic.py
T

106 lines
4.3 KiB
Python
Raw Normal View History

"""Measure detection/readability and field assignment on generated pages only.
Set OCR_MODEL_HOME to a local cache with official classic model weights first.
This script never loads personal documents and prints aggregate counts only.
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import re
import tempfile
from . import server
from .synthetic import TARGET_CODES
def _norm(value: str) -> str:
return re.sub(r"\s+", " ", value.upper()).strip()
def _box(polygon: list[list[int]]) -> tuple[float, float, float, float]:
xs, ys = zip(*polygon)
return min(xs), min(ys), max(xs), max(ys)
def _coverage(reference: tuple[float, float, float, float],
candidate: list[float]) -> float:
x1, y1, x2, y2 = reference
a, b, c, d = candidate
overlap = max(0, min(x2, c) - max(x1, a)) * max(0, min(y2, d) - max(y1, b))
return overlap / max(1, (x2 - x1) * (y2 - y1))
def evaluate(dataset: Path, *, split: str = "val", limit: int | None = None) -> dict:
if split not in ("train", "val", "all"):
raise ValueError("split must be train, val or all")
dataset = dataset.resolve()
manifest = json.loads((dataset / "manifest.json").read_text(encoding="utf-8"))
if manifest.get("source") != "fully synthetic" or manifest.get("training_on_real_documents") is not False:
raise ValueError("this benchmark accepts generated datasets only")
rows = [json.loads(line) for line in (dataset / "ground_truth.jsonl").read_text(encoding="utf-8").splitlines()]
if split != "all":
rows = [row for row in rows if row["split"] == split]
if limit is not None:
if limit < 1:
raise ValueError("limit must be positive")
rows = rows[:limit]
counts = {"split": split, "documents": 0, "failed_documents": 0, "value_lines": 0,
"value_boxes_covered": 0, "value_text_read": 0,
"fields": 0, "fields_exact": 0, "fields_wrong": 0,
"fields_missing": 0,
"per_code": {code: {"value_lines": 0, "read": 0, "assigned": 0}
for code in TARGET_CODES}}
for row in rows:
if row.get("source") != "fully synthetic":
raise ValueError("non-synthetic row")
image = (dataset / row["image"]).resolve()
if not image.is_relative_to(dataset / "pages"):
raise ValueError("image outside generated pages")
with tempfile.TemporaryDirectory(prefix="ocr-synthetic-benchmark-") as directory:
result = server._classic(image.read_bytes(), Path(directory))
counts["documents"] += 1
if result["status"] != "ok":
counts["failed_documents"] += 1
readings = result.get("lines", []) if result["status"] == "ok" else []
for item in row["lines"]:
code = item["code"]
if item["part"] != "value" or code not in TARGET_CODES:
continue
counts["value_lines"] += 1
counts["per_code"][code]["value_lines"] += 1
overlaps = [line for line in readings
if line.get("box") and _coverage(_box(item["polygon"]), line["box"]) >= .45]
if overlaps:
counts["value_boxes_covered"] += 1
if any(_norm(item["text"]) in _norm(line["text"]) for line in overlaps):
counts["value_text_read"] += 1
counts["per_code"][code]["read"] += 1
assigned = ({item["code"]: _norm(item["value"]) for item in result["fields"]}
if result["status"] == "ok" else {})
for code, reference in row["fields"].items():
counts["fields"] += 1
if assigned.get(code) == _norm(reference):
counts["fields_exact"] += 1
counts["per_code"][code]["assigned"] += 1
elif code in assigned:
counts["fields_wrong"] += 1
else:
counts["fields_missing"] += 1
return counts
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--dataset", type=Path, required=True)
parser.add_argument("--split", choices=("train", "val", "all"), default="val")
parser.add_argument("--limit", type=int)
args = parser.parse_args()
print(json.dumps(evaluate(args.dataset, split=args.split, limit=args.limit), indent=2))
if __name__ == "__main__":
main()