Improvements to mujoco.rollout:
- `mjSTATE_FULLPHYSICS` as state spec, enabling divergence detection by inspecting time. - User-defined control spec. - Stop squeezing: outputs always have dim=3. PiperOrigin-RevId: 600445256 Change-Id: I4466e88929cb7081e1c94968a5cfe10485bb7475
This commit is contained in:
committed by
Copybara-Service
parent
5cbaa23388
commit
aceb52bd09
+301
-199
@@ -14,15 +14,16 @@
|
||||
# ==============================================================================
|
||||
"""tests for rollout function."""
|
||||
|
||||
import concurrent.futures
|
||||
import threading
|
||||
|
||||
from absl.testing import absltest
|
||||
from absl.testing import parameterized
|
||||
import mujoco
|
||||
import numpy as np
|
||||
import concurrent.futures
|
||||
import threading
|
||||
from mujoco import rollout
|
||||
import numpy as np
|
||||
|
||||
#--------------------------- models used for testing ---------------------------
|
||||
# -------------------------- models used for testing ---------------------------
|
||||
|
||||
TEST_XML = r"""
|
||||
<mujoco>
|
||||
@@ -96,7 +97,7 @@ TEST_XML_MOCAP = r"""
|
||||
</worldbody>
|
||||
<sensor>
|
||||
<framepos objtype="xbody" objname="1"/>
|
||||
<framequat objtype="xbody" objname="2"/>
|
||||
<framequat objtype="xbody" objname="1"/>
|
||||
</sensor>
|
||||
</mujoco>
|
||||
"""
|
||||
@@ -106,12 +107,33 @@ TEST_XML_EMPTY = r"""
|
||||
</mujoco>
|
||||
"""
|
||||
|
||||
TEST_XML_DIVERGE = r"""
|
||||
<mujoco>
|
||||
<option>
|
||||
<flag gravity="disable"/>
|
||||
</option>
|
||||
|
||||
<worldbody>
|
||||
<geom type="plane" size="5 5 .1"/>
|
||||
<body pos="0 0 -.3" euler="30 45 90">
|
||||
<freejoint/>
|
||||
<geom type="box" size=".1 .2 .4"/>
|
||||
</body>
|
||||
</worldbody>
|
||||
|
||||
<keyframe>
|
||||
<key name="non-diverging" qpos="0 0 .5 1 0 0 0"/>
|
||||
</keyframe>
|
||||
</mujoco>
|
||||
"""
|
||||
|
||||
ALL_MODELS = {'TEST_XML': TEST_XML,
|
||||
'TEST_XML_NO_SENSORS': TEST_XML_NO_SENSORS,
|
||||
'TEST_XML_NO_ACTUATORS': TEST_XML_NO_ACTUATORS,
|
||||
'TEST_XML_EMPTY': TEST_XML_EMPTY}
|
||||
|
||||
#------------------------------- tests -----------------------------------------
|
||||
# ------------------------------ tests -----------------------------------------
|
||||
|
||||
|
||||
class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
|
||||
@@ -119,179 +141,209 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
super().setUp()
|
||||
np.random.seed(42)
|
||||
|
||||
#----------------------------- test basic operation
|
||||
# ----------------------------- test basic operation
|
||||
|
||||
@parameterized.parameters(ALL_MODELS.keys())
|
||||
def test_single_step(self, model_name):
|
||||
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
initial_state = np.random.randn(model.nq + model.nv + model.na)
|
||||
ctrl = np.random.randn(model.nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, ctrl)
|
||||
initial_state = np.random.randn(nstate)
|
||||
control = np.random.randn(model.nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
mujoco.mj_resetData(model, data)
|
||||
py_state, py_sensordata = step(model, data, initial_state, ctrl=ctrl)
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
np.testing.assert_array_equal(state, py_state)
|
||||
np.testing.assert_array_equal(sensordata, py_sensordata)
|
||||
|
||||
|
||||
|
||||
@parameterized.parameters(ALL_MODELS.keys())
|
||||
def test_single_rollout(self, model_name):
|
||||
def test_one_rollout(self, model_name):
|
||||
nstep = 3 # number of timesteps
|
||||
|
||||
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
initial_state = np.random.randn(model.nq + model.nv + model.na)
|
||||
ctrl = np.random.randn(nstep, model.nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, ctrl)
|
||||
initial_state = np.random.randn(nstate)
|
||||
control = np.random.randn(nstep, model.nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
py_state, py_sensordata = single_rollout(model, data, initial_state,
|
||||
ctrl=ctrl)
|
||||
np.testing.assert_array_equal(state, np.asarray(py_state))
|
||||
np.testing.assert_array_equal(sensordata, np.asarray(py_sensordata))
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
np.testing.assert_array_equal(state, py_state)
|
||||
np.testing.assert_array_equal(sensordata, py_sensordata)
|
||||
|
||||
@parameterized.parameters(ALL_MODELS.keys())
|
||||
def test_multi_step(self, model_name):
|
||||
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nstate = 5 # number of initial states
|
||||
nroll = 5 # number of rollouts
|
||||
nstep = 1 # number of steps
|
||||
|
||||
initial_state = np.random.randn(nstate, model.nq + model.nv + model.na)
|
||||
ctrl = np.random.randn(nstate, 1, model.nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, ctrl)
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
control = np.random.randn(nroll, nstep, model.nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
mujoco.mj_resetData(model, data)
|
||||
py_state, py_sensordata = multi_rollout(model, data, initial_state,
|
||||
ctrl=ctrl)
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
np.testing.assert_array_equal(state, py_state)
|
||||
np.testing.assert_array_equal(sensordata, py_sensordata)
|
||||
|
||||
@parameterized.parameters(ALL_MODELS.keys())
|
||||
def test_single_rollout_fixed_ctrl(self, model_name):
|
||||
nstep = 3
|
||||
def test_one_rollout_fixed_ctrl(self, model_name):
|
||||
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
initial_state = np.random.randn(model.nq + model.nv + model.na)
|
||||
ctrl = np.random.randn(model.nu)
|
||||
state = np.empty((nstep, model.nq + model.nv + model.na))
|
||||
sensordata = np.empty((nstep, model.nsensordata))
|
||||
rollout.rollout(model, data, initial_state, ctrl,
|
||||
nroll = 1 # number of rollouts
|
||||
nstep = 3 # number of steps
|
||||
|
||||
initial_state = np.random.randn(nstate)
|
||||
control = np.random.randn(model.nu)
|
||||
state = np.empty((nroll, nstep, nstate))
|
||||
sensordata = np.empty((nroll, nstep, model.nsensordata))
|
||||
rollout.rollout(model, data, initial_state, control,
|
||||
state=state, sensordata=sensordata)
|
||||
|
||||
ctrl = np.tile(ctrl, (nstep, 1)) # repeat??
|
||||
py_state, py_sensordata = single_rollout(model, data, initial_state,
|
||||
ctrl=ctrl)
|
||||
control = np.tile(control, (nstep, 1))
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
np.testing.assert_array_equal(state, py_state)
|
||||
np.testing.assert_array_equal(sensordata, py_sensordata)
|
||||
|
||||
@parameterized.parameters(ALL_MODELS.keys())
|
||||
def test_multi_rollout(self, model_name):
|
||||
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nstate = 2 # number of initial states
|
||||
nroll = 2 # number of initial states
|
||||
nstep = 3 # number of timesteps
|
||||
|
||||
initial_state = np.random.randn(nstate, model.nq + model.nv + model.na)
|
||||
ctrl = np.random.randn(nstate, nstep, model.nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, ctrl)
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
control = np.random.randn(nroll, nstep, model.nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
py_state, py_sensordata = multi_rollout(model, data, initial_state,
|
||||
ctrl=ctrl)
|
||||
np.testing.assert_array_equal(py_state, py_state)
|
||||
np.testing.assert_array_equal(py_sensordata, py_sensordata)
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
np.testing.assert_array_equal(state, py_state)
|
||||
np.testing.assert_array_equal(sensordata, py_sensordata)
|
||||
|
||||
@parameterized.parameters(ALL_MODELS.keys())
|
||||
def test_multi_rollout_fixed_ctrl_infer_from_output(self, model_name):
|
||||
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nstate = 2 # number of initial states
|
||||
nroll = 2 # number of rollouts
|
||||
nstep = 3 # number of timesteps
|
||||
|
||||
initial_state = np.random.randn(nstate, model.nq + model.nv + model.na)
|
||||
ctrl = np.random.randn(nstate, 1, model.nu) # 1 control in the time dimension
|
||||
state = np.empty((nstate, nstep, model.nq + model.nv + model.na))
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, ctrl,
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
control = np.random.randn(nroll, 1, model.nu)
|
||||
state = np.empty((nroll, nstep, nstate))
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control,
|
||||
state=state)
|
||||
|
||||
ctrl = np.repeat(ctrl, nstep, axis=1)
|
||||
py_state, py_sensordata = multi_rollout(model, data, initial_state,
|
||||
ctrl=ctrl)
|
||||
control = np.repeat(control, nstep, axis=1)
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
np.testing.assert_array_equal(state, py_state)
|
||||
np.testing.assert_array_equal(sensordata, py_sensordata)
|
||||
|
||||
@parameterized.product(arg_nstep=[[3, 1, 1], [3, 3, 1], [3, 1, 3]],
|
||||
model_name=list(ALL_MODELS.keys()))
|
||||
def test_multi_rollout_multiple_inputs(self, arg_nstep, model_name):
|
||||
@parameterized.parameters(ALL_MODELS.keys())
|
||||
def test_py_rollout_generalized_control(self, model_name):
|
||||
model = mujoco.MjModel.from_xml_string(ALL_MODELS[model_name])
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nstate = 4 # number of initial states
|
||||
nroll = 4 # number of rollouts
|
||||
nstep = 3 # number of timesteps
|
||||
|
||||
initial_state = np.random.randn(nstate, model.nq + model.nv + model.na)
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
|
||||
# arg_nstep is the horizon for {ctrl, qfrc_applied, xfrc_applied}, respectively
|
||||
ctrl = np.random.randn(nstate, arg_nstep[0], model.nu)
|
||||
qfrc_applied = np.random.randn(nstate, arg_nstep[1], model.nv)
|
||||
xfrc_applied = np.random.randn(nstate, arg_nstep[2], model.nbody*6)
|
||||
control_spec = (mujoco.mjtState.mjSTATE_CTRL |
|
||||
mujoco.mjtState.mjSTATE_QFRC_APPLIED |
|
||||
mujoco.mjtState.mjSTATE_XFRC_APPLIED)
|
||||
ncontrol = mujoco.mj_stateSize(model, control_spec)
|
||||
control = np.random.randn(nroll, nstep, ncontrol)
|
||||
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, ctrl,
|
||||
qfrc_applied=qfrc_applied,
|
||||
xfrc_applied=xfrc_applied)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control,
|
||||
control_spec=control_spec)
|
||||
|
||||
# tile singleton arguments
|
||||
nstep = max(arg_nstep)
|
||||
if arg_nstep[0] == 1:
|
||||
ctrl = np.repeat(ctrl, nstep, axis=1)
|
||||
if arg_nstep[1] == 1:
|
||||
qfrc_applied = np.repeat(qfrc_applied, nstep, axis=1)
|
||||
if arg_nstep[2] == 1:
|
||||
xfrc_applied = np.repeat(xfrc_applied, nstep, axis=1)
|
||||
|
||||
py_state, py_sensordata = multi_rollout(model, data, initial_state,
|
||||
ctrl=ctrl,
|
||||
qfrc_applied=qfrc_applied,
|
||||
xfrc_applied=xfrc_applied)
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control,
|
||||
control_spec=control_spec)
|
||||
np.testing.assert_array_equal(state, py_state)
|
||||
np.testing.assert_array_equal(sensordata, py_sensordata)
|
||||
|
||||
#----------------------------- test threaded operation
|
||||
def test_detect_divergence(self):
|
||||
model = mujoco.MjModel.from_xml_string(TEST_XML_DIVERGE)
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nroll = 4 # number of rollouts
|
||||
initial_state = np.empty((nroll, nstate))
|
||||
|
||||
# get diverging (0, 2) and non-diverging (1, 3) states
|
||||
mujoco.mj_getState(model, data, initial_state[0],
|
||||
mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
mujoco.mj_getState(model, data, initial_state[2],
|
||||
mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
mujoco.mj_resetDataKeyframe(model, data, 0) # keyframe 0 does not diverge
|
||||
mujoco.mj_getState(model, data, initial_state[1],
|
||||
mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
mujoco.mj_getState(model, data, initial_state[3],
|
||||
mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
|
||||
nstep = 10000 # divergence after ~15s, timestep = 2e-3
|
||||
|
||||
state = np.random.randn(nroll, nstep, nstate)
|
||||
|
||||
rollout.rollout(model, data, initial_state, state=state)
|
||||
|
||||
# initial_state[0,2] diverged, final timesteps are identical
|
||||
assert state[0][-1][0] == state[0][-2][0]
|
||||
assert state[2][-1][0] == state[2][-2][0]
|
||||
|
||||
# initial_state[1,3] did not diverge, final timesteps are different
|
||||
assert state[1][-1][0] != state[1][-2][0]
|
||||
assert state[3][-1][0] != state[3][-2][0]
|
||||
|
||||
# ----------------------------- test threaded operation
|
||||
|
||||
def test_threading(self):
|
||||
model = mujoco.MjModel.from_xml_string(TEST_XML)
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
num_workers = 32
|
||||
nstate = 10000
|
||||
nroll = 10000
|
||||
nstep = 5
|
||||
initial_state = np.random.randn(nstate, model.nq+model.nv+model.na)
|
||||
state = np.zeros((nstate, nstep, model.nq+model.nv+model.na))
|
||||
sensordata = np.zeros((nstate, nstep, model.nsensordata))
|
||||
ctrl = np.random.randn(nstate, nstep, model.nu)
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
state = np.empty((nroll, nstep, nstate))
|
||||
sensordata = np.empty((nroll, nstep, model.nsensordata))
|
||||
control = np.random.randn(nroll, nstep, model.nu)
|
||||
|
||||
thread_local = threading.local()
|
||||
|
||||
def thread_initializer():
|
||||
thread_local.data = mujoco.MjData(model)
|
||||
|
||||
def call_rollout(initial_state, ctrl, state):
|
||||
rollout.rollout(model, thread_local.data, skip_checks=True,
|
||||
nstate=initial_state.shape[0], nstep=nstep,
|
||||
initial_state=initial_state, ctrl=ctrl, state=state)
|
||||
def call_rollout(initial_state, control, state, sensordata):
|
||||
rollout.rollout(model, thread_local.data, initial_state, control,
|
||||
skip_checks=True, nroll=initial_state.shape[0],
|
||||
nstep=nstep, state=state, sensordata=sensordata)
|
||||
|
||||
n = initial_state.shape[0] // num_workers # integer division
|
||||
n = nroll // num_workers # integer division
|
||||
chunks = [] # a list of tuples, one per worker
|
||||
for i in range(num_workers-1):
|
||||
chunks.append(
|
||||
(initial_state[i*n:(i+1)*n], ctrl[i*n:(i+1)*n], state[i*n:(i+1)*n]))
|
||||
chunks.append((initial_state[i*n:(i+1)*n],
|
||||
control[i*n:(i+1)*n],
|
||||
state[i*n:(i+1)*n],
|
||||
sensordata[i*n:(i+1)*n]))
|
||||
|
||||
# last chunk, absorbing the remainder:
|
||||
chunks.append(
|
||||
(initial_state[(num_workers-1)*n:], ctrl[(num_workers-1)*n:],
|
||||
state[(num_workers-1)*n:]))
|
||||
chunks.append((initial_state[(num_workers-1)*n:],
|
||||
control[(num_workers-1)*n:],
|
||||
state[(num_workers-1)*n:],
|
||||
sensordata[(num_workers-1)*n:]))
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor(
|
||||
max_workers=num_workers, initializer=thread_initializer) as executor:
|
||||
@@ -302,187 +354,237 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
future.result()
|
||||
|
||||
data = mujoco.MjData(model)
|
||||
py_state, py_sensordata = multi_rollout(model, data, initial_state,
|
||||
ctrl=ctrl)
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
np.testing.assert_array_equal(state, py_state)
|
||||
np.testing.assert_array_equal(sensordata, py_sensordata)
|
||||
|
||||
#----------------------------- test advanced operation
|
||||
|
||||
def test_time(self):
|
||||
model = mujoco.MjModel.from_xml_string(TEST_XML)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
nstate = 1
|
||||
nstep = 3
|
||||
|
||||
initial_time = np.array([[2.]])
|
||||
initial_state = np.random.randn(nstate, model.nq + model.nv + model.na)
|
||||
ctrl = np.random.randn(nstate, nstep, model.nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, ctrl,
|
||||
initial_time=initial_time)
|
||||
|
||||
self.assertAlmostEqual(data.time, 2 + nstep*model.opt.timestep)
|
||||
# ---------------------------- test advanced operation
|
||||
|
||||
def test_warmstart(self):
|
||||
model = mujoco.MjModel.from_xml_string(TEST_XML)
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
state0 = np.zeros(model.nq + model.nv + model.na)
|
||||
ctrl = np.zeros(model.nu)
|
||||
state1, _ = step(model, data, state0, ctrl=ctrl)
|
||||
# take one step, save the state
|
||||
state0 = np.zeros(nstate)
|
||||
control = np.zeros(model.nu)
|
||||
state1, _ = step(model, data, state0, control)
|
||||
|
||||
# save qacc_warmstart
|
||||
initial_warmstart = data.qacc_warmstart.copy()
|
||||
|
||||
state2, _ = step(model, data, state1, ctrl=ctrl)
|
||||
# take one more step (uses correct warmstart)
|
||||
state2, _ = step(model, data, state1[0], control)
|
||||
|
||||
state, _ = rollout.rollout(model, data, state1, ctrl)
|
||||
assert np.linalg.norm(state-state2) > 0
|
||||
# take step using rollout, don't take warmstart into account
|
||||
state, _ = rollout.rollout(model, data, state1[0], control)
|
||||
|
||||
state, _ = rollout.rollout(model, data, state1, ctrl,
|
||||
# assert that stepping without warmstarts is not exact
|
||||
np.testing.assert_raises(AssertionError,
|
||||
np.testing.assert_array_equal, state, state2)
|
||||
|
||||
# take step using rollout, take warmstart into account
|
||||
state, _ = rollout.rollout(model, data, state1, control,
|
||||
initial_warmstart=initial_warmstart)
|
||||
np.testing.assert_array_equal(state, state2)
|
||||
|
||||
# assert exact equality
|
||||
np.testing.assert_array_equal(state, np.expand_dims(state2, axis=0))
|
||||
|
||||
def test_mocap(self):
|
||||
model = mujoco.MjModel.from_xml_string(TEST_XML_MOCAP)
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
initial_state = np.zeros(model.nq + model.nv + model.na)
|
||||
initial_state = np.zeros(nstate)
|
||||
|
||||
control_spec = (mujoco.mjtState.mjSTATE_MOCAP_POS |
|
||||
mujoco.mjtState.mjSTATE_MOCAP_QUAT)
|
||||
|
||||
pos1 = np.array((1., 2., 3.))
|
||||
quat1 = np.array((1., 2., 3., 4.))
|
||||
quat1 /= np.linalg.norm(quat1)
|
||||
pos2 = np.array((2., 3., 4.))
|
||||
quat2 = np.array((2., 3., 4., 5.))
|
||||
quat2 /= np.linalg.norm(quat2)
|
||||
mocap = np.hstack((pos1, quat1, pos2, quat2))
|
||||
control = np.hstack((pos1, pos2, quat1, quat2))
|
||||
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, mocap=mocap)
|
||||
_, sensordata = rollout.rollout(model, data, initial_state, control,
|
||||
control_spec=control_spec)
|
||||
|
||||
np.testing.assert_array_almost_equal(sensordata[:3], pos1)
|
||||
np.testing.assert_array_almost_equal(sensordata[3:], quat2)
|
||||
np.testing.assert_array_almost_equal(sensordata[0][0][:3], pos1)
|
||||
np.testing.assert_array_almost_equal(sensordata[0][0][3:], quat1)
|
||||
|
||||
#----------------------------- test correctness
|
||||
# ---------------------------- test correctness
|
||||
|
||||
def test_intercept_mj_errors(self):
|
||||
model = mujoco.MjModel.from_xml_string(TEST_XML)
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
initial_state = np.zeros(model.nq + model.nv + model.na)
|
||||
ctrl = np.zeros((3, model.nu))
|
||||
nroll = 1
|
||||
nstep = 3
|
||||
|
||||
initial_state = np.zeros((nroll, nstate))
|
||||
ctrl = np.zeros((nroll, nstep, model.nu))
|
||||
|
||||
model.opt.solver = 10 # invalid solver type
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
mujoco.FatalError, 'mj_fwdConstraint: unknown solver type 10'):
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, ctrl)
|
||||
rollout.rollout(model, data, initial_state, ctrl)
|
||||
|
||||
def test_invalid(self):
|
||||
model = mujoco.MjModel.from_xml_string(TEST_XML)
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
initial_state = np.zeros(model.nq + model.nv + model.na)
|
||||
nroll = 1
|
||||
|
||||
ctrl = 'string'
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'ctrl must be a numpy array or float'):
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, ctrl)
|
||||
initial_state = np.zeros((nroll, nstate))
|
||||
|
||||
qfrc_applied = np.zeros((2, 3, 4, 5))
|
||||
control = 'string'
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'qfrc_applied can have at most 3 dimensions'):
|
||||
state, sensordata = rollout.rollout(model, data, initial_state,
|
||||
qfrc_applied=qfrc_applied)
|
||||
ValueError, 'control must be a numpy array or float'):
|
||||
rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
control = np.zeros((2, 3, 4, 5))
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'control can have at most 3 dimensions'):
|
||||
rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
def test_bad_sizes(self):
|
||||
model = mujoco.MjModel.from_xml_string(TEST_XML)
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
initial_state = np.random.randn(model.nq + model.nv + model.na+1)
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'trailing dimension of initial_state must be 5, got 6'):
|
||||
state, sensordata = rollout.rollout(model, data, initial_state)
|
||||
nroll = 1
|
||||
nstep = 3
|
||||
|
||||
initial_state = np.random.randn(model.nq + model.nv + model.na)
|
||||
ctrl = np.random.randn(model.nu+1)
|
||||
initial_state = np.random.randn(nroll, nstate + 1)
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'trailing dimension of ctrl must be 2, got 3'):
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, ctrl)
|
||||
ValueError, 'trailing dimension of initial_state must be 6, got 7'):
|
||||
rollout.rollout(model, data, initial_state)
|
||||
|
||||
ctrl = np.random.randn(2, model.nu)
|
||||
qfrc_applied = np.random.randn(3, model.nv) # incompatible horizon
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
control = np.random.randn(1, nstep, model.nu + 1)
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'dimension 1 inferred as 2 but qfrc_applied has 3'):
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, ctrl,
|
||||
qfrc_applied=qfrc_applied)
|
||||
ValueError, 'trailing dimension of control must be 2, got 3'):
|
||||
rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
control = np.random.randn(nroll, nstep, model.nu)
|
||||
state = np.random.randn(nroll, nstep+1, nstate) # incompatible nstep
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'dimension 1 inferred as 3 but state has 4'):
|
||||
rollout.rollout(model, data, initial_state, control, state=state)
|
||||
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
control = np.random.randn(nroll, nstep, model.nu)
|
||||
bad_spec = mujoco.mjtState.mjSTATE_ACT
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'control_spec can only contain bits in mjSTATE_USER'):
|
||||
rollout.rollout(model, data, initial_state, control,
|
||||
control_spec=bad_spec)
|
||||
|
||||
def test_stateless(self):
|
||||
model = mujoco.MjModel.from_xml_string(TEST_XML)
|
||||
model.opt.disableflags |= mujoco.mjtDisableBit.mjDSBL_WARMSTART.value
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
data = mujoco.MjData(model)
|
||||
|
||||
# call step with a clean mjData
|
||||
initial_state = np.random.randn(model.nq + model.nv + model.na)
|
||||
ctrl = np.random.randn(model.nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, ctrl)
|
||||
# step with a clean mjData
|
||||
initial_state = np.random.randn(nstate)
|
||||
control = np.random.randn(3, 3, model.nu)
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
# fill mjData with some debug value, see that we still get the same outputs
|
||||
mujoco.mj_resetDataDebug(model, data, 255)
|
||||
debug_state, debug_sensordata = rollout.rollout(model, data, initial_state,
|
||||
ctrl)
|
||||
# fill user fields with random values
|
||||
for attr in [
|
||||
'ctrl',
|
||||
'qfrc_applied',
|
||||
'xfrc_applied',
|
||||
'mocap_pos',
|
||||
'mocap_quat',
|
||||
]:
|
||||
setattr(data, attr, np.random.randn(*getattr(data, attr).shape))
|
||||
|
||||
np.testing.assert_array_equal(state, debug_state)
|
||||
np.testing.assert_array_equal(sensordata, debug_sensordata)
|
||||
# roll out again
|
||||
state2, sensordata2 = rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
# assert that we still get the same outputs
|
||||
np.testing.assert_array_equal(state, state2)
|
||||
np.testing.assert_array_equal(sensordata, sensordata2)
|
||||
|
||||
|
||||
#--------------- Python implementation of rollout functionality ----------------
|
||||
# -------------- Python implementation of rollout functionality ----------------
|
||||
|
||||
def get_state(data):
|
||||
return np.hstack((data.qpos, data.qvel, data.act))
|
||||
|
||||
def set_state(model, data, state):
|
||||
data.qpos = state[:model.nq]
|
||||
data.qvel = state[model.nq:model.nq+model.nv]
|
||||
data.act = state[model.nq+model.nv:model.nq+model.nv+model.na]
|
||||
def get_state(model, data):
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
state = np.empty(nstate)
|
||||
mujoco.mj_getState(model, data, state, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
return state.reshape((1, nstate))
|
||||
|
||||
def step(model, data, state, **kwargs):
|
||||
|
||||
def step(model, data, state, control,
|
||||
control_spec=mujoco.mjtState.mjSTATE_CTRL):
|
||||
if state is not None:
|
||||
set_state(model, data, state)
|
||||
for key, value in kwargs.items():
|
||||
if value is not None:
|
||||
setattr(data, key, np.reshape(value, getattr(data, key).shape))
|
||||
mujoco.mj_setState(model, data, state, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
mujoco.mj_setState(model, data, control, control_spec)
|
||||
mujoco.mj_step(model, data)
|
||||
return (get_state(data), data.sensordata)
|
||||
return (get_state(model, data), data.sensordata)
|
||||
|
||||
def single_rollout(model, data, initial_state, **kwargs):
|
||||
arg_nstep = set([a.shape[0] for a in kwargs.values()])
|
||||
assert len(arg_nstep) == 1 # nstep dimensions must match
|
||||
nstep = arg_nstep.pop()
|
||||
|
||||
state = np.empty((nstep, model.nq + model.nv + model.na))
|
||||
def one_rollout(model, data, initial_state, control,
|
||||
control_spec=mujoco.mjtState.mjSTATE_CTRL):
|
||||
nstep = control.shape[0]
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
state = np.empty((nstep, nstate))
|
||||
sensordata = np.empty((nstep, model.nsensordata))
|
||||
|
||||
mujoco.mj_resetData(model, data)
|
||||
for t in range(nstep):
|
||||
kwargs_t = {}
|
||||
for key, value in kwargs.items():
|
||||
kwargs_t[key] = value[0 if value.ndim == 1 else t]
|
||||
state[t], sensordata[t] = step(model, data,
|
||||
initial_state if t==0 else None,
|
||||
**kwargs_t)
|
||||
initial_state if t == 0 else None,
|
||||
control[t], control_spec)
|
||||
return state, sensordata
|
||||
|
||||
def multi_rollout(model, data, initial_state, **kwargs):
|
||||
nstate = initial_state.shape[0]
|
||||
arg_nstep = set([a.shape[1] for a in kwargs.values()])
|
||||
assert len(arg_nstep) == 1 # nstep dimensions must match
|
||||
nstep = arg_nstep.pop()
|
||||
|
||||
state = np.empty((nstate, nstep, model.nq + model.nv + model.na))
|
||||
sensordata = np.empty((nstate, nstep, model.nsensordata))
|
||||
for s in range(nstate):
|
||||
kwargs_s = {key : value[s] for key, value in kwargs.items()}
|
||||
state_s, sensordata_s = single_rollout(model, data, initial_state[s],
|
||||
**kwargs_s)
|
||||
state[s] = state_s
|
||||
sensordata[s] = sensordata_s
|
||||
return state.squeeze(), sensordata.squeeze()
|
||||
def ensure_2d(arg):
|
||||
if arg is None:
|
||||
return None
|
||||
else:
|
||||
return np.ascontiguousarray(np.atleast_2d(arg), dtype=np.float64)
|
||||
|
||||
|
||||
def ensure_3d(arg):
|
||||
if arg is None:
|
||||
return None
|
||||
else:
|
||||
# np.atleast_3d adds both leading and trailing dims, we want only leading
|
||||
if arg.ndim == 0:
|
||||
arg = arg[np.newaxis, np.newaxis, np.newaxis, ...]
|
||||
elif arg.ndim == 1:
|
||||
arg = arg[np.newaxis, np.newaxis, ...]
|
||||
elif arg.ndim == 2:
|
||||
arg = arg[np.newaxis, ...]
|
||||
return np.ascontiguousarray(arg, dtype=np.float64)
|
||||
|
||||
|
||||
def py_rollout(model, data, initial_state, control,
|
||||
control_spec=mujoco.mjtState.mjSTATE_CTRL):
|
||||
initial_state = ensure_2d(initial_state)
|
||||
control = ensure_3d(control)
|
||||
nroll = initial_state.shape[0]
|
||||
nstep = control.shape[1]
|
||||
nstate = mujoco.mj_stateSize(model, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
|
||||
state = np.empty((nroll, nstep, nstate))
|
||||
sensordata = np.empty((nroll, nstep, model.nsensordata))
|
||||
for r in range(nroll):
|
||||
state_r, sensordata_r = one_rollout(
|
||||
model, data, initial_state[r], control[r], control_spec
|
||||
)
|
||||
state[r] = state_r
|
||||
sensordata[r] = sensordata_r
|
||||
return state, sensordata
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
absltest.main()
|
||||
|
||||
Reference in New Issue
Block a user