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