admesh/agentic-intent-classifier
245
1from __future__ import annotations2 3import json4from dataclasses import dataclass5from functools import lru_cache6from pathlib import Path7 8import torch9from transformers import AutoTokenizer10 11try:12 from .config import ( # type: ignore13 CALIBRATION_ARTIFACTS_DIR,14 DECISION_PHASE_HEAD_CONFIG,15 INTENT_HEAD_CONFIG,16 MULTITASK_INTENT_MODEL_DIR,17 SUBTYPE_HEAD_CONFIG,18 )19 from .multitask_model import MultiTaskIntentModel, MultiTaskLabelSizes # type: ignore20except ImportError:21 from config import (22 CALIBRATION_ARTIFACTS_DIR,23 DECISION_PHASE_HEAD_CONFIG,24 INTENT_HEAD_CONFIG,25 MULTITASK_INTENT_MODEL_DIR,26 SUBTYPE_HEAD_CONFIG,27 )28 from multitask_model import MultiTaskIntentModel, MultiTaskLabelSizes29 30 31def round_score(value: float) -> float:32 return round(float(value), 4)33 34 35TASK_TO_CONFIG = {36 "intent_type": INTENT_HEAD_CONFIG,37 "intent_subtype": SUBTYPE_HEAD_CONFIG,38 "decision_phase": DECISION_PHASE_HEAD_CONFIG,39}40 41TASK_TO_LOGIT_KEY = {42 "intent_type": "intent_type_logits",43 "intent_subtype": "intent_subtype_logits",44 "decision_phase": "decision_phase_logits",45}46 47 48@dataclass(frozen=True)49class CalibrationState:50 calibrated: bool51 temperature: float52 confidence_threshold: float53 54 55class MultiTaskRuntime:56 def __init__(self, model_dir: Path):57 self.model_dir = model_dir58 self._tokenizer = None59 self._model = None60 self._metadata = None61 self._predict_batch_size = 3262 63 @property64 def metadata(self) -> dict:65 if self._metadata is None:66 metadata_path = self.model_dir / "metadata.json"67 if not metadata_path.exists():68 raise FileNotFoundError(69 f"Missing multitask metadata at {metadata_path}. Run python3 training/train_multitask_intent.py first."70 )71 self._metadata = json.loads(metadata_path.read_text(encoding="utf-8"))72 return self._metadata73 74 @property75 def tokenizer(self):76 if self._tokenizer is None:77 self._tokenizer = AutoTokenizer.from_pretrained(str(self.model_dir))78 return self._tokenizer79 80 @property81 def model(self) -> MultiTaskIntentModel:82 if self._model is None:83 weights_path = self.model_dir / "multitask_model.pt"84 if not weights_path.exists():85 raise FileNotFoundError(86 f"Missing multitask weights at {weights_path}. Run python3 training/train_multitask_intent.py first."87 )88 payload = torch.load(weights_path, map_location="cpu")89 label_sizes = MultiTaskLabelSizes(90 intent_type=len(TASK_TO_CONFIG["intent_type"].labels),91 intent_subtype=len(TASK_TO_CONFIG["intent_subtype"].labels),92 decision_phase=len(TASK_TO_CONFIG["decision_phase"].labels),93 )94 model = MultiTaskIntentModel(self.metadata["base_model_name"], label_sizes)95 model.load_state_dict(payload["state_dict"], strict=True)96 model.eval()97 self._model = model98 return self._model99 100 def _encode(self, texts: list[str], max_length: int) -> dict[str, torch.Tensor]:101 encoded = self.tokenizer(102 texts,103 return_tensors="pt",104 truncation=True,105 padding=True,106 max_length=max_length,107 )108 return {"input_ids": encoded["input_ids"], "attention_mask": encoded["attention_mask"]}109 110 def _predict_logits(self, task: str, texts: list[str]) -> torch.Tensor:111 config = TASK_TO_CONFIG[task]112 inputs = self._encode(texts, config.max_length)113 with torch.inference_mode():114 outputs = self.model(**inputs)115 return outputs[TASK_TO_LOGIT_KEY[task]]116 117 def predict_all_heads_batch(118 self, texts: list[str]119 ) -> dict[str, torch.Tensor]:120 """Single encoder pass returning logits for all three heads at once.121 122 This is the hot-path entry point. Compared with calling123 ``_predict_logits`` once per head it cuts the number of DistilBERT124 forward passes from 3 → 1, roughly halving CPU latency for a single125 query.126 127 Returns128 -------129 dict with keys ``intent_type_logits``, ``intent_subtype_logits``,130 ``decision_phase_logits`` — raw (pre-softmax) float tensors of shape131 ``(len(texts), n_classes_for_head)``.132 """133 # Use the maximum of the three head max_lengths so all heads see the134 # same truncation boundary.135 max_len = max(cfg.max_length for cfg in TASK_TO_CONFIG.values())136 inputs = self._encode(texts, max_len)137 with torch.inference_mode():138 outputs = self.model(**inputs)139 return {140 "intent_type_logits": outputs["intent_type_logits"],141 "intent_subtype_logits": outputs["intent_subtype_logits"],142 "decision_phase_logits": outputs["decision_phase_logits"],143 }144 145 146class MultiTaskHeadProxy:147 def __init__(self, task: str):148 if task not in TASK_TO_CONFIG:149 raise ValueError(f"Unsupported multitask head: {task}")150 self.task = task151 self.config = TASK_TO_CONFIG[task]152 self.runtime = get_multitask_runtime()153 self._calibration = None154 155 @property156 def tokenizer(self):157 return self.runtime.tokenizer158 159 @property160 def model(self):161 proxy = self162 163 class _TaskModelView:164 config = type("ConfigView", (), {"id2label": proxy.config.id2label})()165 166 def forward(self, input_ids=None, attention_mask=None, **kwargs):167 with torch.inference_mode():168 outputs = proxy.runtime.model(input_ids=input_ids, attention_mask=attention_mask)169 logits = outputs[TASK_TO_LOGIT_KEY[proxy.task]]170 return type("OutputView", (), {"logits": logits})()171 172 __call__ = forward173 174 return _TaskModelView()175 176 @property177 def forward_arg_names(self) -> set[str]:178 return {"input_ids", "attention_mask"}179 180 @property181 def calibration(self) -> CalibrationState:182 if self._calibration is None:183 calibrated = False184 temperature = 1.0185 confidence_threshold = self.config.default_confidence_threshold186 calibration_path = CALIBRATION_ARTIFACTS_DIR / f"{self.task}.json"187 if calibration_path.exists():188 payload = json.loads(calibration_path.read_text(encoding="utf-8"))189 calibrated = bool(payload.get("calibrated", True))190 temperature = float(payload.get("temperature", 1.0))191 confidence_threshold = float(payload.get("confidence_threshold", confidence_threshold))192 self._calibration = CalibrationState(193 calibrated=calibrated,194 temperature=max(temperature, 1e-3),195 confidence_threshold=min(max(confidence_threshold, 0.0), 1.0),196 )197 return self._calibration198 199 def _predict_probs(self, texts: list[str]) -> tuple[torch.Tensor, torch.Tensor]:200 logits = self.runtime._predict_logits(self.task, texts)201 with torch.inference_mode():202 raw_probs = torch.softmax(logits, dim=-1)203 calibrated_probs = torch.softmax(logits / self.calibration.temperature, dim=-1)204 return raw_probs, calibrated_probs205 206 def predict_probs_from_logits(207 self, logits: torch.Tensor208 ) -> tuple[torch.Tensor, torch.Tensor]:209 """Compute calibrated probs from pre-computed logits (hot-path helper).210 211 Called by ``classify_query_fused`` after a single shared encoder pass212 so that each ``MultiTaskHeadProxy`` does not re-run the encoder.213 """214 with torch.inference_mode():215 raw_probs = torch.softmax(logits, dim=-1)216 calibrated_probs = torch.softmax(logits / self.calibration.temperature, dim=-1)217 return raw_probs, calibrated_probs218 219 def predict_from_logits(220 self, logits: torch.Tensor, confidence_threshold: float | None = None221 ) -> dict:222 """Return a single prediction dict from pre-computed logits."""223 effective_threshold = (224 self.calibration.confidence_threshold225 if confidence_threshold is None226 else min(max(float(confidence_threshold), 0.0), 1.0)227 )228 raw_probs, calibrated_probs = self.predict_probs_from_logits(logits.unsqueeze(0))229 raw_row = raw_probs[0]230 calibrated_row = calibrated_probs[0]231 pred_id = int(torch.argmax(calibrated_row).item())232 confidence = float(calibrated_row[pred_id].item())233 raw_confidence = float(raw_row[pred_id].item())234 return {235 "label": self.config.id2label[pred_id],236 "confidence": round_score(confidence),237 "raw_confidence": round_score(raw_confidence),238 "confidence_threshold": round_score(effective_threshold),239 "calibrated": self.calibration.calibrated,240 "meets_confidence_threshold": confidence >= effective_threshold,241 }242 243 def predict_probs_batch(self, texts: list[str]) -> tuple[torch.Tensor, torch.Tensor]:244 if not texts:245 empty = torch.empty((0, len(self.config.labels)), dtype=torch.float32)246 return empty, empty247 raw_chunks: list[torch.Tensor] = []248 calibrated_chunks: list[torch.Tensor] = []249 for start in range(0, len(texts), self.runtime._predict_batch_size):250 batch = texts[start : start + self.runtime._predict_batch_size]251 raw, calibrated = self._predict_probs(batch)252 raw_chunks.append(raw.detach().cpu())253 calibrated_chunks.append(calibrated.detach().cpu())254 return torch.cat(raw_chunks, dim=0), torch.cat(calibrated_chunks, dim=0)255 256 def predict_batch(self, texts: list[str], confidence_threshold: float | None = None) -> list[dict]:257 if not texts:258 return []259 effective_threshold = (260 self.calibration.confidence_threshold261 if confidence_threshold is None262 else min(max(float(confidence_threshold), 0.0), 1.0)263 )264 predictions: list[dict] = []265 for start in range(0, len(texts), self.runtime._predict_batch_size):266 batch = texts[start : start + self.runtime._predict_batch_size]267 raw_probs, calibrated_probs = self._predict_probs(batch)268 for raw_row, calibrated_row in zip(raw_probs, calibrated_probs):269 pred_id = int(torch.argmax(calibrated_row).item())270 confidence = float(calibrated_row[pred_id].item())271 raw_confidence = float(raw_row[pred_id].item())272 predictions.append(273 {274 "label": self.config.id2label[pred_id],275 "confidence": round_score(confidence),276 "raw_confidence": round_score(raw_confidence),277 "confidence_threshold": round_score(effective_threshold),278 "calibrated": self.calibration.calibrated,279 "meets_confidence_threshold": confidence >= effective_threshold,280 }281 )282 return predictions283 284 def predict(self, text: str, confidence_threshold: float | None = None) -> dict:285 return self.predict_batch([text], confidence_threshold=confidence_threshold)[0]286 287 def status(self) -> dict:288 return {289 "head": self.task,290 "model_path": str(self.runtime.model_dir),291 "calibration_path": str(CALIBRATION_ARTIFACTS_DIR / f"{self.task}.json"),292 "ready": (self.runtime.model_dir / "multitask_model.pt").exists(),293 "calibrated": self.calibration.calibrated,294 }295 296 297@lru_cache(maxsize=1)298def get_multitask_runtime() -> MultiTaskRuntime:299 return MultiTaskRuntime(MULTITASK_INTENT_MODEL_DIR)300 