Add com_pos to Mujoco Warp shims.
PiperOrigin-RevId: 896969879 Change-Id: Idc3eb067a8473251a98404525df0c8c134ac2032
This commit is contained in:
committed by
Copybara-Service
parent
05e26e961c
commit
cc0933af8b
@@ -483,3 +483,125 @@ def tendon(m: types.Model, d: types.Data):
|
||||
def tendon_vmap(unused_axis_size, is_batched, m: types.Model, d: types.Data):
|
||||
d = tendon(m, d)
|
||||
return d, is_batched[1]
|
||||
|
||||
|
||||
@ffi.format_args_for_warp
|
||||
def _com_pos_shim(
|
||||
# Model
|
||||
nworld: int,
|
||||
body_inertia: wp.array2d(dtype=wp.vec3),
|
||||
body_mass: wp.array2d(dtype=float),
|
||||
body_parentid: wp.array(dtype=int),
|
||||
body_rootid: wp.array(dtype=int),
|
||||
body_subtreemass: wp.array2d(dtype=float),
|
||||
body_tree: tuple[wp.array(dtype=int), ...],
|
||||
jnt_bodyid: wp.array(dtype=int),
|
||||
jnt_dofadr: wp.array(dtype=int),
|
||||
jnt_type: wp.array(dtype=int),
|
||||
nbody: int,
|
||||
njnt: int,
|
||||
# Data
|
||||
cdof: wp.array2d(dtype=wp.spatial_vector),
|
||||
cinert: wp.array2d(dtype=mjwp_types.vec10),
|
||||
subtree_com: wp.array2d(dtype=wp.vec3),
|
||||
xanchor: wp.array2d(dtype=wp.vec3),
|
||||
xaxis: wp.array2d(dtype=wp.vec3),
|
||||
ximat: wp.array2d(dtype=wp.mat33),
|
||||
xipos: wp.array2d(dtype=wp.vec3),
|
||||
xmat: wp.array2d(dtype=wp.mat33),
|
||||
):
|
||||
_m.stat = _s
|
||||
_m.opt = _o
|
||||
_m.callback = _cb
|
||||
_d.efc = _e
|
||||
_d.contact = _c
|
||||
_m.body_inertia = body_inertia
|
||||
_m.body_mass = body_mass
|
||||
_m.body_parentid = body_parentid
|
||||
_m.body_rootid = body_rootid
|
||||
_m.body_subtreemass = body_subtreemass
|
||||
_m.body_tree = body_tree
|
||||
_m.jnt_bodyid = jnt_bodyid
|
||||
_m.jnt_dofadr = jnt_dofadr
|
||||
_m.jnt_type = jnt_type
|
||||
_m.nbody = nbody
|
||||
_m.njnt = njnt
|
||||
_d.cdof = cdof
|
||||
_d.cinert = cinert
|
||||
_d.subtree_com = subtree_com
|
||||
_d.xanchor = xanchor
|
||||
_d.xaxis = xaxis
|
||||
_d.ximat = ximat
|
||||
_d.xipos = xipos
|
||||
_d.xmat = xmat
|
||||
_d.nworld = nworld
|
||||
mjwarp.com_pos(_m, _d)
|
||||
|
||||
|
||||
def _com_pos_jax_impl(m: types.Model, d: types.Data):
|
||||
output_dims = {
|
||||
'cdof': d.cdof.shape,
|
||||
'cinert': d._impl.cinert.shape,
|
||||
'subtree_com': d.subtree_com.shape,
|
||||
}
|
||||
jf = ffi.jax_callable_variadic_tuple(
|
||||
_com_pos_shim,
|
||||
num_outputs=3,
|
||||
output_dims=output_dims,
|
||||
vmap_method=None,
|
||||
in_out_argnames=set(['cdof', 'cinert', 'subtree_com']),
|
||||
stage_in_argnames=set([
|
||||
'body_inertia',
|
||||
'body_mass',
|
||||
'body_subtreemass',
|
||||
'cdof',
|
||||
'subtree_com',
|
||||
'xanchor',
|
||||
'xaxis',
|
||||
'ximat',
|
||||
'xipos',
|
||||
'xmat',
|
||||
]),
|
||||
stage_out_argnames=set(['cdof', 'subtree_com']),
|
||||
graph_mode=m.opt._impl.graph_mode,
|
||||
has_side_effect=False,
|
||||
)
|
||||
out = jf(
|
||||
d.qpos.shape[0],
|
||||
m.body_inertia,
|
||||
m.body_mass,
|
||||
m.body_parentid,
|
||||
m.body_rootid,
|
||||
m.body_subtreemass,
|
||||
m._impl.body_tree,
|
||||
m.jnt_bodyid,
|
||||
m.jnt_dofadr,
|
||||
m.jnt_type,
|
||||
m.nbody,
|
||||
m.njnt,
|
||||
d.cdof,
|
||||
d._impl.cinert,
|
||||
d.subtree_com,
|
||||
d.xanchor,
|
||||
d.xaxis,
|
||||
d.ximat,
|
||||
d.xipos,
|
||||
d.xmat,
|
||||
)
|
||||
d = d.tree_replace(
|
||||
{'cdof': out[0], '_impl.cinert': out[1], 'subtree_com': out[2]}
|
||||
)
|
||||
return d
|
||||
|
||||
|
||||
@jax.custom_batching.custom_vmap
|
||||
@ffi.marshal_jax_warp_callable
|
||||
def com_pos(m: types.Model, d: types.Data):
|
||||
return _com_pos_jax_impl(m, d)
|
||||
|
||||
|
||||
@com_pos.def_vmap
|
||||
@ffi.marshal_custom_vmap
|
||||
def com_pos_vmap(unused_axis_size, is_batched, m: types.Model, d: types.Data):
|
||||
d = com_pos(m, d)
|
||||
return d, is_batched[1]
|
||||
|
||||
@@ -100,6 +100,37 @@ class SmoothTest(parameterized.TestCase):
|
||||
tu.assert_attr_eq(d, dx, 'site_xpos')
|
||||
tu.assert_eq(d.site_xmat.reshape((-1, 3, 3)), dx.site_xmat, 'site_xmat')
|
||||
|
||||
def test_com_pos(self):
|
||||
"""Tests com_pos with unbatched data."""
|
||||
if not _FORCE_TEST:
|
||||
if not mjxw.WARP_INSTALLED:
|
||||
self.skipTest('Warp not installed.')
|
||||
if not io.has_cuda_gpu_device():
|
||||
self.skipTest('No CUDA GPU device available.')
|
||||
|
||||
m = tu.load_test_file('pendula.xml')
|
||||
|
||||
d = mujoco.MjData(m)
|
||||
mx = mjx.put_model(m, impl='warp')
|
||||
mx = mx.replace(_impl=mx._impl.replace(qM_tiles=()))
|
||||
|
||||
rng = jax.random.PRNGKey(0)
|
||||
dx = mjx.make_data(m, impl='warp')
|
||||
_, key = jax.random.split(rng)
|
||||
qpos = jax.random.uniform(key, (m.nq,))
|
||||
dx = dx.replace(qpos=qpos)
|
||||
|
||||
dx = jax.jit(smooth.kinematics)(mx, dx)
|
||||
dx = jax.jit(smooth.com_pos)(mx, dx)
|
||||
|
||||
d.qpos[:] = qpos
|
||||
mujoco.mj_kinematics(m, d)
|
||||
mujoco.mj_comPos(m, d)
|
||||
|
||||
tu.assert_attr_eq(d, dx, 'subtree_com')
|
||||
tu.assert_attr_eq(d, dx, 'cdof')
|
||||
tu.assert_eq(d.cinert, dx._impl.cinert, 'cinert')
|
||||
|
||||
def test_kinematics_vmap(self):
|
||||
"""Tests kinematics with batched data."""
|
||||
if not mjxw.WARP_INSTALLED:
|
||||
@@ -144,6 +175,57 @@ class SmoothTest(parameterized.TestCase):
|
||||
tu.assert_attr_eq(d, dx, 'site_xpos')
|
||||
tu.assert_eq(d.site_xmat.reshape((-1, 3, 3)), dx.site_xmat, 'site_xmat')
|
||||
|
||||
def test_com_pos_vmap(self):
|
||||
"""Tests com_pos with batched data."""
|
||||
if not mjxw.WARP_INSTALLED:
|
||||
self.skipTest('Warp not installed.')
|
||||
if not io.has_cuda_gpu_device():
|
||||
self.skipTest('No CUDA GPU device available.')
|
||||
|
||||
m = tu.load_test_file('pendula.xml')
|
||||
|
||||
batch_size = 7
|
||||
d = mujoco.MjData(m)
|
||||
mx = mjx.put_model(m, impl='warp')
|
||||
mx = mx.replace(_impl=mx._impl.replace(qM_tiles=()))
|
||||
|
||||
worldids = jp.arange(batch_size)
|
||||
dx_batch = jax.vmap(functools.partial(tu.make_data, m))(worldids)
|
||||
for f in (
|
||||
'xanchor',
|
||||
'xaxis',
|
||||
'xpos',
|
||||
'xipos',
|
||||
'site_xpos',
|
||||
'site_xmat',
|
||||
'subtree_com',
|
||||
'cdof',
|
||||
):
|
||||
dx_batch = dx_batch.replace(**{f: jp.zeros_like(getattr(dx_batch, f))})
|
||||
dx_batch = dx_batch.tree_replace(
|
||||
{'_impl.cinert': jp.zeros_like(dx_batch._impl.cinert)}
|
||||
)
|
||||
|
||||
dx_batch = jax.jit(jax.vmap(smooth.kinematics, in_axes=(None, 0)))(
|
||||
mx, dx_batch
|
||||
)
|
||||
dx_batch = jax.jit(jax.vmap(smooth.com_pos, in_axes=(None, 0)))(
|
||||
mx, dx_batch
|
||||
)
|
||||
|
||||
for i in range(batch_size):
|
||||
dx = dx_batch[i]
|
||||
|
||||
d.qpos[:] = dx.qpos
|
||||
d.mocap_pos[:] = dx.mocap_pos
|
||||
d.mocap_quat[:] = dx.mocap_quat
|
||||
mujoco.mj_kinematics(m, d)
|
||||
mujoco.mj_comPos(m, d)
|
||||
|
||||
tu.assert_attr_eq(d, dx, 'subtree_com')
|
||||
tu.assert_attr_eq(d, dx, 'cdof')
|
||||
tu.assert_eq(d.cinert, dx._impl.cinert, 'cinert')
|
||||
|
||||
def test_kinematics_nested_vmap(self):
|
||||
"""Tests kinematics with nested batch data."""
|
||||
if not _FORCE_TEST:
|
||||
|
||||
Reference in New Issue
Block a user