validops-east-3's picture
Deploy 72701a1d9ec33b1842d0a516ffdee753aaa754a9
efff1ff verified
Raw History Blame Contribute Delete
8.78 kB
"""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