Files
zulassung-ocr-compare/tests/test_synthetic.py
T

118 lines
5.8 KiB
Python

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