Improve performance of getting and putting MJX device data by reducing cost of ncon.
PiperOrigin-RevId: 605454760 Change-Id: I5c52a84c12ab67d042fe43f6ada284b51a762294
This commit is contained in:
committed by
Copybara-Service
parent
3e2d62c248
commit
4933a2c7b6
+6
-1
@@ -8,7 +8,12 @@ Upcoming version (not yet released)
|
||||
MJX
|
||||
^^^
|
||||
|
||||
1. Improved performance of getting and putting device data by using a faster numpy array serialization method.
|
||||
1. Improved performance of getting and putting device data.
|
||||
|
||||
- Use ``tobytes()`` for numpy array serialization, which is orders of magnitude faster than converting to tuples.
|
||||
- Avoid reallocating host ``mjData`` arrays when array shapes are unchanged.
|
||||
- Speed up calculation of ``mjx.ncon`` for models with many geoms.
|
||||
- Avoid calling ``mjx.ncon`` in ``mjx.get_data_into`` when ``nc`` can be derived from ``mjx.Data``.
|
||||
|
||||
Version 3.1.2 (February 05, 2024)
|
||||
-----------------------------------
|
||||
|
||||
@@ -313,8 +313,16 @@ def collision_candidates(m: Union[Model, mujoco.MjModel]) -> CandidateSet:
|
||||
|
||||
body_pairs = []
|
||||
exclude_signature = set(m.exclude_signature)
|
||||
geom_con = m.geom_contype | m.geom_conaffinity
|
||||
b_start = m.body_geomadr
|
||||
b_end = b_start + m.body_geomnum
|
||||
|
||||
for b1 in range(m.nbody):
|
||||
if not geom_con[b_start[b1]:b_end[b1]].any():
|
||||
continue
|
||||
for b2 in range(b1, m.nbody):
|
||||
if not geom_con[b_start[b2]:b_end[b2]].any():
|
||||
continue
|
||||
signature = (b1 << 16) + (b2)
|
||||
if signature in exclude_signature:
|
||||
continue
|
||||
@@ -323,12 +331,12 @@ def collision_candidates(m: Union[Model, mujoco.MjModel]) -> CandidateSet:
|
||||
body_pairs.append((b1, b2))
|
||||
|
||||
for b1, b2 in body_pairs:
|
||||
start1 = m.body_geomadr[b1]
|
||||
end1 = m.body_geomadr[b1] + m.body_geomnum[b1]
|
||||
for g1 in range(start1, end1):
|
||||
start2 = m.body_geomadr[b2]
|
||||
end2 = m.body_geomadr[b2] + m.body_geomnum[b2]
|
||||
for g2 in range(start2, end2):
|
||||
for g1 in range(b_start[b1], b_end[b1]):
|
||||
if not geom_con[g1]:
|
||||
continue
|
||||
for g2 in range(b_start[b2], b_end[b2]):
|
||||
if not geom_con[g2]:
|
||||
continue
|
||||
mask = m.geom_contype[g1] & m.geom_conaffinity[g2]
|
||||
mask |= m.geom_contype[g2] & m.geom_conaffinity[g1]
|
||||
if mask != 0:
|
||||
|
||||
@@ -315,7 +315,7 @@ def _instantiate_contact(m: Model, d: Data) -> Optional[_Efc]:
|
||||
|
||||
|
||||
def count_constraints(
|
||||
m: Union[Model, mujoco.MjModel]
|
||||
m: Union[Model, mujoco.MjModel], d: Optional[Data] = None
|
||||
) -> Tuple[int, int, int, int]:
|
||||
"""Returns equality, friction, limit, and contact constraint counts."""
|
||||
if m.opt.disableflags & DisableBit.CONSTRAINT:
|
||||
@@ -336,7 +336,10 @@ def count_constraints(
|
||||
else:
|
||||
nl = int(m.jnt_limited.sum())
|
||||
|
||||
nc = collision_driver.ncon(m) * 4
|
||||
if d is None:
|
||||
nc = collision_driver.ncon(m) * 4
|
||||
else:
|
||||
nc = d.efc_J.shape[-2] - ne - nf - nl
|
||||
|
||||
return ne, nf, nl, nc
|
||||
|
||||
|
||||
@@ -264,7 +264,7 @@ def get_data_into(
|
||||
d = jax.device_get(d)
|
||||
|
||||
batch_size = d.qpos.shape[0] if batched else 1
|
||||
ne, nf, nl, nc = constraint.count_constraints(m)
|
||||
ne, nf, nl, nc = constraint.count_constraints(m, d)
|
||||
efc_type = np.array([
|
||||
mujoco.mjtConstraint.mjCNSTR_EQUALITY,
|
||||
mujoco.mjtConstraint.mjCNSTR_FRICTION_DOF,
|
||||
@@ -288,7 +288,8 @@ def get_data_into(
|
||||
efc_con = efc_type == mujoco.mjtConstraint.mjCNSTR_CONTACT_PYRAMIDAL
|
||||
nefc, nc = int(efc_active.sum()), int((efc_active & efc_con).sum())
|
||||
result_i.nnzJ = nefc * m.nv
|
||||
mujoco._functions._realloc_con_efc(result_i, ncon=ncon, nefc=nefc) # pylint: disable=protected-access
|
||||
if ncon != result_i.ncon or nefc != result_i.nefc:
|
||||
mujoco._functions._realloc_con_efc(result_i, ncon=ncon, nefc=nefc) # pylint: disable=protected-access
|
||||
result_i.efc_J_rownnz[:] = np.repeat(m.nv, nefc)
|
||||
result_i.efc_J_rowadr[:] = np.arange(0, nefc * m.nv, m.nv)
|
||||
result_i.efc_J_colind[:] = np.tile(np.arange(m.nv), nefc)
|
||||
|
||||
Reference in New Issue
Block a user