Spaces:
Running
Running
Download src/predictor.py from asriel14/article_classifier: direct link, hf CLI and curl.
- Browser
- Download file 2.62 kB
-
https://huggingface.co/spaces/asriel14/article_classifier/resolve/main/src/predictor.py
- Command line
-
hf download hf://spaces/asriel14/article_classifier/src/predictor.py
-
curl -L -o predictor.py https://huggingface.co/spaces/asriel14/article_classifier/resolve/main/src/predictor.py
2.62 kB
| 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() | |
| 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, | |
| } | |