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)
|
||||
|
||||
Reference in New Issue
Block a user