Defer device puts to the end of make_data. Fixes #2461.

PiperOrigin-RevId: 744875869
Change-Id: I8d88d4f73407b44f761d01309e2cc28b43aa305f
This commit is contained in:
Baruch Tabanpour
2025-04-07 15:42:51 -07:00
committed by Copybara-Service
parent 606f00f802
commit 4e0a4f4d39
+146 -144
View File
@@ -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