#!/usr/bin/env python3 """DreamWaQ rsl_rl play — 加载 CENetActorModel checkpoint 渲染。 用法: uv run scripts/play_dreamwaq_rsl.py # 自动找最新 uv run scripts/play_dreamwaq_rsl.py --checkpoint runs/.../model_100.pt --vx 0.5 uv run scripts/play_dreamwaq_rsl.py --terrain --level 5 --num-envs 1 """ import argparse, glob, os, sys, time os.environ.setdefault("XLA_PYTHON_CLIENT_PREALLOCATE", "false") os.environ.setdefault("JAX_PLATFORMS", "cpu") terrain_type = "pyramid" if "--flat" in sys.argv: os.environ["DREAMWAQ_TERRAIN"] = "flat" terrain_type = "flat" elif "--flat-stairs" in sys.argv: os.environ["DREAMWAQ_TERRAIN"] = "flat_stairs" elif "--stairs" in sys.argv: os.environ["DREAMWAQ_TERRAIN"] = "stairs" elif "--terrain" in sys.argv: os.environ.setdefault("DREAMWAQ_TERRAIN", "pyramid") else: os.environ.setdefault("DREAMWAQ_TERRAIN", "flat") sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) import numpy as np import torch from tensordict import TensorDict import motrix_envs.locomotion.go1.dreamwaq # noqa: F401 from motrix_envs import registry as env_registry from motrix_envs.np.renderer import NpRenderer from motrix_rl.rslrl.torch.models.cenet_actor import CENetActorModel PROJECT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) NUM_OBS, NUM_PRIV, NUM_HIST, NUM_ACT, CENET_OUT = 45, 247, 5, 12, 19 CLIP_ACT = 100.0 def find_latest(): models = glob.glob(os.path.join(PROJECT, "runs", "go1-dreamwaq-walk", "rslrl", "*", "model_*.pt")) if not models: print("[ERROR] no rslrl checkpoints found"); sys.exit(1) def _iter_of(p): try: return int(os.path.basename(p).split("_")[1].split(".")[0]) except: return -1 return max(models, key=_iter_of) def main(): p = argparse.ArgumentParser() p.add_argument("--checkpoint", default=None) p.add_argument("--num-envs", type=int, default=4) p.add_argument("--vx", type=float, default=0.5) p.add_argument("--vy", type=float, default=0.0) p.add_argument("--wz", type=float, default=0.0) p.add_argument("--terrain", action="store_true") p.add_argument("--flat", action="store_true") p.add_argument("--flat-stairs", action="store_true") p.add_argument("--stairs", action="store_true") p.add_argument("--level", type=int, default=None) p.add_argument("--spawn-height", type=float, default=None) args = p.parse_args() ckpt = args.checkpoint or find_latest() device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") print(f"[Play] checkpoint: {ckpt}") # 创建模型(与训练时一致) dummy_obs = TensorDict({ "policy": torch.zeros(1, NUM_OBS), "obs_history": torch.zeros(1, NUM_HIST * NUM_OBS), "privileged_obs": torch.zeros(1, NUM_PRIV), }, batch_size=[1]) model = CENetActorModel( dummy_obs, {"actor": ["policy", "obs_history"]}, "actor", NUM_ACT, hidden_dims=[512, 256, 128], activation="elu", stochastic=True, init_noise_std=0.5, cenet_in_dim=NUM_HIST * NUM_OBS, cenet_out_dim=CENET_OUT, ) ckpt_data = torch.load(ckpt, map_location=device) # OnPolicyRunner 保存格式: actor_state_dict, critic_state_dict if "actor_state_dict" in ckpt_data: model.load_state_dict(ckpt_data["actor_state_dict"], strict=False) elif "model_state_dict" in ckpt_data: model.load_state_dict(ckpt_data["model_state_dict"], strict=False) else: model.load_state_dict(ckpt_data, strict=False) model.to(device) model.eval() print(f"[Play] model loaded, device={device}") # 创建环境 env = env_registry.make("go1-dreamwaq-walk", num_envs=args.num_envs) if args.level is not None: env._force_level = args.level print(f"[Play] forcing level={args.level}") if args.spawn_height is not None: env._spawn_absolute = args.spawn_height print(f"[Play] spawn z={args.spawn_height}m") renderer = NpRenderer(env) env.init_state() n = env._num_envs cmd = np.array([args.vx, args.vy, args.wz], dtype=np.float32) print(f"[Play] {n} envs | cmd=(vx={args.vx}, vy={args.vy}, wz={args.wz}) | R=reset Ctrl+C=stop") @torch.no_grad() def act_fn(obs_np, hist_np): obs_t = torch.from_numpy(obs_np).float().to(device) hist_t = torch.from_numpy(hist_np.reshape(n, -1)).float().to(device) td = TensorDict({"policy": obs_t, "obs_history": hist_t}, batch_size=[n], device=device) return model(td).cpu().numpy() from motrixsim.render import RenderClosedError show_heights = False try: while True: if renderer._render.input.is_key_just_pressed("r"): env.init_state() print("[R] Reset") if renderer._render.input.is_key_just_pressed("h"): show_heights = not show_heights print(f"[H] Heights: {'ON' if show_heights else 'OFF'}") env._state.info["commands"][:] = cmd obs = env._state.obs.astype(np.float32) hist = env._state.info.get("obs_history", np.zeros((n, NUM_HIST, NUM_OBS), np.float32)) act = act_fn(obs, hist) env.step(np.clip(act, -CLIP_ACT, CLIP_ACT).astype(np.float32)) if show_heights: from motrix_envs.math import quaternion pose = env._body.get_pose(env._state.data) bp = pose[0, :3]; yaw = quaternion.get_yaw(pose[0:1, 3:7])[0] cos_y, sin_y = np.cos(yaw), np.sin(yaw) for gy in env._hy: for gx in env._hx: wx = bp[0]+cos_y*gx-sin_y*gy; wy = bp[1]+sin_y*gx+cos_y*gy try: wz = float(env._sample_terrain_height(np.array([[wx,wy]]))[0]) renderer._render.gizmos.draw_sphere(0.02, (np.float32(wx), np.float32(wy), np.float32(wz))) except: pass renderer.render() time.sleep(0.01) except (KeyboardInterrupt, RenderClosedError): pass try: renderer.close() except: pass print("[Play] done") if __name__ == "__main__": main()