Spaces:
Paused
Paused
Download server.py from hugging-science/storage: direct link, hf CLI and curl.
- Browser
- Download file 15.1 kB
-
https://huggingface.co/spaces/hugging-science/storage/resolve/main/server.py
- Command line
-
hf download hf://spaces/hugging-science/storage/server.py
-
curl -L -o server.py https://huggingface.co/spaces/hugging-science/storage/resolve/main/server.py
15.1 kB
| 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("<I", len(h)) + h + body | |
| def unpack(b: bytes): | |
| (hl,) = struct.unpack("<I", b[:4]); return json.loads(b[4:4 + hl]), b[4 + hl:] | |
| # ---------------- datasets ---------------- | |
| def _fetch(url): | |
| with urllib.request.urlopen(url, timeout=180) as r: return r.read() | |
| def load_mnist(): | |
| p = f"{DS_DIR}/mnist.npz" | |
| if not os.path.exists(p): | |
| base = "https://storage.googleapis.com/cvdf-datasets/mnist/" | |
| xi = gzip.decompress(_fetch(base + "train-images-idx3-ubyte.gz")) | |
| yi = gzip.decompress(_fetch(base + "train-labels-idx1-ubyte.gz")) | |
| np.savez(p, x=np.frombuffer(xi, "u1", offset=16).reshape(-1, 784), y=np.frombuffer(yi, "u1", offset=8)) | |
| d = np.load(p); return d["x"], d["y"], None | |
| def load_shakespeare(ctx): | |
| p = f"{DS_DIR}/tinyshakespeare.txt" | |
| if not os.path.exists(p): | |
| open(p, "wb").write(_fetch("https://raw.githubusercontent.com/karpathy/char-rnn/master/data/tinyshakespeare/input.txt")) | |
| txt = open(p, encoding="utf-8").read() | |
| chars = sorted(set(txt)); stoi = {c: i for i, c in enumerate(chars)} | |
| ids = np.array([stoi[c] for c in txt], "u1") | |
| x = np.ascontiguousarray(np.lib.stride_tricks.sliding_window_view(ids[:-1], ctx)) # (N-ctx, ctx) | |
| return x, ids[ctx:], chars | |
| def load_dataset(template): | |
| t = TEMPLATES[template] | |
| return load_mnist() if t["dataset"] == "mnist" else load_shakespeare(t["ctx"]) | |
| # ---------------- model helpers (numpy, for init + validation only) ---------------- | |
| def init_weights(layers, seed): | |
| rng = np.random.default_rng(seed); W = [] | |
| for a, b in zip(layers[:-1], layers[1:]): | |
| W += [(rng.standard_normal((a, b)) * np.sqrt(2 / a)).astype("<f4"), np.zeros(b, "<f4")] | |
| return W | |
| def encode(x, j): | |
| if j["encoding"] == "pixels": return x.astype("f4") / 255.0 | |
| return np.eye(j["vocab"], dtype="f4")[x].reshape(len(x), -1) | |
| def forward(W, X): | |
| h = X | |
| for i in range(0, len(W), 2): | |
| h = h @ W[i] + W[i + 1] | |
| if i < len(W) - 2: h = np.maximum(h, 0) | |
| return h | |
| def evaluate(W, X, Y): | |
| z = forward(W, X); m = z.max(1, keepdims=True) | |
| lse = (m + np.log(np.exp(z - m).sum(1, keepdims=True)))[:, 0] | |
| loss = float(np.mean(lse - z[np.arange(len(Y)), Y])); acc = float(np.mean(z.argmax(1) == Y)) | |
| return loss, acc | |
| # ---------------- persistence ---------------- | |
| def job_dir(jid): return f"{JOBS_DIR}/{jid}" | |
| def save(job): | |
| d = job_dir(job["j"]["id"]); os.makedirs(d, exist_ok=True) | |
| tmp = f"{d}/state.json.tmp"; json.dump(job["j"], open(tmp, "w")); os.replace(tmp, f"{d}/state.json") | |
| open(f"{d}/weights.bin", "wb").write(pack({"shapes": [list(w.shape) for w in job["W"]], "round": job["j"]["round"]}, | |
| b"".join(w.tobytes() for w in job["W"]))) | |
| def hub_api(): | |
| from huggingface_hub import HfApi; return HfApi(token=HF_TOKEN) | |
| def hub_save(jid): | |
| if not (HF_TOKEN and CKPT_REPO): return | |
| try: | |
| api = hub_api(); d = job_dir(jid); r = JOBS[jid]["j"]["round"] | |
| api.upload_file(path_or_fileobj=f"{d}/state.json", path_in_repo=f"jobs/{jid}/state.json", repo_id=CKPT_REPO) | |
| api.upload_file(path_or_fileobj=f"{d}/weights.bin", path_in_repo=f"jobs/{jid}/weights.bin", repo_id=CKPT_REPO) | |
| api.upload_file(path_or_fileobj=f"{d}/weights.bin", path_in_repo=f"jobs/{jid}/history/round_{r:05d}.bin", repo_id=CKPT_REPO) | |
| except Exception as e: print("hub_save failed:", e) | |
| def hub_restore(): | |
| if not (HF_TOKEN and CKPT_REPO): return | |
| try: | |
| from huggingface_hub import hf_hub_download | |
| api = hub_api(); api.create_repo(CKPT_REPO, exist_ok=True, private=False) | |
| for f in api.list_repo_files(CKPT_REPO): | |
| if f.startswith("jobs/") and f.endswith("/state.json"): | |
| jid = f.split("/")[1] | |
| if os.path.exists(f"{job_dir(jid)}/state.json"): continue | |
| hf_hub_download(CKPT_REPO, f"jobs/{jid}/state.json", local_dir=DATA_DIR) | |
| hf_hub_download(CKPT_REPO, f"jobs/{jid}/weights.bin", local_dir=DATA_DIR) | |
| print("restored", jid, "from Hub") | |
| except Exception as e: print("hub_restore failed:", e) | |
| def build_shards(j, x, y): | |
| rng = np.random.default_rng(j["seed"]); perm = rng.permutation(len(x)); x, y = x[perm], y[perm] | |
| d = job_dir(j["id"]); os.makedirs(f"{d}/shards", exist_ok=True) | |
| np.savez(f"{d}/val.npz", x=x[:VAL_N], y=y[:VAL_N]) | |
| for i, (xs, ys) in enumerate(zip(np.array_split(x[VAL_N:], j["shards"]), np.array_split(y[VAL_N:], j["shards"]))): | |
| hdr = {"n": len(xs), "x_shape": list(xs.shape), "encoding": j["encoding"], "vocab": j["vocab"]} | |
| open(f"{d}/shards/{i}.bin", "wb").write(pack(hdr, xs.tobytes() + ys.tobytes())) | |
| def load_job(jid): | |
| d = job_dir(jid); j = json.load(open(f"{d}/state.json")) | |
| hdr, body = unpack(open(f"{d}/weights.bin", "rb").read()) | |
| arr = np.frombuffer(body, "<f4"); W, off = [], 0 | |
| for s in hdr["shapes"]: | |
| n = int(np.prod(s)); W.append(arr[off:off + n].reshape(s).copy()); off += n | |
| if not os.path.exists(f"{d}/val.npz") or len(glob.glob(f"{d}/shards/*.bin")) != j["shards"]: | |
| x, y, _ = load_dataset(j["template"]); build_shards(j, x, y) # shards are rebuildable | |
| v = np.load(f"{d}/val.npz") | |
| JOBS[jid] = {"j": j, "W": W, "pending": [], "val": (encode(v["x"], j), v["y"])} | |
| def startup(): | |
| hub_restore() | |
| for p in glob.glob(f"{JOBS_DIR}/*/state.json"): | |
| try: load_job(p.split("/")[-2]) | |
| except Exception as e: print("load_job failed:", p, e) | |
| print(f"DATA_DIR={DATA_DIR} jobs={list(JOBS)}") | |
| # ---------------- helpers ---------------- | |
| def public(j): | |
| keys = ["id", "name", "owner", "template", "encoding", "ctx", "vocab", "layers", "round", "status", | |
| "target_rounds", "local_steps", "batch", "lr", "shards", "k", "val_loss", "created"] | |
| return {k: j.get(k) for k in keys} | |
| def state_receipt(j): | |
| s = {"kind": "state", "job": j["id"], "round": j["round"], "status": j["status"], "ts": int(time.time())} | |
| return {"state": s, "sig": sign(s)} | |
| def get(jid): | |
| job = JOBS.get(jid) | |
| if not job: raise HTTPException(404, "unknown job") | |
| return job | |
| def aggregate(job): | |
| j, P = job["j"], job["pending"] | |
| job["W"] = [np.mean([p["w"][i] for p in P], axis=0).astype("<f4") for i in range(len(job["W"]))] | |
| job["pending"] = []; j["round"] += 1 | |
| loss, acc = evaluate(job["W"], *job["val"]); j["val_loss"] = loss | |
| j["metrics"].append({"round": j["round"], "val_loss": round(loss, 4), "val_acc": round(acc, 4), | |
| "updates": len(P), "ts": int(time.time())}) | |
| j["metrics"] = j["metrics"][-500:] | |
| if j["round"] >= 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 ---------------- | |
| def root(): return {"ok": True, "data_dir": DATA_DIR, "persistent": DATA_DIR == "/data", "jobs": len(JOBS)} | |
| def list_jobs(): return [public(job["j"]) for job in JOBS.values()] | |
| 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) | |
| def job_info(jid: str): return public(get(jid)["j"]) | |
| def job_status(jid: str): return state_receipt(get(jid)["j"]) | |
| def job_metrics(jid: str): | |
| j = get(jid)["j"]; return {"metrics": j["metrics"], "round": j["round"], "status": j["status"], "chars": j.get("chars")} | |
| 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"'}) | |
| 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") | |
| 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, "<f4"); new, off = [], 0 | |
| for s in shapes: | |
| n = int(np.prod(s)); new.append(arr[off:off + n].reshape(s)); off += n | |
| if not all(np.isfinite(w).all() for w in new): raise HTTPException(400, "non-finite values") | |
| delta = float(np.sqrt(sum(float(((a - b) ** 2).sum()) for a, b in zip(new, job["W"])))) | |
| wnorm = float(np.sqrt(sum(float((b ** 2).sum()) for b in job["W"]))) | |
| if delta > 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)} | |
| 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"]) | |
| 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} |