Remove synthetic training path and add fine-text OCR option

This commit is contained in:
OCR Team
2026-10-07 22:07:30 +02:00
parent e28f50bb88
commit e4adf70db0
10 changed files with 82 additions and 526 deletions
+46
View File
@@ -2,7 +2,9 @@ from io import BytesIO
import json
from pathlib import Path
import subprocess
import sys
import tempfile
from types import SimpleNamespace
import unittest
from unittest.mock import patch
@@ -10,6 +12,7 @@ from PIL import Image
from starlette.testclient import TestClient
from ocr_compare import server
from ocr_compare import classic_worker
from ocr_compare.fields import assign_fields
from ocr_compare.vision_worker import _missing_neighbor_box, _retry_table, extract
@@ -205,6 +208,49 @@ class LocalComparisonTests(unittest.TestCase):
self.assertEqual(response.json()["results"]["classic"], classic)
self.assertEqual(response.json()["results"]["vision"], vision)
def test_fine_text_choice_reaches_classic_engine_only(self):
classic = {"status": "error", "error": "classic_failed", "elapsed_ms": 0}
vision = {"status": "error", "error": "vision_disabled", "elapsed_ms": 0}
with patch.object(server, "_classic", return_value=classic) as run_classic, patch.object(
server, "_vision", return_value=vision) as run_vision:
with TestClient(server.app) as client:
response = client.post("/api/compare", content=photo(), headers={
"content-type": "image/jpeg", "x-ocr-detail": "fine"})
invalid = client.post("/api/compare", content=photo(), headers={
"content-type": "image/jpeg", "x-ocr-detail": "unknown"})
self.assertEqual(response.status_code, 200)
self.assertTrue(run_classic.call_args.kwargs["fine_text"])
self.assertEqual(run_vision.call_count, 1)
self.assertEqual(invalid.status_code, 422)
self.assertEqual(invalid.json()["error"], "invalid_detail")
def test_classic_worker_uses_requested_detector_resolution(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
for name in ("PP-OCRv5_mobile_det", "latin_PP-OCRv5_mobile_rec"):
model = root / name
model.mkdir()
(model / "inference.pdiparams").touch()
seen = []
class FakePaddleOCR:
def __init__(self, **kwargs):
seen.append(kwargs)
def predict(self, _image):
return [SimpleNamespace(json={"res": {"rec_texts": ["B"],
"rec_polys": [[[1, 1], [2, 1], [2, 2], [1, 2]]]}})]
for options, expected in (([], 960), (["--fine-text"], 2048)):
output = root / "result.json"
with patch.object(classic_worker, "MODELS", root), patch.object(
sys, "argv", ["classic_worker", str(root / "input.jpg"), str(output), *options]), patch.dict(
sys.modules, {"paddleocr": SimpleNamespace(PaddleOCR=FakePaddleOCR)}):
classic_worker.main()
self.assertEqual(seen[-1]["text_det_limit_side_len"], expected)
self.assertEqual(seen[-1]["text_det_limit_type"], "max")
self.assertEqual(json.loads(output.read_text())["lines"][0]["text"], "B")
def test_worker_launch_errors_are_reported_per_engine(self):
with tempfile.TemporaryDirectory() as directory:
model_home = Path(directory)
-117
View File
@@ -1,117 +0,0 @@
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()