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
Upstream: https://github.com/michaelgillett/mjlab Upstream-Commit: c19f713c415a699a79d71cd96aa13c3104a05047 Upstream-Branch: main
189 lines
6.6 KiB
Python
189 lines
6.6 KiB
Python
"""Tests for sim.py."""
|
|
|
|
import mujoco
|
|
import mujoco_warp as mjwarp
|
|
import numpy as np
|
|
import pytest
|
|
import torch
|
|
from conftest import get_test_device
|
|
|
|
from mjlab.sim import MujocoCfg, Simulation, SimulationCfg
|
|
|
|
|
|
@pytest.fixture
|
|
def device():
|
|
"""Test device fixture."""
|
|
return get_test_device()
|
|
|
|
|
|
@pytest.fixture
|
|
def robot_xml():
|
|
"""Simple robot with geoms and joints."""
|
|
return """
|
|
<mujoco>
|
|
<worldbody>
|
|
<body name="base" pos="0 0 1">
|
|
<freejoint name="free_joint"/>
|
|
<geom name="base_geom" type="box" size="0.1 0.1 0.1" mass="1.0"
|
|
friction="0.5 0.01 0.005"/>
|
|
<body name="foot1" pos="0.2 0 0">
|
|
<joint name="joint1" type="hinge" axis="0 0 1" range="0 1.57"/>
|
|
<geom name="foot1_geom" type="box" size="0.05 0.05 0.05" mass="0.1"
|
|
friction="0.5 0.01 0.005"/>
|
|
</body>
|
|
<body name="foot2" pos="-0.2 0 0">
|
|
<joint name="joint2" type="hinge" axis="0 0 1" range="0 1.57"/>
|
|
<geom name="foot2_geom" type="box" size="0.05 0.05 0.05" mass="0.1"
|
|
friction="0.5 0.01 0.005"/>
|
|
</body>
|
|
</body>
|
|
</worldbody>
|
|
</mujoco>
|
|
"""
|
|
|
|
|
|
def test_simulation_config_is_piped(robot_xml, device):
|
|
"""Test that SimulationCfg values are applied to both mj_model and wp_model."""
|
|
model = mujoco.MjModel.from_xml_string(robot_xml)
|
|
|
|
cfg = SimulationCfg(
|
|
contact_sensor_maxmatch=128,
|
|
broadphase="sap_tile",
|
|
broadphase_filter=("plane", "aabb"),
|
|
mujoco=MujocoCfg(
|
|
timestep=0.02,
|
|
integrator="euler",
|
|
solver="cg",
|
|
iterations=7,
|
|
ls_iterations=14,
|
|
ccd_iterations=20,
|
|
gravity=(0, 0, 7.5),
|
|
enableflags=("energy",),
|
|
),
|
|
)
|
|
|
|
sim = Simulation(num_envs=1, cfg=cfg, model=model, device=device)
|
|
|
|
# MujocoCfg should be applied to mj_model.
|
|
assert sim.mj_model.opt.timestep == cfg.mujoco.timestep
|
|
assert sim.mj_model.opt.integrator == mujoco.mjtIntegrator.mjINT_EULER
|
|
assert sim.mj_model.opt.solver == mujoco.mjtSolver.mjSOL_CG
|
|
assert sim.mj_model.opt.iterations == cfg.mujoco.iterations
|
|
assert sim.mj_model.opt.ls_iterations == cfg.mujoco.ls_iterations
|
|
assert sim.mj_model.opt.ccd_iterations == cfg.mujoco.ccd_iterations
|
|
assert tuple(sim.mj_model.opt.gravity) == cfg.mujoco.gravity
|
|
assert sim.mj_model.opt.enableflags & mujoco.mjtEnableBit.mjENBL_ENERGY
|
|
|
|
# MujocoCfg should be inherited by wp_model via put_model.
|
|
np.testing.assert_almost_equal(
|
|
sim.model.opt.timestep[0].cpu().numpy(), cfg.mujoco.timestep
|
|
)
|
|
np.testing.assert_almost_equal(
|
|
sim.model.opt.gravity[0].cpu().numpy(), cfg.mujoco.gravity
|
|
)
|
|
assert sim.model.opt.integrator == mujoco.mjtIntegrator.mjINT_EULER
|
|
assert sim.model.opt.solver == mujoco.mjtSolver.mjSOL_CG
|
|
assert sim.model.opt.iterations == cfg.mujoco.iterations
|
|
assert sim.model.opt.enableflags & mujoco.mjtEnableBit.mjENBL_ENERGY
|
|
|
|
# SimulationCfg's warp-only settings should be applied to wp_model.opt.
|
|
assert sim.wp_model.opt.contact_sensor_maxmatch == cfg.contact_sensor_maxmatch
|
|
assert sim.wp_model.opt.broadphase == mjwarp.BroadphaseType.SAP_TILE
|
|
assert sim.wp_model.opt.broadphase_filter == (
|
|
mjwarp.BroadphaseFilter.PLANE | mjwarp.BroadphaseFilter.AABB
|
|
)
|
|
|
|
|
|
def test_default_broadphase_keeps_put_model_heuristic(robot_xml, device):
|
|
"""Unset broadphase settings should not override put_model's own heuristic."""
|
|
model = mujoco.MjModel.from_xml_string(robot_xml)
|
|
heuristic_opt = mjwarp.put_model(model).opt
|
|
|
|
sim = Simulation(num_envs=1, cfg=SimulationCfg(), model=model, device=device)
|
|
|
|
assert sim.wp_model.opt.broadphase == heuristic_opt.broadphase
|
|
assert sim.wp_model.opt.broadphase_filter == heuristic_opt.broadphase_filter
|
|
|
|
|
|
def test_ls_parallel_is_deprecated():
|
|
"""Setting the removed ls_parallel option warns instead of erroring."""
|
|
with pytest.warns(DeprecationWarning, match="ls_parallel"):
|
|
SimulationCfg(ls_parallel=True)
|
|
|
|
|
|
def test_sim_reset_restores_initial_state(robot_xml, device):
|
|
"""Test that sim.reset() restores qpos/qvel to initial values."""
|
|
model = mujoco.MjModel.from_xml_string(robot_xml)
|
|
sim = Simulation(num_envs=2, cfg=SimulationCfg(), model=model, device=device)
|
|
|
|
qpos0 = sim.data.qpos.clone()
|
|
qvel0 = sim.data.qvel.clone()
|
|
|
|
# Run simulation to modify state.
|
|
for _ in range(10):
|
|
sim.step()
|
|
|
|
assert not torch.allclose(sim.data.qpos, qpos0)
|
|
assert not torch.allclose(sim.data.qvel, qvel0)
|
|
|
|
# Reset should restore initial state.
|
|
sim.reset()
|
|
|
|
torch.testing.assert_close(sim.data.qpos[:], qpos0)
|
|
torch.testing.assert_close(sim.data.qvel[:], qvel0)
|
|
# qacc_warmstart should be zeroed.
|
|
assert (sim.data.qacc_warmstart == 0).all()
|
|
|
|
|
|
@pytest.mark.skipif(not torch.cuda.is_available(), reason="Likely bug on CPU MjWarp")
|
|
def test_sim_reset_selective(robot_xml, device):
|
|
"""Test that sim.reset() only affects specified environments."""
|
|
model = mujoco.MjModel.from_xml_string(robot_xml)
|
|
sim = Simulation(num_envs=4, cfg=SimulationCfg(), model=model, device=device)
|
|
|
|
qpos0 = sim.data.qpos.clone()
|
|
|
|
# Run simulation to modify state.
|
|
for _ in range(10):
|
|
sim.step()
|
|
|
|
qpos_after_sim = sim.data.qpos.clone()
|
|
|
|
# Reset only env 1 and 3.
|
|
sim.reset(torch.tensor([1, 3], device=device))
|
|
|
|
# Envs 1 and 3 should be reset.
|
|
torch.testing.assert_close(sim.data.qpos[1], qpos0[1])
|
|
torch.testing.assert_close(sim.data.qpos[3], qpos0[3])
|
|
# Envs 0 and 2 should be unchanged.
|
|
torch.testing.assert_close(sim.data.qpos[0], qpos_after_sim[0])
|
|
torch.testing.assert_close(sim.data.qpos[2], qpos_after_sim[2])
|
|
|
|
|
|
def test_xpos_matches_qpos_after_forward(robot_xml, device):
|
|
"""sim.step() leaves xpos stale; sim.forward() makes it match qpos.
|
|
|
|
In MuJoCo, mj_step = mj_step1 (forward kinematics + forces) + mj_step2
|
|
(integration). After mj_step, qpos/qvel are post-integration but xpos is
|
|
from the pre-integration forward pass. sim.forward() recomputes xpos from
|
|
the current qpos.
|
|
"""
|
|
model = mujoco.MjModel.from_xml_string(robot_xml)
|
|
cfg = SimulationCfg(mujoco=MujocoCfg(timestep=0.01)) # Large dt for clear signal
|
|
sim = Simulation(num_envs=2, cfg=cfg, model=model, device=device)
|
|
|
|
# Step enough for significant velocity -> large staleness gap.
|
|
for _ in range(50):
|
|
sim.step()
|
|
|
|
# xpos is stale: reflects pre-integration state of last step.
|
|
# For the freejoint body (body 1), qpos[:3] is the true position.
|
|
xpos_stale = sim.data.xpos[:, 1].clone()
|
|
qpos_pos = sim.data.qpos[:, :3].clone()
|
|
assert not torch.allclose(xpos_stale, qpos_pos, atol=1e-4)
|
|
|
|
# forward() refreshes derived quantities from current qpos.
|
|
sim.forward()
|
|
xpos_fresh = sim.data.xpos[:, 1].clone()
|
|
torch.testing.assert_close(xpos_fresh, qpos_pos, atol=1e-5, rtol=0)
|