Sample symmetric DreamWaQ speed commands
This commit is contained in:
@@ -99,7 +99,7 @@ class DreamWaQCfg(Go1WalkNpEnvCfg):
|
|||||||
cell_size: float = 8.0
|
cell_size: float = 8.0
|
||||||
border_size: float = 5.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
|
@dataclass
|
||||||
class Commands:
|
class Commands:
|
||||||
forward_speed_range = (0.5, 1.0)
|
forward_speed_range = (0.5, 1.0)
|
||||||
@@ -205,6 +205,10 @@ class DreamWaQTask(Go1WalkTask):
|
|||||||
cmds[:, 0] = np.random.uniform(
|
cmds[:, 0] = np.random.uniform(
|
||||||
speed_min, speed_max, size=num_envs
|
speed_min, speed_max, size=num_envs
|
||||||
).astype(np.float32)
|
).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
|
stand = np.random.random(num_envs) < self.cfg.commands.stand_probability
|
||||||
cmds[stand, 0] = 0.0
|
cmds[stand, 0] = 0.0
|
||||||
return cmds
|
return cmds
|
||||||
|
|||||||
@@ -59,7 +59,7 @@ def test_apply_action_advances_three_frame_action_history():
|
|||||||
np.testing.assert_array_equal(state.info["current_actions"], 3.0)
|
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 = DreamWaQTask.__new__(DreamWaQTask)
|
||||||
task._cfg = SimpleNamespace(
|
task._cfg = SimpleNamespace(
|
||||||
commands=SimpleNamespace(
|
commands=SimpleNamespace(
|
||||||
@@ -71,6 +71,11 @@ def test_dreamwaq_commands_are_zero_or_forward_between_half_and_one(monkeypatch)
|
|||||||
"uniform",
|
"uniform",
|
||||||
lambda low, high, size: np.array([0.5, 0.75, 1.0], dtype=np.float32),
|
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(
|
monkeypatch.setattr(
|
||||||
np.random,
|
np.random,
|
||||||
"random",
|
"random",
|
||||||
@@ -79,7 +84,7 @@ def test_dreamwaq_commands_are_zero_or_forward_between_half_and_one(monkeypatch)
|
|||||||
|
|
||||||
commands = task.resample_commands(3)
|
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)
|
np.testing.assert_array_equal(commands[:, 1:], 0.0)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user