Add private local OCR comparison package
This commit is contained in:
@@ -0,0 +1,304 @@
|
||||
"""Local-only, privacy-preserving comparison of classic OCR and PaddleOCR-VL."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from io import BytesIO
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import shlex
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
import warnings
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
from PIL import Image, ImageOps, UnidentifiedImageError
|
||||
from starlette.applications import Starlette
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import FileResponse, JSONResponse
|
||||
from starlette.routing import Route
|
||||
|
||||
from .fields import FIELD_CODES, assign_fields
|
||||
|
||||
|
||||
HERE = Path(__file__).resolve().parent
|
||||
PROJECT = HERE.parent
|
||||
MODEL_HOME = Path(os.environ.get("OCR_MODEL_HOME", PROJECT / ".cache")).expanduser().resolve()
|
||||
|
||||
|
||||
def _remote_command(seconds: int = 180, *, tiles_only: bool = False) -> str:
|
||||
remote_dir = os.environ["OCR_VISION_REMOTE_DIR"]
|
||||
remote_python = os.environ.get("OCR_VISION_REMOTE_PYTHON", "./.venv-vision/bin/python")
|
||||
runner = f'{shlex.quote(remote_python)} -m ocr_compare.vision_worker'
|
||||
if tiles_only:
|
||||
runner += ' --tiles-only'
|
||||
isolation = 'unshare -Urn ' if os.environ.get("OCR_VISION_REMOTE_UNSHARE") == "1" else ""
|
||||
device = shlex.quote(os.environ.get("OCR_VISION_REMOTE_DEVICE", "gpu:0"))
|
||||
model_home = os.environ.get("OCR_VISION_REMOTE_MODEL_HOME")
|
||||
model_env = f'OCR_MODEL_HOME={shlex.quote(model_home)} ' if model_home else ""
|
||||
return (
|
||||
f'set -eu; cd {shlex.quote(remote_dir)}; directory=$(mktemp -d); '
|
||||
'trap \'rm -rf -- "$directory"\' EXIT; '
|
||||
f'OCR_PRIVATE_DIR="$directory" OCR_VISION_DEVICE={device} {model_env}'
|
||||
f'timeout -k 5s {seconds}s {isolation}{runner}'
|
||||
)
|
||||
|
||||
|
||||
MAX_UPLOAD = 16 * 1024 * 1024
|
||||
MAX_PIXELS = 24_000_000
|
||||
MAX_SIDE = 10_000
|
||||
OCR_SIDE = 4000
|
||||
Image.MAX_IMAGE_PIXELS = MAX_PIXELS
|
||||
_slot = asyncio.Semaphore(1)
|
||||
|
||||
|
||||
class InputError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
def _geometry(image: Image.Image, header: str | None) -> Image.Image:
|
||||
try:
|
||||
spec = json.loads(header) if header else {}
|
||||
rotation = int(spec.get("rotation", 0))
|
||||
if rotation not in (0, 90, 180, 270):
|
||||
raise InputError("invalid_rotation")
|
||||
points = spec.get("corners")
|
||||
if points is not None:
|
||||
if (not isinstance(points, list) or len(points) != 4 or
|
||||
any(not isinstance(p, list) or len(p) != 2 or
|
||||
any(not isinstance(v, (int, float)) or isinstance(v, bool) or
|
||||
not 0 <= v <= 1 for v in p) for p in points)):
|
||||
raise InputError("invalid_corners")
|
||||
w, h = image.size
|
||||
source = np.array([[x * (w - 1), y * (h - 1)] for x, y in points], np.float32)
|
||||
area = cv2.contourArea(source)
|
||||
edges = [float(np.linalg.norm(source[(i + 1) % 4] - source[i])) for i in range(4)]
|
||||
turns = []
|
||||
for i in range(4):
|
||||
ax, ay = source[(i + 1) % 4] - source[i]
|
||||
bx, by = source[(i + 2) % 4] - source[(i + 1) % 4]
|
||||
turns.append(float(ax * by - ay * bx))
|
||||
if area < .02 * w * h or min(edges) < 30 or min(turns) <= 0:
|
||||
raise InputError("invalid_corners")
|
||||
out_w = round(max(edges[0], edges[2]))
|
||||
out_h = round(max(edges[1], edges[3]))
|
||||
if out_w < 30 or out_h < 30 or out_w * out_h > MAX_PIXELS:
|
||||
raise InputError("invalid_corners")
|
||||
# The four positions are always TL, TR, BR, BL. This avoids the
|
||||
# previous fixed-template aspect-ratio and rotation ambiguity.
|
||||
target = np.array([[0, 0], [out_w - 1, 0],
|
||||
[out_w - 1, out_h - 1], [0, out_h - 1]], np.float32)
|
||||
matrix = cv2.getPerspectiveTransform(source, target)
|
||||
warped = cv2.warpPerspective(np.asarray(image), matrix, (out_w, out_h),
|
||||
flags=cv2.INTER_LINEAR,
|
||||
borderMode=cv2.BORDER_CONSTANT,
|
||||
borderValue=(255, 255, 255))
|
||||
image = Image.fromarray(warped)
|
||||
if rotation:
|
||||
image = image.rotate(-rotation, expand=True)
|
||||
if max(image.size) > OCR_SIDE:
|
||||
image.thumbnail((OCR_SIDE, OCR_SIDE), Image.Resampling.LANCZOS)
|
||||
return image
|
||||
except (ValueError, TypeError, KeyError, json.JSONDecodeError) as exc:
|
||||
if isinstance(exc, InputError):
|
||||
raise
|
||||
raise InputError("invalid_corners") from None
|
||||
|
||||
|
||||
def _prepare(data: bytes, header: str | None) -> tuple[bytes, tuple[int, int]]:
|
||||
try:
|
||||
with warnings.catch_warnings():
|
||||
warnings.simplefilter("error", Image.DecompressionBombWarning)
|
||||
with Image.open(BytesIO(data)) as source:
|
||||
if source.format not in ("JPEG", "PNG", "MPO"):
|
||||
raise InputError("unsupported_encoding")
|
||||
w, h = source.size
|
||||
if min(w, h) < 1 or max(w, h) > MAX_SIDE or w * h > MAX_PIXELS:
|
||||
raise InputError("image_too_large")
|
||||
source.load()
|
||||
image = ImageOps.exif_transpose(source).convert("RGB")
|
||||
image = _geometry(image, header)
|
||||
output = BytesIO()
|
||||
image.save(output, format="JPEG", quality=92, subsampling=0)
|
||||
return output.getvalue(), image.size
|
||||
except InputError:
|
||||
raise
|
||||
except (OSError, ValueError, UnidentifiedImageError, Image.DecompressionBombError,
|
||||
Image.DecompressionBombWarning):
|
||||
raise InputError("invalid_image") from None
|
||||
|
||||
|
||||
def _classic(data: bytes, directory: Path) -> dict:
|
||||
started = time.monotonic()
|
||||
models = MODEL_HOME / "official_models"
|
||||
if any(not (models / name / "inference.pdiparams").is_file() for name in
|
||||
("PP-OCRv5_mobile_det", "latin_PP-OCRv5_mobile_rec")):
|
||||
return {"status": "error", "error": "classic_models_missing", "elapsed_ms": 0}
|
||||
image_path = directory / "same-input.jpg"
|
||||
output_path = directory / "classic.json"
|
||||
image_path.write_bytes(data)
|
||||
env = os.environ.copy()
|
||||
env.update(PADDLE_PDX_CACHE_HOME=str(MODEL_HOME),
|
||||
HF_HOME=str(MODEL_HOME / "hf"), HF_HUB_OFFLINE="1",
|
||||
TRANSFORMERS_OFFLINE="1", PADDLE_PDX_DISABLE_MODEL_SOURCE_CHECK="True",
|
||||
OMP_NUM_THREADS="1", OPENBLAS_NUM_THREADS="1", OMP_THREAD_LIMIT="1")
|
||||
try:
|
||||
process = subprocess.run(
|
||||
[sys.executable, "-m", "ocr_compare.classic_worker", str(image_path), str(output_path)],
|
||||
cwd=PROJECT, env=env, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL,
|
||||
timeout=120, check=False,
|
||||
)
|
||||
except subprocess.TimeoutExpired:
|
||||
return {"status": "error", "error": "classic_timeout", "elapsed_ms": round((time.monotonic()-started)*1000)}
|
||||
except OSError:
|
||||
return {"status": "error", "error": "classic_failed", "elapsed_ms": round((time.monotonic()-started)*1000)}
|
||||
elapsed = round((time.monotonic() - started) * 1000)
|
||||
if process.returncode or not output_path.is_file():
|
||||
return {"status": "error", "error": "classic_failed", "elapsed_ms": elapsed}
|
||||
try:
|
||||
raw = json.loads(output_path.read_text())
|
||||
lines = [{"text": row["text"], "box": row["bbox"], "source": "OCR-Zeile"}
|
||||
for row in raw["lines"] if row.get("bbox")]
|
||||
except (OSError, ValueError, KeyError, TypeError):
|
||||
return {"status": "error", "error": "classic_invalid", "elapsed_ms": elapsed}
|
||||
return {"status": "ok", "model": "PP-OCRv5 mobile (Latin), lokal",
|
||||
"elapsed_ms": elapsed, "lines": lines, "fields": assign_fields(lines),
|
||||
"regions": []}
|
||||
|
||||
|
||||
def _vision(data: bytes) -> dict:
|
||||
started = time.monotonic()
|
||||
mode = os.environ.get("OCR_VISION_MODE", "local").lower()
|
||||
if mode == "off":
|
||||
return {"status": "error", "error": "vision_disabled", "elapsed_ms": 0}
|
||||
if mode not in ("local", "ssh"):
|
||||
return {"status": "error", "error": "vision_config_error", "elapsed_ms": 0}
|
||||
if mode == "local" and any(not path.is_file() for path in (
|
||||
MODEL_HOME / "official_models" / "PP-DocLayoutV3" / "inference.pdiparams",
|
||||
MODEL_HOME / "official_models" / "PaddleOCR-VL-1.6" / "model.safetensors")):
|
||||
return {"status": "error", "error": "vision_models_missing", "elapsed_ms": 0}
|
||||
if mode == "ssh" and (not os.environ.get("OCR_VISION_SSH_TARGET") or
|
||||
not os.environ.get("OCR_VISION_REMOTE_DIR")):
|
||||
return {"status": "error", "error": "vision_config_error", "elapsed_ms": 0}
|
||||
with Image.open(BytesIO(data)) as page:
|
||||
wide = page.width > 2500 and page.width / page.height > 1.65
|
||||
limit = 90 if wide else 180
|
||||
|
||||
def run_worker(seconds: int, *, tiles_only: bool = False):
|
||||
if mode == "ssh":
|
||||
command = ["ssh", "-T", "-o", "BatchMode=yes", "-o", "StrictHostKeyChecking=yes",
|
||||
"-o", "ConnectTimeout=10"]
|
||||
identity = os.environ.get("OCR_VISION_SSH_IDENTITY")
|
||||
if identity:
|
||||
command += ["-i", str(Path(identity).expanduser()), "-o", "IdentitiesOnly=yes"]
|
||||
command += [os.environ["OCR_VISION_SSH_TARGET"],
|
||||
_remote_command(seconds, tiles_only=tiles_only)]
|
||||
env = None
|
||||
else:
|
||||
command = [os.environ.get("OCR_VISION_PYTHON", sys.executable),
|
||||
"-m", "ocr_compare.vision_worker"]
|
||||
if tiles_only:
|
||||
command.append("--tiles-only")
|
||||
env = os.environ.copy()
|
||||
env.update(OCR_MODEL_HOME=str(MODEL_HOME),
|
||||
OCR_VISION_DEVICE=os.environ.get("OCR_VISION_DEVICE", "cpu"))
|
||||
try:
|
||||
with tempfile.TemporaryDirectory(prefix="ocr-vision-") as private:
|
||||
if env is not None:
|
||||
env["OCR_PRIVATE_DIR"] = private
|
||||
return subprocess.run(command, input=data, stdout=subprocess.PIPE,
|
||||
stderr=subprocess.DEVNULL, timeout=seconds + 15,
|
||||
check=False, cwd=PROJECT if mode == "local" else None, env=env)
|
||||
except subprocess.TimeoutExpired:
|
||||
return None
|
||||
except OSError:
|
||||
return subprocess.CompletedProcess(command, 127, b"")
|
||||
|
||||
process = run_worker(limit)
|
||||
if wide and (process is None or process.returncode == 124):
|
||||
# Wide documents can form one prohibitively expensive layout block.
|
||||
# The three generic views cover the same uploaded page without a
|
||||
# document-specific field template.
|
||||
process = run_worker(90, tiles_only=True)
|
||||
elapsed = round((time.monotonic() - started) * 1000)
|
||||
if process is None or process.returncode == 124:
|
||||
return {"status": "error", "error": "vision_timeout", "elapsed_ms": elapsed}
|
||||
if process.returncode or len(process.stdout) > 2_000_000:
|
||||
return {"status": "error", "error": "vision_failed", "elapsed_ms": elapsed}
|
||||
try:
|
||||
raw = json.loads(process.stdout)
|
||||
if raw.get("error"):
|
||||
code = "vision_models_missing" if raw["error"] == "models_missing" else "vision_failed"
|
||||
return {"status": "error", "error": code, "elapsed_ms": elapsed}
|
||||
lines = raw["lines"]
|
||||
regions = raw["regions"]
|
||||
if not isinstance(lines, list) or not isinstance(regions, list):
|
||||
raise ValueError("invalid collections")
|
||||
except (ValueError, TypeError, KeyError):
|
||||
return {"status": "error", "error": "vision_invalid", "elapsed_ms": elapsed}
|
||||
return {"status": "ok", "model": raw.get("model", "PaddleOCR-VL"),
|
||||
"elapsed_ms": elapsed, "lines": lines,
|
||||
"fields": assign_fields(lines, allow_adjacent=False),
|
||||
"regions": regions, "layout_note": raw.get("layout_note", ""),
|
||||
"retry_used": bool(raw.get("retry_used", False))}
|
||||
|
||||
|
||||
async def index(_: Request):
|
||||
return FileResponse(HERE / "index.html", media_type="text/html; charset=utf-8")
|
||||
|
||||
|
||||
async def script(_: Request):
|
||||
return FileResponse(HERE / "app.js", media_type="text/javascript; charset=utf-8")
|
||||
|
||||
|
||||
async def style(_: Request):
|
||||
return FileResponse(HERE / "style.css", media_type="text/css; charset=utf-8")
|
||||
|
||||
|
||||
async def health(_: Request):
|
||||
return JSONResponse({"ready": True, "scope": "localhost only",
|
||||
"vision_mode": os.environ.get("OCR_VISION_MODE", "local")})
|
||||
|
||||
|
||||
async def compare(request: Request):
|
||||
if _slot.locked():
|
||||
return JSONResponse({"error": "busy"}, status_code=429)
|
||||
if request.headers.get("content-type", "").split(";", 1)[0] not in ("image/jpeg", "image/png"):
|
||||
return JSONResponse({"error": "unsupported_media_type"}, status_code=415)
|
||||
async with _slot:
|
||||
data = bytearray()
|
||||
try:
|
||||
async with asyncio.timeout(30):
|
||||
async for chunk in request.stream():
|
||||
data.extend(chunk)
|
||||
if len(data) > MAX_UPLOAD:
|
||||
return JSONResponse({"error": "upload_too_large"}, status_code=413)
|
||||
except TimeoutError:
|
||||
return JSONResponse({"error": "upload_timeout"}, status_code=408)
|
||||
try:
|
||||
image, (width, height) = await asyncio.to_thread(
|
||||
_prepare, bytes(data), request.headers.get("x-ocr-geometry"))
|
||||
except InputError as exc:
|
||||
return JSONResponse({"error": str(exc)}, status_code=422)
|
||||
with tempfile.TemporaryDirectory(prefix="ocr-compare-") as path:
|
||||
directory = Path(path)
|
||||
classic, vision = await asyncio.gather(
|
||||
asyncio.to_thread(_classic, image, directory),
|
||||
asyncio.to_thread(_vision, image),
|
||||
)
|
||||
# Exact JPEG bytes sent to both pipelines are returned for evidence
|
||||
# highlighting. They never enter a report, log, or external OCR API.
|
||||
import base64
|
||||
return JSONResponse({"image": "data:image/jpeg;base64," + base64.b64encode(image).decode(),
|
||||
"image_size": {"width": width, "height": height},
|
||||
"field_codes": FIELD_CODES,
|
||||
"results": {"classic": classic, "vision": vision}})
|
||||
|
||||
|
||||
app = Starlette(routes=[Route("/", index), Route("/app.js", script),
|
||||
Route("/style.css", style), Route("/api/health", health),
|
||||
Route("/api/compare", compare, methods=["POST"])])
|
||||
Reference in New Issue
Block a user