storage / server.py
Bc-AI's picture
Update server.py
0b2b232 verified
Raw History Blame Contribute Delete
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"])}
@app.on_event("startup")
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 ----------------
@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, "<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)}
@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}