fix: boundary termination, physics crash handler, action_scale 0.25, play script update

This commit is contained in:
8x54zj-m
2026-06-30 19:03:57 +08:00
parent 7affccc695
commit cb62f35ef6
3 changed files with 115 additions and 74 deletions

View File

@@ -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]:

View File

@@ -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

View File

@@ -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__":