"""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")