""" 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")) @app.get("/login") 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)) @app.get("/auth/callback", name="auth_callback") 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 @app.get("/logout") 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) @app.get("/", response_class=HTMLResponse) 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()) @app.get("/train", response_class=HTMLResponse) 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) @app.get("/models", response_class=HTMLResponse) async def models_page(request: Request): if not current_user(request): return RedirectResponse("/login?next=/models", status_code=303) return render(request, "models.html") @app.get("/models/{job_id}", response_class=HTMLResponse) 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"])) @app.get("/explore", response_class=HTMLResponse) async def explore_page(request: Request): return render(request, "explore.html", tiers=tiers_dto()) @app.get("/u/{username}", response_class=HTMLResponse) 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)) @app.get("/queue", response_class=HTMLResponse) async def queue_page(request: Request): return render(request, "queue.html") @app.get("/profile", response_class=HTMLResponse) 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]]) @app.get("/about", response_class=HTMLResponse) 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 -------- @app.get("/api/me") 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"]} @app.get("/api/tiers") async def api_tiers(): return {"tiers": tiers_dto(), "min_tokens": MIN_TOKENS, "max_tokens": MAX_TOKENS, "stops_m": TOKEN_STOPS_M} @app.get("/api/estimate") 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)} @app.get("/api/jobs") 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)]} @app.get("/api/jobs/{job_id}") 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) @app.get("/api/jobs/{job_id}/metrics") async def api_metrics(job_id: str): return {"metrics": db.get_metrics(job_id)} @app.get("/api/jobs/{job_id}/logs") async def api_logs(job_id: str): return {"logs": db.get_logs(job_id, 300)} @app.post("/api/jobs") 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)} @app.post("/api/jobs/{job_id}/cancel") 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} @app.post("/api/jobs/{job_id}/publish") 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) @app.post("/api/jobs/{job_id}/generate") 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), } @app.get("/api/explore") 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()} @app.get("/api/leaderboard") 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)]} @app.get("/api/users/{username}") 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")]} @app.get("/api/trainers") async def api_trainers(): return {"trainers": db.top_trainers()} @app.get("/api/queue") 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}, } @app.get("/api/health") async def health(): return {"ok": True, "devices": DEVICES, "data": ndata.state(), "stats": db.stats(), "store": store.state(), "bucket": store.BUCKET_ID} @app.exception_handler(HTTPException) 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)))