Defer device puts to the end of make_data. Fixes #2461.
PiperOrigin-RevId: 744875869 Change-Id: I8d88d4f73407b44f761d01309e2cc28b43aa305f
This commit is contained in:
committed by
Copybara-Service
parent
606f00f802
commit
4e0a4f4d39
+146
-144
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user