#!/usr/bin/env python3 """DreamWaQ play — renders env with trained policy in MotrixSim. Usage: uv run scripts/play_dreamwaq.py uv run scripts/play_dreamwaq.py --num-envs 16 """ import argparse, os, time, sys # CRITICAL: disable JAX GPU memory preallocation BEFORE importing jax. # Otherwise JAX grabs 75% of GPU memory and starves the MotrixSim (Vulkan) # renderer → "Couldn't get swap chain texture" crash. Must be set first. os.environ.setdefault("XLA_PYTHON_CLIENT_PREALLOCATE", "false") sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) import motrix_envs.locomotion.go1.dreamwaq # noqa import motrix_rl.tasks.go1_dreamwaq # noqa import numpy as np import jax, jax.numpy as jnp import pickle, msgpack from motrix_envs import registry as env_registry from motrix_envs.np.renderer import NpRenderer from motrix_rl.skrl.jax.train.dreamwaq_ppo import DreamWaQWrapper, CENet def _decode_arr(ext): if not hasattr(ext, "code"): return None parts = msgpack.unpackb(ext.data, raw=False) if not isinstance(parts, list) or len(parts) < 3: return None shape = [] def _flatten(s): if isinstance(s, list): for x in s: _flatten(x) elif isinstance(s, int): shape.append(s) _flatten(parts[0]) return np.frombuffer(parts[2], dtype=np.dtype(parts[1])).reshape(shape) def load_params(ckpt_path): with open(ckpt_path, 'rb') as f: ckpt = pickle.load(f) raw = msgpack.unpackb(ckpt['policy'])['params'] params = {} for name, val in raw.items(): if isinstance(val, dict): params[name] = {k: _decode_arr(v) for k, v in val.items()} else: params[name] = _decode_arr(val) # State-preprocessor stats for first 64 dims (REQUIRED: policy trained on normalized obs) mean64 = std64 = None if 'state_preprocessor' in ckpt: sp = msgpack.unpackb(ckpt['state_preprocessor'], raw=False) mean64 = _decode_arr(sp['running_mean'])[:64].astype(np.float32) std64 = np.sqrt(_decode_arr(sp['running_variance'])[:64]).astype(np.float32) return params, mean64, std64 CLIP_ACT = 23.7 CLIP_OBS = 100.0 def policy_forward(x, p, mean64=None, std64=None): x = jnp.array(x[:, :64]) # Apply state-preprocessor normalization (clip((x-mean)/(std+eps), -5, 5)) if mean64 is not None: x = jnp.clip((x - jnp.array(mean64)) / (jnp.array(std64) + 1e-8), -5.0, 5.0) else: x = jnp.clip(x, -CLIP_OBS, CLIP_OBS) x = jax.nn.elu(x @ jnp.array(p['Dense_0']['kernel']) + jnp.array(p['Dense_0']['bias'])) x = jax.nn.elu(x @ jnp.array(p['Dense_1']['kernel']) + jnp.array(p['Dense_1']['bias'])) x = jax.nn.elu(x @ jnp.array(p['Dense_2']['kernel']) + jnp.array(p['Dense_2']['bias'])) return np.clip(np.array(x @ jnp.array(p['Dense_3']['kernel']) + jnp.array(p['Dense_3']['bias'])), -CLIP_ACT, CLIP_ACT) def main(): p = argparse.ArgumentParser() p.add_argument("--num-envs", type=int, default=9) p.add_argument("--checkpoint", default=None) args = p.parse_args() # Auto-find checkpoint if args.checkpoint is None: run_dir = "runs/go1-dreamwaq-walk/skrl" runs = sorted([d for d in os.listdir(run_dir) if os.path.isdir(os.path.join(run_dir, d)) and d.startswith("26-")]) args.checkpoint = os.path.join(run_dir, runs[-1], "checkpoints", "best_agent.pickle") # Load policy + state-preprocessor normalization policy_params, mean64, std64 = load_params(args.checkpoint) print(f"[Play] Policy: {args.checkpoint}") print(f"[Play] State normalization: {'ON' if mean64 is not None else 'OFF'}") # Load VAE (saved in skrl/ base dir, not run subdir) run_dir = os.path.dirname(os.path.dirname(os.path.dirname(args.checkpoint))) # skrl/ base vae_path = os.path.join(run_dir, "cenet_params.pkl") if not os.path.exists(vae_path): vae_files = sorted([f for f in os.listdir(run_dir) if f.startswith("vae_")], key=lambda x: int(x.split("_")[1].split(".")[0])) if vae_files: vae_path = os.path.join(run_dir, vae_files[-1]) with open(vae_path, 'rb') as f: vae_params = pickle.load(f) print(f"[Play] VAE: {vae_path}") # Create env + renderer (like view.py) raw_env = env_registry.make("go1-dreamwaq-walk", num_envs=args.num_envs) renderer = NpRenderer(raw_env) # CENet for policy inference cenet = CENet() rng = jax.random.PRNGKey(42) wrapper = DreamWaQWrapper(raw_env, cenet, vae_params, rng=rng) # Init env raw_env.init_state() wrapper._vae_buf = [] n = raw_env._num_envs print(f"[Play] {n} envs, Ctrl+C to stop") from motrixsim.render import RenderClosedError try: while True: # CENet inference (mean mode) hist = jnp.array(raw_env._state.info.get("obs_history", np.zeros((n, 5, 45), dtype=np.float32))) z, vel = cenet.apply(vae_params, hist, method=cenet.inference) code = np.concatenate([np.array(vel), np.array(z)], axis=-1) obs_arr = raw_env._state.obs priv = raw_env._state.info.get("privileged_obs", np.zeros((n, 235), dtype=np.float32)) heights = priv[:, 48:] if priv.shape[1] > 48 else np.zeros((n, 187), dtype=np.float32) base_vel = raw_env._state.info.get("base_vel", np.zeros((n, 3), dtype=np.float32)) base_vel_n = base_vel * np.array([2.0, 2.0, 1.0], dtype=np.float32) aug_obs = np.concatenate([code, obs_arr, base_vel_n, heights], axis=-1) actions = policy_forward(aug_obs, policy_params, mean64, std64) wrapper.step(actions) renderer.render() time.sleep(0.01) except (KeyboardInterrupt, RenderClosedError): pass try: renderer.close() except: pass print("[Play] Done") if __name__ == "__main__": main()