Add rne_postconstraint function to MJX. This function matches mj_rnePostConstraint and computes cacc, cfrc_ext, and cfrc_int.
PiperOrigin-RevId: 671336319 Change-Id: I263b46c56a7609b9a8a30895c012798579d3dffa
This commit is contained in:
committed by
Copybara-Service
parent
cd8ff4401d
commit
7b073b60c4
@@ -43,6 +43,7 @@ from mujoco.mjx._src.smooth import crb
|
||||
from mujoco.mjx._src.smooth import factor_m
|
||||
from mujoco.mjx._src.smooth import kinematics
|
||||
from mujoco.mjx._src.smooth import rne
|
||||
from mujoco.mjx._src.smooth import rne_postconstraint
|
||||
from mujoco.mjx._src.smooth import subtree_vel
|
||||
from mujoco.mjx._src.smooth import tendon
|
||||
from mujoco.mjx._src.smooth import transmission
|
||||
|
||||
@@ -40,7 +40,7 @@ def _assert_attr_eq(a, b, attr):
|
||||
|
||||
class SensorTest(parameterized.TestCase):
|
||||
|
||||
@parameterized.parameters('no_sensor.xml', 'sensor.xml')
|
||||
@parameterized.parameters('sensor/model.xml', 'sensor/sensor.xml')
|
||||
def test_sensor(self, filename):
|
||||
"""Tests MJX sensor functions match MuJoCo sensor functions."""
|
||||
m = test_util.load_test_file(filename)
|
||||
@@ -65,7 +65,7 @@ class SensorTest(parameterized.TestCase):
|
||||
|
||||
def test_disable_sensor(self):
|
||||
"""Tests disabling sensor."""
|
||||
m = test_util.load_test_file('sensor.xml')
|
||||
m = test_util.load_test_file('sensor/sensor.xml')
|
||||
# disable sensors
|
||||
m.opt.disableflags = m.opt.disableflags | mjx.DisableBit.SENSOR
|
||||
d = mujoco.MjData(m)
|
||||
@@ -84,7 +84,7 @@ class SensorTest(parameterized.TestCase):
|
||||
|
||||
def test_unsupported_sensor(self):
|
||||
"""Tests MJX sensor functions do not break for unsupported sensors."""
|
||||
m = test_util.load_test_file('unsupported_sensor.xml')
|
||||
m = test_util.load_test_file('sensor/unsupported.xml')
|
||||
mx = mjx.put_model(m)
|
||||
dx = jax.jit(mjx.forward)(mx, mjx.put_data(m, mujoco.MjData(m)))
|
||||
_assert_eq(np.zeros(m.nsensordata), dx.sensordata, 'sensordata')
|
||||
|
||||
@@ -558,6 +558,129 @@ def rne(m: Model, d: Data) -> Data:
|
||||
return d
|
||||
|
||||
|
||||
def rne_postconstraint(m: Model, d: Data) -> Data:
|
||||
"""RNE with complete data: compute cacc, cfrc_ext, cfrc_int."""
|
||||
|
||||
def _transform_force(frc, offset):
|
||||
force, torque = jp.split(frc, 2)
|
||||
torque -= jp.cross(offset, force)
|
||||
# spatial motion vector layout is flipped: (torque, force)
|
||||
return jp.concatenate([torque, force])
|
||||
|
||||
# cfrc_ext = perturb
|
||||
cfrc_ext = jp.vstack([
|
||||
jp.zeros((1, 6)), # world body
|
||||
jax.vmap(_transform_force)(
|
||||
d.xfrc_applied[1:], d.subtree_com[m.body_rootid][1:] - d.xipos[1:]
|
||||
),
|
||||
])
|
||||
|
||||
# cfrc_ext += contacts
|
||||
|
||||
# compute contact forces for each condim
|
||||
forces = []
|
||||
condim_idx = []
|
||||
for dim in set(d.contact.dim):
|
||||
force, idx = support.contact_force_dim(m, d, dim)
|
||||
forces.append(force)
|
||||
condim_idx.append(idx)
|
||||
|
||||
# update cfrc_ext with contact forces
|
||||
if forces:
|
||||
|
||||
@jax.vmap
|
||||
def _contact_force_to_cfrc_ext(force, pos, frame, id1, id2, com1, com2):
|
||||
# force: contact to world frame
|
||||
force = force.reshape((-1, 3)) @ frame
|
||||
force = force.reshape(-1)
|
||||
|
||||
# contact force on bodies
|
||||
cfrc_com1 = _transform_force(force, com1 - pos)
|
||||
cfrc_com2 = _transform_force(force, com2 - pos)
|
||||
|
||||
# mask
|
||||
mask1 = id1 != 0
|
||||
mask2 = id2 != 0
|
||||
|
||||
return jp.vstack([-1 * cfrc_com1 * mask1, cfrc_com2 * mask2]), jp.array(
|
||||
[id1, id2]
|
||||
)
|
||||
|
||||
condim_idx = jp.concatenate(condim_idx)
|
||||
frame = d.contact.frame[condim_idx]
|
||||
pos = d.contact.pos[condim_idx]
|
||||
id1 = jp.array(m.geom_bodyid)[d.contact.geom[condim_idx, 0]]
|
||||
id2 = jp.array(m.geom_bodyid)[d.contact.geom[condim_idx, 1]]
|
||||
com1 = d.subtree_com[jp.array(m.body_rootid)][id1]
|
||||
com2 = d.subtree_com[jp.array(m.body_rootid)][id2]
|
||||
|
||||
cfrc_contact, cfrc_idx = _contact_force_to_cfrc_ext(
|
||||
jp.concatenate(forces), pos, frame, id1, id2, com1, com2
|
||||
)
|
||||
|
||||
cfrc_ext = cfrc_ext.at[cfrc_idx.reshape(-1)].add(
|
||||
cfrc_contact.reshape((-1, 6))
|
||||
)
|
||||
|
||||
# TODO(taylorhowell): connect and weld constraints
|
||||
|
||||
# forward pass over bodies: compute cacc, cfrc_int
|
||||
def _forward(carry, cfrc_ext, cinert, cvel, body_dofadr, body_dofnum):
|
||||
if carry is None:
|
||||
if m.opt.disableflags & DisableBit.GRAVITY:
|
||||
cacc0 = jp.zeros(6)
|
||||
else:
|
||||
cacc0 = jp.concatenate((jp.zeros(3), -m.opt.gravity))
|
||||
return cacc0, jp.zeros(6)
|
||||
else:
|
||||
cacc_parent, _ = carry
|
||||
|
||||
# create dof mask
|
||||
indices = jp.arange(m.nv)
|
||||
mask = jp.logical_and(
|
||||
indices >= body_dofadr, indices < body_dofadr + body_dofnum
|
||||
)
|
||||
|
||||
# cacc = cacc_parent + cdofdot * qvel + cdof * qacc
|
||||
cacc_vel = d.cdof_dot.T @ (mask * d.qvel)
|
||||
cacc_acc = d.cdof.T @ (mask * d.qacc)
|
||||
cacc = cacc_parent + cacc_vel + cacc_acc
|
||||
|
||||
# cfrc_body = cinert * cacc + cvel x (cinert * cvel)
|
||||
cfrc_body = math.inert_mul(cinert, cacc)
|
||||
cfrc_corr = math.inert_mul(cinert, cvel)
|
||||
cfrc = math.motion_cross_force(cvel, cfrc_corr)
|
||||
cfrc_body = cfrc_body + cfrc
|
||||
cfrc_int = cfrc_body - cfrc_ext
|
||||
|
||||
return cacc, cfrc_int
|
||||
|
||||
cacc, cfrc_int = scan.body_tree(
|
||||
m,
|
||||
_forward,
|
||||
'bbbbb',
|
||||
'bb',
|
||||
cfrc_ext,
|
||||
d.cinert,
|
||||
d.cvel,
|
||||
jp.array(m.body_dofadr),
|
||||
jp.array(m.body_dofnum),
|
||||
)
|
||||
|
||||
# backward pass over bodies: accumulate cfrc_int from children
|
||||
cfrc_int = scan.body_tree(
|
||||
m,
|
||||
lambda c, p: p + c if c is not None else p, # add child to parent
|
||||
'b',
|
||||
'b',
|
||||
cfrc_int,
|
||||
reverse=True,
|
||||
)
|
||||
|
||||
# update data
|
||||
return d.replace(cacc=cacc, cfrc_int=cfrc_int, cfrc_ext=cfrc_ext)
|
||||
|
||||
|
||||
def tendon(m: Model, d: Data) -> Data:
|
||||
"""Computes tendon lengths and moments."""
|
||||
if not m.ntendon:
|
||||
|
||||
@@ -191,6 +191,44 @@ class SmoothTest(absltest.TestCase):
|
||||
_assert_attr_eq(d, dx, 'subtree_linvel')
|
||||
_assert_attr_eq(d, dx, 'subtree_angmom')
|
||||
|
||||
def test_rnepostconstraint(self):
|
||||
"""Tests MJX rne_postconstraint function to match MuJoCo mj_rnePostConstraint."""
|
||||
|
||||
m = mujoco.MjModel.from_xml_string("""
|
||||
<mujoco>
|
||||
<worldbody>
|
||||
<geom name="floor" size="0 0 .05" type="plane"/>
|
||||
<body pos="0 0 1">
|
||||
<joint type="ball" damping="1"/>
|
||||
<geom type="capsule" size="0.1 0.5" fromto="0 0 0 0.5 0 0"/>
|
||||
<body pos="0.5 0 0">
|
||||
<joint type="ball" damping="1"/>
|
||||
<geom type="capsule" size="0.1 0.5" fromto="0 0 0 0.5 0 0"/>
|
||||
</body>
|
||||
</body>
|
||||
</worldbody>
|
||||
<keyframe>
|
||||
<key qpos='0.424577 0.450592 0.451703 -0.642391 0.729379 0.545151 0.407756 0.0674697'/>
|
||||
</keyframe>
|
||||
</mujoco>
|
||||
""")
|
||||
d = mujoco.MjData(m)
|
||||
mujoco.mj_resetDataKeyframe(m, d, 0)
|
||||
# apply external forces
|
||||
d.xfrc_applied = 0.001 * np.ones(d.xfrc_applied.shape)
|
||||
mujoco.mj_step(m, d, 2)
|
||||
mujoco.mj_forward(m, d)
|
||||
mx = mjx.put_model(m)
|
||||
dx = mjx.put_data(m, d)
|
||||
|
||||
# rne postconstraint
|
||||
mujoco.mj_rnePostConstraint(m, d)
|
||||
dx = jax.jit(mjx.rne_postconstraint)(mx, dx)
|
||||
|
||||
_assert_eq(d.cacc, dx.cacc, 'cacc')
|
||||
_assert_eq(d.cfrc_ext, dx.cfrc_ext, 'cfrc_ext')
|
||||
_assert_eq(d.cfrc_int, dx.cfrc_int, 'cfrc_int')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
absltest.main()
|
||||
|
||||
@@ -319,3 +319,28 @@ def contact_force(
|
||||
force = force.reshape(-1)
|
||||
|
||||
return force * (efc_address >= 0)
|
||||
|
||||
|
||||
def contact_force_dim(
|
||||
m: Model, d: Data, dim: int
|
||||
) -> Tuple[jax.Array, np.ndarray]:
|
||||
"""Extract 6D force:torque for contacts with dimension dim."""
|
||||
# valid contact and condim indices
|
||||
idx_dim = (d.contact.efc_address >= 0) & (d.contact.dim == dim)
|
||||
|
||||
# contact force from efc
|
||||
if m.opt.cone == mujoco.mjtCone.mjCONE_PYRAMIDAL:
|
||||
efc_address = (
|
||||
d.contact.efc_address[idx_dim, None]
|
||||
+ np.arange(np.where(dim == 1, 1, 2 * (dim - 1)))[None]
|
||||
)
|
||||
efc_force = d.efc_force[efc_address]
|
||||
force = jax.vmap(_decode_pyramid, in_axes=(0, 0, None))(
|
||||
efc_force, d.contact.friction[idx_dim], dim
|
||||
)
|
||||
return force, np.where(idx_dim)[0]
|
||||
elif m.opt.cone == mujoco.mjtCone.mjCONE_ELLIPTIC:
|
||||
# TODO(taylorhowell): add support for elliptic cone
|
||||
raise NotImplementedError('Elliptic cone force is not implemented yet.')
|
||||
else:
|
||||
raise ValueError(f'Unknown cone type: {m.opt.cone}.')
|
||||
|
||||
@@ -1,9 +0,0 @@
|
||||
<!-- For validating model with no sensor -->
|
||||
<mujoco model="no_sensor">
|
||||
<worldbody>
|
||||
<body>
|
||||
<joint type="hinge"/>
|
||||
<geom size="1"/>
|
||||
</body>
|
||||
</worldbody>
|
||||
</mujoco>
|
||||
@@ -0,0 +1,58 @@
|
||||
<mujoco model="model">
|
||||
<asset>
|
||||
<material name="material"/>
|
||||
</asset>
|
||||
<worldbody>
|
||||
<!-- tree 0 -->
|
||||
<body name="body0" pos="1 2 3">
|
||||
<joint name="hinge0" type="hinge" axis="1 0 0"/>
|
||||
<geom size="0.1" material="material"/>
|
||||
<site name="site_rangefinder0" pos="-1e-3 0 0.2"/>
|
||||
<site name="site_rangefinder1" pos="-1e-3 0 0.175"/>
|
||||
<site name="site0" pos=".1 .2 .3"/>
|
||||
<body name="body1" pos="0.1 0.2 0.3">
|
||||
<joint name="hinge1" type="hinge" axis="0 1 0"/>
|
||||
<geom size="0.25"/>
|
||||
<site name="site1" pos=".2 .4 .6"/>
|
||||
</body>
|
||||
</body>
|
||||
|
||||
<!-- body 2 -->
|
||||
<body name="body2" pos=".1 .1 .1">
|
||||
<joint name="ballquat2" type="ball" pos="0.1 0.1 0.1"/>
|
||||
<geom name="geom2" size="1"/>
|
||||
</body>
|
||||
|
||||
<!-- body 3 -->
|
||||
<body name="body3" pos="-.1 -.1 -.1">
|
||||
<joint name="ballquat3" type="ball" pos="0.1 0.2 0.3"/>
|
||||
<geom size="1"/>
|
||||
<site name="site3"/>
|
||||
</body>
|
||||
|
||||
<!-- bodies for camera projection -->
|
||||
<body pos="11.1 0 1">
|
||||
<geom type="box" size=".1 .6 .375"/>
|
||||
<site name="frontorigin" pos="-.1 .6 .375"/>
|
||||
<site name="frontcenter" pos="-.1 0 0"/>
|
||||
</body>
|
||||
<body pos="10 0 0">
|
||||
<joint axis="0 0 1" range="-180 180" limited="false"/>
|
||||
<geom type="sphere" size=".2" pos="0 0 0.9"/>
|
||||
<camera pos="0 0 1" xyaxes="0 -1 0 0 0 1" fovy="41.11209"
|
||||
resolution="1920 1200" name="fixedcamera"/>
|
||||
</body>
|
||||
|
||||
<!-- body for rangefinder -->
|
||||
<body name="body_rangefinder" pos="1 2 4">
|
||||
<geom size="0.01" material="material"/>
|
||||
</body>
|
||||
</worldbody>
|
||||
|
||||
<actuator>
|
||||
<motor name="motor0" joint="hinge0" ctrlrange="-1 1" gear="10"
|
||||
ctrllimited="true"/>
|
||||
<motor name="motor1" joint="hinge1" ctrlrange="-1 1" gear="10"
|
||||
ctrllimited="true"/>
|
||||
</actuator>
|
||||
</mujoco>
|
||||
@@ -0,0 +1,65 @@
|
||||
<!-- For validating model with unsupported sensors -->
|
||||
<mujoco model="unsupported_sensor">
|
||||
<asset>
|
||||
<material name="material"/>
|
||||
</asset>
|
||||
<worldbody>
|
||||
<!-- tree 0 -->
|
||||
<body name="body0" pos="1 2 3">
|
||||
<joint name="hinge0" type="hinge" axis="1 0 0"/>
|
||||
<geom name="geom0" size="0.1" material="material"/>
|
||||
<site name="site_rangefinder0" pos="-1e-3 0 0.2"/>
|
||||
<site name="site_rangefinder1" pos="-1e-3 0 0.175"/>
|
||||
<site name="site0" pos=".1 .2 .3"/>
|
||||
<body name="body1" pos="0.1 0.2 0.3">
|
||||
<joint name="hinge1" type="hinge" axis="0 1 0"/>
|
||||
<geom name="geom1" size="0.25"/>
|
||||
<site name="site1" pos=".2 .4 .6"/>
|
||||
</body>
|
||||
</body>
|
||||
|
||||
<!-- body 2 -->
|
||||
<body name="body2" pos=".1 .1 .1">
|
||||
<joint name="ballquat2" type="ball" pos="0.1 0.1 0.1"/>
|
||||
<geom name="geom2" size="1"/>
|
||||
</body>
|
||||
|
||||
<!-- body 3 -->
|
||||
<body name="body3" pos="-.1 -.1 -.1">
|
||||
<joint name="ballquat3" type="ball" pos="0.1 0.2 0.3"/>
|
||||
<geom size="1"/>
|
||||
<site name="site3"/>
|
||||
</body>
|
||||
|
||||
<!-- bodies for camera projection -->
|
||||
<body pos="11.1 0 1">
|
||||
<geom type="box" size=".1 .6 .375"/>
|
||||
<site name="frontorigin" pos="-.1 .6 .375"/>
|
||||
<site name="frontcenter" pos="-.1 0 0"/>
|
||||
</body>
|
||||
<body pos="10 0 0">
|
||||
<joint axis="0 0 1" range="-180 180" limited="false"/>
|
||||
<geom type="sphere" size=".2" pos="0 0 0.9"/>
|
||||
<camera pos="0 0 1" xyaxes="0 -1 0 0 0 1" fovy="41.11209"
|
||||
resolution="1920 1200" name="fixedcamera"/>
|
||||
</body>
|
||||
|
||||
<!-- body for rangefinder -->
|
||||
<body name="body_rangefinder" pos="1 2 4">
|
||||
<geom size="0.01" material="material"/>
|
||||
</body>
|
||||
</worldbody>
|
||||
|
||||
<actuator>
|
||||
<motor name="motor0" joint="hinge0" ctrlrange="-1 1" gear="10"
|
||||
ctrllimited="true"/>
|
||||
<motor name="motor1" joint="hinge1" ctrlrange="-1 1" gear="10"
|
||||
ctrllimited="true"/>
|
||||
</actuator>
|
||||
|
||||
<sensor>
|
||||
<distance name="distance" geom1="geom0" geom2="geom1"/>
|
||||
<framelinvel name="framelinvel" objtype="site" objname="site0"/>
|
||||
<touch name="touch" site="site0"/>
|
||||
</sensor>
|
||||
</mujoco>
|
||||
@@ -1,19 +0,0 @@
|
||||
<!-- For validating model with unsupported sensor -->
|
||||
<mujoco model="unsupported_sensor">
|
||||
<worldbody>
|
||||
<body>
|
||||
<joint type="hinge"/>
|
||||
<geom name="geom0" size="1"/>
|
||||
<site name="site"/>
|
||||
</body>
|
||||
<body>
|
||||
<joint type="hinge"/>
|
||||
<geom name="geom1" size="1"/>
|
||||
</body>
|
||||
</worldbody>
|
||||
<sensor>
|
||||
<distance name="distance" geom1="geom0" geom2="geom1"/>
|
||||
<framelinvel name="framelinvel" objtype="site" objname="site"/>
|
||||
<touch name="touch" site="site"/>
|
||||
</sensor>
|
||||
</mujoco>
|
||||
Reference in New Issue
Block a user