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:
Taylor Howell
2024-09-05 05:34:27 -07:00
committed by Copybara-Service
parent cd8ff4401d
commit 7b073b60c4
10 changed files with 313 additions and 31 deletions
+1
View File
@@ -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
+3 -3
View File
@@ -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')
+123
View File
@@ -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:
+38
View File
@@ -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()
+25
View File
@@ -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}.')
-9
View File
@@ -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>
+58
View File
@@ -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>