diff --git a/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py b/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py index 7bca6e3..94fab4b 100644 --- a/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py +++ b/motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py @@ -30,6 +30,19 @@ from motrix_envs.math import quaternion _SCENE_PRINTED = False +def _sanitize_dof_pos(dof_pos: np.ndarray) -> np.ndarray: + """Make free-joint positions finite and guarantee normalized xyzw quaternions.""" + clean = np.nan_to_num( + np.array(dof_pos, copy=True), nan=0.0, posinf=0.0, neginf=0.0 + ) + quat = clean[:, 3:7] + quat_norm = np.linalg.norm(quat, axis=1) + valid_quat = np.isfinite(quat_norm) & (quat_norm > 1e-6) + quat[valid_quat] /= quat_norm[valid_quat, np.newaxis] + quat[~valid_quat] = [0.0, 0.0, 0.0, 1.0] + return clean + + def _scene_file(): """选择地形场景。DREAMWAQ_TERRAIN=flat|pyramid|flat_stairs。默认 FLAT。""" global _SCENE_PRINTED @@ -353,22 +366,12 @@ class DreamWaQTask(Go1WalkTask): dv_clean[:, 6:] = np.nan_to_num(dv_clean[:, 6:], nan=0.0, posinf=0.0, neginf=0.0) data.set_dof_vel(dv_clean) dp = data.dof_pos - if np.any(~np.isfinite(dp)): - dp_clean = np.array(dp) - # 四元数 NaN → 整行替换为单位四元数 [0,0,0,1] - quat_nan_row = np.any(~np.isfinite(dp_clean[:, 3:7]), axis=1) - if np.any(quat_nan_row): - dp_clean[quat_nan_row, 3:7] = [0.0, 0.0, 0.0, 1.0] - # 关节位置 NaN → 0 - dp_clean[:, 7:] = np.nan_to_num(dp_clean[:, 7:], nan=0.0, posinf=0.0, neginf=0.0) - data.set_dof_pos(dp_clean, self._model) - else: - quat_norm = np.linalg.norm(dp[:, 3:7], axis=1) - bad_quat = quat_norm < 1e-6 - if np.any(bad_quat): - dp_clean = np.array(dp) - dp_clean[bad_quat, 3:7] = [0.0, 0.0, 0.0, 1.0] - data.set_dof_pos(dp_clean, self._model) + quat_norm = np.linalg.norm(dp[:, 3:7], axis=1) + needs_clean = np.any(~np.isfinite(dp)) or np.any( + ~np.isfinite(quat_norm) | (np.abs(quat_norm - 1.0) > 1e-4) + ) + if needs_clean: + data.set_dof_pos(_sanitize_dof_pos(dp), self._model) obs = self._get_obs(data, state.info) obs = np.nan_to_num(obs, nan=0.0, posinf=0.0, neginf=0.0) @@ -459,6 +462,10 @@ class DreamWaQTask(Go1WalkTask): else: done_idx = np.arange(num_reset) + if hasattr(self, "_action_buffer"): + self._action_buffer[done_idx] = 0.0 + self._latency_steps[done_idx] = np.random.randint(0, 4, size=num_reset) + # 游戏启发式地形课程(per-env,与上游一致) if not hasattr(self, '_terrain_origins'): all_origins = np.zeros((self._num_rows, self._num_cols, 2), dtype=np.float32) diff --git a/motrix_envs/src/motrix_envs/locomotion/go1/walk_np.py b/motrix_envs/src/motrix_envs/locomotion/go1/walk_np.py index 1694a98..39112bd 100644 --- a/motrix_envs/src/motrix_envs/locomotion/go1/walk_np.py +++ b/motrix_envs/src/motrix_envs/locomotion/go1/walk_np.py @@ -81,6 +81,19 @@ class Go1WalkTask(NpEnv): def get_dof_vel(self, data: mtx.SceneModel): return self._body.get_joint_dof_vel(data) + def _invalid_physics_state_mask(self, state: NpEnvState) -> np.ndarray: + invalid = super()._invalid_physics_state_mask(state) + dof_pos = np.asarray(state.data.dof_pos) + dof_vel = np.asarray(state.data.dof_vel) + + quat_norm = np.linalg.norm(dof_pos[:, 3:7], axis=1) + invalid |= ~np.isfinite(quat_norm) + invalid |= np.abs(quat_norm - 1.0) > 1e-3 + invalid |= np.any(np.abs(dof_pos[:, :3]) > 1e4, axis=1) + invalid |= np.any(np.abs(dof_pos[:, 7:]) > 20.0, axis=1) + invalid |= np.any(np.abs(dof_vel) > 1e3, axis=1) + return invalid + def _init_buffer(self): cfg = self._cfg assert isinstance(cfg, Go1WalkNpEnvCfg) diff --git a/motrix_envs/src/motrix_envs/np/env.py b/motrix_envs/src/motrix_envs/np/env.py index 050d472..3cf6b75 100644 --- a/motrix_envs/src/motrix_envs/np/env.py +++ b/motrix_envs/src/motrix_envs/np/env.py @@ -201,6 +201,40 @@ class NpEnv(ABEnv): if self._physics_crash_count <= 3: print(f"[WARN] physics crash #{self._physics_crash_count}: {e} — resetting {n} envs") + def _invalid_physics_state_mask(self, state: NpEnvState) -> np.ndarray: + """Return environments whose simulator state cannot be consumed safely.""" + invalid = np.zeros(self._num_envs, dtype=bool) + for values in (state.data.dof_pos, state.data.dof_vel, state.data.actuator_ctrls): + array = np.asarray(values) + if array.ndim == 1: + array = array[:, np.newaxis] + invalid |= ~np.isfinite(array).all(axis=1) + return invalid + + def _reset_invalid_physics_states(self) -> bool: + invalid = self._invalid_physics_state_mask(self._state) + if not np.any(invalid): + return False + + self._state.terminated[invalid] = True + self._state.reward[invalid] = 0.0 + # Make the backing SceneData safe before task-specific reset logic reads + # poses or curriculum statistics from the full batch. + self._state.data[invalid].reset(self._model) + if not hasattr(self, "_invalid_physics_state_count"): + self._invalid_physics_state_count = 0 + self._invalid_physics_reset_events = 0 + self._invalid_physics_state_count += int(invalid.sum()) + self._invalid_physics_reset_events += 1 + if self._invalid_physics_reset_events <= 10: + print( + f"[WARN] reset {int(invalid.sum())} invalid physics states " + f"(event={self._invalid_physics_reset_events}, " + f"total={self._invalid_physics_state_count})" + ) + self._reset_done_envs() + return True + def _prev_physics_step(self): state = self._state state.reward.fill(0.0) @@ -218,6 +252,8 @@ class NpEnv(ABEnv): if getattr(self, "_physics_crashed_this_step", False): self._reset_done_envs() return self._state + if self._reset_invalid_physics_states(): + return self._state self._state = self.update_state(self._state) self._state.info["steps"] += 1 self._update_truncate() diff --git a/motrix_envs/tests/test_dreamwaq_state_safety.py b/motrix_envs/tests/test_dreamwaq_state_safety.py new file mode 100644 index 0000000..761a6a0 --- /dev/null +++ b/motrix_envs/tests/test_dreamwaq_state_safety.py @@ -0,0 +1,19 @@ +import numpy as np + +from motrix_envs.locomotion.go1.dreamwaq import _sanitize_dof_pos + + +def test_sanitize_dof_pos_handles_zero_quaternion_with_nonfinite_joints(): + dof_pos = np.zeros((2, 19), dtype=np.float32) + dof_pos[0, 3:7] = 0.0 + dof_pos[0, 7] = np.inf + dof_pos[1, 3:7] = [0.0, 0.0, 0.5, 0.5] + dof_pos[1, 8] = np.nan + + clean = _sanitize_dof_pos(dof_pos) + + assert np.isfinite(clean).all() + np.testing.assert_allclose(clean[0, 3:7], [0.0, 0.0, 0.0, 1.0]) + np.testing.assert_allclose(np.linalg.norm(clean[:, 3:7], axis=1), 1.0) + assert clean[0, 7] == 0.0 + assert clean[1, 8] == 0.0