mjlab/tests/test_tracking_task.py
Upstream Snapshot 32a241c28f
Some checks failed
nightly / Test against latest dependencies (py3.10) (push) Has been cancelled
nightly / Test against latest dependencies (py3.13) (push) Has been cancelled
tests / tests (3.13, locked) (push) Has been cancelled
tests / tests (3.13, unlocked) (push) Has been cancelled
tests / pyright (3.10) (push) Has been cancelled
tests / lint-format (push) Has been cancelled
tests / tests (3.10, locked) (push) Has been cancelled
tests / tests (3.11, locked) (push) Has been cancelled
tests / tests (3.12, locked) (push) Has been cancelled
tests / pyright (3.11) (push) Has been cancelled
tests / pyright (3.12) (push) Has been cancelled
tests / pyright (3.13) (push) Has been cancelled
tests / ty-check (3.10) (push) Has been cancelled
tests / ty-check (3.11) (push) Has been cancelled
tests / ty-check (3.12) (push) Has been cancelled
tests / ty-check (3.13) (push) Has been cancelled
tests / stubs (push) Has been cancelled
tests / smoke-test (push) Has been cancelled
Docker / check_paths (push) Has been cancelled
docs / build (push) Has been cancelled
Docker / build (push) Has been cancelled
Import upstream snapshot c19f713c415a699a79d71cd96aa13c3104a05047
Upstream: https://github.com/michaelgillett/mjlab
Upstream-Commit: c19f713c415a699a79d71cd96aa13c3104a05047
Upstream-Branch: main
2026-08-28 15:42:17 +08:00

135 lines
4.6 KiB
Python

"""Tests specific to motion tracking tasks."""
import pytest
from mjlab.asset_zoo.robots import G1_ACTION_SCALE
from mjlab.envs.mdp.actions import JointPositionActionCfg
from mjlab.tasks.registry import list_tasks, load_env_cfg
from mjlab.tasks.tracking.mdp import MotionCommandCfg
@pytest.fixture(scope="module")
def tracking_task_ids() -> list[str]:
"""Get all tracking task IDs."""
return [t for t in list_tasks() if "Tracking" in t]
@pytest.fixture(scope="module")
def g1_tracking_task_ids(tracking_task_ids: list[str]) -> list[str]:
"""Get all G1 tracking task IDs."""
return [t for t in tracking_task_ids if "G1" in t]
def test_tracking_tasks_have_motion_command(tracking_task_ids: list[str]) -> None:
"""All tracking tasks should have a 'motion' command of type MotionCommandCfg."""
for task_id in tracking_task_ids:
cfg = load_env_cfg(task_id)
assert "motion" in cfg.commands, f"Task {task_id} missing 'motion' command"
motion_cmd = cfg.commands["motion"]
assert isinstance(motion_cmd, MotionCommandCfg), (
f"Task {task_id} motion command is not MotionCommandCfg"
)
def test_tracking_tasks_have_self_collision_sensor(
tracking_task_ids: list[str],
) -> None:
"""All tracking tasks should have a self_collision sensor."""
for task_id in tracking_task_ids:
cfg = load_env_cfg(task_id)
assert cfg.scene.sensors is not None, f"Task {task_id} has no sensors"
sensor_names = {s.name for s in cfg.scene.sensors}
assert "self_collision" in sensor_names, (
f"Task {task_id} missing self_collision sensor"
)
def test_tracking_no_state_estimation_observations() -> None:
"""No-state-estimation tasks remove observations that depend on state estimation."""
task_id = "Mjlab-Tracking-Flat-Unitree-G1-No-State-Estimation"
# Test both training and play modes
for play_mode in [False, True]:
cfg = load_env_cfg(task_id, play=play_mode)
mode_str = "play mode" if play_mode else "training mode"
assert "actor" in cfg.observations, (
f"Task {task_id} ({mode_str}) missing policy observations"
)
actor_terms = cfg.observations["actor"].terms
assert "motion_anchor_pos_b" not in actor_terms, (
f"Task {task_id} ({mode_str}) has motion_anchor_pos_b in policy, "
"expected it to be removed for no-state-estimation variant"
)
assert "base_lin_vel" not in actor_terms, (
f"Task {task_id} ({mode_str}) has base_lin_vel in policy, "
"expected it to be removed for no-state-estimation variant"
)
def test_tracking_play_disables_rsi_randomization() -> None:
"""Tracking play tasks should disable RSI randomization."""
tracking_tasks = [
"Mjlab-Tracking-Flat-Unitree-G1",
"Mjlab-Tracking-Flat-Unitree-G1-No-State-Estimation",
]
for task_id in tracking_tasks:
cfg = load_env_cfg(task_id, play=True)
motion_cmd = cfg.commands["motion"]
assert isinstance(motion_cmd, MotionCommandCfg), (
f"Task {task_id} (play mode) motion command is not MotionCommandCfg"
)
assert motion_cmd.pose_range == {}, (
f"Task {task_id} (play mode) has non-empty pose_range={motion_cmd.pose_range}, "
"expected empty dict for disabled RSI"
)
assert motion_cmd.velocity_range == {}, (
f"Task {task_id} (play mode) has non-empty velocity_range={motion_cmd.velocity_range}, "
"expected empty dict for disabled RSI"
)
def test_tracking_play_uses_start_sampling_mode() -> None:
"""Tracking play tasks should use sampling_mode='start'."""
tracking_tasks = [
"Mjlab-Tracking-Flat-Unitree-G1",
"Mjlab-Tracking-Flat-Unitree-G1-No-State-Estimation",
]
for task_id in tracking_tasks:
cfg = load_env_cfg(task_id, play=True)
motion_cmd = cfg.commands["motion"]
assert isinstance(motion_cmd, MotionCommandCfg), (
f"Task {task_id} (play mode) motion command is not MotionCommandCfg"
)
assert motion_cmd.sampling_mode == "start", (
f"Task {task_id} (play mode) sampling_mode={motion_cmd.sampling_mode}, expected 'start'"
)
def test_g1_tracking_has_correct_action_scale(g1_tracking_task_ids: list[str]) -> None:
"""G1 tracking tasks should use G1_ACTION_SCALE."""
for task_id in g1_tracking_task_ids:
cfg = load_env_cfg(task_id)
assert "joint_pos" in cfg.actions, f"Task {task_id} missing 'joint_pos' action"
joint_pos_action = cfg.actions["joint_pos"]
assert isinstance(joint_pos_action, JointPositionActionCfg), (
f"Task {task_id} joint_pos action is not JointPositionActionCfg"
)
assert joint_pos_action.scale == G1_ACTION_SCALE, (
f"Task {task_id} action scale mismatch, expected G1_ACTION_SCALE"
)