breakdown-risk-demo / app /predictors.py
Cyprien
Resolve refs to commits so nothing serves a stale model or split
9668975
Raw
History Blame Contribute Delete
1.73 kB
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