import os, json, time, hmac, hashlib, struct, threading, gzip, urllib.request, glob import numpy as np from fastapi import FastAPI, HTTPException, Request, Header from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import Response # ---------------- config ---------------- SECRET = "secret-e".encode() # same value as on PythonAnywhere ADMIN_KEY = "admin-ke!" HF_TOKEN = os.environ.get("HF_TOKEN") # write token (optional but recommended) CKPT_REPO = "hugging-science/train-browser-model" # e.g. "yourname/train-browser-ckpts" ORIGINS = [o.strip() for o in os.environ.get("ALLOWED_ORIGINS", "*").split(",")] DATA_DIR = os.environ.get("DATA_DIR") or ("/data" if os.path.isdir("/data") else "/app/data") JOBS_DIR, DS_DIR = f"{DATA_DIR}/jobs", f"{DATA_DIR}/datasets" os.makedirs(JOBS_DIR, exist_ok=True); os.makedirs(DS_DIR, exist_ok=True) VAL_N = 2000 # held-out validation samples per job (never sent to browsers) AGG_WAIT = 180 # seconds: aggregate even if fewer than k updates are pending CKPT_EVERY = 5 # rounds between Hub checkpoints MAX_BODY = 64 * 1024 * 1024 TEMPLATES = { "mnist_mlp": dict(dataset="mnist", encoding="pixels", ctx=1), "char_lm": dict(dataset="tinyshakespeare", encoding="onehot", ctx=8), } LIMITS = dict(shards=(2, 200), k=(1, 20), target_rounds=(1, 500), local_steps=(5, 500), batch=(8, 128)) app = FastAPI(title="Train Browser Space") app.add_middleware(CORSMiddleware, allow_origins=ORIGINS, allow_methods=["*"], allow_headers=["*"]) LOCK = threading.Lock() JOBS = {} # id -> {"j": state dict, "W": [np arrays], "pending": [...], "val": (Xenc, Y)} # ---------------- signing / binary format ---------------- def canon(d): return json.dumps(d, sort_keys=True, separators=(",", ":")).encode() def sign(d): return hmac.new(SECRET, canon(d), hashlib.sha256).hexdigest() def verify(d, sig): return bool(sig) and hmac.compare_digest(sign(d), sig) def check_ticket(t, sig, kind, job_id=None): if not isinstance(t, dict) or not verify(t, sig) or t.get("kind") != kind \ or float(t.get("exp", 0)) < time.time() or (job_id and t.get("job") != job_id): raise HTTPException(401, "bad or expired ticket") def pack(hdr, body: bytes) -> bytes: h = json.dumps(hdr).encode(); return struct.pack("= j["target_rounds"]: j["status"] = "done" save(job) if j["round"] % CKPT_EVERY == 0 or j["status"] == "done": threading.Thread(target=hub_save, args=(j["id"],), daemon=True).start() # ---------------- routes ---------------- @app.get("/") def root(): return {"ok": True, "data_dir": DATA_DIR, "persistent": DATA_DIR == "/data", "jobs": len(JOBS)} @app.get("/jobs") def list_jobs(): return [public(job["j"]) for job in JOBS.values()] @app.post("/jobs") async def create_job(req: Request): d = await req.json(); t = d.get("ticket"); check_ticket(t, d.get("sig"), "create_job") jid = t["job"] if jid in JOBS: return state_receipt(JOBS[jid]["j"]) if t["template"] not in TEMPLATES: raise HTTPException(400, "unknown template") for key, (lo, hi) in LIMITS.items(): if not lo <= int(t[key]) <= hi: raise HTTPException(400, f"{key} out of range") hidden = [int(h) for h in t["hidden"]][:3] if any(not 4 <= h <= 1024 for h in hidden): raise HTTPException(400, "hidden out of range") tp = TEMPLATES[t["template"]] x, y, chars = load_dataset(t["template"]) vocab = len(chars) if chars else 0 in_dim, out_dim = (784, 10) if tp["encoding"] == "pixels" else (tp["ctx"] * vocab, vocab) j = {"id": jid, "name": str(t.get("name", jid))[:60], "owner": t["owner"], "template": t["template"], "encoding": tp["encoding"], "ctx": tp["ctx"], "vocab": vocab, "chars": chars, "layers": [in_dim] + hidden + [out_dim], "shards": int(t["shards"]), "k": int(t["k"]), "target_rounds": int(t["target_rounds"]), "local_steps": int(t["local_steps"]), "batch": int(t["batch"]), "lr": float(t["lr"]), "seed": int(t.get("seed", 0)), "round": 0, "status": "running", "metrics": [], "val_loss": None, "created": int(time.time())} with LOCK: build_shards(j, x, y) W = init_weights(j["layers"], j["seed"]) v = np.load(f"{job_dir(jid)}/val.npz"); val = (encode(v["x"], j), v["y"]) j["val_loss"] = evaluate(W, *val)[0] JOBS[jid] = {"j": j, "W": W, "pending": [], "val": val}; save(JOBS[jid]) threading.Thread(target=hub_save, args=(jid,), daemon=True).start() return state_receipt(j) @app.get("/jobs/{jid}") def job_info(jid: str): return public(get(jid)["j"]) @app.get("/jobs/{jid}/status") def job_status(jid: str): return state_receipt(get(jid)["j"]) @app.get("/jobs/{jid}/metrics") def job_metrics(jid: str): j = get(jid)["j"]; return {"metrics": j["metrics"], "round": j["round"], "status": j["status"], "chars": j.get("chars")} @app.get("/jobs/{jid}/weights") def job_weights(jid: str): job = get(jid) with LOCK: hdr = public(job["j"]) | {"shapes": [list(w.shape) for w in job["W"]]} body = b"".join(w.tobytes() for w in job["W"]) return Response(pack(hdr, body), media_type="application/octet-stream", headers={"Content-Disposition": f'attachment; filename="{jid}_round{job["j"]["round"]}.bin"'}) @app.get("/jobs/{jid}/shards/{i}") def job_shard(jid: str, i: int): get(jid); p = f"{job_dir(jid)}/shards/{i}.bin" if not os.path.exists(p): raise HTTPException(404, "no such shard") return Response(open(p, "rb").read(), media_type="application/octet-stream") @app.post("/jobs/{jid}/submit") async def submit(jid: str, req: Request, x_ticket: str = Header(None), x_sig: str = Header(None)): job = get(jid); j = job["j"] try: t = json.loads(x_ticket or "") except Exception: raise HTTPException(401, "missing ticket") check_ticket(t, x_sig, "task", jid) if int(req.headers.get("content-length", "0")) > MAX_BODY: raise HTTPException(413) hdr, body = unpack(await req.body()) with LOCK: if j["status"] != "running": raise HTTPException(409, f"job is {j['status']}") if int(hdr.get("round", -1)) != j["round"]: raise HTTPException(409, "stale round, refetch weights") shapes = [w.shape for w in job["W"]] if len(body) != sum(int(np.prod(s)) for s in shapes) * 4: raise HTTPException(400, "wrong payload size") arr = np.frombuffer(body, " 2.0 * wnorm + 20: raise HTTPException(400, f"implausible update (delta={delta:.1f})") loss, acc = evaluate(new, *job["val"]) if not np.isfinite(loss) or loss > j["val_loss"] * 1.15 + 0.25: raise HTTPException(400, f"update hurts validation loss ({loss:.3f})") job["pending"].append({"w": new, "ts": time.time(), "user": t["user"]}) rnd = j["round"] if len(job["pending"]) >= j["k"] or time.time() - job["pending"][0]["ts"] > AGG_WAIT: aggregate(job) receipt = {"kind": "task", "task_id": t["task_id"], "user": t["user"], "job": jid, "round": rnd, "new_round": j["round"], "status": j["status"], "val_loss": round(loss, 4), "val_acc": round(acc, 4), "ts": int(time.time())} return {"receipt": receipt, "sig": sign(receipt)} @app.post("/jobs/{jid}/set_status") async def set_status(jid: str, req: Request): d = await req.json(); t = d.get("ticket"); check_ticket(t, d.get("sig"), "set_status", jid) job = get(jid) with LOCK: if t["status"] in ("running", "paused", "cancelled") and job["j"]["status"] != "done": job["j"]["status"] = t["status"]; save(job) return state_receipt(job["j"]) @app.delete("/jobs/{jid}") def delete_job(jid: str, x_admin_key: str = Header(None)): if not ADMIN_KEY or x_admin_key != ADMIN_KEY: raise HTTPException(401) import shutil with LOCK: JOBS.pop(jid, None); shutil.rmtree(job_dir(jid), ignore_errors=True) return {"ok": True}