Spaces:
Runtime error
Runtime error
| """WASD-controlled push_cube world model (AR shortcut-forcing DiT-S, 5-step inference). | |
| HF Space edition: loads a stripped fp16 checkpoint + Wan-VAE from the repo, initial frames from | |
| frames/*.png. Same stdlib HTTP server + HTML page as the cluster demo. CPU or GPU (auto). | |
| """ | |
| import os, sys, io, json, time, threading, socket, base64, glob | |
| from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer | |
| import numpy as np | |
| import torch | |
| import yaml | |
| ROOT = os.path.dirname(os.path.abspath(__file__)) | |
| sys.path.insert(0, ROOT) | |
| # weights live in a separate public model repo (Space repo has a 1 GB cap) | |
| from huggingface_hub import hf_hub_download | |
| WREPO = "lamamkh/push-cube-wasd-worldmodel-weights" | |
| CKPT_PATH = hf_hub_download(WREPO, "model_fp16.pt") | |
| VAE_PATH = hf_hub_download(WREPO, "Wan2.1_VAE.pth") | |
| os.environ["WAN_VAE_PATH"] = VAE_PATH | |
| from acwm.model.interface import get_dynamics_class | |
| DEVICE = "cuda" if torch.cuda.is_available() else "cpu" | |
| if DEVICE == "cpu": | |
| torch.set_num_threads(max(1, os.cpu_count() or 2)) | |
| PORT = int(os.environ.get("PORT", "7860")) | |
| K_STEPS = 5 | |
| MAX_LATENT = 10 | |
| print(f"[wasd] device={DEVICE}", flush=True) | |
| cfg = yaml.safe_load(open(os.path.join(ROOT, "push_cube_ar.yaml"))) | |
| mc = cfg["model_config"]; mc["action_dim"] = 2 | |
| mc["use_flash_attn"] = torch.cuda.is_available() | |
| WM = get_dynamics_class(cfg["dynamics_class"])(cfg["model_name"], mc).to(DEVICE).eval() | |
| _ck = torch.load(CKPT_PATH, map_location="cpu", weights_only=False) | |
| _sd = {k: (v.float() if v.is_floating_point() else v) for k, v in _ck["model_state_dict"].items()} | |
| WM.load_state_dict(_sd, strict=False) | |
| WM = WM.float() | |
| print(f"[wasd] loaded step={_ck.get('step','?')}", flush=True) | |
| ACR = WM.model.action_compress_rate | |
| KMAX = WM.k_max | |
| FRAMES = sorted(glob.glob(os.path.join(ROOT, "frames", "*.png"))) | |
| def load_frame(idx): | |
| import imageio.v2 as iio | |
| img = iio.imread(FRAMES[idx % len(FRAMES)]).astype(np.float32) / 255.0 | |
| return torch.from_numpy(img) # [H,W,3] | |
| def ar_one_frame(z_all, a_curr): | |
| h, w, D = z_all.shape[2], z_all.shape[3], z_all.shape[4] | |
| t_len = z_all.shape[1] + 1 | |
| K = K_STEPS; d = 1.0 / K; d_min = 1.0 / KMAX | |
| z = torch.randn(1, 1, h, w, D, device=DEVICE) | |
| for i in range(K): | |
| seq = torch.cat([z_all, z], dim=1) | |
| sig = torch.ones(1, t_len, device=DEVICE); sig[:, -1] = i / K | |
| ds = torch.full((1, t_len), d_min, device=DEVICE); ds[:, -1] = d | |
| b = WM.model(seq, sig, ds, a_curr) | |
| z = z + b[:, -1:] * d | |
| return z | |
| def decode_last(z_all): | |
| vid = WM.vae.decode_to_pixel(z_all.permute(0, 1, 4, 2, 3).contiguous()) | |
| vid = ((vid + 1) / 2).clamp(0, 1)[0].permute(0, 2, 3, 1).cpu().numpy() | |
| return vid[-1] | |
| def png_b64(img01): | |
| import imageio.v2 as iio | |
| buf = io.BytesIO() | |
| iio.imwrite(buf, (np.clip(img01, 0, 1) * 255).astype(np.uint8), format="png") | |
| return "data:image/png;base64," + base64.b64encode(buf.getvalue()).decode() | |
| LOCK = threading.Lock() | |
| S = {"z": None, "acts": []} | |
| def reset(idx): | |
| f0 = load_frame(int(idx)) | |
| o0 = f0.unsqueeze(0).to(DEVICE) # [1,H,W,3] | |
| S["z"] = WM.encode_obs(o0.unsqueeze(1)); S["acts"] = [] | |
| return png_b64(f0.numpy()), 1 | |
| def step(dx, dy): | |
| v = np.array([dx, dy], np.float32); n = float(np.linalg.norm(v)) | |
| a = v / n if n > 1e-6 else np.array([1.0, 0.0], np.float32) | |
| S["acts"].extend([a] * ((ACR + 1) if len(S["acts"]) == 0 else ACR)) | |
| a_curr = torch.tensor(np.stack(S["acts"]), device=DEVICE).unsqueeze(0).float() | |
| S["z"] = torch.cat([S["z"], ar_one_frame(S["z"], a_curr)], dim=1) | |
| if S["z"].shape[1] > MAX_LATENT: | |
| S["z"] = S["z"][:, 1:]; S["acts"] = S["acts"][ACR:] | |
| return png_b64(decode_last(S["z"])), S["z"].shape[1], [float(a[0]), float(a[1])] | |
| PAGE = """<!doctype html><html><head><meta charset="utf-8"><title>push_cube · WASD world model</title><style> | |
| :root{--bg:#0e0f13;--panel:#16181f;--ink:#ecebe6;--soft:#9aa0ad;--amber:#f59e42;--line:#24262f; | |
| --mono:ui-monospace,Menlo,Consolas,monospace;--sans:ui-sans-serif,system-ui,sans-serif} | |
| *{box-sizing:border-box}body{margin:0;background:var(--bg);color:var(--ink);font-family:var(--sans); | |
| display:flex;min-height:100vh;align-items:center;justify-content:center} | |
| .app{display:flex;gap:28px;padding:28px;flex-wrap:wrap;justify-content:center} | |
| .stage{background:var(--panel);border:1px solid var(--line);border-radius:14px;padding:16px} | |
| #frame{width:min(480px,88vw);aspect-ratio:1;border-radius:8px;display:block;background:#000} | |
| #frame:not([src]){opacity:0} | |
| .hud{font-family:var(--mono);font-size:12.5px;color:var(--soft);margin-top:10px;display:flex; | |
| justify-content:space-between} | |
| .side{max-width:290px;display:flex;flex-direction:column;gap:16px} | |
| h1{font-size:20px;margin:0}p{color:var(--soft);font-size:13.5px;line-height:1.5;margin:0} | |
| .keys{display:grid;grid-template-columns:repeat(3,54px);gap:6px;justify-content:center} | |
| .key{height:54px;border:1px solid var(--line);border-radius:9px;display:flex;align-items:center; | |
| justify-content:center;font-family:var(--mono);font-size:16px;color:var(--soft);background:var(--panel); | |
| transition:all .08s;cursor:pointer;user-select:none}.key.on{background:var(--amber);color:#14161c;border-color:var(--amber)} | |
| select,button{background:var(--panel);color:var(--ink);border:1px solid var(--line);border-radius:8px; | |
| padding:8px 12px;font-family:var(--mono);font-size:13px;cursor:pointer}button:hover{border-color:var(--amber)} | |
| .row{display:flex;gap:8px;align-items:center}.busy #frame{opacity:.7} | |
| #spin{display:none;color:var(--amber)}.busy #spin{display:inline}</style></head><body> | |
| <div class="app" id="app"> | |
| <div class="stage"><img id="frame" alt="world model frame"> | |
| <div class="hud"><span id="stat">press R or Reset to start</span><span><span id="spin">thinking…</span> <span id="lat"></span></span></div></div> | |
| <div class="side"><h1>Drive the world model</h1> | |
| <p><b style="color:var(--amber)">WASD</b> (or tap the keys) pushes the disk. Each press imagines ~4 | |
| frames with 5-step denoising in an AR shortcut DiT — everything is generated, no simulator. | |
| On the free CPU tier a step takes a while; on GPU it's ~0.3 s.</p> | |
| <div class="keys"><span></span><div class="key" data-k="w">W</div><span></span> | |
| <div class="key" data-k="a">A</div><div class="key" data-k="s">S</div><div class="key" data-k="d">D</div></div> | |
| <div class="row"><select id="ep"></select><button id="reset">⟲ Reset (R)</button></div> | |
| <p id="msg"></p></div></div><script> | |
| const $=i=>document.getElementById(i);const DIR={w:[0,-1],s:[0,1],a:[-1,0],d:[1,0]}; | |
| let held=new Set(),busy=false,alive=false; | |
| for(let i=0;i<16;i++){const o=document.createElement('option');o.value=i;o.textContent='scene '+i;$('ep').appendChild(o)} | |
| async function post(u,b){try{const r=await fetch(u,{method:'POST',headers:{'Content-Type':'application/json'},body:JSON.stringify(b)});if(!r.ok){const t=await r.text();$('msg').textContent='server error: '+t.slice(0,200);throw new Error(t)}return r.json()}catch(e){$('msg').textContent='request failed: '+e.message;busy=false;$('app').classList.remove('busy');throw e}} | |
| async function doReset(){busy=true;$('app').classList.add('busy');$('stat').textContent='resetting…'; | |
| const j=await post('reset',{idx:+$('ep').value});$('frame').src=j.frame; | |
| $('stat').textContent='latent 1 · ready';alive=true;busy=false;$('app').classList.remove('busy')} | |
| async function doStep(dx,dy){if(busy||!alive)return;busy=true;$('app').classList.add('busy'); | |
| const t0=performance.now();const j=await post('step',{dx,dy});$('frame').src=j.frame; | |
| $('stat').textContent=`latent ${j.n} · a=[${j.a[0].toFixed(2)},${j.a[1].toFixed(2)}]`; | |
| $('lat').textContent=`${Math.round(performance.now()-t0)} ms`;busy=false;$('app').classList.remove('busy'); | |
| if(held.size){let dx2=0,dy2=0;for(const k of held){dx2+=DIR[k][0];dy2+=DIR[k][1]}if(dx2||dy2)doStep(dx2,dy2)}} | |
| addEventListener('keydown',e=>{const k=e.key.toLowerCase();if(k==='r'){doReset();return} | |
| if(DIR[k]&&!held.has(k)){held.add(k);mark();let dx=0,dy=0;for(const kk of held){dx+=DIR[kk][0];dy+=DIR[kk][1]}doStep(dx,dy)}}); | |
| addEventListener('keyup',e=>{const k=e.key.toLowerCase();if(DIR[k]){held.delete(k);mark()}}); | |
| function mark(){document.querySelectorAll('.key').forEach(el=>el.classList.toggle('on',held.has(el.dataset.k)))} | |
| document.querySelectorAll('.key').forEach(el=>el.onclick=()=>{const d=DIR[el.dataset.k];doStep(d[0],d[1])}); | |
| $('reset').onclick=doReset; | |
| doReset();</script></body></html>""" | |
| class H(BaseHTTPRequestHandler): | |
| def log_message(self, *a): pass | |
| def _json(self, obj, code=200): | |
| b = json.dumps(obj).encode() | |
| self.send_response(code); self.send_header("Content-Type", "application/json") | |
| self.send_header("Content-Length", str(len(b))); self.end_headers(); self.wfile.write(b) | |
| def do_GET(self): | |
| b = PAGE.encode() | |
| self.send_response(200); self.send_header("Content-Type", "text/html; charset=utf-8") | |
| self.send_header("Content-Length", str(len(b))); self.end_headers(); self.wfile.write(b) | |
| def do_POST(self): | |
| n = int(self.headers.get("Content-Length", 0)) | |
| req = json.loads(self.rfile.read(n) or b"{}") | |
| try: | |
| with LOCK: | |
| if self.path.endswith("/reset"): | |
| frame, nlat = reset(req.get("idx", 0)); self._json({"frame": frame, "n": nlat}) | |
| elif self.path.endswith("/step"): | |
| if S["z"] is None: self._json({"error": "reset first"}, 400); return | |
| frame, nlat, a = step(float(req.get("dx", 0)), float(req.get("dy", 0))) | |
| self._json({"frame": frame, "n": nlat, "a": a}) | |
| else: self._json({"error": "unknown"}, 404) | |
| except Exception as e: | |
| import traceback; traceback.print_exc() | |
| self._json({"error": str(e)}, 500) | |
| if __name__ == "__main__": | |
| print(f"[wasd] serving :{PORT}", flush=True) | |
| ThreadingHTTPServer(("0.0.0.0", PORT), H).serve_forever() | |