breakdown-risk-demo / app /handlers.py
Cyprien
Put the loading indicator beside Revision
6a2ee52
Raw
History Blame Contribute Delete
4.19 kB
import json
import time
from functools import partial
from pathlib import Path
import gradio as gr
from . import evaluation
from .config import display
from .hub import model_repos, newest, revisions
from .predictors import load
from .text import (
CLASSIFY,
EVALUATE,
LOADING,
NO_MODEL,
SCORING,
STATUS_FAILED,
STATUS_LOADING,
WARM_FAILED,
WARMING,
pushed_at,
status_line,
status_ready,
)
from .turns import window
EXAMPLES = json.loads((Path(__file__).parent / "examples.json").read_text(encoding="utf-8"))
def classify(repo: str, revision: str, transcript: str) -> tuple[dict[str, float], str, str]:
"""Score a caller's turns for breakdown risk."""
text = window(transcript)
if not repo or not text:
return {}, "", ""
predictor, sha = load(repo, revision)
started = time.perf_counter() # after load, so the reading is inference alone
labels, probabilities = predictor([text])
elapsed = (time.perf_counter() - started) * 1000
scores = {display(name): float(p) for name, p in zip(labels, probabilities[0])}
return scores, text, f"`{sha[:7]}` · `{elapsed:.0f} ms`"
def _actions(ready: bool, busy: str = "") -> tuple[gr.Button, gr.Button]:
"""Both actions at once: they wait on the same model, so they move together."""
return (
gr.Button(busy or CLASSIFY, interactive=ready),
gr.Button(busy or EVALUATE, interactive=ready),
)
# What each session is waiting on, so a pick that got superseded lets go quietly.
_awaited: dict[str, tuple[str, str]] = {}
def warm(repo: str, revision: str, request: gr.Request):
"""Fetch the picked revision up front, so an action only ever runs inference."""
if not repo:
yield (*_actions(ready=False), status_line(NO_MODEL))
return
pick = (repo, revision)
_awaited[request.session_hash] = pick
yield (*_actions(ready=False, busy=WARMING), status_line(STATUS_LOADING))
try:
predictor, _ = load(repo, revision)
sample = [EXAMPLES[0]["text"]]
predictor(sample) # the first call pays torch's lazy init; pay it here
started = time.perf_counter()
predictor(sample) # so the reading matches what a click will cost
ms = (time.perf_counter() - started) * 1000
except Exception as failure:
if _awaited.get(request.session_hash) != pick:
return
# a bad revision must hand the buttons back, not dead-end the interface
yield (*_actions(ready=True), status_line(STATUS_FAILED))
raise gr.Error(f"{WARM_FAILED}: {failure}") from failure
if _awaited.get(request.session_hash) == pick: # else a newer pick owns the buttons
yield (*_actions(ready=True), status_line(status_ready(ms)))
def summary(report: evaluation.Report) -> str:
return (
f"### {report.hits}/{report.total} — accuracy {report.hits / report.total:.1%}\n"
f"`{report.repo}` @ `{report.model_sha[:7]}` · data `{report.dataset_sha[:7]}`"
f" · {report.ms_per_case:.0f} ms per case"
)
def evaluate(repo: str, revision: str, progress=gr.Progress()):
if not repo:
return NO_MODEL, [], []
progress(0, desc=LOADING)
predictor, sha = load(repo, revision)
track = partial(progress.tqdm, desc=SCORING)
report = evaluation.run(repo, sha, predict=predictor, track=track)
return summary(report), report.confusion, report.cases
def pick_revision(repo: str) -> tuple[gr.Dropdown, str]:
choices = revisions(repo)
return (
gr.Dropdown(choices=choices, value=choices[0][1] if choices else None),
pushed_at(model_repos().get(repo)),
)
def refresh(repo: str, revision: str) -> tuple[gr.Dropdown, gr.Dropdown, str]:
pushes = model_repos()
chosen = repo if repo in pushes else newest(pushes)
choices = revisions(chosen) if chosen else []
shas = [sha for _, sha in choices]
keep = revision if revision in shas else (shas[0] if shas else None)
return (
gr.Dropdown(choices=list(pushes), value=chosen),
gr.Dropdown(choices=choices, value=keep),
pushed_at(pushes[chosen] if chosen else None),
)