"""LayoutLM inference: batched, thread-safe, warmed. Design notes ------------ **Batching is the main optimisation.** The naive pattern -- one pipeline call per key -- re-encodes the same page for every key. LayoutLM's encoder is bidirectional and the question is interleaved with the document, so nothing is shareable across calls. Passing N keys in one call instead lets the CPU GEMMs amortise weight reads. Measured gains are in ``bench/bench_batching.py``. **The model has no null-answer head.** It will always return *some* span for any input, including questions whose answer is not on the page. Absence therefore has to be decided by the caller using the returned score, which is what ``confidence_threshold`` in the service layer does. This is a genuine limitation of the architecture, not something that can be tuned away. **Thread safety.** A single module instance is guarded by a lock. Torch's CPU kernels are internally parallel, so concurrent Python threads around one module buy nothing and risk races in the tokenizer; batching is the correct lever. """ from __future__ import annotations import os import threading import time from dataclasses import dataclass from typing import Sequence from .config import Settings from .logging import get_logger from .schema import ModelUnavailableError log = get_logger(__name__) @dataclass(frozen=True, slots=True) class Span: """A candidate answer span for one key on one page.""" key: str answer: str confidence: float page: int start: int end: int class Engine: """Thread-safe wrapper around the HF document-question-answering pipeline.""" def __init__(self, settings: Settings): self.settings = settings self._lock = threading.Lock() self._pipe = None self.load_seconds: float = 0.0 self.warm: bool = False # ------------------------------------------------------------------ def load(self) -> None: if self._pipe is not None: return try: import torch from transformers import pipeline torch.set_num_threads(self.settings.torch_threads) t0 = time.perf_counter() pipe = pipeline( "document-question-answering", model=self.settings.model_id, device=self.settings.device, ) pipe.model.eval() self.load_seconds = time.perf_counter() - t0 self._pipe = pipe log.info("engine loaded", extra={ "detail": f"{self.settings.model_id} in {self.load_seconds:.2f}s", }) except Exception as exc: # noqa: BLE001 raise ModelUnavailableError("model could not be loaded", detail=str(exc)) from exc def warmup(self) -> None: """Run one throwaway inference so the first real request is not slow. Without this, the first request pays for lazy kernel selection and allocator growth, inflating its latency and skewing p95. """ if self.warm: return self.load() with self._lock: self._pipe(image=None, question="What is the date?", word_boxes=[["Invoice", [0, 0, 100, 20]], ["Date", [110, 0, 190, 20]], ["2025", [200, 0, 260, 20]]], top_k=1) self.warm = True log.info("engine warmed") # ------------------------------------------------------------------ @staticmethod def _unwrap(obj): """Flatten the pipeline's nested list output to a list of dicts. transformers 4.50 returns `[[{...}]]` for batched input even though the docstring advertises a bare dict, and the nesting depth has changed between releases. Normalising here keeps callers version-agnostic. """ out: list[dict] = [] stack = [obj] while stack: item = stack.pop(0) if isinstance(item, dict): out.append(item) elif isinstance(item, (list, tuple)): stack.extend(item) return out def answer_batch(self, page_words: Sequence[str], page_boxes: Sequence[Sequence[int]], keys: Sequence[str], page_index: int = 1) -> list[Span]: """Answer every key against one page in a batched forward pass. Returns one Span per key, in the order the keys were given. """ self.load() word_boxes = [[w, list(b)] for w, b in zip(page_words, page_boxes)] # The pipeline is given `image=None` so it uses our pre-computed # word_boxes instead of running its own Tesseract pass. `image` is a # required positional parameter in transformers' implementation. payload = [{"image": None, "question": k, "word_boxes": word_boxes} for k in keys] batch = self.settings.batch_size results: list[dict] = [] with self._lock: for start in range(0, len(payload), batch): chunk = payload[start:start + batch] out = self._pipe(chunk, top_k=1) results.extend(self._unwrap(out)) spans: list[Span] = [] for key, res in zip(keys, results): spans.append(Span( key=key, answer=str(res.get("answer", "")).strip(), confidence=float(res.get("score", 0.0) or 0.0), page=page_index, start=int(res.get("start", -1)), end=int(res.get("end", -1)), )) return spans # ------------------------------------------------------------------ def extract(self, pages, keys: Sequence[str], max_pages: int) -> list[Span]: """Answer every key, choosing the best-scoring page for each. Multi-page documents are scored page by page and the highest-confidence span wins. Keys that only appear on page 3 would be missed by a first-page-only strategy. """ best: dict[str, Span] = {} limit = min(max_pages, len(pages)) for page_index in range(limit): page = pages[page_index] if not page.words: # An empty page cannot contain the answer; asking the model # anyway only invites fabrication. for key in keys: best.setdefault(key, Span(key, "", 0.0, page_index + 1, -1, -1)) continue for span in self.answer_batch(page.words, page.boxes, keys, page_index + 1): current = best.get(span.key) if current is None or span.confidence > current.confidence: best[span.key] = span return [best[k] for k in keys if k in best] _engine: Engine | None = None _engine_lock = threading.Lock() def get_engine(settings: Settings) -> Engine: """Return the configured engine, preferring ONNX Runtime when available. Set ``DOCX_BACKEND=torch`` to force PyTorch. The ONNX path is selected automatically only if its artefacts exist; it was measured at ~30% lower latency with identical answers at fp32. """ global _engine if _engine is None: with _engine_lock: if _engine is None: backend = os.environ.get("DOCX_BACKEND", "auto").lower() if backend in ("auto", "onnx"): try: from .onnx_engine import OnnxEngine, default_onnx_path if default_onnx_path(settings.model_id).exists(): _engine = OnnxEngine(settings) # type: ignore[assignment] log.info("using ONNX Runtime backend") return _engine if backend == "onnx": raise ModelUnavailableError( "ONNX backend requested but artefacts are missing", detail=f"expected {default_onnx_path(settings.model_id)}", ) except FileNotFoundError as exc: if backend == "onnx": raise ModelUnavailableError(str(exc)) from exc log.info("ONNX artefacts absent, falling back to torch") _engine = Engine(settings) return _engine def reset_engine() -> None: """Drop the engine. Used by tests that swap models or settings.""" global _engine with _engine_lock: _engine = None