Spaces:
Running
Running
| """ONNX Runtime-backed GLiNER2 extraction engine (no PyTorch at inference). | |
| Runs the monolithic ``lion-ai/gliner2-base-v1-onnx`` export -- the GLiNER2 | |
| encoder + span head fused into a single ``model.onnx`` -- on ONNX Runtime. | |
| The companion ``tokenizer.json`` is loaded through the ``tokenizers`` library, | |
| so inference does not depend on torch/transformers. | |
| This engine reproduces the GLiNER2 schema semantics on top of the ONNX span | |
| head for the ``/json/feature-extract`` route: | |
| * zero-shot entity extraction (mode='entities') | |
| * structured-JSON field extraction with ``str``/``list`` dtypes and choice | |
| ("classification") fields (mode='json') | |
| Each schema group (the entity list, or each structure parent) runs as its own | |
| ONNX forward pass, because the graph expects a single ``[P]`` schema per call. | |
| Known behavioural differences vs. the PyTorch model (``fastino/gliner2-base-v1``): | |
| the export fixes the count dimension to 1, so a structure with multiple | |
| repeating objects yields only the single best object. For documents that | |
| contain exactly one object (the common case for this route) the output matches | |
| the PyTorch model. | |
| """ | |
| from __future__ import annotations | |
| import os | |
| import re | |
| from typing import Any, Dict, List, Optional, Tuple | |
| import numpy as np | |
| import onnxruntime as ort | |
| from huggingface_hub import snapshot_download | |
| from tokenizers import Tokenizer | |
| from gliner2.inference.schema import Schema | |
| from app.core.logger import get_logger | |
| logger = get_logger(__name__) | |
| MAX_WIDTH = 8 | |
| SEP_TEXT = "[SEP_TEXT]" | |
| DESC_TOKEN = "[DESCRIPTION]" | |
| P_TOKEN = "[P]" | |
| C_TOKEN = "[C]" | |
| E_TOKEN = "[E]" | |
| R_TOKEN = "[R]" | |
| L_TOKEN = "[L]" | |
| SPECIAL_TOKENS = frozenset((P_TOKEN, C_TOKEN, E_TOKEN, R_TOKEN, L_TOKEN)) | |
| # Identical to gliner2.processor.WhitespaceTokenSplitter. | |
| _WORD_RE = re.compile( | |
| r"""(?:https?://[^\s]+|www\.[^\s]+) | |
| |[a-z0-9._%+-]+@[a-z0-9.-]+\.[a-z]{2,} | |
| |@[a-z0-9_]+ | |
| |\w+(?:[-_]\w+)* | |
| |\S""", | |
| re.VERBOSE | re.IGNORECASE, | |
| ) | |
| def _parse_field_spec(spec: Any) -> Tuple[str, str, Optional[List[str]], Optional[str]]: | |
| """Parse a structure field spec into ``(name, dtype, choices, description)``. | |
| Mirrors ``gliner2``'s ``GLiNER2._parse_field_spec``. Field specs look like | |
| ``"name::dtype::choices::description"`` where all parts after ``name`` are | |
| optional; dict specs with ``name``/``dtype``/``choices``/``description`` | |
| keys are accepted too. | |
| """ | |
| if isinstance(spec, dict): | |
| return ( | |
| spec.get("name", ""), | |
| spec.get("dtype", "list"), | |
| spec.get("choices"), | |
| spec.get("description"), | |
| ) | |
| parts = spec.split("::") | |
| name = parts[0] | |
| dtype: str = "list" | |
| choices: Optional[List[str]] = None | |
| desc: Optional[str] = None | |
| dtype_explicitly_set = False | |
| for part in parts[1:]: | |
| if part in ("str", "list"): | |
| dtype = part | |
| dtype_explicitly_set = True | |
| elif part.startswith("[") and part.endswith("]"): | |
| choices = [c.strip() for c in part[1:-1].split("|")] | |
| if not dtype_explicitly_set: | |
| dtype = "str" | |
| else: | |
| desc = part | |
| return name, dtype, choices, desc | |
| class GLiNER2ONNXEngine: | |
| """Cached wrapper around the monolithic GLiNER2 ONNX model. | |
| The model and tokenizer are loaded once and reused for the life of the | |
| process. ``onnxruntime``'s ``InferenceSession.run`` is thread-safe, so the | |
| engine can serve concurrent requests (the caller bounds concurrency). | |
| """ | |
| def __init__(self, model_id: str, device: str = "cpu") -> None: | |
| self._model_id = model_id | |
| self._device = device | |
| self._session: Optional[ort.InferenceSession] = None | |
| self._tokenizer: Optional[Tokenizer] = None | |
| self._load_error: Optional[str] = None | |
| # ------------------------------------------------------------------ # | |
| # Lifecycle | |
| # ------------------------------------------------------------------ # | |
| def load(self) -> None: | |
| """Load and cache the ONNX session + tokenizer. Idempotent.""" | |
| if self._session is not None: | |
| return | |
| try: | |
| model_dir = self._resolve_model_dir() | |
| tokenizer_path = os.path.join(model_dir, "tokenizer.json") | |
| model_path = os.path.join(model_dir, "model.onnx") | |
| if not os.path.exists(tokenizer_path): | |
| raise FileNotFoundError(f"tokenizer.json not found in {model_dir}") | |
| if not os.path.exists(model_path): | |
| raise FileNotFoundError(f"model.onnx not found in {model_dir}") | |
| providers = ["CPUExecutionProvider"] | |
| if self._device == "cuda": | |
| providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] | |
| self._tokenizer = Tokenizer.from_file(tokenizer_path) | |
| self._session = ort.InferenceSession(model_path, providers=providers) | |
| self._load_error = None | |
| except Exception as exc: # noqa: BLE001 | |
| self._load_error = str(exc) | |
| logger.exception("GLiNER2 ONNX model load failed") | |
| raise | |
| def unload(self) -> None: | |
| """Release the cached session + tokenizer.""" | |
| self._session = None | |
| self._tokenizer = None | |
| self._load_error = None | |
| def is_loaded(self) -> bool: | |
| return self._session is not None | |
| def load_error(self) -> Optional[str]: | |
| return self._load_error | |
| def model_id(self) -> str: | |
| return self._model_id | |
| def device(self) -> str: | |
| return self._device | |
| def _resolve_model_dir(self) -> str: | |
| """Return a local directory containing ``model.onnx`` + ``tokenizer.json``.""" | |
| if os.path.isdir(self._model_id): | |
| return self._model_id | |
| return snapshot_download(self._model_id) | |
| def _ensure_loaded(self) -> None: | |
| if self._session is None: | |
| self.load() | |
| # ------------------------------------------------------------------ # | |
| # Inference | |
| # ------------------------------------------------------------------ # | |
| def extract_entities(self, text: str, labels: List[str], threshold: float = 0.5) -> Dict[str, Any]: | |
| """Zero-shot entity extraction. Returns ``{"entities": {label: [texts]}}``.""" | |
| self._ensure_loaded() | |
| result: Dict[str, List[str]] = {label: [] for label in labels} | |
| text = self._normalize_text(text) | |
| words = self._split_words(text) | |
| if not words: | |
| return {"entities": result} | |
| word_strings = [w for w, _, _ in words] | |
| start_map = [w[1] for w in words] | |
| end_map = [w[2] for w in words] | |
| schema_tokens = self._build_schema_tokens("entities", list(labels), E_TOKEN, {}) | |
| feeds = self._build_feeds(schema_tokens, word_strings) | |
| grid = self._session.run(None, feeds)[0][0] | |
| text_len = len(word_strings) | |
| for i, label in enumerate(labels): | |
| spans = self._find_spans(grid[i], threshold, text_len, text, start_map, end_map) | |
| result[label] = self._format_spans(spans) | |
| return {"entities": result} | |
| def extract_json(self, text: str, structure: Dict[str, Any], threshold: float = 0.5) -> Dict[str, Any]: | |
| """Structured JSON extraction. Returns ``{parent: [object]}``.""" | |
| self._ensure_loaded() | |
| if not structure: | |
| return {} | |
| schema_obj = self._build_structure_schema(structure) | |
| schema = schema_obj.build() | |
| prefix = self._build_classification_prefix(schema) | |
| text = self._normalize_text(text) | |
| words = self._split_words(text) | |
| if not words: | |
| return {parent: [] for parent in schema.get("json_structures", []) for parent in parent} | |
| word_strings = [w for w, _, _ in words] | |
| start_map = [w[1] for w in words] | |
| end_map = [w[2] for w in words] | |
| text_tokens = prefix + word_strings | |
| num_prefix = len(prefix) | |
| text_len = len(word_strings) | |
| json_descs = schema.get("json_descriptions", {}) | |
| result: Dict[str, Any] = {} | |
| for struct in schema.get("json_structures", []): | |
| for parent, fields in struct.items(): | |
| field_names = list(fields.keys()) | |
| schema_tokens = self._build_schema_tokens( | |
| parent, field_names, C_TOKEN, json_descs.get(parent, {}) | |
| ) | |
| feeds = self._build_feeds(schema_tokens, text_tokens) | |
| grid = self._session.run(None, feeds)[0][0] | |
| obj: Dict[str, Any] = {} | |
| for fname in schema_obj._field_orders.get(parent, field_names): | |
| fidx = field_names.index(fname) | |
| full = grid[fidx] | |
| meta = schema_obj._field_metadata.get(f"{parent}.{fname}", {}) | |
| dtype = meta.get("dtype", "list") | |
| field_threshold = ( | |
| meta.get("threshold") if meta.get("threshold") is not None else threshold | |
| ) | |
| fval = fields[fname] | |
| if isinstance(fval, dict) and "choices" in fval: | |
| obj[fname] = self._decode_choice_field( | |
| full, prefix, fval["choices"], num_prefix, field_threshold, dtype | |
| ) | |
| else: | |
| scores = full[num_prefix:num_prefix + text_len] | |
| spans = self._find_spans(scores, field_threshold, text_len, text, start_map, end_map) | |
| if dtype == "list": | |
| obj[fname] = self._dedup_texts(self._format_spans(spans)) | |
| else: | |
| obj[fname] = spans[0][0] if spans else None | |
| instances = [obj] if any(v is not None and v != [] for v in obj.values()) else [] | |
| result[parent] = instances | |
| return result | |
| # ------------------------------------------------------------------ # | |
| # Schema construction (mirrors gliner2.processor / gliner2 engine) | |
| # ------------------------------------------------------------------ # | |
| def _normalize_text(text: str) -> str: | |
| """Append a period so the model sees the same input gliner2 would.""" | |
| if text and not text.endswith((".", "!", "?")): | |
| return text + "." | |
| return text if text else "." | |
| def _split_words(text: str) -> List[Tuple[str, int, int]]: | |
| """Split text into lowercased ``(word, char_start, char_end)`` tuples.""" | |
| lowered = text.lower() | |
| return [(m.group(), m.start(), m.end()) for m in _WORD_RE.finditer(lowered)] | |
| def _build_structure_schema(self, structure: Dict[str, Any]) -> Schema: | |
| """Build a gliner2 ``Schema`` object from a route ``structure`` dict.""" | |
| schema = Schema() | |
| for parent, specs in structure.items(): | |
| builder = schema.structure(parent) | |
| for spec in specs: | |
| name, dtype, choices, desc = _parse_field_spec(spec) | |
| builder.field(name, dtype=dtype, choices=choices, description=desc) | |
| builder._auto_finish() | |
| return schema | |
| def _build_classification_prefix(schema: Dict[str, Any]) -> List[str]: | |
| """Build the classification (choice) prefix tokens prepended to the text.""" | |
| prefix_tokens: List[str] = [] | |
| for struct in schema.get("json_structures", []): | |
| for parent, fields in struct.items(): | |
| cls_fields = [ | |
| (fname, fval) for fname, fval in fields.items() | |
| if isinstance(fval, dict) and "value" in fval and "choices" in fval | |
| ] | |
| inner: List[str] = [] | |
| for fname, fval in cls_fields: | |
| choices = fval["choices"] | |
| choice_tokens: List[str] = [] | |
| for i, c in enumerate(choices): | |
| if i > 0: | |
| choice_tokens.append("|") | |
| choice_tokens.append(c) | |
| inner.extend([fname, "("] + choice_tokens + [")", ","]) | |
| if inner: | |
| inner = inner[:-1] | |
| prefix_tokens.extend(["(", f"{parent}:", *inner, ")"]) | |
| return prefix_tokens | |
| def _build_schema_tokens( | |
| parent: str, | |
| fields: List[str], | |
| child_token: str, | |
| descriptions: Dict[str, str], | |
| ) -> List[str]: | |
| """Build a schema token sequence like gliner2's ``_transform_schema``.""" | |
| prompt_str = parent | |
| descs = [(lbl, d) for lbl, d in descriptions.items() if lbl in fields] | |
| for lbl, d in descs: | |
| prompt_str += f" {DESC_TOKEN} {lbl}: {d}" | |
| tokens = ["(", P_TOKEN, prompt_str, "("] | |
| for field in fields: | |
| tokens.extend([child_token, field]) | |
| tokens.extend([")", ")"]) | |
| return tokens | |
| # ------------------------------------------------------------------ # | |
| # ONNX inputs | |
| # ------------------------------------------------------------------ # | |
| def _build_feeds(self, schema_tokens: List[str], text_tokens: List[str]) -> Dict[str, np.ndarray]: | |
| """Build the five ONNX input tensors for one schema + text pair.""" | |
| combined = schema_tokens + [SEP_TEXT] + text_tokens | |
| encoding = self._tokenizer.encode(combined, is_pretokenized=True, add_special_tokens=False) | |
| token_ids = encoding.ids | |
| word_ids = encoding.word_ids | |
| input_ids = np.array([token_ids], dtype=np.int64) | |
| attention_mask = np.ones((1, len(token_ids)), dtype=np.int64) | |
| # Number of words before the text segment: schema words + [SEP_TEXT]. | |
| num_schema_words = len(schema_tokens) + 1 | |
| word_first: Dict[int, int] = {} | |
| for tok_pos, wid in enumerate(word_ids): | |
| if wid is not None and wid not in word_first: | |
| word_first[wid] = tok_pos | |
| text_positions = np.array( | |
| [word_first[num_schema_words + i] for i in range(len(text_tokens))], dtype=np.int64 | |
| ) | |
| schema_positions = np.array( | |
| [word_first[i] for i, tok in enumerate(schema_tokens) if tok in SPECIAL_TOKENS], | |
| dtype=np.int64, | |
| ) | |
| num_words = len(text_positions) | |
| spans = [] | |
| for start in range(num_words): | |
| for width in range(1, MAX_WIDTH + 1): | |
| end = start + width | |
| if end <= num_words: | |
| spans.append((start, end - 1)) | |
| else: | |
| spans.append((0, 0)) | |
| span_idx = np.array(spans, dtype=np.int64).reshape(1, -1, 2) | |
| return { | |
| "input_ids": input_ids, | |
| "attention_mask": attention_mask, | |
| "text_positions": text_positions, | |
| "schema_positions": schema_positions, | |
| "span_idx": span_idx, | |
| } | |
| # ------------------------------------------------------------------ # | |
| # Decoding (mirrors gliner2's span/choice decoding on the ONNX scores) | |
| # ------------------------------------------------------------------ # | |
| def _find_spans( | |
| scores: np.ndarray, | |
| threshold: float, | |
| text_len: int, | |
| text: str, | |
| start_map: List[int], | |
| end_map: List[int], | |
| ) -> List[Tuple[str, float, int, int]]: | |
| """Find spans above threshold. Returns ``(text, confidence, start, end)``.""" | |
| valid = np.argwhere(scores >= threshold) | |
| spans: List[Tuple[str, float, int, int]] = [] | |
| for start, width in valid: | |
| end = start + width + 1 | |
| if 0 <= start < text_len and end <= text_len: | |
| try: | |
| char_start = start_map[start] | |
| char_end = end_map[end - 1] | |
| text_span = text[char_start:char_end].strip() | |
| except (IndexError, KeyError): | |
| continue | |
| if text_span: | |
| spans.append((text_span, float(scores[start, width]), char_start, char_end)) | |
| return spans | |
| def _format_spans(spans: List[Tuple[str, float, int, int]]) -> List[str]: | |
| """Score-sort + overlap removal, matching gliner2's ``_format_spans``.""" | |
| if not spans: | |
| return [] | |
| sorted_spans = sorted(spans, key=lambda x: x[1], reverse=True) | |
| selected: List[Tuple[str, float, int, int]] = [] | |
| for text, conf, start, end in sorted_spans: | |
| overlap = any(not (end <= s[2] or start >= s[3]) for s in selected) | |
| if not overlap: | |
| selected.append((text, conf, start, end)) | |
| return [s[0] for s in selected] | |
| def _dedup_texts(values: List[str]) -> List[str]: | |
| """Dedup a list of extracted texts (case-insensitive), matching gliner2.""" | |
| unique: List[str] = [] | |
| seen = set() | |
| for v in values: | |
| if v and v.lower() not in seen: | |
| seen.add(v.lower()) | |
| unique.append(v) | |
| return unique | |
| def _find_choice_idx(choice: str, tokens: List[str]) -> int: | |
| """Index of a choice token inside the classification prefix tokens.""" | |
| choice_lower = choice.lower() | |
| for i, tok in enumerate(tokens): | |
| if tok.lower() == choice_lower or choice_lower in tok.lower(): | |
| return i | |
| return -1 | |
| def _decode_choice_field( | |
| self, | |
| full: np.ndarray, | |
| prefix: List[str], | |
| choices: List[str], | |
| num_prefix: int, | |
| threshold: float, | |
| dtype: str, | |
| ) -> Any: | |
| """Score the classification-prefix choice tokens (gliner2 semantics).""" | |
| prefix_scores = full[:num_prefix, 0] | |
| if dtype == "list": | |
| selected: List[str] = [] | |
| seen = set() | |
| for choice in choices: | |
| if choice in seen: | |
| continue | |
| idx = self._find_choice_idx(choice, prefix) | |
| if 0 <= idx < len(prefix_scores): | |
| score = float(prefix_scores[idx]) | |
| if score >= threshold: | |
| selected.append(choice) | |
| seen.add(choice) | |
| return selected | |
| best: Optional[str] = None | |
| best_score = -1.0 | |
| for choice in choices: | |
| idx = self._find_choice_idx(choice, prefix) | |
| if 0 <= idx < len(prefix_scores): | |
| score = float(prefix_scores[idx]) | |
| if score > best_score: | |
| best_score = score | |
| best = choice | |
| return best if best and best_score >= threshold else None |