Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
61 changes: 45 additions & 16 deletions src_rbd_shaders/dynamics/multibody/gravity_and_lu.rs
Original file line number Diff line number Diff line change
Expand Up @@ -21,9 +21,7 @@ use crate::utils::linalg::{MAX_MB_DOFS, MatSlice, fill_par, gemv_tr_spatial_spli
use crate::utils::{BatchIndices, Slice};
use crate::{AngVector, Vector, gcross_av};

use super::lu::{
LANES, lu_apply_pivots, lu_factor_in_shared, lu_triangular_solve_in_place, sm_idx,
};
use super::lu::{LANES, NO_PARENT, ltdl_factor_in_shared, ltdl_solve_in_shared, sm_idx};
use super::types::{MultibodyInfo, MultibodyLinkStatic, MultibodyLinkWorkspace};

/// Fused gravity / Coriolis-force assembly + LU factor + LU solve.
Expand Down Expand Up @@ -235,30 +233,61 @@ pub fn gpu_mb_gravity_and_lu(
}
x.write(lane as usize, gen_forces.read(rhs_offset + lane as usize));
}

// Per-DOF parent array (tree metadata for the sparse LᵀDL factor/solves),
// written into the old pivots buffer so downstream solves read the tree
// from the binding they already have. A joint's DOFs chain internally;
// its first DOF hangs off the last DOF of the nearest ancestor link that
// has any (welded links contribute none).
if lane == 0 {
let mut k = 0u32;
while k < num_links {
let stat = stat_slice[k as usize];
if stat.ndofs > 0 {
let mut p = NO_PARENT;
if k != 0 {
let mut a = stat.parent_link_id;
for _ in 0..MAX_MB_DOFS as u32 {
let s = stat_slice[a as usize];
if s.ndofs > 0 {
p = s.assembly_id + s.ndofs - 1;
break;
}
if a == 0 {
break;
}
a = s.parent_link_id;
}
}
lu_pivots.write(piv_offset + stat.assembly_id as usize, p);
for t in 1..stat.ndofs {
lu_pivots.write(
piv_offset + (stat.assembly_id + t) as usize,
stat.assembly_id + t - 1,
);
}
}
k += 1;
}
}
workgroup_memory_barrier_with_group_sync();

lu_factor_in_shared(
ndofs,
lane,
mat,
lu_pivots,
piv_offset,
pivot_row_shared,
inv_akk_shared,
);
ltdl_factor_in_shared(ndofs, lane, mat, lu_pivots, piv_offset);

// Persist LU factors to global memory (joint / contact constraint init
// reuses them for unit-RHS solves).
// Persist the LᵀDL factors to global memory (joint / contact constraint
// init reuses them for unit-RHS solves).
if lane < ndofs {
for r in 0..ndofs {
mass_matrices.write(m_view.idx(r, lane), mat.read(sm_idx(r, lane)));
}
}

// ---- Phase 4: solve M·x = τ for the gravity rhs. ----
lu_apply_pivots(ndofs, lane, lu_pivots, piv_offset, x);
lu_triangular_solve_in_place(ndofs, lane, mat, x, partial);
ltdl_solve_in_shared(ndofs, lane, mat, lu_pivots, piv_offset, x);

let _ = partial;
let _ = pivot_row_shared;
let _ = inv_akk_shared;
if lane < ndofs {
gen_forces.write(rhs_offset + lane as usize, x.read(lane as usize));
}
Expand Down
244 changes: 96 additions & 148 deletions src_rbd_shaders/dynamics/multibody/lu.rs
Original file line number Diff line number Diff line change
@@ -1,18 +1,38 @@
//! LU decomposition + solve.
//! Tree-sparse LᵀDL factorization + solve for the multibody mass matrix.
//!
//! Split into two kernels so the factorization can be reused across multiple
//! right-hand sides within a frame (e.g. gravity τ, contact impulses, …).
//! Dispatched one workgroup per `(multibody, batch)` pair.
//! The mass matrix of a kinematic tree has branch-induced sparsity:
//! `M[i, j] != 0` only when DOFs `i` and `j` lie on the same root-to-leaf
//! path. Eliminating leaves-to-root (descending DOF order, since parents are
//! assembled before children) factors `M = Lᵀ·D·L` with ZERO fill-in —
//! `L[k, i]` is nonzero only for `i` on `k`'s ancestor chain (Featherstone,
//! RBDA §6; MuJoCo's `mj_factorM`). Factor cost drops from O(n³) to
//! O(n·depth²) and every solve from O(n²) to O(n·depth).
//!
//! The per-DOF parent index (`u32::MAX` at roots) is stored in the buffer
//! that used to hold the LU pivots, so every downstream solve reads the tree
//! from a binding it already had. Factors live in the dense `mass_matrices`
//! storage: strict-lower ancestor-chain entries hold `L` (unit diagonal
//! implied), the diagonal holds `D`, and the upper triangle keeps stale
//! symmetric copies of M that no solve reads.
//!
//! The factor/solve loops are data-dependent but contain NO barriers, so a
//! single lane runs them serially — legal under WGSL uniformity rules (the
//! old dense kernels needed fixed 64-iteration loops only because they
//! barriered inside). At n ≤ 64, depth ≈ 7, the serial sparse path does a
//! few hundred FMAs versus the dense path's ~192 workgroup barriers.

