From 29d6e103757b4caeacdca2328fbf50940204cb90 Mon Sep 17 00:00:00 2001 From: Jaroslav Sevcik Date: Sun, 25 Feb 2024 08:48:07 -0800 Subject: [PATCH 1/2] Update independent elements in solve_m in parallel --- mjx/mujoco/mjx/_src/smooth.py | 20 ++++++++++++-------- 1 file changed, 12 insertions(+), 8 deletions(-) diff --git a/mjx/mujoco/mjx/_src/smooth.py b/mjx/mujoco/mjx/_src/smooth.py index 632a225d..310aa3ae 100644 --- a/mjx/mujoco/mjx/_src/smooth.py +++ b/mjx/mujoco/mjx/_src/smooth.py @@ -272,6 +272,10 @@ def solve_m(m: Model, d: Data, x: jax.Array) -> jax.Array: if not support.is_sparse(m): return jax.scipy.linalg.cho_solve((d.qLD, False), x) + depth = [] + for i in range(m.nv): + depth.append(depth[m.dof_parentid[i]] + 1 if m.dof_parentid[i] != -1 else 0) + updates_i, updates_j = {}, {} for i in range(m.nv): madr_ij, j = m.dof_Madr[i], i @@ -279,21 +283,21 @@ def solve_m(m: Model, d: Data, x: jax.Array) -> jax.Array: madr_ij, j = madr_ij + 1, m.dof_parentid[j] if j == -1: break - updates_i.setdefault(i, []).append((madr_ij, j)) - updates_j.setdefault(j, []).append((madr_ij, i)) + updates_i.setdefault(depth[i], []).append((i, madr_ij, j)) + updates_j.setdefault(depth[j], []).append((j, madr_ij, i)) # x <- inv(L') * x - for j, vals in sorted(updates_j.items(), reverse=True): - madr_ij, i = jp.array(vals).T - x = x.at[j].add(-jp.sum(d.qLD[madr_ij] * x[i])) + for _, vals in sorted(updates_j.items(), reverse=True): + j, madr_ij, i = np.array(vals).T + x = x.at[j].add(-(d.qLD[madr_ij] * x[i])) # x <- inv(D) * x x = x * d.qLDiagInv # x <- inv(L) * x - for i, vals in sorted(updates_i.items()): - madr_ij, j = jp.array(vals).T - x = x.at[i].add(-jp.sum(d.qLD[madr_ij] * x[j])) + for _, vals in sorted(updates_i.items()): + i, madr_ij, j = np.array(vals).T + x = x.at[i].add(-(d.qLD[madr_ij] * x[j])) return x From 7688e03eea9e471413abcb640bb30356c53d98e5 Mon Sep 17 00:00:00 2001 From: Jaroslav Sevcik Date: Tue, 27 Feb 2024 00:10:21 -0800 Subject: [PATCH 2/2] Remove unnecessary parens --- mjx/mujoco/mjx/_src/smooth.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/mjx/mujoco/mjx/_src/smooth.py b/mjx/mujoco/mjx/_src/smooth.py index 310aa3ae..496c5dbb 100644 --- a/mjx/mujoco/mjx/_src/smooth.py +++ b/mjx/mujoco/mjx/_src/smooth.py @@ -289,7 +289,7 @@ def solve_m(m: Model, d: Data, x: jax.Array) -> jax.Array: # x <- inv(L') * x for _, vals in sorted(updates_j.items(), reverse=True): j, madr_ij, i = np.array(vals).T - x = x.at[j].add(-(d.qLD[madr_ij] * x[i])) + x = x.at[j].add(-d.qLD[madr_ij] * x[i]) # x <- inv(D) * x x = x * d.qLDiagInv @@ -297,7 +297,7 @@ def solve_m(m: Model, d: Data, x: jax.Array) -> jax.Array: # x <- inv(L) * x for _, vals in sorted(updates_i.items()): i, madr_ij, j = np.array(vals).T - x = x.at[i].add(-(d.qLD[madr_ij] * x[j])) + x = x.at[i].add(-d.qLD[madr_ij] * x[j]) return x