fix: clamp std before distribution, lower init_noise to 0.5, NaN guard
This commit is contained in:
147
scripts/play_dreamwaq.py
Normal file
147
scripts/play_dreamwaq.py
Normal file
@@ -0,0 +1,147 @@
|
||||
#!/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()
|
||||
Reference in New Issue
Block a user