Spaces:
Sleeping
Sleeping
| 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 | |
| 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 | |