fix: clamp std before distribution, lower init_noise to 0.5, NaN guard

This commit is contained in:
8x54zj-m
2026-06-30 15:35:58 +08:00
parent f2a8e0e2ff
commit b0f1da4596
61 changed files with 6522 additions and 45 deletions

147
scripts/play_dreamwaq.py Normal file
View 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()