Download main.py from SLM-Archive/hyperdex-trainer: direct link, hf CLI and curl.
- Browser
- Download file 25.1 kB
-
https://huggingface.co/SLM-Archive/hyperdex-trainer/resolve/main/main.py
- Command line
-
hf download hf://SLM-Archive/hyperdex-trainer/main.py
-
curl -L -o main.py https://huggingface.co/SLM-Archive/hyperdex-trainer/resolve/main/main.py
25.1 kB
| """ | |
| NanoDex — a place to pre-train small decoder-only language models from scratch. | |
| FastAPI multi-page app: | |
| / landing (logged out) / dashboard (logged in) | |
| /train multi-step run wizard | |
| /models your runs: training, queued, finished | |
| /models/{id} one run in detail — live loss, logs, publish, playground | |
| /profile account + stats | |
| /login /logout Hugging Face OAuth | |
| /api/* JSON for live polling | |
| """ | |
| import os | |
| import secrets | |
| import sys | |
| import threading | |
| import time | |
| from typing import Optional | |
| # Printed before anything heavy so a stalled boot is always diagnosable from | |
| # the Space's run log. | |
| print("[boot] NanoDex starting", flush=True) | |
| from fastapi import FastAPI, Request, HTTPException, Body | |
| from fastapi.responses import HTMLResponse, RedirectResponse, JSONResponse | |
| from fastapi.staticfiles import StaticFiles | |
| from fastapi.templating import Jinja2Templates | |
| from starlette.middleware.sessions import SessionMiddleware | |
| from nanodex import db, worker, hub, estimate, store | |
| from nanodex import data as ndata | |
| from nanodex.config import (TIERS, TIER_ORDER, MIN_TOKENS, MAX_TOKENS, | |
| TOKEN_STOPS_M, SEQ_LEN, VOCAB_SIZE) | |
| from nanodex.trainer import total_steps_for | |
| BASE = os.path.dirname(os.path.abspath(__file__)) | |
| TOKENIZER_DIR = os.path.join(BASE, "tokenizer") | |
| MAX_ACTIVE_PER_USER = int(os.environ.get("MAX_ACTIVE_PER_USER", "2")) | |
| OAUTH_CLIENT_ID = os.environ.get("OAUTH_CLIENT_ID") | |
| OAUTH_CLIENT_SECRET = os.environ.get("OAUTH_CLIENT_SECRET") | |
| OAUTH_SCOPES = os.environ.get("OAUTH_SCOPES", "openid profile read-repos write-repos manage-repos") | |
| OPENID_PROVIDER_URL = os.environ.get("OPENID_PROVIDER_URL", "https://huggingface.co") | |
| SPACE_HOST = os.environ.get("SPACE_HOST") | |
| OAUTH_READY = bool(OAUTH_CLIENT_ID and OAUTH_CLIENT_SECRET) | |
| print(f"[boot] data dir: {db.DATA_DIR}", flush=True) | |
| db.init() | |
| print("[boot] database ready", flush=True) | |
| from tokenizers import Tokenizer as _FastTok | |
| TOKENIZER = _FastTok.from_file(os.path.join(TOKENIZER_DIR, "tokenizer.json")) | |
| print(f"[boot] tokenizer loaded ({TOKENIZER.get_vocab_size()} tokens)", flush=True) | |
| DEVICES = worker.devices() | |
| print(f"[boot] devices: {DEVICES}", flush=True) | |
| def _restore_history(): | |
| """ | |
| Pull previously finished runs back out of the bucket. Done in a thread so | |
| a slow or unreachable bucket never delays the app coming up. | |
| """ | |
| try: | |
| if not store.available(): | |
| return | |
| known = db.known_job_ids() | |
| ids = [i for i in store.list_runs() if i not in known] | |
| if not ids: | |
| print("[store] no archived runs to restore", flush=True) | |
| return | |
| n = 0 | |
| for jid in ids: | |
| meta = store.fetch_run_meta(jid) | |
| if meta and db.import_run(meta): | |
| n += 1 | |
| store._state["restored"] = n | |
| print(f"[store] restored {n} run(s) from durable storage", flush=True) | |
| except Exception as exc: | |
| print(f"[store] history restore failed: {exc!r}", flush=True) | |
| threading.Thread(target=_restore_history, name="restore-history", | |
| daemon=True).start() | |
| ndata.start_background_build(TOKENIZER) | |
| worker.start_workers(TOKENIZER_DIR) | |
| ON_GPU = DEVICES[0].startswith("cuda") | |
| INFER_DEVICE = DEVICES[0] if ON_GPU else "cpu" | |
| app = FastAPI(title="NanoDex") | |
| IN_SPACE = bool(SPACE_HOST) | |
| app.add_middleware( | |
| SessionMiddleware, | |
| secret_key=os.environ.get("SESSION_SECRET") or (OAUTH_CLIENT_SECRET or secrets.token_hex(32)), | |
| # A Space is served inside an iframe on huggingface.co, which makes the | |
| # session cookie third-party: "lax" gets dropped on the redirect back from | |
| # the OAuth provider and the CSRF state check fails. "none" + Secure is the | |
| # only combination that survives that round trip. | |
| same_site="none" if IN_SPACE else "lax", | |
| https_only=IN_SPACE, | |
| max_age=8 * 3600, | |
| ) | |
| app.mount("/static", StaticFiles(directory=os.path.join(BASE, "web/static")), name="static") | |
| templates = Jinja2Templates(directory=os.path.join(BASE, "web/templates")) | |
| # --------------------------------------------------------------- oauth ------ | |
| oauth = None | |
| if OAUTH_READY: | |
| from authlib.integrations.starlette_client import OAuth | |
| oauth = OAuth() | |
| oauth.register( | |
| name="huggingface", | |
| client_id=OAUTH_CLIENT_ID, | |
| client_secret=OAUTH_CLIENT_SECRET, | |
| server_metadata_url=f"{OPENID_PROVIDER_URL}/.well-known/openid-configuration", | |
| client_kwargs={"scope": OAUTH_SCOPES}, | |
| ) | |
| def current_user(request: Request) -> Optional[dict]: | |
| u = request.session.get("user") | |
| if u and u.get("expires_at", 0) > time.time(): | |
| return u | |
| if u: | |
| request.session.pop("user", None) | |
| return None | |
| def require_user(request: Request) -> dict: | |
| u = current_user(request) | |
| if not u: | |
| raise HTTPException(401, "not signed in") | |
| return u | |
| def _redirect_uri(request: Request) -> str: | |
| if SPACE_HOST: | |
| return f"https://{SPACE_HOST}/auth/callback" | |
| return str(request.url_for("auth_callback")) | |
| async def login(request: Request, next: str = "/"): | |
| request.session["next"] = next | |
| if not OAUTH_READY: | |
| # Local development only — never reachable on a Space. | |
| request.session["user"] = { | |
| "username": os.environ.get("DEV_USER", "localdev"), | |
| "name": "Local Dev", "picture": None, "token": None, | |
| "expires_at": time.time() + 8 * 3600, | |
| } | |
| return RedirectResponse(next, status_code=303) | |
| return await oauth.huggingface.authorize_redirect(request, _redirect_uri(request)) | |
| async def auth_callback(request: Request): | |
| nxt = request.session.get("next", "/") or "/" | |
| try: | |
| token = await oauth.huggingface.authorize_access_token(request) | |
| except Exception as exc: | |
| if "state" in str(exc).lower() and request.cookies.get("nx_retry") != "1": | |
| # The state cookie didn't survive the round trip. Try once more | |
| # with a fresh session; the marker cookie stops this from looping. | |
| request.session.clear() | |
| r = RedirectResponse(f"/login?next={nxt}", status_code=303) | |
| r.set_cookie("nx_retry", "1", max_age=120, path="/", | |
| httponly=True, secure=IN_SPACE, | |
| samesite="none" if IN_SPACE else "lax") | |
| return r | |
| return render(request, "error.html", status_code=400, | |
| code="Sign-in failed", detail=str(exc)[:300]) | |
| info = token.get("userinfo") or await oauth.huggingface.userinfo(token=token) | |
| request.session["user"] = { | |
| "username": info.get("preferred_username") or info.get("name"), | |
| "name": info.get("name") or info.get("preferred_username"), | |
| "picture": info.get("picture"), | |
| "token": token.get("access_token"), | |
| "expires_at": time.time() + 8 * 3600, | |
| } | |
| r = RedirectResponse(nxt, status_code=303) | |
| r.delete_cookie("nx_retry", path="/") | |
| return r | |
| async def logout(request: Request): | |
| request.session.clear() | |
| return RedirectResponse("/", status_code=303) | |
| # ---------------------------------------------------------- serialization --- | |
| def job_dto(j, with_position=False): | |
| # A row written before a tier list change would otherwise KeyError here and | |
| # take down the whole page, so fall back to the first tier for display. | |
| tier = TIERS.get(j["tier"]) or TIERS[TIER_ORDER[0]] | |
| target = max(int(j["target_tokens"]), 1) | |
| # Step counts rarely divide the budget exactly, so a finished run would | |
| # otherwise sit at 99.9% forever. | |
| progress = (1.0 if j["status"] == "done" | |
| else min(1.0, (j["tokens_seen"] or 0) / target)) | |
| d = { | |
| "id": j["id"], | |
| "name": j["model_name"], | |
| "username": j["username"], | |
| "tier": j["tier"], | |
| "tier_label": tier.label, | |
| "n_params": j["n_params"], | |
| "status": j["status"], | |
| "target_tokens": j["target_tokens"], | |
| "tokens_seen": j["tokens_seen"] or 0, | |
| "progress": progress, | |
| "step": j["step"] or 0, | |
| "total_steps": j["total_steps"] or 0, | |
| "loss": j["loss"], | |
| "best_loss": j["best_loss"], | |
| "lr": j["lr"], | |
| "tok_per_s": j["tok_per_s"] or 0, | |
| "eta_s": j["eta_s"], | |
| "gpu_index": j["gpu_index"], | |
| "error": j["error"], | |
| "pushed_repo": j["pushed_repo"], | |
| "live_step": j["live_step"], | |
| "live_at": j["live_at"], | |
| "has_checkpoint": bool(j["live_step"] is not None) or j["status"] == "done", | |
| "created_at": j["created_at"], | |
| "started_at": j["started_at"], | |
| "finished_at": j["finished_at"], | |
| "arch": { | |
| "layers": tier.num_hidden_layers, "hidden": tier.hidden_size, | |
| "heads": tier.num_attention_heads, "kv_heads": tier.num_key_value_heads, | |
| "ffn": tier.intermediate_size, "head_dim": tier.head_dim, | |
| "ctx": SEQ_LEN, "vocab": VOCAB_SIZE, | |
| }, | |
| } | |
| if with_position and j["status"] == "queued": | |
| d["queue_position"] = db.queue_position(j["id"]) | |
| return d | |
| def _tok_label(n): | |
| return f"{n / 1e9:g}B" if n >= 1_000_000_000 else f"{n / 1e6:g}M" | |
| def tiers_dto(): | |
| out = [] | |
| for k in TIER_ORDER: | |
| t = TIERS[k] | |
| out.append({ | |
| "key": k, "label": t.label, "params": t.param_count(), | |
| "layers": t.num_hidden_layers, "hidden": t.hidden_size, | |
| "heads": t.num_attention_heads, "kv_heads": t.num_key_value_heads, | |
| "ffn": t.intermediate_size, "batch_tokens": t.batch_tokens, "lr": t.lr, | |
| }) | |
| return out | |
| # ------------------------------------------------------------- pages -------- | |
| def render(request: Request, template: str, status_code: int = 200, **kw): | |
| base = {"user": current_user(request), | |
| "on_gpu": ON_GPU, "n_devices": len(DEVICES)} | |
| base.update(kw) | |
| return templates.TemplateResponse(request, template, base, | |
| status_code=status_code) | |
| async def index(request: Request): | |
| u = current_user(request) | |
| if not u: | |
| return render(request, "landing.html", stats=db.stats(), | |
| tiers=tiers_dto(), min_tokens=MIN_TOKENS, | |
| max_tokens=MAX_TOKENS, | |
| recent=[job_dto(j) for j in db.list_public(limit=6)]) | |
| jobs = [job_dto(j, True) for j in db.list_jobs(username=u["username"], limit=8)] | |
| return render(request, "home.html", jobs=jobs, stats=db.stats()) | |
| async def train_page(request: Request): | |
| if not current_user(request): | |
| return RedirectResponse("/login?next=/train", status_code=303) | |
| return render(request, "train.html", tiers=tiers_dto(), | |
| min_tokens=MIN_TOKENS, max_tokens=MAX_TOKENS, | |
| stops=TOKEN_STOPS_M, | |
| max_active=MAX_ACTIVE_PER_USER, seq_len=SEQ_LEN, | |
| vocab=VOCAB_SIZE) | |
| async def models_page(request: Request): | |
| if not current_user(request): | |
| return RedirectResponse("/login?next=/models", status_code=303) | |
| return render(request, "models.html") | |
| async def model_detail(request: Request, job_id: str): | |
| # Readable by anyone — browsing the shelf shouldn't need a sign-in. | |
| # Generating, cancelling and publishing still do. | |
| j = db.get_job(job_id) | |
| if not j: | |
| return render(request, "error.html", status_code=404, | |
| code="404", detail="No run with that id.") | |
| me = current_user(request) | |
| return render(request, "detail.html", job=job_dto(j, True), | |
| infer_device=INFER_DEVICE, | |
| owner=bool(me and j["username"] == me["username"])) | |
| async def explore_page(request: Request): | |
| return render(request, "explore.html", tiers=tiers_dto()) | |
| async def public_profile(request: Request, username: str): | |
| prof = db.user_profile(username) | |
| if not prof: | |
| return render(request, "error.html", status_code=404, code="404", | |
| detail=f"@{username} hasn't trained anything here yet.") | |
| me = current_user(request) | |
| return render(request, "user.html", profile=prof, tiers=tiers_dto(), | |
| is_me=bool(me and me["username"] == username)) | |
| async def queue_page(request: Request): | |
| return render(request, "queue.html") | |
| async def profile_page(request: Request): | |
| u = current_user(request) | |
| if not u: | |
| return RedirectResponse("/login?next=/profile", status_code=303) | |
| jobs = db.list_jobs(username=u["username"], limit=500) | |
| done = [j for j in jobs if j["status"] == "done"] | |
| return render(request, "profile.html", summary={ | |
| "total": len(jobs), "done": len(done), | |
| "active": sum(1 for j in jobs if j["status"] in ("queued", "running")), | |
| "tokens": sum(j["tokens_seen"] or 0 for j in jobs), | |
| "params": sum(j["n_params"] for j in done), | |
| "published": sum(1 for j in done if j["pushed_repo"]), | |
| "best": min([j["best_loss"] for j in done if j["best_loss"]], default=None), | |
| }, | |
| jobs=[job_dto(j, True) for j in jobs[:20]]) | |
| async def about_page(request: Request): | |
| return render(request, "about.html", tiers=tiers_dto(), | |
| min_tokens=MIN_TOKENS, max_tokens=MAX_TOKENS, | |
| seq_len=SEQ_LEN, vocab=VOCAB_SIZE, | |
| cache_target=ndata.TARGET_TOKENS, | |
| max_active=MAX_ACTIVE_PER_USER) | |
| # --------------------------------------------------------------- api -------- | |
| async def api_me(request: Request): | |
| u = current_user(request) | |
| if not u: | |
| return {"signed_in": False} | |
| return {"signed_in": True, "username": u["username"], | |
| "name": u["name"], "picture": u["picture"]} | |
| async def api_tiers(): | |
| return {"tiers": tiers_dto(), "min_tokens": MIN_TOKENS, | |
| "max_tokens": MAX_TOKENS, "stops_m": TOKEN_STOPS_M} | |
| async def api_estimate(tier: str, tokens: int): | |
| if tier not in TIERS: | |
| raise HTTPException(400, "unknown tier") | |
| tokens = max(MIN_TOKENS, min(MAX_TOKENS, int(tokens))) | |
| s = estimate.summary(tier, tokens, ON_GPU) | |
| return {**s, "eta_human": estimate.human_time(s["eta_s"]), | |
| "chinchilla_x": tokens / (s["n_params"] * 20)} | |
| async def api_jobs(request: Request): | |
| u = require_user(request) | |
| return {"jobs": [job_dto(j, True) | |
| for j in db.list_jobs(username=u["username"], limit=200)]} | |
| async def api_job(job_id: str): | |
| j = db.get_job(job_id) | |
| if not j: | |
| raise HTTPException(404, "no such run") | |
| return job_dto(j, True) | |
| async def api_metrics(job_id: str): | |
| return {"metrics": db.get_metrics(job_id)} | |
| async def api_logs(job_id: str): | |
| return {"logs": db.get_logs(job_id, 300)} | |
| async def api_create(request: Request, payload: dict = Body(...)): | |
| u = require_user(request) | |
| tier = payload.get("tier") | |
| if tier not in TIERS: | |
| raise HTTPException(400, "Pick one of the four model sizes.") | |
| try: | |
| tokens = int(payload.get("tokens") or 0) | |
| except (TypeError, ValueError): | |
| raise HTTPException(400, "Token budget must be a number.") | |
| if not (MIN_TOKENS <= tokens <= MAX_TOKENS): | |
| raise HTTPException(400, f"Token budget must be between " | |
| f"{_tok_label(MIN_TOKENS)} and {_tok_label(MAX_TOKENS)}.") | |
| if db.user_active_count(u["username"]) >= MAX_ACTIVE_PER_USER: | |
| raise HTTPException(429, f"You already have {MAX_ACTIVE_PER_USER} runs in " | |
| "flight. Wait for one to finish, or cancel it.") | |
| import re | |
| name = re.sub(r"[^A-Za-z0-9._-]+", "-", (payload.get("name") or "").strip()).strip("-._")[:80] | |
| if not name: | |
| name = f"{TIERS[tier].label.lower()}-{tokens // 10**6}m" | |
| job_id = db.create_job( | |
| username=u["username"], display_name=u["name"], avatar_url=u["picture"], | |
| model_name=name, tier=tier, n_params=TIERS[tier].param_count(), | |
| target_tokens=tokens, total_steps=total_steps_for(tier, tokens), | |
| seed=int(time.time()) % 100000, | |
| ) | |
| db.add_log(job_id, f"queued by @{u['username']}") | |
| return {"id": job_id, "position": db.queue_position(job_id)} | |
| async def api_cancel(request: Request, job_id: str): | |
| u = require_user(request) | |
| j = db.get_job(job_id) | |
| if not j: | |
| raise HTTPException(404, "no such run") | |
| if j["username"] != u["username"]: | |
| raise HTTPException(403, "not your run") | |
| worker.request_cancel(job_id) | |
| return {"ok": True} | |
| def api_publish(request: Request, job_id: str, payload: dict = Body(default={})): | |
| u = require_user(request) | |
| j = db.get_job(job_id) | |
| if not j: | |
| raise HTTPException(404, "no such run") | |
| if j["username"] != u["username"]: | |
| raise HTTPException(403, "not your run") | |
| if j["status"] != "done": | |
| raise HTTPException(400, "This run hasn't finished yet.") | |
| if not u.get("token"): | |
| raise HTTPException(401, "Sign in again — your token expired.") | |
| out_dir = os.path.join(db.RUNS_DIR, job_id) | |
| if not os.path.exists(os.path.join(out_dir, "model.safetensors")): | |
| store.fetch_checkpoint(job_id, out_dir) | |
| if not os.path.exists(os.path.join(out_dir, "model.safetensors")): | |
| raise HTTPException(410, "Checkpoint could not be restored from storage.") | |
| import re | |
| name = re.sub(r"[^A-Za-z0-9._-]+", "-", | |
| (payload.get("repo") or j["model_name"]).strip()).strip("-._")[:80] | |
| try: | |
| repo_id = hub.push_run(out_dir, j, u["token"], name or j["model_name"]) | |
| except Exception as exc: | |
| raise HTTPException(500, f"Upload failed: {exc}") | |
| db.update_job(job_id, pushed_repo=repo_id) | |
| return {"repo": repo_id, "url": f"https://huggingface.co/{repo_id}"} | |
| # One cached inference model, keyed by (job, checkpoint) so a newer snapshot | |
| # replaces it automatically. Generation runs on the GPU when there is one — | |
| # a 3M-parameter model costs nothing next to the training job sharing it. | |
| _pg = {"key": None, "model": None} | |
| _pg_lock = threading.Lock() | |
| def _resolve_checkpoint(j): | |
| """Final weights if the run finished, otherwise the newest live snapshot.""" | |
| final = os.path.join(db.RUNS_DIR, j["id"]) | |
| if j["status"] == "done": | |
| if not os.path.exists(os.path.join(final, "model.safetensors")): | |
| # The container was rebuilt since this run finished — pull the | |
| # weights back out of the bucket. | |
| store.fetch_checkpoint(j["id"], final) | |
| if os.path.exists(os.path.join(final, "model.safetensors")): | |
| return final, ("final", j["step"]) | |
| live = os.path.join(final, "live") | |
| if os.path.exists(os.path.join(live, "model.safetensors")): | |
| return live, ("live", j["live_step"]) | |
| return None, (None, None) | |
| def api_generate(request: Request, job_id: str, payload: dict = Body(...)): | |
| # Anyone signed in can generate from anyone's model — that is the point of | |
| # a shared shelf. Cancelling and publishing stay with the owner. | |
| require_user(request) | |
| j = db.get_job(job_id) | |
| if not j: | |
| raise HTTPException(404, "no such run") | |
| path, (kind, step) = _resolve_checkpoint(j) | |
| if not path: | |
| if j["status"] == "queued": | |
| raise HTTPException(409, "This run hasn't started yet — " | |
| "the first checkpoint appears a few steps in.") | |
| if j["status"] == "running": | |
| raise HTTPException(409, "No checkpoint saved yet. The first one lands " | |
| "within the first few percent of the run.") | |
| raise HTTPException(410, "Checkpoint is gone — the Space restarted.") | |
| import torch | |
| from transformers import AutoModelForCausalLM | |
| key = (path, step) | |
| with _pg_lock: | |
| if _pg["key"] != key: | |
| m = AutoModelForCausalLM.from_pretrained(path, dtype=torch.float32) | |
| m.config.use_cache = True | |
| _pg["model"] = m.eval().to(INFER_DEVICE) | |
| _pg["key"] = key | |
| model = _pg["model"] | |
| prompt = payload.get("prompt") or " " | |
| ids = TOKENIZER.encode(prompt).ids[-(SEQ_LEN - 1):] or [0] | |
| temp = float(payload.get("temperature", 0.85)) | |
| x = torch.tensor([ids], dtype=torch.long, device=INFER_DEVICE) | |
| t0 = time.time() | |
| with torch.no_grad(): | |
| out = model.generate( | |
| x, | |
| max_new_tokens=max(1, min(int(payload.get("max_new_tokens", 80)), 256)), | |
| do_sample=temp > 0.01, temperature=max(temp, 1e-4), | |
| top_k=int(payload.get("top_k", 40)), | |
| top_p=float(payload.get("top_p", 0.95)), | |
| pad_token_id=1, eos_token_id=None, | |
| ) | |
| ids_out = out[0].tolist() | |
| return { | |
| "text": TOKENIZER.decode(ids_out), | |
| "prompt": prompt, | |
| "continuation": TOKENIZER.decode(ids_out[len(ids):]), | |
| "checkpoint": kind, | |
| "step": step, | |
| "total_steps": j["total_steps"], | |
| "device": INFER_DEVICE, | |
| "ms": round((time.time() - t0) * 1000), | |
| "new_tokens": len(ids_out) - len(ids), | |
| } | |
| async def api_explore(tier: str = "", sort: str = "recent", | |
| limit: int = 60, offset: int = 0, username: str = ""): | |
| limit = max(1, min(limit, 120)) | |
| rows = db.list_public(tier=tier or None, sort=sort, limit=limit, | |
| offset=max(0, offset), username=username or None) | |
| return {"models": [job_dto(j) for j in rows], | |
| "active": [job_dto(j, True) for j in db.list_active_public()], | |
| "stats": db.stats()} | |
| async def api_leaderboard(): | |
| lb = db.leaderboard() | |
| return {"tiers": [ | |
| {"tier": k, "label": (TIERS.get(k) or TIERS[TIER_ORDER[0]]).label, | |
| "order": TIER_ORDER.index(k) if k in TIER_ORDER else 99, | |
| "models": [job_dto(j) for j in v]} | |
| for k, v in sorted(lb.items(), | |
| key=lambda kv: TIER_ORDER.index(kv[0]) | |
| if kv[0] in TIER_ORDER else 99)]} | |
| async def api_user(username: str): | |
| prof = db.user_profile(username) | |
| if not prof: | |
| raise HTTPException(404, "no such trainer") | |
| return {"profile": prof, | |
| "models": [job_dto(j) for j in | |
| db.list_public(username=username, limit=100)], | |
| "active": [job_dto(j, True) for j in | |
| db.list_jobs(username=username, limit=100) | |
| if j["status"] in ("queued", "running")]} | |
| async def api_trainers(): | |
| return {"trainers": db.top_trainers()} | |
| async def api_queue(): | |
| infos = worker.worker_info() | |
| return { | |
| "workers": [{"index": w["index"], "device": w["device"], | |
| "job": job_dto(w["job"]) if w["job"] else None} | |
| for w in infos], | |
| "running": [job_dto(j) for j in db.list_jobs(status=["running"], limit=64)], | |
| "queued": [job_dto(j, True) for j in db.list_jobs(status=["queued"], limit=64)], | |
| "recent": [job_dto(j) for j in db.list_jobs(status=["done"], limit=12)], | |
| "stats": db.stats(), | |
| "data": {**ndata.state(), "target": ndata.TARGET_TOKENS}, | |
| "on_gpu": ON_GPU, | |
| "store": {**store.state(), "bucket": store.BUCKET_ID}, | |
| } | |
| async def health(): | |
| return {"ok": True, "devices": DEVICES, "data": ndata.state(), | |
| "stats": db.stats(), "store": store.state(), | |
| "bucket": store.BUCKET_ID} | |
| async def http_error(request: Request, exc: HTTPException): | |
| if request.url.path.startswith("/api/"): | |
| return JSONResponse({"error": exc.detail}, status_code=exc.status_code) | |
| if exc.status_code == 401: | |
| return RedirectResponse(f"/login?next={request.url.path}", status_code=303) | |
| return render(request, "error.html", status_code=exc.status_code, | |
| code=str(exc.status_code), detail=exc.detail) | |
| if __name__ == "__main__": | |
| import uvicorn | |
| uvicorn.run(app, host="0.0.0.0", port=int(os.environ.get("PORT", 7860))) | |