# file: hybrid_ner/llm_disambiguator.py
from __future__ import annotations
import requests
import subprocess
import json
import time
import re
from typing import List, Optional, Dict, Any
from dataclasses import dataclass
from urllib.parse import urljoin
# Import your project's LabelHypothesis, CandidateSpan, SchemaIndex
from .models import LabelHypothesis, CandidateSpan
from .schema import SchemaIndex
@dataclass
[docs]
class LLMConfig:
[docs]
use_cli: bool = False # Use Ollama CLI (subprocess) instead of HTTP
[docs]
cli_binary: str = "ollama" # CLI binary name
[docs]
http_url: str = "http://localhost:11434/v1/chat/completions" # Example Ollama HTTP endpoint
[docs]
model: str = "ollama/gpt-oss:20B" # Model name to call
[docs]
temperature: float = 0.0 # deterministic
[docs]
stop_sequences: Optional[List[str]] = None
# Controls
[docs]
min_confidence: float = 0.15 # attach hypotheses with at least this soft score (if LLM emits a score)
[docs]
cache_ttl_seconds: int = 3600
[docs]
dry_run: bool = False # if true, do not call the model; return NO_LABEL
# Confidence-band trigger for disambiguation.
# Spans with a best hypothesis score in [uncertain_score_floor, uncertain_score_ceiling]
# are sent to the LLM; outside this range they are accepted or dropped without LLM involvement.
[docs]
uncertain_score_floor: float = 0.40 # below this the span is too weak even for LLM
[docs]
uncertain_score_ceiling: float = 0.65 # above this the embed result is confident enough
[docs]
high_conf_bypass_score: float = 0.85 # gazetteer/high-conf hit at or above this bypasses LLM
[docs]
class LLMDisambiguator:
def __init__(self, schema: SchemaIndex, config: LLMConfig = LLMConfig()):
[docs]
self._cache: Dict[str, Dict[str, Any]] = {} # key -> {"time":ts, "resp":...}
[docs]
self.llm_ok = self.health_check()
if not self.llm_ok:
return
[docs]
def should_call(self, c: CandidateSpan) -> bool:
"""Decide whether to invoke the LLM for this candidate span.
Decision logic (in priority order):
1. High-confidence gazetteer hit exists → bypass LLM entirely (trust deterministic match).
2. No proposed labels at all → ask LLM (span has no evidence).
3. All hypotheses lack a known schema group → ask LLM (labels are unrecognised).
4. Multiple competing schema groups → ask LLM (genuine ambiguity).
5. Best score in the uncertain confidence band → ask LLM (embed result is marginal).
6. Best score below the floor → do NOT ask LLM (span is too weak; LLM would speculate).
7. Best score above the ceiling → do NOT ask LLM (embed is confident enough).
"""
proposed = list(getattr(c, "proposed_labels", None) or [])
# 1. High-confidence bypass: at least one deterministic (gazetteer) hit above threshold.
high_conf = self.config.high_conf_bypass_score
if any(
getattr(h, "score", 0.0) >= high_conf
and "gazetteer" in str(getattr(h, "rationale", ""))
for h in proposed
):
return False
# 2. No labels at all.
if not proposed:
return True
# 3. All hypotheses lack a known schema group.
if all(getattr(h, "group", None) is None for h in proposed):
return True
# 4. Ambiguity across multiple schema groups.
groups = {getattr(h, "group", None) for h in proposed if getattr(h, "group", None)}
if len(groups) > 1:
return True
# 5 & 6 & 7. Confidence-band gate on the best available score.
best_score = max((getattr(h, "score", 0.0) for h in proposed), default=0.0)
floor = self.config.uncertain_score_floor
ceiling = self.config.uncertain_score_ceiling
return floor <= best_score <= ceiling
# Build a concise prompt. Keep it short and constrained.
[docs]
def _build_prompt(self, doc_text: str, c: CandidateSpan, candidate_labels: List[str]) -> str:
ctx = doc_text[max(0, c.start - 200): c.end + 200] # short surrounding context
label_info = "\n".join([f"- {lbl}: {self._label_short(lbl)}" for lbl in candidate_labels])
prompt = f"""
You are a domain expert restricted to choosing one label from a provided list.
Context: "{ctx}"
Span: "{c.text}"
Possible labels:
{label_info}
Instruction: Choose exactly one label id from the list above that best fits the span in this context, or reply EXACTLY "NO_LABEL" if none apply. Answer with a JSON object only:
{{"label":"<label_id_or_NO_LABEL>", "score":<0.0-1.0 optional>, "rationale":"<short explanation>"}}
Do NOT invent labels. Do NOT provide additional text.
"""
return prompt.strip()
[docs]
def _label_short(self, lbl: str) -> str:
# Provide a short description from schema if available (fallback to label)
# The schema in your repo stores descriptions in the original json loaded by DescriptionEmbedGenerator.
# If not present, return label.
entry = getattr(self.schema, "label_descriptions", {}) or {}
if lbl in entry:
return entry[lbl].get("short_description") or entry[lbl].get("name") or lbl
return lbl
[docs]
def _cache_key(self, c: CandidateSpan) -> str:
return f"{c.doc_id}:{c.start}:{c.end}:{c.text}"
# Main entry: disambiguate a list of candidates, augment their proposed_labels with llm hypotheses
[docs]
def disambiguate(self, doc_text: str, candidates: List[CandidateSpan]) -> None:
if not self.llm_ok:
return
for c in candidates:
if not self.should_call(c):
continue
key = self._cache_key(c)
now = time.time()
if key in self._cache and now - self._cache[key]["time"] < self.config.cache_ttl_seconds:
resp = self._cache[key]["resp"]
else:
resp = self._call_llm_for_candidate(doc_text, c)
self._cache[key] = {"time": now, "resp": resp}
# parse & attach
hypothesis = self._parse_llm_response(resp, c)
if hypothesis:
# Only attach if label is known in schema and clears min_confidence
if (hypothesis.label in self.schema.label_to_group
and getattr(hypothesis, "score", 0.0) >= self.config.min_confidence):
# assign group
hypothesis.group = self.schema.label_to_group[hypothesis.label]
c.proposed_labels.append(hypothesis)
[docs]
def _call_llm_for_candidate(self, doc_text: str, c: CandidateSpan) -> Optional[Dict[str, Any]]:
if self.config.dry_run:
return {"label": "NO_LABEL", "score": 0.0, "rationale": "dry_run"}
# Candidate label set: use schema labels that are plausible (all labels in schema)
candidate_labels = list(self.schema.label_to_group.keys()) # or narrow set if desired
prompt = self._build_prompt(doc_text, c, candidate_labels)
if self.config.use_cli:
# `ollama run MODEL PROMPT` prints the model output to stdout.
cmd = [self.config.cli_binary, "run", self.config.model, prompt]
try:
p = subprocess.run(cmd, capture_output=True, text=True, timeout=self.config.timeout)
text = p.stdout.strip()
parsed = self._extract_json_object(text)
return parsed if parsed is not None else {"label": "NO_LABEL", "score": 0.0, "rationale": "empty"}
except Exception as e:
return {"label": "NO_LABEL", "score": 0.0, "rationale": f"error:{e}"}
else:
try:
url = (self.config.http_url or "").strip()
# OpenAI-compatible chat endpoint
if url.endswith("/chat/completions"):
payload = {
"model": self.config.model,
"messages": [
{"role": "system", "content": "Return ONLY valid JSON. No extra text."},
{"role": "user", "content": prompt},
],
"temperature": self.config.temperature,
"max_tokens": self.config.max_tokens,
}
r = requests.post(url, json=payload, timeout=self.config.timeout)
r.raise_for_status()
raw = r.json()
content = (
raw.get("choices", [{}])[0]
.get("message", {})
.get("content", "")
)
parsed = self._extract_json_object(content)
return parsed if parsed is not None else {"label": "NO_LABEL", "score": 0.0, "rationale": "unparseable_chat_json"}
# OpenAI-compatible legacy completions endpoint
if url.endswith("/completions"):
payload = {
"model": self.config.model,
"prompt": prompt,
"temperature": self.config.temperature,
"max_tokens": self.config.max_tokens,
}
r = requests.post(url, json=payload, timeout=self.config.timeout)
r.raise_for_status()
raw = r.json()
text = (raw.get("choices", [{}])[0].get("text", "") or "")
parsed = self._extract_json_object(text)
return parsed if parsed is not None else {"label": "NO_LABEL", "score": 0.0, "rationale": "unparseable_completion_json"}
# Fallback: unknown endpoint; try prompt-style (your old behavior)
payload = {
"model": self.config.model,
"prompt": prompt,
"temperature": self.config.temperature,
"max_tokens": self.config.max_tokens,
}
r = requests.post(url, json=payload, timeout=self.config.timeout)
r.raise_for_status()
return r.json()
except Exception as e:
return {"label": "NO_LABEL", "score": 0.0, "rationale": f"error:{e}"}
[docs]
def _parse_llm_response(self, resp: Any, c: CandidateSpan) -> Optional[LabelHypothesis]:
# Expected resp is a dict containing {"label":..., "score":..., "rationale":...}
if not isinstance(resp, dict):
# attempt hydration if resp is string
try:
resp = json.loads(str(resp))
except Exception:
return None
label = resp.get("label") or resp.get("label_id")
if not label or label == "NO_LABEL":
return None
try:
score = float(resp.get("score", 0.5))
except (TypeError, ValueError):
# malformed score from the LLM (e.g. null, "high"); skip this hypothesis
return None
rationale = str(resp.get("rationale", "")).strip()
# Build a LabelHypothesis (adapt to your models.py constructor)
lh = LabelHypothesis(label=label, score=score, rationale=rationale)
return lh
[docs]
def health_check(self) -> bool:
try:
# derive base API URL from configured endpoint
base = self.config.http_url.rsplit("/", 2)[0] # -> http://localhost:11434/v1
models_url = urljoin(base + "/", "models")
r = requests.get(models_url, timeout=3)
r.raise_for_status()
models = [m["id"] for m in r.json().get("data", [])]
if self.config.model not in models:
raise RuntimeError(
f"Model '{self.config.model}' not found. Available: {models}"
)
return True
except Exception as e:
import logging
logging.error(f"LLM health check failed: {e}")
return False