From cb62f35ef61953460851b8bd88664b8835aa71d3 Mon Sep 17 00:00:00 2001 From: 8x54zj-m <8x54zj-m@motrixlab.local> Date: Tue, 30 Jun 2026 19:03:57 +0800 Subject: [PATCH] fix: boundary termination, physics crash handler, action_scale 0.25, play script update --- .../motrix_envs/locomotion/go1/dreamwaq.py | 16 +- motrix_envs/src/motrix_envs/np/env.py | 15 +- scripts/play_dreamwaq_rsl.py | 158 ++++++++++-------- 3 files changed, 115 insertions(+), 74 deletions(-) diff --git a/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py b/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py index 2463c5d..647b6a5 100644 --- a/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py +++ b/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py @@ -109,7 +109,7 @@ class DreamWaQCfg(Go1WalkNpEnvCfg): self.commands = DreamWaQCfg.Commands() self.control_config.stiffness = 28.0 self.control_config.damping = 0.7 - self.control_config.action_scale = 0.25 # 上游原值(基类默认 0.05) + self.control_config.action_scale = 0.25 # 上游原值 self.noise_config.scale_joint_angle = 0.01 # 上游原值(基类默认 0.03) self._apply_dreamwaq() @@ -530,6 +530,20 @@ class DreamWaQTask(Go1WalkTask): return obs, info + def update_terminated(self, state): + """基类接触终止 + hfield 边界终止(防止走出地形导致 NaN)。""" + state = super().update_terminated(state) + pose = self._body.get_pose(state.data) + base_xy = pose[:, :2] + hf = self._model.get_hfield(0) + b = hf.bound + out_of_bounds = ( + (base_xy[:, 0] < b[0]) | (base_xy[:, 0] > b[3]) | + (base_xy[:, 1] < b[1]) | (base_xy[:, 1] > b[4]) + ) + state.terminated = state.terminated | out_of_bounds + return state + # ── 奖励(与上游对齐,update_reward 中 × dt)── def _get_reward(self, data: mtx.SceneData, info: dict) -> dict[str, np.ndarray]: diff --git a/motrix_envs/src/motrix_envs/np/env.py b/motrix_envs/src/motrix_envs/np/env.py index 6d58a71..ea18335 100644 --- a/motrix_envs/src/motrix_envs/np/env.py +++ b/motrix_envs/src/motrix_envs/np/env.py @@ -183,8 +183,19 @@ class NpEnv(ABEnv): def physics_step(self): # motrixsim.SceneModel.step only supports single step, so we loop - for _ in range(self._cfg.sim_substeps): - self._model.step(self._state.data) + try: + for _ in range(self._cfg.sim_substeps): + self._model.step(self._state.data) + except Exception as e: + # Rust panic / MotrixSim 物理崩溃 → 标记所有 env 为终止 + n = self._state.data.shape[0] + self._state.terminated[:] = True + self._state.reward[:] = 0.0 + if not hasattr(self, '_physics_crash_count'): + self._physics_crash_count = 0 + self._physics_crash_count += 1 + if self._physics_crash_count <= 3: + print(f"[WARN] physics crash #{self._physics_crash_count}: {e} — resetting {n} envs") def _prev_physics_step(self): state = self._state diff --git a/scripts/play_dreamwaq_rsl.py b/scripts/play_dreamwaq_rsl.py index 972a1ac..6710c00 100644 --- a/scripts/play_dreamwaq_rsl.py +++ b/scripts/play_dreamwaq_rsl.py @@ -1,99 +1,119 @@ #!/usr/bin/env python3 -"""DreamWaQ rsl_rl play — render the trained ActorCritic_DWAQ policy in NATIVE MotrixSim. +"""DreamWaQ rsl_rl play — 加载 CENetActorModel checkpoint 渲染。 -Loads the PyTorch checkpoint DIRECTLY (no ONNX). ONNX is only for cross-sim -deployment (e.g. MuJoCo sim2sim); the native MotrixSim env runs the torch policy. -Deterministic inference: mean CENet code + actor. - -Usage: - uv run scripts/play_dreamwaq_rsl.py # auto-find latest, walk forward - uv run scripts/play_dreamwaq_rsl.py --checkpoint runs/.../model_1100.pt --vx 0.5 - uv run scripts/play_dreamwaq_rsl.py --vx 0 --num-envs 1 # stand still, single robot +用法: + 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 -# avoid JAX grabbing GPU memory and starving the MotrixSim (Vulkan) renderer os.environ.setdefault("XLA_PYTHON_CLIENT_PREALLOCATE", "false") os.environ.setdefault("JAX_PLATFORMS", "cpu") -# --terrain / --level / --flat-stairs / --stairs: pick hfield scene (before import). + terrain_type = "pyramid" -if "--flat-stairs" in sys.argv: terrain_type = "flat_stairs" -elif "--stairs" in sys.argv: terrain_type = "stairs" -if "--terrain" in sys.argv or "--level" in sys.argv or "--flat-stairs" in sys.argv or "--stairs" in sys.argv: - os.environ["DREAMWAQ_TERRAIN"] = terrain_type +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 register env +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.dwaq_rsl.actor_critic_dwaq import ActorCritic_DWAQ +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, 235, 5, 12, 19 -CLIP_ACT = 23.7 - - -def _iter_of(path): - try: - return int(os.path.basename(path).split("_")[1].split(".")[0]) - except Exception: - return -1 +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", "rsl_dwaq", "*", "model_*.pt")) + models = glob.glob(os.path.join(PROJECT, "runs", "go1-dreamwaq-walk", "rslrl", "*", "model_*.pt")) if not models: - print("[ERROR] no rsl_dwaq checkpoints found"); sys.exit(1) - return max(models, key=_iter_of) # highest iteration (flat model_1100 > terrain early 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, help="forward velocity command [m/s]") - p.add_argument("--vy", type=float, default=0.0, help="lateral velocity command [m/s]") - p.add_argument("--wz", type=float, default=0.0, help="yaw rate command [rad/s]") - p.add_argument("--terrain", action="store_true", - help="view the training pyramid hfield (default: flat plane)") - p.add_argument("--flat-stairs", action="store_true", - help="view the 2-level flat+stairs terrain (implies terrain)") - p.add_argument("--stairs", action="store_true", - help="view the stairs terrain scene (implies terrain)") - p.add_argument("--level", type=int, default=None, - help="force ALL spawns at this terrain level (implies --terrain)") - p.add_argument("--spawn-height", type=float, default=None, - help="spawn clearance above terrain in meters (default 0.45; try 1-2 to experiment)") + 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() - ac = ActorCritic_DWAQ(NUM_OBS + CENET_OUT, NUM_PRIV, NUM_ACT, NUM_HIST * NUM_OBS, CENET_OUT) - ac.load_state_dict(torch.load(ckpt, map_location="cpu")["model_state_dict"]) - ac.eval() - print(f"[Play-rsl] policy (native torch): {ckpt}") + 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 # pin all spawns to this level (read in reset) - print(f"[Play-rsl] forcing ALL spawns at terrain level {args.level}") + 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 # absolute world z, no offset - print(f"[Play-rsl] spawn absolute z = {args.spawn_height}m") + 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-rsl] {n} envs | cmd=(vx={args.vx}, vy={args.vy}, wz={args.wz}) | Ctrl+C to stop") + 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, hist): - obs_t = torch.from_numpy(obs) - h = ac.encoder(torch.from_numpy(hist).reshape(obs.shape[0], -1)) # (n,225)->(n,64) - code = torch.cat([ac.encode_mean_vel(h), ac.encode_mean_latent(h)], dim=-1) # (n,19) - return ac.actor(torch.cat([code, obs_t], dim=-1)).numpy() # (n,12) + 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 @@ -101,40 +121,36 @@ def main(): while True: if renderer._render.input.is_key_just_pressed("r"): env.init_state() - print("[R] Reset all envs") + print("[R] Reset") if renderer._render.input.is_key_just_pressed("h"): show_heights = not show_heights - print(f"[H] Height points: {'ON' if show_heights else 'OFF'}") + 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)).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] + 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 - wz = float(env._sample_terrain_height(np.array([[wx,wy]]))[0]) - g = renderer._render.gizmos - g.draw_sphere(0.02, (np.float32(wx), np.float32(wy), np.float32(wz))) - + 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 Exception: - pass - print("[Play-rsl] done") + except: pass + print("[Play] done") if __name__ == "__main__":