Add pyink and isort config. Reformat.
PiperOrigin-RevId: 704533915 Change-Id: I37e9fd51261bd166b725c7460fc65d02fed2b391
This commit is contained in:
committed by
Copybara-Service
parent
6f6244b739
commit
f3b3024291
+125
-74
@@ -127,10 +127,12 @@ TEST_XML_DIVERGE = r"""
|
||||
</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}
|
||||
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 -----------------------------------------
|
||||
|
||||
@@ -242,8 +244,9 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
initial_state = np.random.randn(nstate)
|
||||
control = np.random.randn(nstep, model.nu)
|
||||
initial_warmstart = np.tile(data.qacc_warmstart.copy(), (nroll, 1))
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control,
|
||||
initial_warmstart=initial_warmstart)
|
||||
state, sensordata = rollout.rollout(
|
||||
model, data, initial_state, control, initial_warmstart=initial_warmstart
|
||||
)
|
||||
|
||||
mujoco.mj_resetData(model, data)
|
||||
initial_state = np.tile(initial_state, (nroll, 1))
|
||||
@@ -264,8 +267,9 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
initial_state = np.random.randn(nstate)
|
||||
control = np.random.randn(nstep, model.nu)
|
||||
state = np.empty((nroll, nstep, nstate))
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control,
|
||||
state=state)
|
||||
state, sensordata = rollout.rollout(
|
||||
model, data, initial_state, control, state=state
|
||||
)
|
||||
|
||||
mujoco.mj_resetData(model, data)
|
||||
initial_state = np.tile(initial_state, (nroll, 1))
|
||||
@@ -286,8 +290,9 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
initial_state = np.random.randn(nstate)
|
||||
control = np.random.randn(nstep, model.nu)
|
||||
sensordata = np.empty((nroll, nstep, model.nsensordata))
|
||||
state, sensordata = rollout.rollout(model, data, initial_state, control,
|
||||
sensordata=sensordata)
|
||||
state, sensordata = rollout.rollout(
|
||||
model, data, initial_state, control, sensordata=sensordata
|
||||
)
|
||||
|
||||
mujoco.mj_resetData(model, data)
|
||||
initial_state = np.tile(initial_state, (nroll, 1))
|
||||
@@ -309,8 +314,9 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
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)
|
||||
rollout.rollout(
|
||||
model, data, initial_state, control, state=state, sensordata=sensordata
|
||||
)
|
||||
|
||||
control = np.tile(control, (nstep, 1))
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
@@ -374,8 +380,9 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
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)
|
||||
state, sensordata = rollout.rollout(
|
||||
model, data, initial_state, control, state=state
|
||||
)
|
||||
|
||||
control = np.repeat(control, nstep, axis=1)
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control)
|
||||
@@ -393,17 +400,21 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
|
||||
control_spec = (mujoco.mjtState.mjSTATE_CTRL |
|
||||
mujoco.mjtState.mjSTATE_QFRC_APPLIED |
|
||||
mujoco.mjtState.mjSTATE_XFRC_APPLIED)
|
||||
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, control,
|
||||
control_spec=control_spec)
|
||||
state, sensordata = rollout.rollout(
|
||||
model, data, initial_state, control, control_spec=control_spec
|
||||
)
|
||||
|
||||
py_state, py_sensordata = py_rollout(model, data, initial_state, control,
|
||||
control_spec=control_spec)
|
||||
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)
|
||||
|
||||
@@ -416,15 +427,19 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
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_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)
|
||||
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
|
||||
|
||||
@@ -459,27 +474,40 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
thread_local.data = mujoco.MjData(model)
|
||||
|
||||
model_list = [model] * nroll
|
||||
|
||||
def call_rollout(initial_state, control, state, sensordata):
|
||||
rollout.rollout(model_list, thread_local.data, initial_state, control,
|
||||
skip_checks=True,
|
||||
nstep=nstep, state=state, sensordata=sensordata)
|
||||
rollout.rollout(
|
||||
model_list,
|
||||
thread_local.data,
|
||||
initial_state,
|
||||
control,
|
||||
skip_checks=True,
|
||||
nstep=nstep,
|
||||
state=state,
|
||||
sensordata=sensordata,
|
||||
)
|
||||
|
||||
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],
|
||||
control[i*n:(i+1)*n],
|
||||
state[i*n:(i+1)*n],
|
||||
sensordata[i*n:(i+1)*n]))
|
||||
for i in range(num_workers - 1):
|
||||
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:],
|
||||
control[(num_workers-1)*n:],
|
||||
state[(num_workers-1)*n:],
|
||||
sensordata[(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:
|
||||
max_workers=num_workers, initializer=thread_initializer
|
||||
) as executor:
|
||||
futures = []
|
||||
for chunk in chunks:
|
||||
futures.append(executor.submit(call_rollout, *chunk))
|
||||
@@ -513,12 +541,14 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
state, _ = rollout.rollout(model, data, state1[0], control)
|
||||
|
||||
# assert that stepping without warmstarts is not exact
|
||||
np.testing.assert_raises(AssertionError,
|
||||
np.testing.assert_array_equal, state, state2)
|
||||
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)
|
||||
state, _ = rollout.rollout(
|
||||
model, data, state1, control, initial_warmstart=initial_warmstart
|
||||
)
|
||||
|
||||
# assert exact equality
|
||||
np.testing.assert_array_equal(state, np.expand_dims(state2, axis=0))
|
||||
@@ -530,19 +560,21 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
|
||||
initial_state = np.zeros(nstate)
|
||||
|
||||
control_spec = (mujoco.mjtState.mjSTATE_MOCAP_POS |
|
||||
mujoco.mjtState.mjSTATE_MOCAP_QUAT)
|
||||
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.))
|
||||
pos1 = np.array((1.0, 2.0, 3.0))
|
||||
quat1 = np.array((1.0, 2.0, 3.0, 4.0))
|
||||
quat1 /= np.linalg.norm(quat1)
|
||||
pos2 = np.array((2., 3., 4.))
|
||||
quat2 = np.array((2., 3., 4., 5.))
|
||||
pos2 = np.array((2.0, 3.0, 4.0))
|
||||
quat2 = np.array((2.0, 3.0, 4.0, 5.0))
|
||||
quat2 /= np.linalg.norm(quat2)
|
||||
control = np.hstack((pos1, pos2, quat1, quat2))
|
||||
|
||||
_, sensordata = rollout.rollout(model, data, initial_state, control,
|
||||
control_spec=control_spec)
|
||||
_, sensordata = rollout.rollout(
|
||||
model, data, initial_state, control, control_spec=control_spec
|
||||
)
|
||||
|
||||
np.testing.assert_array_almost_equal(sensordata[0][0][:3], pos1)
|
||||
np.testing.assert_array_almost_equal(sensordata[0][0][3:], quat1)
|
||||
@@ -562,7 +594,8 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
|
||||
model.opt.solver = 10 # invalid solver type
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
mujoco.FatalError, 'mj_fwdConstraint: unknown solver type 10'):
|
||||
mujoco.FatalError, 'mj_fwdConstraint: unknown solver type 10'
|
||||
):
|
||||
rollout.rollout(model, data, initial_state, ctrl)
|
||||
|
||||
def test_invalid(self):
|
||||
@@ -576,12 +609,14 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
|
||||
control = 'string'
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'control must be a numpy array or float'):
|
||||
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'):
|
||||
ValueError, 'control can have at most 3 dimensions'
|
||||
):
|
||||
rollout.rollout(model, data, initial_state, control)
|
||||
|
||||
def test_bad_sizes(self):
|
||||
@@ -594,28 +629,33 @@ class MuJoCoRolloutTest(parameterized.TestCase):
|
||||
|
||||
initial_state = np.random.randn(nroll, nstate + 1)
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'trailing dimension of initial_state must be 6, got 7'):
|
||||
ValueError, 'trailing dimension of initial_state must be 6, got 7'
|
||||
):
|
||||
rollout.rollout(model, data, initial_state)
|
||||
|
||||
initial_state = np.random.randn(nroll, nstate)
|
||||
control = np.random.randn(1, nstep, model.nu + 1)
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'trailing dimension of control must be 2, got 3'):
|
||||
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
|
||||
state = np.random.randn(nroll, nstep + 1, nstate) # incompatible nstep
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError, 'dimension 1 inferred as 3 but state has 4'):
|
||||
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)
|
||||
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)
|
||||
@@ -655,8 +695,9 @@ def get_state(model, data):
|
||||
return state.reshape((1, nstate))
|
||||
|
||||
|
||||
def step(model, data, state, control,
|
||||
control_spec=mujoco.mjtState.mjSTATE_CTRL):
|
||||
def step(
|
||||
model, data, state, control, control_spec=mujoco.mjtState.mjSTATE_CTRL
|
||||
):
|
||||
if state is not None:
|
||||
mujoco.mj_setState(model, data, state, mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
mujoco.mj_setState(model, data, control, control_spec)
|
||||
@@ -664,8 +705,13 @@ def step(model, data, state, control,
|
||||
return (get_state(model, data), data.sensordata)
|
||||
|
||||
|
||||
def one_rollout(model, data, initial_state, control,
|
||||
control_spec=mujoco.mjtState.mjSTATE_CTRL):
|
||||
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))
|
||||
@@ -673,9 +719,9 @@ def one_rollout(model, data, initial_state, control,
|
||||
|
||||
mujoco.mj_resetData(model, data)
|
||||
for t in range(nstep):
|
||||
state[t], sensordata[t] = step(model, data,
|
||||
initial_state if t == 0 else None,
|
||||
control[t], control_spec)
|
||||
state[t], sensordata[t] = step(
|
||||
model, data, initial_state if t == 0 else None, control[t], control_spec
|
||||
)
|
||||
return state, sensordata
|
||||
|
||||
|
||||
@@ -700,15 +746,20 @@ def ensure_3d(arg):
|
||||
return np.ascontiguousarray(arg, dtype=np.float64)
|
||||
|
||||
|
||||
def py_rollout(model, data, initial_state, control,
|
||||
control_spec=mujoco.mjtState.mjSTATE_CTRL):
|
||||
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]
|
||||
|
||||
if isinstance(model, mujoco.MjModel):
|
||||
model = [model]*nroll
|
||||
model = [model] * nroll
|
||||
|
||||
nstate = mujoco.mj_stateSize(model[0], mujoco.mjtState.mjSTATE_FULLPHYSICS)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user