from __future__ import annotations import json from pathlib import Path import numpy as np import torch from transformers import AutoModelForSequenceClassification, AutoTokenizer from src.data_utils import combine_text class ArticleTopicPredictor: def __init__(self, model_dir: str | Path, device: str | None = None) -> None: self.model_dir = Path(model_dir) if not self.model_dir.exists(): raise FileNotFoundError(f"Model directory does not exist: {self.model_dir}") self.tokenizer = AutoTokenizer.from_pretrained(self.model_dir) self.model = AutoModelForSequenceClassification.from_pretrained(self.model_dir) mapping_path = self.model_dir / "label_mapping.json" if mapping_path.exists(): payload = json.loads(mapping_path.read_text(encoding="utf-8")) self.id2label = {int(k): v for k, v in payload["id2label"].items()} else: self.id2label = {int(k): v for k, v in self.model.config.id2label.items()} if device is None: device = "cuda" if torch.cuda.is_available() else "cpu" self.device = torch.device(device) self.model.to(self.device) self.model.eval() @torch.inference_mode() def predict(self, title: str = "", abstract: str = "", top95_threshold: float = 0.95) -> dict: text = combine_text(title, abstract) if not text.strip(): raise ValueError("Provide at least title or abstract.") encoded = self.tokenizer( text, truncation=True, padding=False, return_tensors="pt", ) encoded = {k: v.to(self.device) for k, v in encoded.items()} logits = self.model(**encoded).logits[0] probs = torch.softmax(logits, dim=-1).detach().cpu().numpy() order = np.argsort(probs)[::-1] sorted_probs = probs[order] all_probs = [ {"label": self.id2label[int(idx)], "probability": float(probs[idx])} for idx in order ] cumulative = 0.0 top95 = [] for idx, prob in zip(order, sorted_probs): cumulative += float(prob) top95.append( { "label": self.id2label[int(idx)], "probability": float(prob), "cumulative_probability": float(cumulative), } ) if cumulative >= top95_threshold: break return { "input_text": text, "top95": top95, "all_probs": all_probs, }