Spaces:
Sleeping
Sleeping
File size: 1,730 Bytes
759bf41 9668975 759bf41 9668975 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 | 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
|