Reset invalid DreamWaQ physics states safely

This commit is contained in:
8x54zj-m
2026-07-22 04:30:13 +08:00
parent 660393a9ee
commit 1a84cf69ab
4 changed files with 91 additions and 16 deletions

View File

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

View File

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

View File

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

View File

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