ZenMan67's picture
Add files using upload-large-folder tool
81e8ada verified
Raw
History Blame Contribute Delete
3.99 kB
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,
}