Import google-deepmind/mujoco_warp from GitHub.

PiperOrigin-RevId: 831416862
Change-Id: I9ffe9698a3cd4a173189a804998ce8afc16b91f2
This commit is contained in:
Silvia Cruciani
2025-11-12 09:17:25 -08:00
committed by Copybara-Service
parent 4722925ff4
commit 727347813f
+118 -111
View File
@@ -1304,7 +1304,7 @@ def put_data(
else:
nconvar = condim if mjm.opt.cone == mujoco.mjtCone.mjCONE_ELLIPTIC else 2 * (condim - 1)
for k in range(nconvar):
contact_efc_address[i * mjd.ncon + j, k] = mjd.nefc * i + efc_address + k
contact_efc_address[i * mjd.ncon + j, k] = efc_address + k
contact_worldid = np.pad(np.repeat(np.arange(nworld), mjd.ncon), (0, naconmax - nworld * mjd.ncon))
@@ -1517,6 +1517,7 @@ def get_data_into(
result: mujoco.MjData,
mjm: mujoco.MjModel,
d: types.Data,
world_id: int = 0,
):
"""Gets data from a device into an existing mujoco.MjData.
@@ -1524,41 +1525,43 @@ def get_data_into(
result: The data object containing the current state and output arrays (host).
mjm: The model containing kinematic and dynamic information (host).
d: The data object containing the current state and output arrays (device).
world_id: The id of the world to get the data from.
"""
if d.nworld > 1:
raise NotImplementedError("only nworld == 1 supported for now")
# nacon and nefc can overflow. in that case, only pull up to the max contacts and constraints
nacon = min(d.nacon.numpy()[0], d.naconmax)
nefc = min(d.nefc.numpy()[0], d.njmax)
nefc = min(d.nefc.numpy()[world_id], d.njmax)
if nacon != result.ncon or nefc != result.nefc:
ncon_filter = np.zeros_like(d.contact.worldid.numpy(), dtype=bool)
ncon_filter[:nacon] = d.contact.worldid.numpy()[:nacon] == world_id
ncon = ncon_filter.sum()
if ncon != result.ncon or nefc != result.nefc:
# TODO(team): if sparse, set nJ based on sparse efc_J
mujoco._functions._realloc_con_efc(result, ncon=nacon, nefc=nefc, nJ=nefc * mjm.nv)
mujoco._functions._realloc_con_efc(result, ncon=ncon, nefc=nefc, nJ=nefc * mjm.nv)
ne = d.ne.numpy()[0]
nf = d.nf.numpy()[0]
nl = d.nl.numpy()[0]
ne = d.ne.numpy()[world_id]
nf = d.nf.numpy()[world_id]
nl = d.nl.numpy()[world_id]
# efc indexing
# mujoco expects contiguous efc ordering for contacts
# this ordering is not guaranteed with mujoco warp, we enforce order here
if nacon > 0:
if ncon > 0:
efc_idx_efl = np.arange(ne + nf + nl)
contact_dim = d.contact.dim.numpy()
contact_efc_address = d.contact.efc_address.numpy()
contact_dim = d.contact.dim.numpy()[ncon_filter]
contact_efc_address = d.contact.efc_address.numpy()[ncon_filter]
efc_idx_c = []
contact_efc_address_ordered = [ne + nf + nl]
for i in range(nacon):
for i in range(ncon):
dim = contact_dim[i]
if mjm.opt.cone == mujoco.mjtCone.mjCONE_PYRAMIDAL:
ndim = np.maximum(1, 2 * (dim - 1))
else:
ndim = dim
efc_idx_c.append(contact_efc_address[i, :ndim])
if i < nacon - 1:
if i < ncon - 1:
contact_efc_address_ordered.append(contact_efc_address_ordered[-1] + ndim)
efc_idx = np.concatenate((efc_idx_efl, *efc_idx_c))
contact_efc_address_ordered = np.array(contact_efc_address_ordered)
@@ -1568,130 +1571,134 @@ def get_data_into(
efc_idx = efc_idx[:nefc] # dont emit indices for overflow constraints
result.solver_niter[0] = d.solver_niter.numpy()[0]
result.ncon = nacon
result.solver_niter[0] = d.solver_niter.numpy()[world_id]
result.ncon = ncon
result.ne = ne
result.nf = nf
result.nl = nl
result.time = d.time.numpy()[0]
result.energy[:] = d.energy.numpy()[0]
result.qpos[:] = d.qpos.numpy()[0]
result.qvel[:] = d.qvel.numpy()[0]
result.act[:] = d.act.numpy()[0]
result.qacc_warmstart[:] = d.qacc_warmstart.numpy()[0]
result.ctrl[:] = d.ctrl.numpy()[0]
result.qfrc_applied[:] = d.qfrc_applied.numpy()[0]
result.xfrc_applied[:] = d.xfrc_applied.numpy()[0]
result.eq_active[:] = d.eq_active.numpy()[0]
result.mocap_pos[:] = d.mocap_pos.numpy()[0]
result.mocap_quat[:] = d.mocap_quat.numpy()[0]
result.qacc[:] = d.qacc.numpy()[0]
result.act_dot[:] = d.act_dot.numpy()[0]
result.xpos[:] = d.xpos.numpy()[0]
result.xquat[:] = d.xquat.numpy()[0]
result.xmat[:] = d.xmat.numpy().reshape((-1, 9))
result.xipos[:] = d.xipos.numpy()[0]
result.ximat[:] = d.ximat.numpy().reshape((-1, 9))
result.xanchor[:] = d.xanchor.numpy()[0]
result.xaxis[:] = d.xaxis.numpy()[0]
result.geom_xpos[:] = d.geom_xpos.numpy()[0]
result.geom_xmat[:] = d.geom_xmat.numpy().reshape((-1, 9))
result.site_xpos[:] = d.site_xpos.numpy()[0]
result.site_xmat[:] = d.site_xmat.numpy().reshape((-1, 9))
result.cam_xpos[:] = d.cam_xpos.numpy()[0]
result.cam_xmat[:] = d.cam_xmat.numpy().reshape((-1, 9))
result.light_xpos[:] = d.light_xpos.numpy()[0]
result.light_xdir[:] = d.light_xdir.numpy()[0]
result.subtree_com[:] = d.subtree_com.numpy()[0]
result.cdof[:] = d.cdof.numpy()[0]
result.cinert[:] = d.cinert.numpy()[0]
result.flexvert_xpos[:] = d.flexvert_xpos.numpy()[0]
result.flexedge_length[:] = d.flexedge_length.numpy()[0]
result.flexedge_velocity[:] = d.flexedge_velocity.numpy()[0]
result.actuator_length[:] = d.actuator_length.numpy()[0]
result.time = d.time.numpy()[world_id]
result.energy[:] = d.energy.numpy()[world_id]
result.qpos[:] = d.qpos.numpy()[world_id]
result.qvel[:] = d.qvel.numpy()[world_id]
result.act[:] = d.act.numpy()[world_id]
result.qacc_warmstart[:] = d.qacc_warmstart.numpy()[world_id]
result.ctrl[:] = d.ctrl.numpy()[world_id]
result.qfrc_applied[:] = d.qfrc_applied.numpy()[world_id]
result.xfrc_applied[:] = d.xfrc_applied.numpy()[world_id]
result.eq_active[:] = d.eq_active.numpy()[world_id]
result.mocap_pos[:] = d.mocap_pos.numpy()[world_id]
result.mocap_quat[:] = d.mocap_quat.numpy()[world_id]
result.qacc[:] = d.qacc.numpy()[world_id]
result.act_dot[:] = d.act_dot.numpy()[world_id]
result.xpos[:] = d.xpos.numpy()[world_id]
result.xquat[:] = d.xquat.numpy()[world_id]
result.xmat[:] = d.xmat.numpy()[world_id].reshape((-1, 9))
result.xipos[:] = d.xipos.numpy()[world_id]
result.ximat[:] = d.ximat.numpy()[world_id].reshape((-1, 9))
result.xanchor[:] = d.xanchor.numpy()[world_id]
result.xaxis[:] = d.xaxis.numpy()[world_id]
result.geom_xpos[:] = d.geom_xpos.numpy()[world_id]
result.geom_xmat[:] = d.geom_xmat.numpy()[world_id].reshape((-1, 9))
result.site_xpos[:] = d.site_xpos.numpy()[world_id]
result.site_xmat[:] = d.site_xmat.numpy()[world_id].reshape((-1, 9))
result.cam_xpos[:] = d.cam_xpos.numpy()[world_id]
result.cam_xmat[:] = d.cam_xmat.numpy()[world_id].reshape((-1, 9))
result.light_xpos[:] = d.light_xpos.numpy()[world_id]
result.light_xdir[:] = d.light_xdir.numpy()[world_id]
result.subtree_com[:] = d.subtree_com.numpy()[world_id]
result.cdof[:] = d.cdof.numpy()[world_id]
result.cinert[:] = d.cinert.numpy()[world_id]
result.flexvert_xpos[:] = d.flexvert_xpos.numpy()[world_id]
result.flexedge_length[:] = d.flexedge_length.numpy()[world_id]
result.flexedge_velocity[:] = d.flexedge_velocity.numpy()[world_id]
result.actuator_length[:] = d.actuator_length.numpy()[world_id]
mujoco.mju_dense2sparse(
result.actuator_moment, d.actuator_moment.numpy()[0], result.moment_rownnz, result.moment_rowadr, result.moment_colind
result.actuator_moment,
d.actuator_moment.numpy()[world_id],
result.moment_rownnz,
result.moment_rowadr,
result.moment_colind,
)
result.crb[:] = d.crb.numpy()[0]
result.qLDiagInv[:] = d.qLDiagInv.numpy()[0]
result.ten_velocity[:] = d.ten_velocity.numpy()[0]
result.actuator_velocity[:] = d.actuator_velocity.numpy()[0]
result.cvel[:] = d.cvel.numpy()[0]
result.cdof_dot[:] = d.cdof_dot.numpy()[0]
result.qfrc_bias[:] = d.qfrc_bias.numpy()[0]
result.qfrc_spring[:] = d.qfrc_spring.numpy()[0]
result.qfrc_damper[:] = d.qfrc_damper.numpy()[0]
result.qfrc_gravcomp[:] = d.qfrc_gravcomp.numpy()[0]
result.qfrc_fluid[:] = d.qfrc_fluid.numpy()[0]
result.qfrc_passive[:] = d.qfrc_passive.numpy()[0]
result.subtree_linvel[:] = d.subtree_linvel.numpy()[0]
result.subtree_angmom[:] = d.subtree_angmom.numpy()[0]
result.actuator_force[:] = d.actuator_force.numpy()[0]
result.qfrc_actuator[:] = d.qfrc_actuator.numpy()[0]
result.qfrc_smooth[:] = d.qfrc_smooth.numpy()[0]
result.qacc_smooth[:] = d.qacc_smooth.numpy()[0]
result.qfrc_constraint[:] = d.qfrc_constraint.numpy()[0]
result.qfrc_inverse[:] = d.qfrc_inverse.numpy()[0]
result.crb[:] = d.crb.numpy()[world_id]
result.qLDiagInv[:] = d.qLDiagInv.numpy()[world_id]
result.ten_velocity[:] = d.ten_velocity.numpy()[world_id]
result.actuator_velocity[:] = d.actuator_velocity.numpy()[world_id]
result.cvel[:] = d.cvel.numpy()[world_id]
result.cdof_dot[:] = d.cdof_dot.numpy()[world_id]
result.qfrc_bias[:] = d.qfrc_bias.numpy()[world_id]
result.qfrc_spring[:] = d.qfrc_spring.numpy()[world_id]
result.qfrc_damper[:] = d.qfrc_damper.numpy()[world_id]
result.qfrc_gravcomp[:] = d.qfrc_gravcomp.numpy()[world_id]
result.qfrc_fluid[:] = d.qfrc_fluid.numpy()[world_id]
result.qfrc_passive[:] = d.qfrc_passive.numpy()[world_id]
result.subtree_linvel[:] = d.subtree_linvel.numpy()[world_id]
result.subtree_angmom[:] = d.subtree_angmom.numpy()[world_id]
result.actuator_force[:] = d.actuator_force.numpy()[world_id]
result.qfrc_actuator[:] = d.qfrc_actuator.numpy()[world_id]
result.qfrc_smooth[:] = d.qfrc_smooth.numpy()[world_id]
result.qacc_smooth[:] = d.qacc_smooth.numpy()[world_id]
result.qfrc_constraint[:] = d.qfrc_constraint.numpy()[world_id]
result.qfrc_inverse[:] = d.qfrc_inverse.numpy()[world_id]
# contact
result.contact.dist[:] = d.contact.dist.numpy()[:nacon]
result.contact.pos[:] = d.contact.pos.numpy()[:nacon]
result.contact.frame[:] = d.contact.frame.numpy()[:nacon].reshape((-1, 9))
result.contact.includemargin[:] = d.contact.includemargin.numpy()[:nacon]
result.contact.friction[:] = d.contact.friction.numpy()[:nacon]
result.contact.solref[:] = d.contact.solref.numpy()[:nacon]
result.contact.solreffriction[:] = d.contact.solreffriction.numpy()[:nacon]
result.contact.solimp[:] = d.contact.solimp.numpy()[:nacon]
result.contact.dim[:] = d.contact.dim.numpy()[:nacon]
result.contact.geom[:] = d.contact.geom.numpy()[:nacon]
result.contact.efc_address[:] = contact_efc_address_ordered[:nacon]
result.contact.dist[:ncon] = d.contact.dist.numpy()[ncon_filter]
result.contact.pos[:ncon] = d.contact.pos.numpy()[ncon_filter]
result.contact.frame[:ncon] = d.contact.frame.numpy()[ncon_filter].reshape((-1, 9))
result.contact.includemargin[:ncon] = d.contact.includemargin.numpy()[ncon_filter]
result.contact.friction[:ncon] = d.contact.friction.numpy()[ncon_filter]
result.contact.solref[:ncon] = d.contact.solref.numpy()[ncon_filter]
result.contact.solreffriction[:ncon] = d.contact.solreffriction.numpy()[ncon_filter]
result.contact.solimp[:ncon] = d.contact.solimp.numpy()[ncon_filter]
result.contact.dim[:ncon] = d.contact.dim.numpy()[ncon_filter]
result.contact.geom[:ncon] = d.contact.geom.numpy()[ncon_filter]
result.contact.efc_address[:ncon] = contact_efc_address_ordered[:ncon]
if mujoco.mj_isSparse(mjm):
result.qM[:] = d.qM.numpy()[0, 0]
result.qLD[:] = d.qLD.numpy()[0, 0]
result.qM[:] = d.qM.numpy()[world_id, 0]
result.qLD[:] = d.qLD.numpy()[world_id, 0]
if nefc > 0:
efc_J = d.efc.J.numpy()[0, efc_idx, : mjm.nv]
efc_J = d.efc.J.numpy()[world_id, efc_idx, : mjm.nv]
mujoco.mju_dense2sparse(result.efc_J, efc_J, result.efc_J_rownnz, result.efc_J_rowadr, result.efc_J_colind)
else:
qM = d.qM.numpy()
qM = d.qM.numpy()[world_id]
adr = 0
for i in range(mjm.nv):
j = i
while j >= 0:
result.qM[adr] = qM[0, i, j]
result.qM[adr] = qM[i, j]
j = mjm.dof_parentid[j]
adr += 1
mujoco.mj_factorM(mjm, result)
if nefc > 0:
result.efc_J[: nefc * mjm.nv] = d.efc.J.numpy()[0, :nefc, : mjm.nv].flatten()
result.efc_J[: nefc * mjm.nv] = d.efc.J.numpy()[world_id, :nefc, : mjm.nv].flatten()
# efc
result.efc_type[:] = d.efc.type.numpy()[0, efc_idx]
result.efc_id[:] = d.efc.id.numpy()[0, efc_idx]
result.efc_pos[:] = d.efc.pos.numpy()[0, efc_idx]
result.efc_margin[:] = d.efc.margin.numpy()[0, efc_idx]
result.efc_D[:] = d.efc.D.numpy()[0, efc_idx]
result.efc_vel[:] = d.efc.vel.numpy()[0, efc_idx]
result.efc_aref[:] = d.efc.aref.numpy()[0, efc_idx]
result.efc_frictionloss[:] = d.efc.frictionloss.numpy()[0, efc_idx]
result.efc_state[:] = d.efc.state.numpy()[0, efc_idx]
result.efc_force[:] = d.efc.force.numpy()[0, efc_idx]
result.efc_type[:] = d.efc.type.numpy()[world_id, efc_idx]
result.efc_id[:] = d.efc.id.numpy()[world_id, efc_idx]
result.efc_pos[:] = d.efc.pos.numpy()[world_id, efc_idx]
result.efc_margin[:] = d.efc.margin.numpy()[world_id, efc_idx]
result.efc_D[:] = d.efc.D.numpy()[world_id, efc_idx]
result.efc_vel[:] = d.efc.vel.numpy()[world_id, efc_idx]
result.efc_aref[:] = d.efc.aref.numpy()[world_id, efc_idx]
result.efc_frictionloss[:] = d.efc.frictionloss.numpy()[world_id, efc_idx]
result.efc_state[:] = d.efc.state.numpy()[world_id, efc_idx]
result.efc_force[:] = d.efc.force.numpy()[world_id, efc_idx]
# rne_postconstraint
result.cacc[:] = d.cacc.numpy()[0]
result.cfrc_int[:] = d.cfrc_int.numpy()[0]
result.cfrc_ext[:] = d.cfrc_ext.numpy()[0]
result.cacc[:] = d.cacc.numpy()[world_id]
result.cfrc_int[:] = d.cfrc_int.numpy()[world_id]
result.cfrc_ext[:] = d.cfrc_ext.numpy()[world_id]
# tendon
result.ten_length[:] = d.ten_length.numpy()[0]
result.ten_J[:] = d.ten_J.numpy()[0]
result.ten_wrapadr[:] = d.ten_wrapadr.numpy()[0]
result.ten_wrapnum[:] = d.ten_wrapnum.numpy()[0]
result.wrap_obj[:] = d.wrap_obj.numpy()[0]
result.wrap_xpos[:] = d.wrap_xpos.numpy()[0]
result.ten_length[:] = d.ten_length.numpy()[world_id]
result.ten_J[:] = d.ten_J.numpy()[world_id]
result.ten_wrapadr[:] = d.ten_wrapadr.numpy()[world_id]
result.ten_wrapnum[:] = d.ten_wrapnum.numpy()[world_id]
result.wrap_obj[:] = d.wrap_obj.numpy()[world_id]
result.wrap_xpos[:] = d.wrap_xpos.numpy()[world_id]
# sensors
result.sensordata[:] = d.sensordata.numpy()
result.sensordata[:] = d.sensordata.numpy()[world_id]
def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None):