use khal_std::index::MaybeIndexUnchecked;
use khal_std::sync::workgroup_memory_barrier_with_group_sync;

use crate::utils::linalg::MAX_MB_DOFS;

/// Workgroup width for the parallelised LU kernels. Must match the
/// Workgroup width for the multibody solver kernels. Must match the
/// `threads(N, 1, 1)` attribute and `MB_LU_LANES` on the host side.
pub(super) const LANES: u32 = 64;

/// Parent sentinel for root DOFs in the per-DOF parent array.
pub(super) const NO_PARENT: u32 = u32::MAX;

/// Side length of the workgroup-shared matrix tile. Must equal both the lane
/// count and the maximum supported `ndofs`.
const MAX_DOFS_U32: u32 = MAX_MB_DOFS as u32;
Expand All @@ -23,168 +43,96 @@ pub(super) fn sm_idx(r: u32, c: u32) -> usize {
(c * MAX_DOFS_U32 + r) as usize
}

/// Workgroup-parallel LU factorization on the shared `mat` tile in place.
/// Tree-sparse LᵀDL factorization of the shared `mat` tile in place, run
/// serially by lane 0 (no barriers inside — the loops are data-dependent).
/// `parents` holds the per-DOF parent indices (`NO_PARENT` at roots), read
/// from the pivots buffer at `parents_offset`.
///
/// Ends with a workgroup barrier so every lane sees the factors.
#[inline]
pub(super) fn lu_factor_in_shared(
pub(super) fn ltdl_factor_in_shared(
n: u32,
lane: u32,
mat: &mut [f32; MAX_MB_DOFS * MAX_MB_DOFS],
pivots_dst: &mut [u32],
pivots_offset: usize,
pivot_row_shared: &mut u32,
inv_akk_shared: &mut f32,
mat: &mut impl MaybeIndexUnchecked<f32>,
parents: &[u32],
parents_offset: usize,
) {
// NOTE: fixed number of iterations for uniform control flow.
// TODO(PERF): on non-web platforms we could just use `n` as the upper bound.
for k in 0..MAX_DOFS_U32 {
let active = k < n;
if active && lane == 0 {
let mut pivot_row = k;
let mut pivot_val = {
let v = mat.read(sm_idx(k, k));
if v >= 0.0 { v } else { -v }
};
for i in (k + 1)..n {
let v = mat.read(sm_idx(i, k));
let av = if v >= 0.0 { v } else { -v };
if av > pivot_val {
pivot_val = av;
pivot_row = i;
if lane == 0 {
// Eliminate leaves-to-root: descending k, updating only the
// ancestor-chain entries of rows above.
for step in 0..n {
let k = n - 1 - step;
let d = mat.read(sm_idx(k, k));
let inv_d = if d != 0.0 { 1.0 / d } else { 0.0 };
let mut i = parents.read(parents_offset + k as usize);
// Bounded loops (parents are strictly decreasing) so a corrupt
// parent array can't hang the GPU.
for _ in 0..MAX_DOFS_U32 {
if i == NO_PARENT {
break;
}
}
*pivot_row_shared = pivot_row;
pivots_dst.write(pivots_offset + k as usize, pivot_row);
}
workgroup_memory_barrier_with_group_sync();
let pivot_row = *pivot_row_shared;

if active && pivot_row != k && lane < n {
let c = lane;
let a = mat.read(sm_idx(k, c));
let b = mat.read(sm_idx(pivot_row, c));
mat.write(sm_idx(k, c), b);
mat.write(sm_idx(pivot_row, c), a);
}
workgroup_memory_barrier_with_group_sync();

if active && lane == 0 {
let akk = mat.read(sm_idx(k, k));
*inv_akk_shared = if akk != 0.0 { 1.0 / akk } else { 0.0 };
}
workgroup_memory_barrier_with_group_sync();
let inv_akk = *inv_akk_shared;

if active {
let r = k + 1 + lane;
if r < n {
let v = mat.read(sm_idx(r, k)) * inv_akk;
mat.write(sm_idx(r, k), v);
}
}
workgroup_memory_barrier_with_group_sync();

if active {
let j = k + 1 + lane;
if j < n {
let akj = mat.read(sm_idx(k, j));
for r in (k + 1)..n {
let lik = mat.read(sm_idx(r, k));
let v = mat.read(sm_idx(r, j)) - lik * akj;
mat.write(sm_idx(r, j), v);
let a = mat.read(sm_idx(k, i)) * inv_d;
let mut j = i;
for _ in 0..MAX_DOFS_U32 {
if j == NO_PARENT {
break;
}
let v = mat.read(sm_idx(i, j)) - a * mat.read(sm_idx(k, j));
mat.write(sm_idx(i, j), v);
j = parents.read(parents_offset + j as usize);
}
mat.write(sm_idx(k, i), a);
i = parents.read(parents_offset + i as usize);
}
}
workgroup_memory_barrier_with_group_sync();
}
workgroup_memory_barrier_with_group_sync();
}

/// Inner of the workgroup-parallel triangular solves used by the LU
/// solve/factor-and-solve kernels. Preconditions: `mat`
/// already holds the LU factors and `x` already holds the permuted rhs.
/// Solve `M·x = b` in place on the shared tile holding LᵀDL factors, run
/// serially by lane 0. Ends with a workgroup barrier.
#[inline]
pub(super) fn lu_triangular_solve_in_place(
pub(super) fn ltdl_solve_in_shared(
n: u32,
lane: u32,
mat: &[f32; MAX_MB_DOFS * MAX_MB_DOFS],
x: &mut [f32; MAX_MB_DOFS],
partial: &mut [f32; LANES as usize],
mat: &impl MaybeIndexUnchecked<f32>,
parents: &[u32],
parents_offset: usize,
x: &mut impl MaybeIndexUnchecked<f32>,
) {
// NOTE: fixed number of iterations for uniform control flow.
// TODO(PERF): on non-web platforms we could just use `n` as the upper bound.
for i in 0..MAX_DOFS_U32 {
let active = i < n;
let s = if active && lane < i {
mat.read(sm_idx(i, lane)) * x.read(lane as usize)
} else {
0.0f32
};
partial.write(lane as usize, s);
workgroup_memory_barrier_with_group_sync();
for step in 0..6u32 {
let stride = 1u32 << (5 - step);
if lane < stride {
let v = partial.read(lane as usize) + partial.read((lane + stride) as usize);
partial.write(lane as usize, v);
}
workgroup_memory_barrier_with_group_sync();
}
if active && lane == 0 {
let cur = x.read(i as usize);
x.write(i as usize, cur - partial.read(0));
}
workgroup_memory_barrier_with_group_sync();
}

// NOTE: fixed number of iterations for uniform control flow.
// TODO(PERF): on non-web platforms we could just use `n` as the upper bound.
for step in 0..MAX_DOFS_U32 {
let active = step < n;
// For dummy iterations (step >= n), `i` is not meaningful — guard
// every use of it behind `active`.
let i = if active { n - 1 - step } else { 0 };
let s = if active && lane > i && lane < n {
mat.read(sm_idx(i, lane)) * x.read(lane as usize)
} else {
0.0f32
};
partial.write(lane as usize, s);
workgroup_memory_barrier_with_group_sync();
for r in 0..6u32 {
let stride = 1u32 << (5 - r);
if lane < stride {
let v = partial.read(lane as usize) + partial.read((lane + stride) as usize);
partial.write(lane as usize, v);
if lane == 0 {
// Solve Lᵀ·z = b: scatter descending (descendants before ancestors).
for step in 0..n {
let i = n - 1 - step;
let xi = x.read(i as usize);
let mut j = parents.read(parents_offset + i as usize);
for _ in 0..MAX_DOFS_U32 {
if j == NO_PARENT {
break;
}
let v = x.read(j as usize) - mat.read(sm_idx(i, j)) * xi;
x.write(j as usize, v);
j = parents.read(parents_offset + j as usize);
}
workgroup_memory_barrier_with_group_sync();
}
if active && lane == 0 {
let u = mat.read(sm_idx(i, i));
let cur = x.read(i as usize) - partial.read(0);
x.write(i as usize, if u != 0.0 { cur / u } else { 0.0 });
// z = D⁻¹·z.
for i in 0..n {
let d = mat.read(sm_idx(i, i));
let v = x.read(i as usize);
x.write(i as usize, if d != 0.0 { v / d } else { 0.0 });
}
workgroup_memory_barrier_with_group_sync();
}
}

/// Apply the recorded pivots (sequential — lane 0 only). `n` is small so the
/// extra parallelism wouldn't pay for the barrier.
#[inline]
pub(super) fn lu_apply_pivots(
n: u32,
lane: u32,
buf_pivots: &[u32],
pivots_offset: usize,
x: &mut [f32; MAX_MB_DOFS],
) {
if lane == 0 {
for k in 0..n {
let p = buf_pivots.read(pivots_offset + k as usize);
if p != k {
let a = x.read(k as usize);
let b = x.read(p as usize);
x.write(k as usize, b);
x.write(p as usize, a);
// Solve L·x = z: gather ascending (ancestors before descendants).
for i in 0..n {
let mut s = x.read(i as usize);
let mut j = parents.read(parents_offset + i as usize);
for _ in 0..MAX_DOFS_U32 {
if j == NO_PARENT {
break;
}
s -= mat.read(sm_idx(i, j)) * x.read(j as usize);
j = parents.read(parents_offset + j as usize);
}
x.write(i as usize, s);
}
}
workgroup_memory_barrier_with_group_sync();
Expand Down
Loading
Loading