118 lines
5.8 KiB
Python
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()
|