from __future__ import annotations import json import os from pathlib import Path from typing import Any import torch from huggingface_hub import snapshot_download from safetensors.torch import load_file from transformers import AutoTokenizer from encoder_model import EncoderClassifier DEFAULT_REPO_ID = "ZenMan67/support-ticket-classifiers-minilm" class HubTicketClassifier: """Download public checkpoints at startup and run the two-stage cascade.""" def __init__( self, repo_id: str | None = None, revision: str = "main", device: str | None = None, preload_all: bool = False, ) -> None: self.repo_id = repo_id or os.getenv("HF_MODEL_REPO", DEFAULT_REPO_ID) self.revision = os.getenv("HF_MODEL_REVISION", revision) self.device = torch.device( device or ("mps" if torch.backends.mps.is_available() else "cpu") ) self.model_root = Path(snapshot_download( repo_id=self.repo_id, revision=self.revision, repo_type="model", allow_patterns=[ "tokenizer/*", "handler/*", ], )) self.tokenizer = AutoTokenizer.from_pretrained( self.model_root / "tokenizer", local_files_only=True, ) self.models: dict[str, tuple[EncoderClassifier, dict[str, Any]]] = {} self._load_task("handler") if preload_all: for task in ("human", "llm", "auto"): self._load_task(task) def _load_task( self, task: str, ) -> tuple[EncoderClassifier, dict[str, Any]]: if task in self.models: return self.models[task] task_root = self.model_root / task if not task_root.exists(): self.model_root = Path(snapshot_download( repo_id=self.repo_id, revision=self.revision, repo_type="model", allow_patterns=[f"{task}/*"], )) task_root = self.model_root / task config = json.loads( (task_root / "config.json").read_text(encoding="utf-8") ) model = EncoderClassifier( config["base_model"], config["num_labels"], dropout=config["dropout"], pooling=config["pooling"], pretrained=False, encoder_config=config["encoder_config"], ) model.load_state_dict(load_file(task_root / "model.safetensors")) model.to(self.device).eval() self.models[task] = (model, config) return model, config @torch.inference_mode() def classify( self, text: str, task: str, top_k: int = 3, ) -> dict[str, Any]: model, config = self._load_task(task) encoded = self.tokenizer( text, return_tensors="pt", truncation=True, max_length=config["max_length"], ) logits = model( encoded["input_ids"].to(self.device), encoded["attention_mask"].to(self.device), ) probabilities = logits.softmax(dim=-1)[0].cpu() values, indices = probabilities.topk( min(top_k, len(config["labels"])) ) return { "label": config["labels"][int(indices[0])], "confidence": float(values[0]), "top": [ { "label": config["labels"][int(index)], "confidence": float(value), } for value, index in zip(values, indices) ], } def predict(self, text: str, top_k: int = 3) -> dict[str, Any]: handler = self.classify(text, "handler", top_k) category = self.classify(text, handler["label"], top_k) return { "text": text, "handler": handler, "category": category, }