Team Ai
Modelpublic

admesh/agentic-intent-classifier

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