Spaces:
Running
Running
Download docqa/docxextract/engine.py from validops-east-3/instance-2: direct link, hf CLI and curl.
- Browser
- Download file 8.78 kB
-
https://huggingface.co/spaces/validops-east-3/instance-2/resolve/main/docqa/docxextract/engine.py
- Command line
-
hf download hf://spaces/validops-east-3/instance-2/docqa/docxextract/engine.py
-
curl -L -o engine.py https://huggingface.co/spaces/validops-east-3/instance-2/resolve/main/docqa/docxextract/engine.py
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__) | |
| 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") | |
| # ------------------------------------------------------------------ | |
| 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 | |