From 6d6a805091c4e56736748cf52b5df9e2376da3b9 Mon Sep 17 00:00:00 2001 From: 8x54zj-m <8x54zj-m@motrixlab.local> Date: Wed, 22 Jul 2026 19:39:13 +0800 Subject: [PATCH] Limit DreamWaQ commands to stand or walk --- .../src/motrix_envs/locomotion/go1/dreamwaq.py | 14 +++++++------- motrix_envs/tests/test_dreamwaq_state_safety.py | 13 +++++++++++++ 2 files changed, 20 insertions(+), 7 deletions(-) diff --git a/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py b/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py index f693ab2..79db784 100644 --- a/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py +++ b/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py @@ -94,10 +94,11 @@ class DreamWaQCfg(Go1WalkNpEnvCfg): cell_size: float = 8.0 border_size: float = 5.0 - # 命令范围(上游: [-1,1] 对称) + # Stair training commands: stand or walk straight forward at 0.5 m/s. @dataclass class Commands: - vel_limit = [[-1.0, -1.0, -1.0], [1.0, 1.0, 1.0]] + forward_values = (0.0, 0.5) + stand_probability = 0.2 commands: Commands = field(default_factory=Commands) @@ -191,13 +192,12 @@ class DreamWaQTask(Go1WalkTask): self._hf_cache = None self._hm_cache = None - # ── 命令:全范围,无课程(匹配上游)── + # ── 命令:站立或固定速度直行 ── def resample_commands(self, num_envs: int) -> np.ndarray: - lim = np.array(self.cfg.commands.vel_limit, dtype=np.float32) - cmds = np.random.uniform(lim[0], lim[1], size=(num_envs, 3)).astype(np.float32) - small = np.linalg.norm(cmds[:, :2], axis=1) < 0.2 - cmds[small, :2] = 0.0 + cmds = np.zeros((num_envs, 3), dtype=np.float32) + walk = np.random.random(num_envs) >= self.cfg.commands.stand_probability + cmds[walk, 0] = self.cfg.commands.forward_values[1] return cmds # ── 动作裁剪 + 力矩计算 ── diff --git a/motrix_envs/tests/test_dreamwaq_state_safety.py b/motrix_envs/tests/test_dreamwaq_state_safety.py index 0b4c7b7..f043ca5 100644 --- a/motrix_envs/tests/test_dreamwaq_state_safety.py +++ b/motrix_envs/tests/test_dreamwaq_state_safety.py @@ -44,6 +44,19 @@ def test_apply_action_advances_three_frame_action_history(): np.testing.assert_array_equal(state.info["current_actions"], 3.0) +def test_dreamwaq_commands_are_stand_or_fixed_forward_speed(monkeypatch): + task = DreamWaQTask.__new__(DreamWaQTask) + task._cfg = DreamWaQCfg() + monkeypatch.setattr( + np.random, "random", lambda size: np.array([0.1, 0.2, 0.9], dtype=np.float32) + ) + + commands = task.resample_commands(3) + + np.testing.assert_array_equal(commands[:, 0], [0.0, 0.5, 0.5]) + np.testing.assert_array_equal(commands[:, 1:], 0.0) + + def test_dreamwaq_feet_air_time_matches_upstream_contact_filtering(): task = DreamWaQTask.__new__(DreamWaQTask) task._cfg = SimpleNamespace(ctrl_dt=0.02)