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:
Erik Frey
2024-02-08 16:05:47 -08:00
committed by Copybara-Service
parent 3e2d62c248
commit 4933a2c7b6
4 changed files with 28 additions and 11 deletions
+6 -1
View File
@@ -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)
-----------------------------------
+14 -6
View File
@@ -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:
+5 -2
View File
@@ -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
+3 -2
View File
@@ -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)