Remove synthetic training path and add fine-text OCR option
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user