diff --git a/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py b/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py index 8401583..1cb9980 100644 --- a/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py +++ b/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py @@ -99,7 +99,7 @@ class DreamWaQCfg(Go1WalkNpEnvCfg): cell_size: float = 8.0 border_size: float = 5.0 - # Forward stair commands: stand or walk at 0.5-1.0 m/s. + # Symmetric walking commands with no low-speed samples except standing. @dataclass class Commands: forward_speed_range = (0.5, 1.0) @@ -205,6 +205,10 @@ class DreamWaQTask(Go1WalkTask): cmds[:, 0] = np.random.uniform( speed_min, speed_max, size=num_envs ).astype(np.float32) + signs = np.random.choice( + np.array([-1.0, 1.0], dtype=np.float32), size=num_envs + ) + cmds[:, 0] *= signs stand = np.random.random(num_envs) < self.cfg.commands.stand_probability cmds[stand, 0] = 0.0 return cmds diff --git a/motrix_envs/tests/test_dreamwaq_state_safety.py b/motrix_envs/tests/test_dreamwaq_state_safety.py index 0032d99..345525e 100644 --- a/motrix_envs/tests/test_dreamwaq_state_safety.py +++ b/motrix_envs/tests/test_dreamwaq_state_safety.py @@ -59,7 +59,7 @@ 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_zero_or_forward_between_half_and_one(monkeypatch): +def test_dreamwaq_commands_are_zero_or_symmetric_without_low_speeds(monkeypatch): task = DreamWaQTask.__new__(DreamWaQTask) task._cfg = SimpleNamespace( commands=SimpleNamespace( @@ -71,6 +71,11 @@ def test_dreamwaq_commands_are_zero_or_forward_between_half_and_one(monkeypatch) "uniform", lambda low, high, size: np.array([0.5, 0.75, 1.0], dtype=np.float32), ) + monkeypatch.setattr( + np.random, + "choice", + lambda values, size: np.array([-1.0, -1.0, 1.0], dtype=np.float32), + ) monkeypatch.setattr( np.random, "random", @@ -79,7 +84,7 @@ def test_dreamwaq_commands_are_zero_or_forward_between_half_and_one(monkeypatch) commands = task.resample_commands(3) - np.testing.assert_array_equal(commands[:, 0], [0.0, 0.75, 1.0]) + np.testing.assert_array_equal(commands[:, 0], [0.0, -0.75, 1.0]) np.testing.assert_array_equal(commands[:, 1:], 0.0)