From 4e0a4f4d39eb16adaec47c5642942780521779ea Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Mon, 7 Apr 2025 15:42:51 -0700 Subject: [PATCH] Defer device puts to the end of make_data. Fixes #2461. PiperOrigin-RevId: 744875869 Change-Id: I8d88d4f73407b44f761d01309e2cc28b43aa305f --- mjx/mujoco/mjx/_src/io.py | 290 +++++++++++++++++++------------------- 1 file changed, 146 insertions(+), 144 deletions(-) diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 1e8c2b5e..aac18731 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -237,153 +237,154 @@ def make_data( ne, nf, nl, nc = constraint.counts(efc_type) ncon, nefc = dim.size, ne + nf + nl + nc - with jax.default_device(device): - contact = types.Contact( - dist=jp.zeros((ncon,), dtype=float), - pos=jp.zeros((ncon, 3), dtype=float), - frame=jp.zeros((ncon, 3, 3), dtype=float), - includemargin=jp.zeros((ncon,), dtype=float), - friction=jp.zeros((ncon, 5), dtype=float), - solref=jp.zeros((ncon, mujoco.mjNREF), dtype=float), - solreffriction=jp.zeros((ncon, mujoco.mjNREF), dtype=float), - solimp=jp.zeros((ncon, mujoco.mjNIMP), dtype=float), - dim=dim, - # let jax pick contact.geom int precision, for interop with - # jax_enable_x64 - geom1=jp.full((ncon,), -1, dtype=int), - geom2=jp.full((ncon,), -1, dtype=int), - geom=jp.full((ncon, 2), -1, dtype=int), - efc_address=efc_address, + float_ = jp.zeros(1, float).dtype + int_ = jp.zeros(1, int).dtype + contact = types.Contact( + dist=np.zeros((ncon,), dtype=float_), + pos=np.zeros((ncon, 3), dtype=float_), + frame=np.zeros((ncon, 3, 3), dtype=float_), + includemargin=np.zeros((ncon,), dtype=float_), + friction=np.zeros((ncon, 5), dtype=float_), + solref=np.zeros((ncon, mujoco.mjNREF), dtype=float_), + solreffriction=np.zeros((ncon, mujoco.mjNREF), dtype=float_), + solimp=np.zeros((ncon, mujoco.mjNIMP), dtype=float_), + dim=dim, + # let jax pick contact.geom int precision, for interop with + # jax_enable_x64 + geom1=np.full((ncon,), -1, dtype=int_), + geom2=np.full((ncon,), -1, dtype=int_), + geom=np.full((ncon, 2), -1, dtype=int_), + efc_address=efc_address, + ) + + if m.opt.cone == types.ConeType.ELLIPTIC and np.any(contact.dim == 1): + raise NotImplementedError( + 'condim=1 with ConeType.ELLIPTIC not implemented.' ) - if m.opt.cone == types.ConeType.ELLIPTIC and np.any(contact.dim == 1): - raise NotImplementedError( - 'condim=1 with ConeType.ELLIPTIC not implemented.' - ) + zero_fields = { + 'solver_niter': (int_,), + 'time': (float_,), + 'qvel': (m.nv, float_), + 'act': (m.na, float_), + 'qacc_warmstart': (m.nv, float_), + 'ctrl': (m.nu, float_), + 'qfrc_applied': (m.nv, float_), + 'xfrc_applied': (m.nbody, 6, float_), + 'mocap_pos': (m.nmocap, 3, float_), + 'mocap_quat': (m.nmocap, 4, float_), + 'qacc': (m.nv, float_), + 'act_dot': (m.na, float_), + 'userdata': (m.nuserdata, float_), + 'sensordata': (m.nsensordata, float_), + 'xpos': (m.nbody, 3, float_), + 'xquat': (m.nbody, 4, float_), + 'xmat': (m.nbody, 3, 3, float_), + 'xipos': (m.nbody, 3, float_), + 'ximat': (m.nbody, 3, 3, float_), + 'xanchor': (m.njnt, 3, float_), + 'xaxis': (m.njnt, 3, float_), + 'geom_xpos': (m.ngeom, 3, float_), + 'geom_xmat': (m.ngeom, 3, 3, float_), + 'site_xpos': (m.nsite, 3, float_), + 'site_xmat': (m.nsite, 3, 3, float_), + 'cam_xpos': (m.ncam, 3, float_), + 'cam_xmat': (m.ncam, 3, 3, float_), + 'light_xpos': (m.nlight, 3, float_), + 'light_xdir': (m.nlight, 3, float_), + 'subtree_com': (m.nbody, 3, float_), + 'cdof': (m.nv, 6, float_), + 'cinert': (m.nbody, 10, float_), + 'flexvert_xpos': (m.nflexvert, 3, float_), + 'flexelem_aabb': (m.nflexelem, 6, float_), + 'flexedge_J_rownnz': (m.nflexedge, np.int32), + 'flexedge_J_rowadr': (m.nflexedge, np.int32), + 'flexedge_J_colind': (m.nflexedge, m.nv, np.int32), + 'flexedge_J': (m.nflexedge, m.nv, float_), + 'flexedge_length': (m.nflexedge, float_), + 'ten_wrapadr': (m.ntendon, np.int32), + 'ten_wrapnum': (m.ntendon, np.int32), + 'ten_J_rownnz': (m.ntendon, np.int32), + 'ten_J_rowadr': (m.ntendon, np.int32), + 'ten_J_colind': (m.ntendon, m.nv, np.int32), + 'ten_J': (m.ntendon, m.nv, float_), + 'ten_length': (m.ntendon, float_), + 'wrap_obj': (m.nwrap, 2, np.int32), + 'wrap_xpos': (m.nwrap, 6, float_), + 'actuator_length': (m.nu, float_), + 'moment_rownnz': (m.nu, np.int32), + 'moment_rowadr': (m.nu, np.int32), + 'moment_colind': (m.nJmom, np.int32), + 'actuator_moment': (m.nu, m.nv, float_), + 'crb': (m.nbody, 10, float_), + 'qM': (m.nM, float_) if support.is_sparse(m) else (m.nv, m.nv, float_), + 'qLD': (m.nM, float_) if support.is_sparse(m) else (m.nv, m.nv, float_), + 'qLDiagInv': (m.nv, float_) if support.is_sparse(m) else (0, float_), + 'bvh_aabb_dyn': (m.nbvhdynamic, 6, float_), + 'bvh_active': (m.nbvh, np.uint8), + 'flexedge_velocity': (m.nflexedge, float_), + 'ten_velocity': (m.ntendon, float_), + 'actuator_velocity': (m.nu, float_), + 'cvel': (m.nbody, 6, float_), + 'cdof_dot': (m.nv, 6, float_), + 'qfrc_bias': (m.nv, float_), + 'qfrc_spring': (m.nv, float_), + 'qfrc_damper': (m.nv, float_), + 'qfrc_gravcomp': (m.nv, float_), + 'qfrc_fluid': (m.nv, float_), + 'qfrc_passive': (m.nv, float_), + 'subtree_linvel': (m.nbody, 3, float_), + 'subtree_angmom': (m.nbody, 3, float_), + 'qH': (m.nM, float_) if support.is_sparse(m) else (m.nv, m.nv, float_), + 'qHDiagInv': (m.nv, float_), + 'B_rownnz': (m.nbody, np.int32), + 'B_rowadr': (m.nbody, np.int32), + 'B_colind': (m.nB, np.int32), + 'M_rownnz': (m.nv, np.int32), + 'M_rowadr': (m.nv, np.int32), + 'M_colind': (m.nM, np.int32), + 'mapM2M': (m.nM, np.int32), + 'C_rownnz': (m.nv, np.int32), + 'C_rowadr': (m.nv, np.int32), + 'C_colind': (m.nC, np.int32), + 'mapM2C': (m.nC, np.int32), + 'D_rownnz': (m.nv, np.int32), + 'D_rowadr': (m.nv, np.int32), + 'D_diag': (m.nv, np.int32), + 'D_colind': (m.nD, np.int32), + 'mapM2D': (m.nD, np.int32), + 'mapD2M': (m.nM, np.int32), + 'qDeriv': (m.nD, float_), + 'qLU': (m.nD, float_), + 'actuator_force': (m.nu, float_), + 'qfrc_actuator': (m.nv, float_), + 'qfrc_smooth': (m.nv, float_), + 'qacc_smooth': (m.nv, float_), + 'qfrc_constraint': (m.nv, float_), + 'qfrc_inverse': (m.nv, float_), + 'cacc': (m.nbody, 6, float_), + 'cfrc_int': (m.nbody, 6, float_), + 'cfrc_ext': (m.nbody, 6, float_), + 'efc_J': (nefc, m.nv, float_), + 'efc_pos': (nefc, float_), + 'efc_margin': (nefc, float_), + 'efc_frictionloss': (nefc, float_), + 'efc_D': (nefc, float_), + 'efc_aref': (nefc, float_), + 'efc_force': (nefc, float_), + '_qM_sparse': (m.nM, float_), + '_qLD_sparse': (m.nM, float_), + '_qLDiagInv_sparse': (m.nv, float_), + } - zero_fields = { - 'solver_niter': (int,), - 'time': (float,), - 'qvel': (m.nv, float), - 'act': (m.na, float), - 'qacc_warmstart': (m.nv, float), - 'ctrl': (m.nu, float), - 'qfrc_applied': (m.nv, float), - 'xfrc_applied': (m.nbody, 6, float), - 'mocap_pos': (m.nmocap, 3, float), - 'mocap_quat': (m.nmocap, 4, float), - 'qacc': (m.nv, float), - 'act_dot': (m.na, float), - 'userdata': (m.nuserdata, float), - 'sensordata': (m.nsensordata, float), - 'xpos': (m.nbody, 3, float), - 'xquat': (m.nbody, 4, float), - 'xmat': (m.nbody, 3, 3, float), - 'xipos': (m.nbody, 3, float), - 'ximat': (m.nbody, 3, 3, float), - 'xanchor': (m.njnt, 3, float), - 'xaxis': (m.njnt, 3, float), - 'geom_xpos': (m.ngeom, 3, float), - 'geom_xmat': (m.ngeom, 3, 3, float), - 'site_xpos': (m.nsite, 3, float), - 'site_xmat': (m.nsite, 3, 3, float), - 'cam_xpos': (m.ncam, 3, float), - 'cam_xmat': (m.ncam, 3, 3, float), - 'light_xpos': (m.nlight, 3, float), - 'light_xdir': (m.nlight, 3, float), - 'subtree_com': (m.nbody, 3, float), - 'cdof': (m.nv, 6, float), - 'cinert': (m.nbody, 10, float), - 'flexvert_xpos': (m.nflexvert, 3, float), - 'flexelem_aabb': (m.nflexelem, 6, float), - 'flexedge_J_rownnz': (m.nflexedge, jp.int32), - 'flexedge_J_rowadr': (m.nflexedge, jp.int32), - 'flexedge_J_colind': (m.nflexedge, m.nv, jp.int32), - 'flexedge_J': (m.nflexedge, m.nv, float), - 'flexedge_length': (m.nflexedge, float), - 'ten_wrapadr': (m.ntendon, jp.int32), - 'ten_wrapnum': (m.ntendon, jp.int32), - 'ten_J_rownnz': (m.ntendon, jp.int32), - 'ten_J_rowadr': (m.ntendon, jp.int32), - 'ten_J_colind': (m.ntendon, m.nv, jp.int32), - 'ten_J': (m.ntendon, m.nv, float), - 'ten_length': (m.ntendon, float), - 'wrap_obj': (m.nwrap, 2, jp.int32), - 'wrap_xpos': (m.nwrap, 6, float), - 'actuator_length': (m.nu, float), - 'moment_rownnz': (m.nu, jp.int32), - 'moment_rowadr': (m.nu, jp.int32), - 'moment_colind': (m.nJmom, jp.int32), - 'actuator_moment': (m.nu, m.nv, float), - 'crb': (m.nbody, 10, float), - 'qM': (m.nM, float) if support.is_sparse(m) else (m.nv, m.nv, float), - 'qLD': (m.nM, float) if support.is_sparse(m) else (m.nv, m.nv, float), - 'qLDiagInv': (m.nv, float) if support.is_sparse(m) else (0, float), - 'bvh_aabb_dyn': (m.nbvhdynamic, 6, float), - 'bvh_active': (m.nbvh, jp.uint8), - 'flexedge_velocity': (m.nflexedge, float), - 'ten_velocity': (m.ntendon, float), - 'actuator_velocity': (m.nu, float), - 'cvel': (m.nbody, 6, float), - 'cdof_dot': (m.nv, 6, float), - 'qfrc_bias': (m.nv, float), - 'qfrc_spring': (m.nv, float), - 'qfrc_damper': (m.nv, float), - 'qfrc_gravcomp': (m.nv, float), - 'qfrc_fluid': (m.nv, float), - 'qfrc_passive': (m.nv, float), - 'subtree_linvel': (m.nbody, 3, float), - 'subtree_angmom': (m.nbody, 3, float), - 'qH': (m.nM, float) if support.is_sparse(m) else (m.nv, m.nv, float), - 'qHDiagInv': (m.nv, float), - 'B_rownnz': (m.nbody, jp.int32), - 'B_rowadr': (m.nbody, jp.int32), - 'B_colind': (m.nB, jp.int32), - 'M_rownnz': (m.nv, jp.int32), - 'M_rowadr': (m.nv, jp.int32), - 'M_colind': (m.nM, jp.int32), - 'mapM2M': (m.nM, jp.int32), - 'C_rownnz': (m.nv, jp.int32), - 'C_rowadr': (m.nv, jp.int32), - 'C_colind': (m.nC, jp.int32), - 'mapM2C': (m.nC, jp.int32), - 'D_rownnz': (m.nv, jp.int32), - 'D_rowadr': (m.nv, jp.int32), - 'D_diag': (m.nv, jp.int32), - 'D_colind': (m.nD, jp.int32), - 'mapM2D': (m.nD, jp.int32), - 'mapD2M': (m.nM, jp.int32), - 'qDeriv': (m.nD, float), - 'qLU': (m.nD, float), - 'actuator_force': (m.nu, float), - 'qfrc_actuator': (m.nv, float), - 'qfrc_smooth': (m.nv, float), - 'qacc_smooth': (m.nv, float), - 'qfrc_constraint': (m.nv, float), - 'qfrc_inverse': (m.nv, float), - 'cacc': (m.nbody, 6, float), - 'cfrc_int': (m.nbody, 6, float), - 'cfrc_ext': (m.nbody, 6, float), - 'efc_J': (nefc, m.nv, float), - 'efc_pos': (nefc, float), - 'efc_margin': (nefc, float), - 'efc_frictionloss': (nefc, float), - 'efc_D': (nefc, float), - 'efc_aref': (nefc, float), - 'efc_force': (nefc, float), - '_qM_sparse': (m.nM, float), - '_qLD_sparse': (m.nM, float), - '_qLDiagInv_sparse': (m.nv, float), - } + if not _full_compat: + for f in types.Data.fields(): + if f.metadata.get('restricted_to') in ('mujoco', 'mjx'): + zero_fields[f.name] = (0, zero_fields[f.name][-1]) - if not _full_compat: - for f in types.Data.fields(): - if f.metadata.get('restricted_to') in ('mujoco', 'mjx'): - zero_fields[f.name] = (0, zero_fields[f.name][-1]) - - zero_fields = { - k: jp.zeros(v[:-1], dtype=v[-1]) for k, v in zero_fields.items() - } + zero_fields = { + k: np.zeros(v[:-1], dtype=v[-1]) for k, v in zero_fields.items() + } d = types.Data( ne=ne, @@ -391,12 +392,13 @@ def make_data( nl=nl, nefc=nefc, ncon=ncon, - qpos=jp.array(m.qpos0), + qpos=jp.array(m.qpos0, dtype=float_), contact=contact, efc_type=efc_type, eq_active=m.eq_active0, **zero_fields, ) + d = jax.device_put(d, device=device) return d