Team Ai
Modelpublic

admesh/agentic-intent-classifier

sourceHugging Faceotherupdated 17d agoView on Hugging Face
2likes45downloads
model_runtime.py293 linesDownload Raw Back to root
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