admesh/agentic-intent-classifier
245
1from __future__ import annotations2 3import json4import os5import inspect6from dataclasses import dataclass7from pathlib import Path8from functools import lru_cache9 10os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")11 12import torch13from transformers import AutoModelForSequenceClassification, AutoTokenizer14 15try:16 from .config import HEAD_CONFIGS, HeadConfig, _looks_like_local_hf_model_dir # type: ignore17 from .multitask_runtime import MultiTaskHeadProxy # type: ignore18except ImportError:19 from config import HEAD_CONFIGS, HeadConfig, _looks_like_local_hf_model_dir20 from multitask_runtime import MultiTaskHeadProxy21 22_TRAIN_SCRIPT_HINTS: dict[str, str] = {23 "intent_type": "python3 training/train.py",24 "decision_phase": "python3 training/train_decision_phase.py",25 "intent_subtype": "python3 training/train_subtype.py",26 "iab_content": "python3 training/train_iab.py",27}28 29 30def _resolved_model_dir(config: HeadConfig) -> Path:31 return Path(config.model_dir).expanduser().resolve()32 33 34def _missing_head_weights_message(config: HeadConfig) -> str:35 path = _resolved_model_dir(config)36 train_hint = _TRAIN_SCRIPT_HINTS.get(37 config.slug,38 "See the `training/` directory for the matching `train_*.py` script.",39 )40 return (41 f"Classifier weights for head '{config.slug}' are missing or incomplete at {path}. "42 f"Expected a Hugging Face model directory with config.json and "43 f"model.safetensors (or pytorch_model.bin), plus tokenizer files. "44 f"From the `agentic-intent-classifier` directory, run: {train_hint}. "45 f"Note: training only `train_iab.py` does not populate `model_output`; "46 f"full `classify_query` / evaluation also needs the intent, subtype, and decision-phase heads."47 )48 49 50def round_score(value: float) -> float:51 return round(float(value), 4)52 53 54@dataclass(frozen=True)55class CalibrationState:56 calibrated: bool57 temperature: float58 confidence_threshold: float59 60 61class SequenceClassifierHead:62 def __init__(self, config: HeadConfig):63 self.config = config64 self._tokenizer = None65 self._model = None66 self._calibration = None67 self._predict_batch_size = 3268 self._forward_arg_names = None69 70 def _weights_dir(self) -> Path:71 return _resolved_model_dir(self.config)72 73 def _require_local_weights(self) -> Path:74 weights_dir = self._weights_dir()75 if not _looks_like_local_hf_model_dir(weights_dir):76 raise FileNotFoundError(_missing_head_weights_message(self.config))77 return weights_dir78 79 @property80 def tokenizer(self):81 if self._tokenizer is None:82 weights_dir = self._require_local_weights()83 self._tokenizer = AutoTokenizer.from_pretrained(str(weights_dir))84 return self._tokenizer85 86 @property87 def model(self):88 if self._model is None:89 weights_dir = self._require_local_weights()90 alt = weights_dir / "iab_weights.safetensors"91 canonical = weights_dir / "model.safetensors"92 if alt.exists() and not canonical.exists():93 os.symlink(str(alt), str(canonical))94 self._model = AutoModelForSequenceClassification.from_pretrained(str(weights_dir))95 self._model.eval()96 return self._model97 98 @property99 def forward_arg_names(self) -> set[str]:100 if self._forward_arg_names is None:101 self._forward_arg_names = set(inspect.signature(self.model.forward).parameters)102 return self._forward_arg_names103 104 @property105 def calibration(self) -> CalibrationState:106 if self._calibration is None:107 calibrated = False108 temperature = 1.0109 confidence_threshold = self.config.default_confidence_threshold110 if self.config.calibration_path.exists():111 payload = json.loads(self.config.calibration_path.read_text())112 calibrated = bool(payload.get("calibrated", True))113 temperature = float(payload.get("temperature", 1.0))114 confidence_threshold = float(115 payload.get("confidence_threshold", self.config.default_confidence_threshold)116 )117 self._calibration = CalibrationState(118 calibrated=calibrated,119 temperature=max(temperature, 1e-3),120 confidence_threshold=min(max(confidence_threshold, 0.0), 1.0),121 )122 return self._calibration123 124 def status(self) -> dict:125 weights_dir = self._weights_dir()126 return {127 "head": self.config.slug,128 "model_path": str(weights_dir),129 "calibration_path": str(self.config.calibration_path),130 "ready": _looks_like_local_hf_model_dir(weights_dir),131 "calibrated": self.calibration.calibrated,132 }133 134 def _encode(self, texts: list[str]):135 encoded = self.tokenizer(136 texts,137 return_tensors="pt",138 truncation=True,139 padding=True,140 max_length=self.config.max_length,141 )142 return {143 key: value144 for key, value in encoded.items()145 if key in self.forward_arg_names146 }147 148 def _predict_probs(self, texts: list[str]) -> tuple[torch.Tensor, torch.Tensor]:149 inputs = self._encode(texts)150 with torch.inference_mode():151 outputs = self.model(**inputs)152 raw_probs = torch.softmax(outputs.logits, dim=-1)153 calibrated_probs = torch.softmax(outputs.logits / self.calibration.temperature, dim=-1)154 return raw_probs, calibrated_probs155 156 def predict_probs_batch(self, texts: list[str]) -> tuple[torch.Tensor, torch.Tensor]:157 if not texts:158 empty = torch.empty((0, len(self.config.labels)), dtype=torch.float32)159 return empty, empty160 raw_chunks: list[torch.Tensor] = []161 calibrated_chunks: list[torch.Tensor] = []162 for start in range(0, len(texts), self._predict_batch_size):163 batch_texts = texts[start : start + self._predict_batch_size]164 raw_probs, calibrated_probs = self._predict_probs(batch_texts)165 raw_chunks.append(raw_probs.detach().cpu())166 calibrated_chunks.append(calibrated_probs.detach().cpu())167 return torch.cat(raw_chunks, dim=0), torch.cat(calibrated_chunks, dim=0)168 169 def predict_batch(self, texts: list[str], confidence_threshold: float | None = None) -> list[dict]:170 if not texts:171 return []172 173 effective_threshold = (174 self.calibration.confidence_threshold175 if confidence_threshold is None176 else min(max(float(confidence_threshold), 0.0), 1.0)177 )178 predictions: list[dict] = []179 180 for start in range(0, len(texts), self._predict_batch_size):181 batch_texts = texts[start : start + self._predict_batch_size]182 raw_probs, calibrated_probs = self._predict_probs(batch_texts)183 for raw_row, calibrated_row in zip(raw_probs, calibrated_probs):184 pred_id = int(torch.argmax(calibrated_row).item())185 confidence = float(calibrated_row[pred_id].item())186 raw_confidence = float(raw_row[pred_id].item())187 predictions.append(188 {189 "label": self.model.config.id2label[pred_id],190 "confidence": round_score(confidence),191 "raw_confidence": round_score(raw_confidence),192 "confidence_threshold": round_score(effective_threshold),193 "calibrated": self.calibration.calibrated,194 "meets_confidence_threshold": confidence >= effective_threshold,195 }196 )197 return predictions198 199 def predict_candidate_batch(200 self,201 texts: list[str],202 candidate_labels: list[list[str]],203 confidence_threshold: float | None = None,204 ) -> list[dict]:205 if not texts:206 return []207 if len(texts) != len(candidate_labels):208 raise ValueError("texts and candidate_labels must have the same length")209 210 effective_threshold = (211 self.calibration.confidence_threshold212 if confidence_threshold is None213 else min(max(float(confidence_threshold), 0.0), 1.0)214 )215 predictions: list[dict] = []216 217 for start in range(0, len(texts), self._predict_batch_size):218 batch_texts = texts[start : start + self._predict_batch_size]219 batch_candidates = candidate_labels[start : start + self._predict_batch_size]220 raw_probs, calibrated_probs = self._predict_probs(batch_texts)221 for raw_row, calibrated_row, labels in zip(raw_probs, calibrated_probs, batch_candidates):222 label_ids = [self.config.label2id[label] for label in labels if label in self.config.label2id]223 if not label_ids:224 predictions.append(225 {226 "label": None,227 "confidence": 0.0,228 "raw_confidence": 0.0,229 "candidate_mass": 0.0,230 "confidence_threshold": round_score(effective_threshold),231 "calibrated": self.calibration.calibrated,232 "meets_confidence_threshold": False,233 }234 )235 continue236 237 calibrated_slice = calibrated_row[label_ids]238 raw_slice = raw_row[label_ids]239 calibrated_mass = float(calibrated_slice.sum().item())240 raw_mass = float(raw_slice.sum().item())241 if calibrated_mass <= 0:242 predictions.append(243 {244 "label": labels[0],245 "confidence": 0.0,246 "raw_confidence": 0.0,247 "candidate_mass": 0.0,248 "confidence_threshold": round_score(effective_threshold),249 "calibrated": self.calibration.calibrated,250 "meets_confidence_threshold": False,251 }252 )253 continue254 255 normalized_calibrated = calibrated_slice / calibrated_mass256 normalized_raw = raw_slice / max(raw_mass, 1e-9)257 pred_offset = int(torch.argmax(normalized_calibrated).item())258 pred_id = label_ids[pred_offset]259 confidence = float(normalized_calibrated[pred_offset].item())260 raw_confidence = float(normalized_raw[pred_offset].item())261 predictions.append(262 {263 "label": self.model.config.id2label[pred_id],264 "confidence": round_score(confidence),265 "raw_confidence": round_score(raw_confidence),266 "candidate_mass": round_score(calibrated_mass),267 "confidence_threshold": round_score(effective_threshold),268 "calibrated": self.calibration.calibrated,269 "meets_confidence_threshold": confidence >= effective_threshold,270 }271 )272 return predictions273 274 def predict(self, text: str, confidence_threshold: float | None = None) -> dict:275 return self.predict_batch([text], confidence_threshold=confidence_threshold)[0]276 277 def predict_candidates(278 self,279 text: str,280 candidate_labels: list[str],281 confidence_threshold: float | None = None,282 ) -> dict:283 return self.predict_candidate_batch([text], [candidate_labels], confidence_threshold=confidence_threshold)[0]284 285 286@lru_cache(maxsize=None)287def get_head(head_name: str) -> SequenceClassifierHead:288 if head_name not in HEAD_CONFIGS:289 raise ValueError(f"Unknown head: {head_name}")290 if head_name in {"intent_type", "intent_subtype", "decision_phase"}:291 return MultiTaskHeadProxy(head_name) # type: ignore[return-value]292 return SequenceClassifierHead(HEAD_CONFIGS[head_name])293 