hub-task-tagger / app.py
davanstrien's picture
davanstrien HF Staff
Accept/reject feedback: per-tag votes stored in a private bucket with a pseudonymous id
5605c42 verified
Raw History Blame Contribute Delete
23.6 kB
"""Hub task tagger demo: FastAPI backend for the single-page frontend in static/index.html.
The model input is rebuilt exactly as in training (see tagger_core.py). The scoring path, temperature and
threshold are copied from the Gradio test Space (demo-space/app.py), with the base model only.
"""
import asyncio
import hashlib
import hmac
import json
import math
import os
import random
import re
import time
import uuid
from collections import OrderedDict, deque
from contextlib import asynccontextmanager
from datetime import datetime, timezone
from pathlib import Path
import requests
import torch
from fastapi import FastAPI, Request
from fastapi.concurrency import run_in_threadpool
from fastapi.responses import FileResponse, JSONResponse
from gliner2.classification import (
ClassificationConfig,
ClassificationSchema,
Classifier,
)
from huggingface_hub import HfApi, hf_hub_download
from huggingface_hub.utils import (
GatedRepoError,
HfHubHTTPError,
RepositoryNotFoundError,
)
from pydantic import BaseModel
from tokenizers import Tokenizer
from tagger_core import RETIRED, StateBuilder, fetch_first_rows, render_first_rows
# cpu-basic has 2 vCPUs; os.cpu_count() reports the host, so set the thread count explicitly.
torch.set_num_threads(int(os.environ.get("TORCH_THREADS", "2")))
HERE = Path(__file__).parent
LABELS = json.loads((HERE / "labels.json").read_text())
TASK = "labels" # task name in the trained classification schema
MAX_TEXT_CHARS = 2000 # train-gliner2.py --max-text-chars
MAX_STATE_TOKENS = 370 # prepare_hub_tasks.py --max-state-tokens
BASE_TOKENIZER = "Qwen/Qwen3.5-4B-Base"
MODEL_REPO = "davanstrien/hub-task-tagger-gliner2.5-base"
# Temperature fitted on calibration (top-1 log-loss); tau chosen on calibration micro-F1 (threshold rule, top-1 kept).
T, TAU = 1.300, 0.225
TOP_K = 5
# Baked into the image at build time (see Dockerfile); falls back to the Hub when run elsewhere.
MODEL_DIR = Path(os.environ.get("MODEL_DIR", "/home/user/models"))
TAGS_URL = "https://huggingface.co/api/datasets-tags-by-type?type=task_categories"
CACHE_SIZE = 2000
WARM_IDS = ["hubxrt/LPMusicCapsMTT_a2t", "dougdotcon/douvras-failure-atlas",
"mim-chess-vlas/train_800_dense__mask__overlay_a75__sim__agentview_camera__static"] # the page examples
# ---- startup ---------------------------------------------------------------------------------------------------
STARTUP = {}
t0 = time.perf_counter()
if (MODEL_DIR / "tokenizer.json").exists():
tokenizer_path = MODEL_DIR / "tokenizer.json"
model_path = MODEL_DIR / "base"
schema_path = model_path / "classification_schema.json"
STARTUP["source"] = "image"
else:
tokenizer_path = hf_hub_download(BASE_TOKENIZER, "tokenizer.json")
model_path = MODEL_REPO
schema_path = hf_hub_download(MODEL_REPO, "classification_schema.json")
STARTUP["source"] = "hub"
states = StateBuilder(Tokenizer.from_file(str(tokenizer_path)), MAX_STATE_TOKENS)
STARTUP["tokenizer_load_s"] = round(time.perf_counter() - t0, 2)
MODEL_REVISION = (MODEL_DIR / "revision.txt").read_text().strip() if (MODEL_DIR / "revision.txt").exists() else None
t0 = time.perf_counter()
SCHEMA = ClassificationSchema().multi(TASK, LABELS)
CONFIG = ClassificationConfig(batch_size=1)
CLASSIFIER = Classifier.from_pretrained(str(model_path)).to(device="cpu").eval()
if json.loads(Path(schema_path).read_text())["tasks"][0]["labels"] != LABELS:
raise SystemExit(f"{MODEL_REPO}: trained labels differ from labels.json")
STARTUP["model_load_s"] = round(time.perf_counter() - t0, 2)
STARTUP["torch_threads"] = torch.get_num_threads()
LIVE_TAGS = None
print("startup", STARTUP, flush=True)
# One request at a time on the model: 2 vCPUs do not serve parallel forward passes well.
MODEL_LOCK = asyncio.Lock()
def live_tags():
global LIVE_TAGS
if LIVE_TAGS is None:
LIVE_TAGS = {x["label"] for x in requests.get(TAGS_URL, timeout=30).json()["task_categories"]}
return LIVE_TAGS
def declared_tags(dataset_id):
"""(raw card task_categories, mapped to current tags, gated, sha, error). Anonymous on purpose: public data only."""
try:
info = HfApi(token=False).dataset_info(dataset_id, expand=["cardData", "gated", "sha"])
except Exception as error: # noqa: BLE001 -- classified by the caller
return [], [], False, None, error
card = info.card_data.to_dict() if info.card_data else {}
raw = card.get("task_categories") or []
raw = [raw] if isinstance(raw, str) else [str(t) for t in raw]
live = live_tags()
mapped = sorted({t if t in live else RETIRED[t] for t in raw if t in live or RETIRED.get(t) in live})
return raw, mapped, bool(info.gated), info.sha, None
def softmax(logits, t):
scaled = [z / t for z in logits]
m = max(scaled)
e = [math.exp(z - m) for z in scaled]
total = sum(e)
return [x / total for x in e]
def score(state):
text = state[:MAX_TEXT_CHARS]
with torch.inference_mode():
scores = CLASSIFIER.batch_score([text], SCHEMA, config=CONFIG)
logits = scores[0].tasks[TASK]
probs = softmax([logits[label] for label in LABELS], T)
ranked = sorted(zip(LABELS, probs), key=lambda x: -x[1])
chosen = [label for label, p in ranked if p >= TAU] or [ranked[0][0]]
if ranked[0][0] not in chosen:
chosen.insert(0, ranked[0][0])
return ranked, chosen
# ---- cache ---------------------------------------------------------------------------------------------------
class PredictionCache:
"""LRU of successful responses (without timing), keyed on (dataset_id, revision sha): a new commit misses."""
def __init__(self, size):
self.size, self.items = size, OrderedDict()
def get(self, dataset_id, sha):
key = (dataset_id, sha)
if sha is None or key not in self.items:
return None
self.items.move_to_end(key)
return self.items[key]
def put(self, dataset_id, sha, body):
if sha is None:
return
self.items[(dataset_id, sha)] = body
self.items.move_to_end((dataset_id, sha))
while len(self.items) > self.size:
self.items.popitem(last=False)
CACHE = PredictionCache(CACHE_SIZE)
READY = {"warm": False, "warm_s": None, "warmed": []}
class Failure(Exception):
def __init__(self, status, kind, message, dataset_id=None, timing=None):
self.status, self.kind, self.message, self.dataset_id, self.timing = status, kind, message, dataset_id, timing
def clean_id(value):
"""Accept `owner/name` or a full dataset URL; return None when it is not a dataset id."""
value = (value or "").strip().removeprefix("https://huggingface.co/datasets/").split("?")[0].strip("/")
parts = value.split("/")
if len(parts) != 2 or not all(parts) or " " in value:
return None
return value
async def predict_one(raw_id):
"""The response body for one dataset id; raises Failure for anything the page should show as an error."""
dataset_id = clean_id(raw_id)
if dataset_id is None:
raise Failure(400, "bad_id", "That is not a dataset id. Use the form owner/name, for example "
"ar0s/kuka_heat_reaction_expert, or paste the dataset's URL.")
timing = {}
t0 = time.perf_counter()
raw, declared, gated, sha, error = await run_in_threadpool(declared_tags, dataset_id)
timing["card_ms"] = round(1000 * (time.perf_counter() - t0))
if isinstance(error, GatedRepoError):
gated, error = True, None
if isinstance(error, RepositoryNotFoundError):
raise Failure(404, "not_found", f"No public dataset called {dataset_id}. Check the spelling. "
"Private datasets cannot be read here.", dataset_id, timing)
if error is not None:
status = error.response.status_code if isinstance(error, HfHubHTTPError) and error.response is not None else None
raise Failure(502, "hub_error", f"The Hub did not answer for {dataset_id} (status {status or 'unknown'}). "
"Try again in a minute.", dataset_id, timing)
if gated:
raise Failure(403, "gated", f"{dataset_id} is gated. This page reads only public datasets, "
"so it cannot see the rows. Try a dataset without an access form.", dataset_id, timing)
cached = CACHE.get(dataset_id, sha)
if cached is not None:
return {**cached, "timing": {"cached": True, "card_ms": timing["card_ms"]}}
t0 = time.perf_counter()
fr, cfg, split, fetch_error = await run_in_threadpool(fetch_first_rows, dataset_id)
timing["fetch_ms"] = round(1000 * (time.perf_counter() - t0))
if fr is None:
print("no preview", dataset_id, fetch_error, flush=True)
raise Failure(422, "no_viewer", f"{dataset_id} has no dataset viewer preview, so there are no rows "
"for the model to read. The viewer is off for some datasets and still processing for new "
"ones. Try again later, or try another dataset.", dataset_id, timing)
t0 = time.perf_counter()
state, truncated = states.build(render_first_rows(fr))
timing["render_ms"] = round(1000 * (time.perf_counter() - t0), 1)
async with MODEL_LOCK:
t0 = time.perf_counter()
ranked, chosen = await run_in_threadpool(score, state)
timing["model_ms"] = round(1000 * (time.perf_counter() - t0))
body = {
"dataset_id": dataset_id,
"sha": sha,
"declared": declared,
"declared_raw": raw,
"top": [{"tag": label, "p": round(p, 4)} for label, p in ranked[:TOP_K]],
"set": chosen,
"state": state,
"truncated": truncated,
"preview": f"{cfg}/{split}",
}
CACHE.put(dataset_id, sha, body)
return {**body, "timing": timing}
async def warm_cache():
"""Score the page examples once so the first visitor's clicks are cache hits. Failures are logged, not fatal."""
t0 = time.perf_counter()
for dataset_id in WARM_IDS:
try:
await predict_one(dataset_id)
READY["warmed"].append(dataset_id)
except Exception as error: # noqa: BLE001 -- warm-up must never stop the app
print("warm-up failed", dataset_id, getattr(error, "kind", repr(error)), flush=True)
READY["warm"], READY["warm_s"] = True, round(time.perf_counter() - t0, 2)
print("warm", READY, flush=True)
# ---- random dataset pool --------------------------------------------------------------------------------------
POOL_REFRESH_S = 6 * 3600 # rebuild the pool every 6 hours
POOL_RETRY_S = 10 * 60 # retry sooner when a build fails and the pool is empty
POOL_EXPAND = ["gated", "private", "tags", "cardData", "downloads", "likes"]
RANDOM_TRIES = 5
IS_VALID_URL = "https://datasets-server.huggingface.co/is-valid"
# Any of these words in the id, the tags or the card's tags/pretty_name keeps a dataset out of the pool.
ADULT_WORDS = {"nsfw", "adult", "adults", "porn", "porno", "pornography", "pornographic", "hentai", "xxx",
"sex", "sexual", "sexy", "nude", "nudes", "nudity", "naked", "erotic", "erotica", "explicit",
"r18", "18plus", "lewd", "fetish", "onlyfans", "rule34", "not-for-all-audiences",
# Safety-research sets (jailbreaks, hate speech) are fine to tag by hand, but a random click should
# not put their first row in front of a visitor.
"toxic", "toxicity", "hate", "hateful", "jailbreak", "jailbreaks", "harmful", "offensive", "unsafe",
"redteam", "red-team", "teaming", "abuse", "abusive"}
POOL = {"ids": [], "task_tagged": {}, "size": 0, "built_at": None, "build_s": None, "error": None,
"is_valid_checks": 0, "is_valid_hits": 0}
def words_in(text):
return set(re.split(r"[^a-z0-9+-]+", str(text).lower())) | set(re.split(r"[^a-z0-9]+", str(text).lower()))
def looks_adult(info):
card = info.card_data.to_dict() if info.card_data else {}
card_tags = card.get("tags") or []
card_tags = [card_tags] if isinstance(card_tags, str) else card_tags
texts = [info.id, card.get("pretty_name") or "", *(info.tags or []), *card_tags]
return any(words_in(text) & ADULT_WORDS for text in texts)
def pool_candidates(api):
"""Popular + trending datasets, then recently modified ones with some traction (downloads >= 50 or likes >= 2)."""
yield from api.list_datasets(sort="trendingScore", limit=500, expand=POOL_EXPAND, gated=False)
yield from api.list_datasets(sort="likes", limit=700, expand=POOL_EXPAND, gated=False)
kept = 0
for info in api.list_datasets(sort="last_modified", limit=20000, expand=POOL_EXPAND, gated=False):
if (info.downloads or 0) >= 50 or (info.likes or 0) >= 2:
kept += 1
yield info
if kept >= 1000:
break
def build_pool():
"""(ids, {id: has task_categories tag}) of public, ungated, not-adult datasets. Blocking; run in a thread."""
ids, task_tagged, skipped = [], {}, 0
for info in pool_candidates(HfApi(token=False)):
if info.id in task_tagged:
continue
if info.private or info.gated or looks_adult(info):
skipped += 1
continue
ids.append(info.id)
task_tagged[info.id] = any(tag.startswith("task_categories:") for tag in info.tags or [])
return ids, task_tagged, skipped
async def refresh_pool_forever():
"""Build the pool now and every 6 hours, in a thread. Never blocks requests or readiness."""
while True:
t0 = time.perf_counter()
try:
ids, task_tagged, skipped = await asyncio.to_thread(build_pool)
POOL.update(ids=ids, task_tagged=task_tagged, size=len(ids), built_at=time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
build_s=round(time.perf_counter() - t0, 1), error=None)
print(f"pool built: {len(ids)} datasets ({skipped} skipped) in {POOL['build_s']} s", flush=True)
except Exception as error: # noqa: BLE001 -- keep the old pool, try again later
POOL["error"] = repr(error)
print("pool build failed", repr(error), flush=True)
await asyncio.sleep(POOL_REFRESH_S if POOL["ids"] else POOL_RETRY_S)
def has_preview(dataset_id):
"""True when the dataset viewer has a first-rows preview for this dataset."""
try:
r = requests.get(IS_VALID_URL, params={"dataset": dataset_id}, timeout=10)
return r.ok and r.json().get("preview") is True
except (requests.RequestException, ValueError):
return False
@asynccontextmanager
async def lifespan(app):
# Warm in the background: the port opens at once, and /healthz reports 503 until warm-up ends.
# The random-dataset pool builds in the background too and does not affect /healthz.
tasks = [asyncio.create_task(warm_cache()), asyncio.create_task(refresh_pool_forever())]
yield
for task in tasks:
task.cancel()
# ---- feedback ------------------------------------------------------------------------------------------------
# A Hub bucket mounted read-write on the Space. One JSON file per event; the app never reads the folder back.
FEEDBACK_DIR = Path(os.environ.get("FEEDBACK_DIR", "/data/feedback"))
FEEDBACK_PER_HOUR = 30
MODEL_INFO = {"repo": MODEL_REPO, "revision": MODEL_REVISION, "T": T, "tau": TAU}
TAG_RE = re.compile(r"^[a-z0-9][a-z0-9-]{0,63}$")
SHA_RE = re.compile(r"^[0-9a-f]{40}$")
# Rate limiting keys on a salted hash of the client IP. The salt lives only in memory, and neither the IP nor
# the hash is written anywhere.
RATE_SALT = os.urandom(16)
RECENT = {} # hashed ip -> deque of submit times (monotonic seconds)
# Pseudonymous source id: HMAC of month + client IP with a secret salt from the Space secrets. The month in the
# message rotates the ids, so events cannot be linked across months. The IP itself is never stored.
FEEDBACK_SALT = os.environ.get("FEEDBACK_SALT", "").encode()
if not FEEDBACK_SALT:
print("WARNING: FEEDBACK_SALT is not set; feedback events get source_id=null", flush=True)
def check_feedback_dir():
"""True when FEEDBACK_DIR exists and a file can be written and removed there."""
probe = FEEDBACK_DIR / f".probe-{uuid.uuid4().hex}"
try:
probe.write_text("ok")
probe.unlink()
return True
except OSError as error:
print(f"WARNING: feedback disabled, {FEEDBACK_DIR} is not writable: {error!r}", flush=True)
return False
FEEDBACK = {"enabled": check_feedback_dir(), "dir": str(FEEDBACK_DIR), "written": 0, "write_errors": 0,
"source_ids": bool(FEEDBACK_SALT)}
def source_id(ip, period):
"""First 16 hex characters of HMAC-SHA256(FEEDBACK_SALT, "YYYY-MM|ip"), or None without a salt."""
if not FEEDBACK_SALT:
return None
return hmac.new(FEEDBACK_SALT, f"{period}|{ip}".encode(), hashlib.sha256).hexdigest()[:16]
print("feedback", FEEDBACK, flush=True)
def rate_limited(ip):
key = hashlib.sha256(RATE_SALT + ip.encode()).hexdigest()
now = time.monotonic()
times = RECENT.setdefault(key, deque())
while times and now - times[0] > 3600:
times.popleft()
if len(times) >= FEEDBACK_PER_HOUR:
return True
times.append(now)
return False
def client_ip(request):
"""The address the Space proxy saw. It appends that address last to X-Forwarded-For; earlier entries come from
the client and can be forged, so they are ignored."""
forwarded = [x.strip() for x in request.headers.get("x-forwarded-for", "").split(",") if x.strip()]
return forwarded[-1] if forwarded else (request.client.host if request.client else "unknown")
class Suggestion(BaseModel):
tag: str
p: float
class FeedbackRequest(BaseModel):
dataset_id: str
sha: str | None = None
model_repo: str | None = None
model_revision: str | None = None
T: float | None = None
tau: float | None = None
suggested: list[Suggestion]
set: list[str]
accepted: list[str] = []
rejected: list[str] = []
owner_tags: list[str] = []
ts: str | None = None
def feedback_problem(fb):
"""A short reason when the event is not valid, else None."""
if clean_id(fb.dataset_id) != fb.dataset_id:
return "dataset_id is not an owner/name id"
if fb.sha is not None and not SHA_RE.match(fb.sha):
return "sha is not a commit sha"
suggested = [s.tag for s in fb.suggested]
if not 1 <= len(suggested) <= TOP_K or len(set(suggested)) != len(suggested):
return f"suggested must list 1 to {TOP_K} different tags"
if not all(tag in LABELS for tag in suggested + fb.set + fb.accepted + fb.rejected):
return "a tag is not one of the model's labels"
if not all(0.0 <= s.p <= 1.0 for s in fb.suggested):
return "a probability is outside 0-1"
if not fb.accepted and not fb.rejected:
return "nothing accepted or rejected"
if set(fb.accepted) & set(fb.rejected):
return "a tag is both accepted and rejected"
if not set(fb.accepted + fb.rejected + fb.set) <= set(suggested):
return "accepted, rejected and set tags must come from the suggested tags"
if len(fb.owner_tags) > 30 or not all(TAG_RE.match(tag) for tag in fb.owner_tags):
return "owner_tags are not valid tags"
if fb.ts is not None and len(fb.ts) > 40:
return "ts is too long"
return None
def write_event(record):
"""Write one event as its own file: a temp name in the same folder, then os.replace."""
day = FEEDBACK_DIR / record["received_at"][:10]
day.mkdir(parents=True, exist_ok=True)
final = day / f"{record['id']}.json"
tmp = day / f".{record['id']}.json.tmp"
tmp.write_text(json.dumps(record, ensure_ascii=False, indent=1))
os.replace(tmp, final)
# ---- API -------------------------------------------------------------------------------------------------------
app = FastAPI(title="Hub task tagger demo", docs_url=None, redoc_url=None, lifespan=lifespan)
class PredictRequest(BaseModel):
dataset_id: str
@app.get("/healthz")
def healthz():
pool = {k: v for k, v in POOL.items() if k not in ("ids", "task_tagged")}
body = {"ok": READY["warm"], "startup": STARTUP, "warm": READY, "cache_entries": len(CACHE.items), "pool": pool,
"feedback": {k: v for k, v in FEEDBACK.items() if k != "dir"}}
return JSONResponse(status_code=200 if READY["warm"] else 503, content=body)
@app.post("/api/predict")
async def predict(req: PredictRequest):
try:
body = await predict_one(req.dataset_id)
return {**body, "model": MODEL_INFO, "feedback": FEEDBACK["enabled"]}
except Failure as f:
return JSONResponse(
status_code=f.status,
content={"error": f.kind, "message": f.message, "dataset_id": f.dataset_id, "timing": f.timing or {}},
)
@app.post("/api/feedback")
async def feedback(fb: FeedbackRequest, request: Request):
if not FEEDBACK["enabled"]:
return JSONResponse(status_code=503, content={"error": "feedback_off", "message": "Feedback is switched off."})
problem = feedback_problem(fb)
if problem:
return JSONResponse(status_code=400, content={"error": "invalid", "message": f"Feedback not saved: {problem}."})
ip = client_ip(request)
if rate_limited(ip):
return JSONResponse(status_code=429, content={"error": "rate_limited", "message": "Too much feedback from "
"this address in the last hour. Thank you, and please try again later."})
now = datetime.now(timezone.utc)
period = now.strftime("%Y-%m")
record = {
"id": uuid.uuid4().hex,
"received_at": now.isoformat(timespec="seconds"),
"source_id": source_id(ip, period),
"source_period": period,
**fb.model_dump(),
# The server's own values win over what the page sent.
"model_repo": MODEL_REPO, "model_revision": MODEL_REVISION, "T": T, "tau": TAU,
}
try:
await run_in_threadpool(write_event, record)
except OSError as error:
FEEDBACK["write_errors"] += 1
print("feedback write failed", repr(error), flush=True)
return JSONResponse(status_code=503, content={"error": "write_failed", "message": "Could not save the "
"feedback just now. Please try again later."})
FEEDBACK["written"] += 1
return {"ok": True, "id": record["id"]}
@app.get("/api/random")
async def random_dataset():
"""A random public dataset from the pool that has a viewer preview. The page then calls /api/predict with it."""
if not POOL["ids"]:
return JSONResponse(status_code=503, content={"error": "no_pool", "message": "The list of random datasets "
"is not ready yet. Try one of the examples, or try again in a few minutes."})
for _ in range(RANDOM_TRIES):
dataset_id = random.choice(POOL["ids"])
ok = await run_in_threadpool(has_preview, dataset_id)
POOL["is_valid_checks"] += 1
POOL["is_valid_hits"] += ok
if ok:
return {"dataset_id": dataset_id, "has_task_tags": POOL["task_tagged"][dataset_id]}
return JSONResponse(status_code=503, content={"error": "no_preview", "message": "Could not find a random "
"dataset with a viewer preview just now. Try again, or try one of the examples."})
@app.get("/")
def index():
return FileResponse(HERE / "static" / "index.html")