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

118 lines
5.8 KiB
Python
Raw Normal View History

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