diff --git a/doc/changelog.rst b/doc/changelog.rst index 8cbb2fb3..9e6dc75d 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -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) ----------------------------------- diff --git a/mjx/mujoco/mjx/_src/collision_driver.py b/mjx/mujoco/mjx/_src/collision_driver.py index c54278b2..721b4345 100644 --- a/mjx/mujoco/mjx/_src/collision_driver.py +++ b/mjx/mujoco/mjx/_src/collision_driver.py @@ -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: diff --git a/mjx/mujoco/mjx/_src/constraint.py b/mjx/mujoco/mjx/_src/constraint.py index 50efcad3..85df7e18 100644 --- a/mjx/mujoco/mjx/_src/constraint.py +++ b/mjx/mujoco/mjx/_src/constraint.py @@ -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 diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index b815c32c..2f979d71 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -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)