Reset invalid DreamWaQ physics states safely
This commit is contained in:
@@ -30,6 +30,19 @@ from motrix_envs.math import quaternion
|
|||||||
_SCENE_PRINTED = False
|
_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():
|
def _scene_file():
|
||||||
"""选择地形场景。DREAMWAQ_TERRAIN=flat|pyramid|flat_stairs。默认 FLAT。"""
|
"""选择地形场景。DREAMWAQ_TERRAIN=flat|pyramid|flat_stairs。默认 FLAT。"""
|
||||||
global _SCENE_PRINTED
|
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)
|
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)
|
data.set_dof_vel(dv_clean)
|
||||||
dp = data.dof_pos
|
dp = data.dof_pos
|
||||||
if np.any(~np.isfinite(dp)):
|
quat_norm = np.linalg.norm(dp[:, 3:7], axis=1)
|
||||||
dp_clean = np.array(dp)
|
needs_clean = np.any(~np.isfinite(dp)) or np.any(
|
||||||
# 四元数 NaN → 整行替换为单位四元数 [0,0,0,1]
|
~np.isfinite(quat_norm) | (np.abs(quat_norm - 1.0) > 1e-4)
|
||||||
quat_nan_row = np.any(~np.isfinite(dp_clean[:, 3:7]), axis=1)
|
)
|
||||||
if np.any(quat_nan_row):
|
if needs_clean:
|
||||||
dp_clean[quat_nan_row, 3:7] = [0.0, 0.0, 0.0, 1.0]
|
data.set_dof_pos(_sanitize_dof_pos(dp), self._model)
|
||||||
# 关节位置 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)
|
|
||||||
obs = self._get_obs(data, state.info)
|
obs = self._get_obs(data, state.info)
|
||||||
obs = np.nan_to_num(obs, nan=0.0, posinf=0.0, neginf=0.0)
|
obs = np.nan_to_num(obs, nan=0.0, posinf=0.0, neginf=0.0)
|
||||||
|
|
||||||
@@ -459,6 +462,10 @@ class DreamWaQTask(Go1WalkTask):
|
|||||||
else:
|
else:
|
||||||
done_idx = np.arange(num_reset)
|
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,与上游一致)
|
# 游戏启发式地形课程(per-env,与上游一致)
|
||||||
if not hasattr(self, '_terrain_origins'):
|
if not hasattr(self, '_terrain_origins'):
|
||||||
all_origins = np.zeros((self._num_rows, self._num_cols, 2), dtype=np.float32)
|
all_origins = np.zeros((self._num_rows, self._num_cols, 2), dtype=np.float32)
|
||||||
|
|||||||
@@ -81,6 +81,19 @@ class Go1WalkTask(NpEnv):
|
|||||||
def get_dof_vel(self, data: mtx.SceneModel):
|
def get_dof_vel(self, data: mtx.SceneModel):
|
||||||
return self._body.get_joint_dof_vel(data)
|
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):
|
def _init_buffer(self):
|
||||||
cfg = self._cfg
|
cfg = self._cfg
|
||||||
assert isinstance(cfg, Go1WalkNpEnvCfg)
|
assert isinstance(cfg, Go1WalkNpEnvCfg)
|
||||||
|
|||||||
@@ -201,6 +201,40 @@ class NpEnv(ABEnv):
|
|||||||
if self._physics_crash_count <= 3:
|
if self._physics_crash_count <= 3:
|
||||||
print(f"[WARN] physics crash #{self._physics_crash_count}: {e} — resetting {n} envs")
|
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):
|
def _prev_physics_step(self):
|
||||||
state = self._state
|
state = self._state
|
||||||
state.reward.fill(0.0)
|
state.reward.fill(0.0)
|
||||||
@@ -218,6 +252,8 @@ class NpEnv(ABEnv):
|
|||||||
if getattr(self, "_physics_crashed_this_step", False):
|
if getattr(self, "_physics_crashed_this_step", False):
|
||||||
self._reset_done_envs()
|
self._reset_done_envs()
|
||||||
return self._state
|
return self._state
|
||||||
|
if self._reset_invalid_physics_states():
|
||||||
|
return self._state
|
||||||
self._state = self.update_state(self._state)
|
self._state = self.update_state(self._state)
|
||||||
self._state.info["steps"] += 1
|
self._state.info["steps"] += 1
|
||||||
self._update_truncate()
|
self._update_truncate()
|
||||||
|
|||||||
19
motrix_envs/tests/test_dreamwaq_state_safety.py
Normal file
19
motrix_envs/tests/test_dreamwaq_state_safety.py
Normal 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
|
||||||
Reference in New Issue
Block a user