diff --git a/README.md b/README.md index 9fa9898..5782124 100644 --- a/README.md +++ b/README.md @@ -9,7 +9,7 @@ Eine lokale Testoberfläche für zwei unabhängige Wege auf **demselben Bild**: Die Oberfläche zeigt Rohtext, Feldvorschläge, Bildbelege und Laufzeiten nebeneinander. Die Ecken des Dokuments lassen sich per Maus oder Touch setzen; alternativ wird das ganze Bild verwendet. Besondere Feldcodes: C.1.1 (Name), C.1.2 (Vorname), C.1.3 (Anschrift), B (Erstzulassung), 2.1 (HSN) und 2.2 (TSN). **Jeden Vorschlag am Bild prüfen.** Unleserliche oder unbelegte Werte bleiben offen. Es werden keine künstlichen Konfidenzwerte angezeigt. -Dieses Repository enthält **keine Dokumentbilder, OCR-Ergebnisse, personenbezogenen Testdaten, Modellgewichte, privaten Hostnamen oder Zugangsdaten**. Die Tests nutzen erfundene Textzeilen und generierte Bilder. Es gibt kein Training und keine Datenübertragung an externe OCR- oder KI-APIs. Die offiziellen Modellgewichte werden beim Einrichten heruntergeladen; die Inferenz nutzt anschließend explizite lokale Modellverzeichnisse. Der optionale SSH-Modus überträgt das Bild ausschließlich an eine selbst verwaltete Maschine. Das Projekt ist für interne Tests gedacht und enthält keine öffentliche Lizenz. +Dieses Repository enthält **keine Dokumentbilder, OCR-Ergebnisse, personenbezogenen Testdaten, Modellgewichte, privaten Hostnamen oder Zugangsdaten**. Die Tests nutzen erfundene Textzeilen und generierte Bilder. Ein Generator kann lokal synthetische Trainingsbeispiele erzeugen; dieses Repository startet selbst kein Training. Es gibt keine Datenübertragung an externe OCR- oder KI-APIs. Die offiziellen Modellgewichte werden beim Einrichten heruntergeladen; die Inferenz nutzt anschließend explizite lokale Modellverzeichnisse. Der optionale SSH-Modus überträgt das Bild ausschließlich an eine selbst verwaltete Maschine. Das Projekt ist für interne Tests gedacht und enthält keine öffentliche Lizenz. ## Voraussetzungen @@ -98,6 +98,21 @@ Die klassische OCR liefert präzisere Zeilenboxen, kann aber Text verlesen. Die Die App verarbeitet Uploads im Speicher und in kurzlebigen temporären Dateien. Sie schreibt weder Dokumente noch OCR-Ausgaben in ein dauerhaftes Verzeichnis. Browser und eigener Server sehen die Bilddaten; bei SSH-Betrieb auch der eigene SSH-Zielhost. Vor dem Einsatz mit echten Dokumenten sollten Modelle bereits eingerichtet sein. Ein fehlender Cache wird als Fehler angezeigt, statt beim Upload Modellgewichte nachzuladen. +## Synthetische Daten und Training + +Der Generator erzeugt **sichtbar ungültige, erfundene** Fahrzeugformulare mit variierter Feldreihenfolge, Schrift, Dichte und Bildstörung. Er liest keine Quelldokumente. Ausgabe und Modellartefakte liegen unter dem von Git ignorierten `data/`-Verzeichnis. Niemals echte Dokumente, daraus ausgeschnittene Texte oder OCR-Ausgaben in diesen Trainingssatz mischen. + +```sh +python -m ocr_compare.synthetic --out data/synthetic-v1 --count 100 --seed 20261007 +python -m ocr_compare.benchmark_synthetic --dataset data/synthetic-v1 +``` + +Vor dem Benchmark die **klassischen** Modelle mit `python -m ocr_compare.setup_models --classic` lokal einrichten. Der Benchmark verarbeitet standardmäßig nur den nach ganzen Dokumenten getrennten Validierungsteil; `--split train` und `--split all` sind für die Fehlersuche gedacht. Er gibt ausschließlich Summen aus: an Wertboxen überlappende OCR-Zeilen, dort vollständig gelesene Werte und korrekt, falsch oder gar nicht zugeordnete Zielfelder. Für eine neue Stichprobe ein neues Ausgabeverzeichnis und einen anderen Seed nehmen. Der Generator überschreibt nichts. + +Der Satz enthält `rec_train.txt`/`rec_val.txt` mit Bildpfad und Text für die PaddleOCR-Texterkennung sowie `det_train.txt`/`det_val.txt` mit Seitenpfad und Zeilenpolygonen für die Texterkennung auf der Seite. `ground_truth.jsonl` enthält die exakten synthetischen Sollwerte für C.1.1, C.1.2, C.1.3, B, 2.1 und 2.2. Eine Validierung der Labels und des Auswertungscodes ist in den lokalen Tests enthalten; **ein Modelltraining und dessen Export sind hier noch nicht getestet**. + +Für einen Trainingsexperiment zuerst den Fehler nach Stufe messen: Fehlende Textboxen sprechen für den [Detektor](https://www.paddleocr.ai/main/en/version3.x/module_usage/text_detection.html), falsch gelesene vorhandene Boxen für den [Recognizer](https://www.paddleocr.ai/main/en/version3.x/module_usage/text_recognition.html). Korrekt gelesene Werte mit falschem Feldcode sind ein Zuordnungs- oder Layoutproblem. PaddleOCR-VL ist eine Pipeline aus Layoutanalyse und VLM. Die [offizielle SFT-Anleitung](https://www.paddleocr.ai/main/en/version3.x/pipeline_usage/PaddleOCR-VL.html#5-model-fine-tuning) unterstützt zurzeit nur das VLM; die [ERNIEKit-Beispielkonfiguration](https://github.com/PaddlePaddle/ERNIE/blob/release/v1.4/docs/paddleocr_vl_sft.md) wurde auf einer 80-GB-GPU demonstriert. Das ist keine belastbare Zusage für Training auf einer kleineren lokalen GPU. Synthetische Validierung allein beweist keine Verbesserung bei echten Dokumenten: Diese ausschließlich als unveränderten, lokalen **Test-Holdout** verwenden, nie zum Training. + ## Prüfen und Fehler eingrenzen ```sh diff --git a/ocr_compare/benchmark_synthetic.py b/ocr_compare/benchmark_synthetic.py new file mode 100644 index 0000000..dbe06b7 --- /dev/null +++ b/ocr_compare/benchmark_synthetic.py @@ -0,0 +1,105 @@ +"""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() diff --git a/ocr_compare/synthetic.py b/ocr_compare/synthetic.py new file mode 100644 index 0000000..c52f76a --- /dev/null +++ b/ocr_compare/synthetic.py @@ -0,0 +1,275 @@ +"""Generate labeled, visibly invalid vehicle-form images without source documents. + +The output is ignored by Git. Each page has exact field/line ground truth, +PaddleOCR detection labels, and recognition crops with page-level splits. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path +import random + +import cv2 +import numpy as np +from PIL import Image, ImageDraw, ImageFont + + +TARGET_CODES = ("C.1.1", "C.1.2", "C.1.3", "B", "2.1", "2.2") +CAPTIONS = { + "C.1.1": "Name", "C.1.2": "Vorname", "C.1.3": "Anschrift", + "B": "Erstzulassung", "2.1": "HSN", "2.2": "TSN", + "A": "Kennzeichen", "D.1": "Marke", "D.2": "Typ", + "D.3": "Handelsname", "E": "Fahrzeug-ID", "P.3": "Kraftstoff", + "F.1": "Masse", "F.2": "Gesamtmasse", "J": "Fahrzeugklasse", + "P.1": "Hubraum", "P.2": "Leistung", "S.1": "Sitzplaetze", + "V.7": "CO2", "G": "Leergewicht", "15.1": "Bereifung", + "15.2": "Bereifung hinten", +} +EXTRA_CODES = ("A", "D.1", "D.2", "D.3", "E", "P.3", "F.1", "F.2", + "J", "P.1", "P.2", "S.1", "V.7", "G", "15.1", "15.2") +FONT_CANDIDATES = ( + "/System/Library/Fonts/Supplemental/Arial.ttf", + "/System/Library/Fonts/Supplemental/Arial Narrow.ttf", + "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", + "/usr/share/fonts/truetype/liberation2/LiberationSans-Regular.ttf", +) +SYLLABLES = ("VE", "NO", "RA", "TU", "ME", "LI", "SA", "KO", "DI", "FA", + "ZEN", "BER", "LON", "FEL", "NAR", "TOV", "XEN", "MIR") + + +def _fonts(font_path: Path | None, size: int) -> ImageFont.FreeTypeFont: + candidates = ((str(font_path),) if font_path is not None else FONT_CANDIDATES) + for path in candidates: + try: + return ImageFont.truetype(path, size) + except OSError: + continue + if font_path is not None: + raise ValueError("font unavailable") + return ImageFont.load_default(size=size) + + +def _word(rng: random.Random, *, syllables: int = 3) -> str: + value = "".join(rng.choice(SYLLABLES) for _ in range(syllables)) + if rng.random() < .22: + value = value.replace("A", "Ä", 1) if "A" in value else value.replace("O", "Ö", 1) + return value + + +def _values(rng: random.Random) -> dict[str, list[str]]: + town = _word(rng, syllables=2) + "STADT" + street = _word(rng, syllables=2) + rng.choice(("WEG", "STRASSE", "ALLEE")) + date = f"{rng.randint(1, 28):02d}.{rng.randint(1, 12):02d}.{rng.randint(1995, 2024)}" + return { + "C.1.1": [_word(rng)], + "C.1.2": [_word(rng, syllables=2)], + "C.1.3": [f"{street} {rng.randint(1, 199)}", + f"{rng.randint(10000, 99999)} {town}"], + "B": [date], + "2.1": [f"{rng.randint(1000, 9999)}"], + "2.2": ["".join(rng.choices("ABCDEFGHJKLMNPRSTUVWXYZ", k=3)) + + f"{rng.randint(0, 999):03d}"], + "A": [f"TEST-{rng.randint(1000, 9999)}"], + "D.1": [_word(rng, syllables=2)], + "D.2": [f"TYP-{rng.randint(100, 999)}"], + "D.3": [f"MODELL-{rng.randint(100, 999)}"], + "E": [f"SYNTH-ID-{rng.randint(100000, 999999)}"], + "P.3": [rng.choice(("BENZIN", "DIESEL", "ELEKTRO"))], + "F.1": [f"{rng.randint(1300, 4200)} KG"], + "F.2": [f"{rng.randint(1300, 4200)} KG"], + "J": [rng.choice(("M1", "N1", "L3E"))], + "P.1": [f"{rng.randint(100, 3000)} CCM"], + "P.2": [f"{rng.randint(20, 250)} KW"], + "S.1": [str(rng.randint(1, 7))], + "V.7": [f"{rng.randint(30, 250)} G/KM"], + "G": [f"{rng.randint(900, 2400)} KG"], + "15.1": [f"{rng.randint(145, 265)}/55 R{rng.randint(14, 20)}"], + "15.2": [f"{rng.randint(145, 265)}/55 R{rng.randint(14, 20)}"], + } + + +def _render(seed: int, font_path: Path | None): + rng = random.Random(seed) + if font_path is None: + installed = [Path(path) for path in FONT_CANDIDATES if Path(path).is_file()] + font_path = rng.choice(installed) if installed else None + hardness = seed % 3 + dense = seed % 4 != 0 + width, height = rng.choice(((1800, 1180), (1680, 1260), (1500, 1450))) + page = Image.new("RGB", (width, height), rng.choice(((248, 247, 243), (241, 244, 247)))) + draw = ImageDraw.Draw(page) + value_font = _fonts(font_path, rng.randint(20, 25) if dense else rng.randint(30, 38)) + small_font = _fonts(font_path, rng.randint(15, 19) if dense else rng.randint(22, 27)) + header_font = _fonts(font_path, 27 if dense else 34) + lines: list[dict] = [] + + if hardness or dense: + for y in range(115, height - 95, 17 if dense else 29): + draw.line((35, y, width - 35, y + (y % 13) - 6), + fill=(220, 225, 228), width=1) + if hardness == 2: + for _ in range(1400): + x, y = rng.randrange(width), rng.randrange(height) + draw.point((x, y), fill=(209, 214, 216)) + + def line(text: str, x: int, y: int, font, *, code: str | None = None, + part: str = "decoration", max_width: int | None = None) -> None: + if len(text) > 25: + raise ValueError("recognition label exceeds 25 characters") + chosen = font + if max_width is not None: + while draw.textlength(text, font=chosen) > max_width and chosen.size > 19: + chosen = _fonts(font_path, chosen.size - 1) + ink = ((65, 70, 74), (76, 78, 80), (88, 85, 82)) if hardness == 2 else ( + (28, 33, 40), (36, 39, 45), (52, 48, 47)) + draw.text((x, y), text, font=chosen, fill=rng.choice(ink)) + x1, y1, x2, y2 = draw.textbbox((x, y), text, font=chosen) + lines.append({"text": text, "code": code, "part": part, + "polygon": [[max(0, x1 - 5), max(0, y1 - 5)], + [min(width - 1, x2 + 5), max(0, y1 - 5)], + [min(width - 1, x2 + 5), min(height - 1, y2 + 5)], + [max(0, x1 - 5), min(height - 1, y2 + 5)]]}) + + margin = rng.randint(65, 90) + gap = rng.randint(24, 45) + cell_width = (width - 2 * margin - gap) // 2 + top = 130 if dense else 155 + rows = 8 if dense else 5 + row_height = min(205, (height - top - 100) // rows) + draw.rectangle((margin - 15, 37, width - margin + 15, height - 35), + outline=(120, 127, 134), width=2) + line("SYNTHETISCHER TEST", margin, 54, header_font) + line("FAHRZEUGDATEN", width - margin - (240 if dense else 300), 60, small_font) + + values = _values(rng) + codes = list(TARGET_CODES) + rng.sample(EXTRA_CODES, 10 if dense else 4) + rng.shuffle(codes) + for index, code in enumerate(codes): + column, row = index // rows, index % rows + x = margin + column * (cell_width + gap) + y = top + row * row_height + draw.rectangle((x, y, x + cell_width, y + row_height - 12), + outline=(138, 147, 151), width=rng.choice((1, 2))) + stacked = rng.random() < .65 or code == "C.1.3" + code_x, code_y = x + 14, y + 11 + line(code, code_x, code_y, small_font, code=code, part="code") + caption_x = code_x + round(draw.textlength(code, font=small_font)) + 17 + line(CAPTIONS[code], caption_x, code_y, small_font, code=code, + part="caption", max_width=max(90, cell_width - (caption_x - x) - 15)) + value_x = x + (18 if stacked else round(cell_width * .43)) + first_y = y + (39 if dense else 53) if stacked else y + 10 + for offset, value in enumerate(values[code]): + if not stacked and offset == 0 and value_x < caption_x + (90 if dense else 175): + first_y = y + (43 if dense else 59) + line(value, value_x, first_y + (28 if dense else 40) * offset, value_font, + code=code, part="value", max_width=cell_width - (value_x - x) - 12) + + line("UNGUELTIG / TESTDATEN", margin, height - 82, small_font) + src = np.float32([[0, 0], [width - 1, 0], [width - 1, height - 1], [0, height - 1]]) + jitter = 20 if seed % 3 else 43 + dst = np.float32([[rng.randint(0, jitter), rng.randint(0, jitter)], + [width - 1 - rng.randint(0, jitter), rng.randint(0, jitter)], + [width - 1 - rng.randint(0, jitter), height - 1 - rng.randint(0, jitter)], + [rng.randint(0, jitter), height - 1 - rng.randint(0, jitter)]]) + transform = cv2.getPerspectiveTransform(src, dst) + image = cv2.warpPerspective(np.asarray(page), transform, (width, height), + flags=cv2.INTER_LINEAR, borderValue=(255, 255, 255)) + for item in lines: + points = np.float32(item["polygon"]).reshape(-1, 1, 2) + transformed = cv2.perspectiveTransform(points, transform).reshape(4, 2) + item["polygon"] = [[round(float(np.clip(px, 0, width - 1))), + round(float(np.clip(py, 0, height - 1)))] + for px, py in transformed] + + yy, xx = np.mgrid[0:height, 0:width] + gradient = 1 - (hardness * .07) * (xx / width) - (hardness * .04) * (yy / height) + image = np.uint8(np.clip(image.astype(np.float32) * gradient[:, :, None], 0, 255)) + if hardness == 2: + image = cv2.GaussianBlur(image, (3, 3), .55) + noise = np.random.default_rng(seed).normal(0, 2.5, image.shape) + image = np.uint8(np.clip(image.astype(np.float32) + noise, 0, 255)) + return image, lines, values, hardness + + +def _crop(image: np.ndarray, polygon: list[list[int]]) -> np.ndarray: + points = np.float32(polygon) + width = max(16, round(max(np.linalg.norm(points[1] - points[0]), + np.linalg.norm(points[2] - points[3])))) + height = max(16, round(max(np.linalg.norm(points[3] - points[0]), + np.linalg.norm(points[2] - points[1])))) + target = np.float32([[0, 0], [width - 1, 0], [width - 1, height - 1], + [0, height - 1]]) + matrix = cv2.getPerspectiveTransform(points, target) + return cv2.warpPerspective(image, matrix, (width, height), + flags=cv2.INTER_LINEAR, borderValue=(255, 255, 255)) + + +def generate(output: Path, *, count: int, seed: int, font_path: Path | None = None) -> dict: + """Create a new directory; never read source files or overwrite a dataset.""" + if count < 2: + raise ValueError("count must be at least 2 for a document-level holdout") + if output.exists(): + raise FileExistsError(f"output already exists: {output}") + (output / "pages").mkdir(parents=True) + (output / "crops").mkdir() + validation = set(range(count - max(1, count // 5), count)) + rec = {"train": [], "val": []} + det = {"train": [], "val": []} + metadata = [] + for index in range(count): + page_id = f"synthetic-{index:05d}" + split = "val" if index in validation else "train" + image, lines, values, hardness = _render(seed + index * 1009, font_path) + page_path = f"pages/{page_id}.jpg" + quality = (92, 82, 69)[hardness] + if not cv2.imwrite(str(output / page_path), cv2.cvtColor(image, cv2.COLOR_RGB2BGR), + [cv2.IMWRITE_JPEG_QUALITY, quality]): + raise OSError("could not write synthetic page") + image = cv2.cvtColor(cv2.imread(str(output / page_path)), cv2.COLOR_BGR2RGB) + labels = [] + for number, item in enumerate(lines): + crop_path = f"crops/{page_id}-{number:02d}.jpg" + crop = _crop(image, item["polygon"]) + if not cv2.imwrite(str(output / crop_path), cv2.cvtColor(crop, cv2.COLOR_RGB2BGR), + [cv2.IMWRITE_JPEG_QUALITY, quality]): + raise OSError("could not write synthetic crop") + rec[split].append(f"{crop_path}\t{item['text']}\n") + item["crop"] = crop_path + labels.append({"transcription": item["text"], "points": item["polygon"]}) + det[split].append(page_path + "\t" + json.dumps(labels, ensure_ascii=False) + "\n") + metadata.append({"id": page_id, "split": split, "source": "fully synthetic", + "image": page_path, "width": image.shape[1], "height": image.shape[0], + "hardness": hardness, + "fields": {code: "\n".join(values[code]) for code in TARGET_CODES}, + "lines": lines}) + for split in ("train", "val"): + (output / f"rec_{split}.txt").write_text("".join(rec[split]), encoding="utf-8") + (output / f"det_{split}.txt").write_text("".join(det[split]), encoding="utf-8") + with (output / "ground_truth.jsonl").open("w", encoding="utf-8") as stream: + for row in metadata: + stream.write(json.dumps(row, ensure_ascii=False) + "\n") + manifest = {"source": "fully synthetic", "training_on_real_documents": False, + "seed": seed, "documents": count, "train_documents": count - len(validation), + "validation_documents": len(validation), + "recognition_crops": sum(len(items) for items in rec.values()), + "target_codes": list(TARGET_CODES)} + (output / "manifest.json").write_text(json.dumps(manifest, indent=2) + "\n", encoding="utf-8") + return manifest + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--out", type=Path, default=Path("data/synthetic-v1")) + parser.add_argument("--count", type=int, default=100) + parser.add_argument("--seed", type=int, default=20261007) + parser.add_argument("--font", type=Path) + args = parser.parse_args() + summary = generate(args.out, count=args.count, seed=args.seed, font_path=args.font) + print(f"Generated {summary['documents']} synthetic pages and " + f"{summary['recognition_crops']} labeled crops in {args.out}") + + +if __name__ == "__main__": + main() diff --git a/tests/test_synthetic.py b/tests/test_synthetic.py new file mode 100644 index 0000000..00bf7fa --- /dev/null +++ b/tests/test_synthetic.py @@ -0,0 +1,117 @@ +import json +from pathlib import Path +import tempfile +import unittest +from unittest.mock import patch + +import cv2 +import numpy as np + +from ocr_compare.synthetic import TARGET_CODES, _render, generate +from ocr_compare.benchmark_synthetic import evaluate + + +class SyntheticDatasetTests(unittest.TestCase): + def test_document_level_split_and_exact_labels(self): + with tempfile.TemporaryDirectory() as directory: + output = Path(directory) / "dataset" + manifest = generate(output, count=6, seed=31415) + self.assertEqual(manifest["source"], "fully synthetic") + self.assertEqual(manifest["train_documents"], 5) + self.assertEqual(manifest["validation_documents"], 1) + rows = [json.loads(line) for line in (output / "ground_truth.jsonl").read_text().splitlines()] + self.assertEqual(len(rows), 6) + self.assertEqual({row["id"] for row in rows if row["split"] == "val"}, + {"synthetic-00005"}) + self.assertEqual(len((output / "det_train.txt").read_text().splitlines()), 5) + self.assertEqual(len((output / "det_val.txt").read_text().splitlines()), 1) + crop_count = 0 + for row in rows: + self.assertEqual(set(row["fields"]), set(TARGET_CODES)) + page = cv2.imread(str(output / row["image"])) + self.assertIsNotNone(page) + self.assertEqual(page.shape[:2], (row["height"], row["width"])) + for item in row["lines"]: + self.assertLessEqual(len(item["text"]), 25) + self.assertTrue((output / item["crop"]).is_file()) + self.assertTrue(all(0 <= x < row["width"] and 0 <= y < row["height"] + for x, y in item["polygon"])) + crop_count += 1 + self.assertEqual(crop_count, manifest["recognition_crops"]) + for split in ("train", "val"): + for entry in (output / f"rec_{split}.txt").read_text().splitlines(): + path, label = entry.split("\t", 1) + self.assertTrue((output / path).is_file()) + self.assertEqual("synthetic-00005" in path, split == "val") + self.assertTrue(label) + + def test_seed_repeats_pixels_and_values(self): + first, lines_a, values_a, hardness_a = _render(12345, None) + second, lines_b, values_b, hardness_b = _render(12345, None) + self.assertTrue(np.array_equal(first, second)) + self.assertEqual(lines_a, lines_b) + self.assertEqual(values_a, values_b) + self.assertEqual(hardness_a, hardness_b) + + def test_requires_a_validation_document_and_never_overwrites(self): + with tempfile.TemporaryDirectory() as directory: + output = Path(directory) / "dataset" + with self.assertRaises(ValueError): + generate(output, count=1, seed=1) + self.assertFalse(output.exists()) + generate(output, count=2, seed=1) + with self.assertRaises(FileExistsError): + generate(output, count=2, seed=2) + + def test_benchmark_separates_reading_from_assignment(self): + with tempfile.TemporaryDirectory() as directory: + output = Path(directory) / "dataset" + generate(output, count=2, seed=31415) + row = json.loads((output / "ground_truth.jsonl").read_text().splitlines()[-1]) + readings = [] + for item in row["lines"]: + if item["part"] != "value" or item["code"] not in TARGET_CODES: + continue + xs, ys = zip(*item["polygon"]) + readings.append({"text": item["text"], + "box": [min(xs), min(ys), max(xs), max(ys)]}) + fields = [{"code": code, "value": value} for code, value in row["fields"].items() + if code not in ("B", "2.2")] + fields.append({"code": "B", "value": "wrong"}) + result = {"status": "ok", "lines": readings, "fields": fields} + with patch("ocr_compare.benchmark_synthetic.server._classic", return_value=result): + counts = evaluate(output, limit=1) + self.assertEqual(counts["value_lines"], 7) + self.assertEqual(counts["value_text_read"], 7) + self.assertEqual(counts["fields_exact"], 4) + self.assertEqual(counts["fields_wrong"], 1) + self.assertEqual(counts["fields_missing"], 1) + + def test_benchmark_rejects_non_synthetic_manifest(self): + with tempfile.TemporaryDirectory() as directory: + output = Path(directory) / "dataset" + generate(output, count=2, seed=31415) + manifest_path = output / "manifest.json" + manifest = json.loads(manifest_path.read_text()) + manifest["training_on_real_documents"] = True + manifest_path.write_text(json.dumps(manifest)) + with self.assertRaisesRegex(ValueError, "generated datasets only"): + evaluate(output, limit=1) + + def test_benchmark_counts_failed_documents_in_denominator(self): + with tempfile.TemporaryDirectory() as directory: + output = Path(directory) / "dataset" + generate(output, count=2, seed=31415) + with patch("ocr_compare.benchmark_synthetic.server._classic", + return_value={"status": "error"}): + counts = evaluate(output, limit=1) + self.assertEqual(counts["documents"], 1) + self.assertEqual(counts["failed_documents"], 1) + self.assertEqual(counts["value_lines"], 7) + self.assertEqual(counts["value_text_read"], 0) + self.assertEqual(counts["fields"], 6) + self.assertEqual(counts["fields_missing"], 6) + + +if __name__ == "__main__": + unittest.main()