import time from collections.abc import Callable, Iterable from functools import lru_cache from typing import NamedTuple import numpy as np from datasets import load_dataset from .config import DATASET_REPO, TOKEN, display from .hub import resolve from .predictors import Predictor, load BATCH_SIZE = 8 Track = Callable[[Iterable], Iterable] class Split(NamedTuple): ids: list[str] texts: list[str] gold: list[str] sha: str class Report(NamedTuple): repo: str model_sha: str dataset_sha: str hits: int total: int ms_per_case: float confusion: list[list] cases: list[list] @lru_cache(maxsize=2) def at_sha(sha: str) -> Split: data = load_dataset(DATASET_REPO, split="test", revision=sha, token=TOKEN) names = data.features["label"].names return Split( ids=list(data["id"]), texts=list(data["text"]), gold=[names[i] for i in data["label"]], sha=sha, ) def test_split() -> Split: return at_sha(resolve(DATASET_REPO, "main", repo_type="dataset")) def confusion(labels: list[str], gold: list[str], predicted: list[str]) -> list[list]: counts = dict.fromkeys(((want, got) for want in labels for got in labels), 0) for pair in zip(gold, predicted): counts[pair] += 1 return [ [f"actual {display(want).lower()}", *(counts[want, got] for got in labels)] for want in labels ] def cases(ids, texts, gold, predicted, probabilities: np.ndarray) -> list[list]: rows = [ [ "✓" if want == got else "✗", row_id, display(want), display(got), round(float(confidence), 3), text.replace("\n", " / "), ] for row_id, text, want, got, confidence in zip( ids, texts, gold, predicted, probabilities.max(axis=1) ) ] return sorted(rows, key=lambda row: (row[0] == "✓", row[4])) def batches(texts: list[str], size: int = BATCH_SIZE) -> list[list[str]]: return [texts[start : start + size] for start in range(0, len(texts), size)] def classify_all( predictor: Predictor, texts: list[str], track: Track | None = None ) -> tuple[list[str], np.ndarray]: chunks = batches(texts) labels: list[str] = [] scored = [] for batch in track(chunks) if track else chunks: labels, probabilities = predictor(batch) scored.append(probabilities) return labels, np.concatenate(scored) def run( repo: str, revision: str, predict: Predictor | None = None, track: Track | None = None, ) -> Report: predictor, model_sha = (predict, revision) if predict else load(repo, revision) split = test_split() started = time.perf_counter() labels, probabilities = classify_all(predictor, split.texts, track) elapsed = time.perf_counter() - started predicted = [labels[int(row.argmax())] for row in probabilities] return Report( repo=repo, model_sha=model_sha, dataset_sha=split.sha, hits=sum(want == got for want, got in zip(split.gold, predicted)), total=len(split.gold), ms_per_case=elapsed * 1000 / len(split.texts), confusion=confusion(labels, split.gold, predicted), cases=cases(split.ids, split.texts, split.gold, predicted, probabilities), )