From 85532b74e109fe9ef78d96129e546afeea66a043 Mon Sep 17 00:00:00 2001 From: Erik Frey Date: Thu, 18 Jul 2024 15:00:44 -0700 Subject: [PATCH] Speed up hfield JIT time by reducing XLA code generation overhead. PiperOrigin-RevId: 653766462 Change-Id: I9d7e7fa3a4d33ae0fa9d00babd5c5478eb4eff88 --- mjx/mujoco/mjx/_src/collision_convex.py | 87 +++++++++++++------------ 1 file changed, 46 insertions(+), 41 deletions(-) diff --git a/mjx/mujoco/mjx/_src/collision_convex.py b/mjx/mujoco/mjx/_src/collision_convex.py index ea9c62d9..b6876c9d 100644 --- a/mjx/mujoco/mjx/_src/collision_convex.py +++ b/mjx/mujoco/mjx/_src/collision_convex.py @@ -1079,51 +1079,56 @@ def _hfield_collision( bmask = jp.array([True, True, False]) # process all prisms in sub-grid - prisms = [] - for r in range(subgrid_size[1]): - for c in range(subgrid_size[0]): - ri, ci = rmin + r, cmin + c + rs = jp.repeat(jp.arange(subgrid_size[1]), subgrid_size[0]) + cs = jp.tile(jp.arange(subgrid_size[0]), subgrid_size[1]) - # ensure ri, ci are in the bounds of the hfield - ri = jp.clip(ri, 0, h.nrow - 2) - ci = jp.clip(ci, 0, h.ncol - 2) + @jax.vmap + def make_prisms(r, c): + ri, ci = rmin + r, cmin + c - p1 = [ - dx * ci - h.size[0], - dy * ri - h.size[1], - h.data[ci, ri] * h.size[2], - ] - p2 = [ - dx * (ci + 1) - h.size[0], - dy * (ri + 1) - h.size[1], - h.data[ci + 1, ri + 1] * h.size[2], - ] - p3 = [ - dx * ci - h.size[0], - dy * (ri + 1) - h.size[1], - h.data[ci, ri + 1] * h.size[2], - ] - top = jp.array([p1, p2, p3]) - bottom = jp.array([p1, p3, p2]) * bmask + bvert - vert = jp.concatenate([bottom, top]) - prisms.append(mesh.hfield_prism(vert)) + # ensure ri, ci are in the bounds of the hfield + ri = jp.clip(ri, 0, h.nrow - 2) + ci = jp.clip(ci, 0, h.ncol - 2) - p3 = p2 - p2 = [ - dx * (ci + 1) - h.size[0], - dy * ri - h.size[1], - h.data[ci + 1, ri] * h.size[2], - ] - top = jp.array([p1, p2, p3]) - bottom = jp.array([p1, p3, p2]) * bmask + bvert - vert = jp.concatenate([bottom, top]) - # NB: If the order of verts is updated above, the corresponding - # hfield_prism function must be updated to ensure that all faces have the - # correct winding order. - prisms.append(mesh.hfield_prism(vert)) + p1 = [ + dx * ci - h.size[0], + dy * ri - h.size[1], + h.data[ci, ri] * h.size[2], + ] + p2 = [ + dx * (ci + 1) - h.size[0], + dy * (ri + 1) - h.size[1], + h.data[ci + 1, ri + 1] * h.size[2], + ] + p3 = [ + dx * ci - h.size[0], + dy * (ri + 1) - h.size[1], + h.data[ci, ri + 1] * h.size[2], + ] + top = jp.array([p1, p2, p3]) + bottom = jp.array([p1, p3, p2]) * bmask + bvert + vert = jp.concatenate([bottom, top]) + prism1 = mesh.hfield_prism(vert) - n_prisms = len(prisms) - prisms = jax.tree_util.tree_map(lambda *x: jp.stack(x), *prisms) + p3 = p2 + p2 = [ + dx * (ci + 1) - h.size[0], + dy * ri - h.size[1], + h.data[ci + 1, ri] * h.size[2], + ] + top = jp.array([p1, p2, p3]) + bottom = jp.array([p1, p3, p2]) * bmask + bvert + vert = jp.concatenate([bottom, top]) + # NB: If the order of verts is updated above, the corresponding + # hfield_prism function must be updated to ensure that all faces have the + # correct winding order. + prism2 = mesh.hfield_prism(vert) + + return prism1, prism2 + + prism1, prism2 = make_prisms(rs, cs) + n_prisms = 2 * rs.shape[0] + prisms = jax.tree_util.tree_map(lambda *x: jp.concatenate(x), prism1, prism2) dist, pos, n = jax.vmap(collider_fn, in_axes=[None, 0])( obj.replace(pos=obj_pos, mat=obj_mat), prisms )