106 lines
4.3 KiB
Python
106 lines
4.3 KiB
Python
|
|
"""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()
|