import json from collections.abc import Callable from functools import lru_cache import joblib import numpy as np import torch from sentence_transformers import SentenceTransformer from transformers import AutoModelForSequenceClassification, AutoTokenizer from .config import TOKEN from .hub import download, has_file, resolve Predictor = Callable[[list[str]], tuple[list[str], np.ndarray]] def setfit(repo: str, revision: str) -> Predictor: encoder = SentenceTransformer(repo, revision=revision, token=TOKEN) head = joblib.load(download(repo, "model_head.pkl", revision)) labels = json.loads(download(repo, "config_setfit.json", revision).read_text())["labels"] def predict(texts: list[str]): return labels, head.predict_proba(encoder.encode(texts)) return predict def automodel(repo: str, revision: str) -> Predictor: tokenizer = AutoTokenizer.from_pretrained(repo, revision=revision, token=TOKEN) model = AutoModelForSequenceClassification.from_pretrained( repo, revision=revision, token=TOKEN ).eval() labels = [model.config.id2label[i] for i in range(model.config.num_labels)] def predict(texts: list[str]): batch = tokenizer(texts, truncation=True, padding=True, return_tensors="pt") with torch.no_grad(): logits = model(**batch).logits return labels, torch.softmax(logits, dim=-1).numpy() return predict @lru_cache(maxsize=4) def at_sha(repo: str, sha: str) -> Predictor: if has_file(repo, "model_head.pkl", sha): return setfit(repo, sha) return automodel(repo, sha) def load(repo: str, revision: str) -> tuple[Predictor, str]: sha = resolve(repo, revision) return at_sha(repo, sha), sha