From f56d0de490892eb4a9bab499a69e1eb103dfc05d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?S=C3=A9bastien=20Crozet?= Date: Thu, 17 Sep 2026 16:52:11 +0200 Subject: [PATCH 1/4] feat: add support for cuda-oxide --- Cargo.toml | 10 ++- crates/examples2d/Cargo.toml | 1 + crates/examples2d/all_examples2.rs | 7 ++ crates/examples3d/Cargo.toml | 1 + crates/nexus2d/Cargo.toml | 1 + crates/nexus3d/Cargo.toml | 1 + crates/nexus_mpm2d/Cargo.toml | 4 + crates/nexus_mpm3d/Cargo.toml | 4 + crates/nexus_mpm_shaders2d/Cargo.toml | 2 + crates/nexus_mpm_shaders3d/Cargo.toml | 2 + crates/nexus_python3d/Cargo.toml | 1 + crates/nexus_rbd2d/Cargo.toml | 4 + crates/nexus_rbd3d/Cargo.toml | 4 + crates/nexus_rbd_shaders2d/Cargo.toml | 2 + crates/nexus_rbd_shaders3d/Cargo.toml | 2 + crates/nexus_viewer2d/Cargo.toml | 1 + crates/nexus_viewer3d/Cargo.toml | 1 + src_mpm_shaders/grid/kernel.rs | 101 +++++++++++++++----------- src_mpm_shaders/solver/g2p.rs | 2 +- src_mpm_shaders/solver/g2p_cdf.rs | 4 +- src_viewer/viewer.rs | 29 ++++++-- 21 files changed, 132 insertions(+), 52 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index ba3757b0..174998b8 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -88,11 +88,19 @@ nexus_viewer3d = { version = "0.5.0", path = "crates/nexus_viewer3d" } [workspace.lints] rust.unexpected_cfgs = { level = "warn", check-cfg = [ - 'cfg(feature, values("dim2", "dim3", "cpu", "cuda", "metal"))', + 'cfg(feature, values("dim2", "dim3", "cpu", "cuda", "cuda-oxide", "metal"))', 'cfg(target_arch_is_gpu)', ] } [patch.crates-io] +# TODO: drop once khal 0.3.1 (WebGPU SPIR-V passthrough entry-point fix) and the +# khal/vortx `cuda-oxide` feature are published. +khal = { path = "../khal/crates/khal" } +khal-derive = { path = "../khal/crates/khal-derive" } +khal-std = { path = "../khal/crates/khal-std" } +khal-builder = { path = "../khal/crates/khal-builder" } +vortx = { path = "../vortx" } +vortx-shaders = { path = "../vortx/vortx-shaders" } #rapier2d = { path = "../rapier/crates/rapier2d" } #rapier3d = { path = "../rapier/crates/rapier3d" } #rapier3d-mjcf = { path = "../rapier/crates/rapier3d-mjcf" } diff --git a/crates/examples2d/Cargo.toml b/crates/examples2d/Cargo.toml index 7517291e..6adf93a3 100644 --- a/crates/examples2d/Cargo.toml +++ b/crates/examples2d/Cargo.toml @@ -10,6 +10,7 @@ default = [] cpu = ["nexus_viewer2d/cpu"] cpu-parallel = ["cpu", "nexus_viewer2d/cpu-parallel"] cuda = ["nexus_viewer2d/cuda"] +cuda-oxide = ["cuda", "nexus_viewer2d/cuda-oxide"] metal = ["nexus_viewer2d/metal"] [dependencies] diff --git a/crates/examples2d/all_examples2.rs b/crates/examples2d/all_examples2.rs index 6014f1e0..325e619d 100644 --- a/crates/examples2d/all_examples2.rs +++ b/crates/examples2d/all_examples2.rs @@ -80,6 +80,7 @@ struct CliOptions { example: Option, list: bool, cpu: bool, + cuda: bool, metal: bool, run: bool, } @@ -90,6 +91,7 @@ fn parse_command_line() -> CliOptions { example: None, list: false, cpu: false, + cuda: false, metal: false, run: false, }; @@ -99,6 +101,7 @@ fn parse_command_line() -> CliOptions { "--example" => opts.example = args.next(), "--list" => opts.list = true, "--cpu" => opts.cpu = true, + "--cuda" => opts.cuda = true, "--metal" => opts.metal = true, "--run" => opts.run = true, _ => {} @@ -139,6 +142,10 @@ pub async fn main() { if opts.cpu { viewer = viewer.with_cpu(); } + #[cfg(feature = "cuda")] + if opts.cuda { + viewer = viewer.with_backend(nexus_viewer2d::BackendType::Cuda); + } #[cfg(feature = "metal")] if opts.metal { viewer = viewer.with_backend(nexus_viewer2d::BackendType::Metal); diff --git a/crates/examples3d/Cargo.toml b/crates/examples3d/Cargo.toml index 9d60f77a..05d48912 100644 --- a/crates/examples3d/Cargo.toml +++ b/crates/examples3d/Cargo.toml @@ -10,6 +10,7 @@ default = [] cpu = ["nexus_viewer3d/cpu"] cpu-parallel = ["cpu", "nexus_viewer3d/cpu-parallel"] cuda = ["nexus_viewer3d/cuda"] +cuda-oxide = ["cuda", "nexus_viewer3d/cuda-oxide"] metal = ["nexus_viewer3d/metal"] [dependencies] diff --git a/crates/nexus2d/Cargo.toml b/crates/nexus2d/Cargo.toml index f8b2c299..18b24a83 100644 --- a/crates/nexus2d/Cargo.toml +++ b/crates/nexus2d/Cargo.toml @@ -28,6 +28,7 @@ metal = ["nexus_rbd2d?/metal", "nexus_mpm2d?/metal"] cpu = ["nexus_rbd2d?/cpu", "nexus_mpm2d?/cpu"] cpu-parallel = ["cpu", "nexus_rbd2d?/cpu-parallel", "nexus_mpm2d?/cpu-parallel"] cuda = ["nexus_rbd2d?/cuda", "nexus_mpm2d?/cuda"] +cuda-oxide = ["cuda", "nexus_rbd2d?/cuda-oxide", "nexus_mpm2d?/cuda-oxide"] rbd = ["dep:nexus_rbd2d"] mpm = ["dep:nexus_mpm2d"] diff --git a/crates/nexus3d/Cargo.toml b/crates/nexus3d/Cargo.toml index ef496c2b..53985582 100644 --- a/crates/nexus3d/Cargo.toml +++ b/crates/nexus3d/Cargo.toml @@ -28,6 +28,7 @@ metal = ["nexus_rbd3d?/metal", "nexus_mpm3d?/metal"] cpu = ["nexus_rbd3d?/cpu", "nexus_mpm3d?/cpu"] cpu-parallel = ["cpu", "nexus_rbd3d?/cpu-parallel", "nexus_mpm3d?/cpu-parallel"] cuda = ["nexus_rbd3d?/cuda", "nexus_mpm3d?/cuda"] +cuda-oxide = ["cuda", "nexus_rbd3d?/cuda-oxide", "nexus_mpm3d?/cuda-oxide"] rbd = ["dep:nexus_rbd3d"] mpm = ["dep:nexus_mpm3d"] diff --git a/crates/nexus_mpm2d/Cargo.toml b/crates/nexus_mpm2d/Cargo.toml index c50eced4..27af3981 100644 --- a/crates/nexus_mpm2d/Cargo.toml +++ b/crates/nexus_mpm2d/Cargo.toml @@ -29,6 +29,10 @@ metal = ["khal/metal", "nexus_rbd2d/metal"] cpu = ["nexus_mpm_shaders2d/cpu", "nexus_rbd2d/cpu", "vortx/cpu"] cpu-parallel = ["cpu", "nexus_mpm_shaders2d/cpu-parallel", "nexus_rbd2d/cpu-parallel", "vortx/cpu-parallel"] cuda = ["khal/cuda", "khal-builder/cuda", "nexus_mpm_shaders2d/cuda"] +# Native CUDA with kernels compiled by cuda-oxide instead of rust-cuda (implies `cuda`). +# Not forwarded to the shader crate: khal-builder enables its `cuda-oxide` in +# the separate `cargo oxide` build; on the host it would break the CPU backend. +cuda-oxide = ["cuda", "khal-builder/cuda-oxide"] [dependencies] nexus_mpm_shaders2d = { workspace = true } diff --git a/crates/nexus_mpm3d/Cargo.toml b/crates/nexus_mpm3d/Cargo.toml index 854e1870..702964c8 100644 --- a/crates/nexus_mpm3d/Cargo.toml +++ b/crates/nexus_mpm3d/Cargo.toml @@ -29,6 +29,10 @@ metal = ["khal/metal", "nexus_rbd3d/metal"] cpu = ["nexus_mpm_shaders3d/cpu", "nexus_rbd3d/cpu", "vortx/cpu"] cpu-parallel = ["cpu", "nexus_mpm_shaders3d/cpu-parallel", "nexus_rbd3d/cpu-parallel", "vortx/cpu-parallel"] cuda = ["khal/cuda", "khal-builder/cuda", "nexus_mpm_shaders3d/cuda"] +# Native CUDA with kernels compiled by cuda-oxide instead of rust-cuda (implies `cuda`). +# Not forwarded to the shader crate: khal-builder enables its `cuda-oxide` in +# the separate `cargo oxide` build; on the host it would break the CPU backend. +cuda-oxide = ["cuda", "khal-builder/cuda-oxide"] [dependencies] nexus_mpm_shaders3d = { workspace = true } diff --git a/crates/nexus_mpm_shaders2d/Cargo.toml b/crates/nexus_mpm_shaders2d/Cargo.toml index 9c85829f..6547169c 100644 --- a/crates/nexus_mpm_shaders2d/Cargo.toml +++ b/crates/nexus_mpm_shaders2d/Cargo.toml @@ -32,6 +32,8 @@ web-compat = ["nexus_rbd_shaders2d/web-compat"] cpu = [] cpu-parallel = ["cpu", "vortx-shaders/cpu-parallel"] cuda = [] +# cuda-oxide compiler for the CUDA kernels (implies `cuda`). +cuda-oxide = ["cuda", "khal-std/cuda-oxide", "vortx-shaders/cuda-oxide", "nexus_rbd_shaders2d/cuda-oxide"] [dependencies] nexus_rbd_shaders2d = { workspace = true } diff --git a/crates/nexus_mpm_shaders3d/Cargo.toml b/crates/nexus_mpm_shaders3d/Cargo.toml index b1c0a1ba..d8bd8ac6 100644 --- a/crates/nexus_mpm_shaders3d/Cargo.toml +++ b/crates/nexus_mpm_shaders3d/Cargo.toml @@ -31,6 +31,8 @@ web-compat = ["nexus_rbd_shaders3d/web-compat"] cpu = [] cpu-parallel = ["cpu", "vortx-shaders/cpu-parallel"] cuda = [] +# cuda-oxide compiler for the CUDA kernels (implies `cuda`). +cuda-oxide = ["cuda", "khal-std/cuda-oxide", "vortx-shaders/cuda-oxide", "nexus_rbd_shaders3d/cuda-oxide"] [dependencies] nexus_rbd_shaders3d = { workspace = true } diff --git a/crates/nexus_python3d/Cargo.toml b/crates/nexus_python3d/Cargo.toml index 1114d9e3..b109caa8 100644 --- a/crates/nexus_python3d/Cargo.toml +++ b/crates/nexus_python3d/Cargo.toml @@ -29,6 +29,7 @@ default = ["webgpu", "extension-module"] webgpu = ["nexus_viewer3d/webgpu"] metal = ["nexus_viewer3d/metal"] cuda = ["nexus_viewer3d/cuda"] +cuda-oxide = ["cuda", "nexus_viewer3d/cuda-oxide"] cpu = ["nexus_viewer3d/cpu"] cpu-parallel = ["nexus_viewer3d/cpu-parallel"] # Build as a real Python extension module (no libpython link). diff --git a/crates/nexus_rbd2d/Cargo.toml b/crates/nexus_rbd2d/Cargo.toml index 18c5fe2d..1691a7c4 100644 --- a/crates/nexus_rbd2d/Cargo.toml +++ b/crates/nexus_rbd2d/Cargo.toml @@ -29,6 +29,10 @@ metal = ["khal/metal"] cpu = ["nexus_rbd_shaders2d/cpu", "vortx/cpu"] cpu-parallel = ["cpu", "vortx/cpu-parallel", "nexus_rbd_shaders2d/cpu-parallel"] cuda = ["khal/cuda", "khal-builder/cuda", "nexus_rbd_shaders2d/cuda"] +# Native CUDA with kernels compiled by cuda-oxide instead of rust-cuda (implies `cuda`). +# Not forwarded to the shader crate: khal-builder enables its `cuda-oxide` in +# the separate `cargo oxide` build; on the host it would break the CPU backend. +cuda-oxide = ["cuda", "khal-builder/cuda-oxide"] [dependencies] nexus_rbd_shaders2d = { workspace = true } diff --git a/crates/nexus_rbd3d/Cargo.toml b/crates/nexus_rbd3d/Cargo.toml index 8660425d..a6faeeb4 100644 --- a/crates/nexus_rbd3d/Cargo.toml +++ b/crates/nexus_rbd3d/Cargo.toml @@ -29,6 +29,10 @@ metal = ["khal/metal"] cpu = ["nexus_rbd_shaders3d/cpu", "vortx/cpu"] cpu-parallel = ["cpu", "vortx/cpu-parallel", "nexus_rbd_shaders3d/cpu-parallel"] cuda = ["khal/cuda", "khal-builder/cuda", "nexus_rbd_shaders3d/cuda"] +# Native CUDA with kernels compiled by cuda-oxide instead of rust-cuda (implies `cuda`). +# Not forwarded to the shader crate: khal-builder enables its `cuda-oxide` in +# the separate `cargo oxide` build; on the host it would break the CPU backend. +cuda-oxide = ["cuda", "khal-builder/cuda-oxide"] [dependencies] nexus_rbd_shaders3d = { workspace = true } diff --git a/crates/nexus_rbd_shaders2d/Cargo.toml b/crates/nexus_rbd_shaders2d/Cargo.toml index d8672afd..68decbb4 100644 --- a/crates/nexus_rbd_shaders2d/Cargo.toml +++ b/crates/nexus_rbd_shaders2d/Cargo.toml @@ -34,6 +34,8 @@ web-compat = [] cpu = [] cpu-parallel = ["cpu", "vortx-shaders/cpu-parallel"] cuda = [] +# cuda-oxide compiler for the CUDA kernels (implies `cuda`). +cuda-oxide = ["cuda", "khal-std/cuda-oxide", "vortx-shaders/cuda-oxide"] [dependencies] vortx-shaders = { workspace = true } diff --git a/crates/nexus_rbd_shaders3d/Cargo.toml b/crates/nexus_rbd_shaders3d/Cargo.toml index 1b306203..942bf8f5 100644 --- a/crates/nexus_rbd_shaders3d/Cargo.toml +++ b/crates/nexus_rbd_shaders3d/Cargo.toml @@ -34,6 +34,8 @@ web-compat = [] cpu = [] cpu-parallel = ["cpu", "vortx-shaders/cpu-parallel"] cuda = [] +# cuda-oxide compiler for the CUDA kernels (implies `cuda`). +cuda-oxide = ["cuda", "khal-std/cuda-oxide", "vortx-shaders/cuda-oxide"] [dependencies] vortx-shaders = { workspace = true } diff --git a/crates/nexus_viewer2d/Cargo.toml b/crates/nexus_viewer2d/Cargo.toml index 61a8741d..8ca75a95 100644 --- a/crates/nexus_viewer2d/Cargo.toml +++ b/crates/nexus_viewer2d/Cargo.toml @@ -27,6 +27,7 @@ metal = ["khal/metal", "nexus2d/metal"] cpu = ["nexus2d/cpu", "vortx/cpu"] cpu-parallel = ["cpu", "nexus2d/cpu-parallel", "vortx/cpu-parallel"] cuda = ["nexus2d/cuda"] +cuda-oxide = ["cuda", "nexus2d/cuda-oxide"] [dependencies] nexus2d = { workspace = true, features = ["rbd", "mpm"]} diff --git a/crates/nexus_viewer3d/Cargo.toml b/crates/nexus_viewer3d/Cargo.toml index c5141da7..71886a97 100644 --- a/crates/nexus_viewer3d/Cargo.toml +++ b/crates/nexus_viewer3d/Cargo.toml @@ -27,6 +27,7 @@ metal = ["khal/metal", "nexus3d/metal"] cpu = ["nexus3d/cpu", "vortx/cpu"] cpu-parallel = ["cpu", "nexus3d/cpu-parallel", "vortx/cpu-parallel"] cuda = ["nexus3d/cuda"] +cuda-oxide = ["cuda", "nexus3d/cuda-oxide"] [dependencies] nexus3d = { workspace = true, features = ["rbd", "mpm"]} diff --git a/src_mpm_shaders/grid/kernel.rs b/src_mpm_shaders/grid/kernel.rs index b18fa7e6..908fd73c 100644 --- a/src_mpm_shaders/grid/kernel.rs +++ b/src_mpm_shaders/grid/kernel.rs @@ -20,55 +20,72 @@ pub const NBH_LEN: usize = 9; #[cfg(feature = "dim3")] pub const NBH_LEN: usize = 27; -/// Returns the stencil offset for neighbor `i` as a UVector. +/// Stencil offsets of the neighbors, as `[x, y]` (2D, 3x3 grid) or `[x, y, z]` +/// (3D, 3x3x3 grid) components. Read them through [`nbh_shift`]. /// -/// In 2D, these are UVec2 offsets into a 3x3 grid. -/// In 3D, these are UVec3 offsets into a 3x3x3 grid. +/// Stored as nested scalar arrays rather than `[UVec2/UVec3; N]`: the +/// cuda-oxide backend only materializes constant arrays whose elements are +/// scalars (or nested arrays of scalars), not structs. #[cfg(feature = "dim2")] -pub const NBH_SHIFTS: [UVec2; 9] = [ - UVec2::new(2, 2), - UVec2::new(2, 0), - UVec2::new(2, 1), - UVec2::new(0, 2), - UVec2::new(0, 0), - UVec2::new(0, 1), - UVec2::new(1, 2), - UVec2::new(1, 0), - UVec2::new(1, 1), +pub const NBH_SHIFTS: [[u32; 2]; 9] = [ + [2, 2], + [2, 0], + [2, 1], + [0, 2], + [0, 0], + [0, 1], + [1, 2], + [1, 0], + [1, 1], ]; -/// Returns the stencil offset for neighbor `i` as a UVector. #[cfg(feature = "dim3")] -pub const NBH_SHIFTS: [UVec3; 27] = [ - UVec3::new(2, 2, 2), - UVec3::new(2, 0, 2), - UVec3::new(2, 1, 2), - UVec3::new(0, 2, 2), - UVec3::new(0, 0, 2), - UVec3::new(0, 1, 2), - UVec3::new(1, 2, 2), - UVec3::new(1, 0, 2), - UVec3::new(1, 1, 2), - UVec3::new(2, 2, 0), - UVec3::new(2, 0, 0), - UVec3::new(2, 1, 0), - UVec3::new(0, 2, 0), - UVec3::new(0, 0, 0), - UVec3::new(0, 1, 0), - UVec3::new(1, 2, 0), - UVec3::new(1, 0, 0), - UVec3::new(1, 1, 0), - UVec3::new(2, 2, 1), - UVec3::new(2, 0, 1), - UVec3::new(2, 1, 1), - UVec3::new(0, 2, 1), - UVec3::new(0, 0, 1), - UVec3::new(0, 1, 1), - UVec3::new(1, 2, 1), - UVec3::new(1, 0, 1), - UVec3::new(1, 1, 1), +pub const NBH_SHIFTS: [[u32; 3]; 27] = [ + [2, 2, 2], + [2, 0, 2], + [2, 1, 2], + [0, 2, 2], + [0, 0, 2], + [0, 1, 2], + [1, 2, 2], + [1, 0, 2], + [1, 1, 2], + [2, 2, 0], + [2, 0, 0], + [2, 1, 0], + [0, 2, 0], + [0, 0, 0], + [0, 1, 0], + [1, 2, 0], + [1, 0, 0], + [1, 1, 0], + [2, 2, 1], + [2, 0, 1], + [2, 1, 1], + [0, 2, 1], + [0, 0, 1], + [0, 1, 1], + [1, 2, 1], + [1, 0, 1], + [1, 1, 1], ]; +/// Returns the stencil offset for neighbor `i` as a `UVec2` (2D). +#[cfg(feature = "dim2")] +#[inline(always)] +pub fn nbh_shift(i: usize) -> UVec2 { + use khal_std::index::MaybeIndexUnchecked; + UVec2::from_array(NBH_SHIFTS.read(i)) +} + +/// Returns the stencil offset for neighbor `i` as a `UVec3` (3D). +#[cfg(feature = "dim3")] +#[inline(always)] +pub fn nbh_shift(i: usize) -> UVec3 { + use khal_std::index::MaybeIndexUnchecked; + UVec3::from_array(NBH_SHIFTS.read(i)) +} + /// Flattens the 2D/3D stencil offset of neighbor `i` into a workgroup /// shared-memory index. #[cfg(feature = "dim2")] diff --git a/src_mpm_shaders/solver/g2p.rs b/src_mpm_shaders/solver/g2p.rs index ab927de5..678fbc6b 100644 --- a/src_mpm_shaders/solver/g2p.rs +++ b/src_mpm_shaders/solver/g2p.rs @@ -191,7 +191,7 @@ fn particle_g2p( for i in 0..27 { // For loop unrolling, use the fixed bound (the maximum one between 2D and 3D). if i < NBH_LEN { - let shift = NBH_SHIFTS.read(i); + let shift = nbh_shift(i); let packed_shift = NBH_SHIFT_SHARED.read(i); let shared_id = (packed_cell_index_in_block + packed_shift) as usize; let mut cell_vel = shared_nodes_vel.read(shared_id); diff --git a/src_mpm_shaders/solver/g2p_cdf.rs b/src_mpm_shaders/solver/g2p_cdf.rs index abb2c79a..ac07c1d1 100644 --- a/src_mpm_shaders/solver/g2p_cdf.rs +++ b/src_mpm_shaders/solver/g2p_cdf.rs @@ -194,7 +194,7 @@ fn particle_g2p( for i in 0..27 { // For loop unrolling, use the fixed bound (the maximum one between 2D and 3D). if i < NBH_LEN { - let shift = NBH_SHIFTS.read(i); + let shift = nbh_shift(i); let packed_shift = NBH_SHIFT_SHARED.read(i); let cell_data = shared_nodes[(packed_cell_index_in_block + packed_shift) as usize]; particle_affinity.set_unsigned_bits(cell_data.affinities); @@ -253,7 +253,7 @@ fn particle_g2p( for i in 0..27 { // For loop unrolling, use the fixed bound (the maximum one between 2D and 3D). if i < NBH_LEN { - let shift = NBH_SHIFTS.read(i); + let shift = nbh_shift(i); let packed_shift = NBH_SHIFT_SHARED.read(i); let cell_data = shared_nodes[(packed_cell_index_in_block + packed_shift) as usize]; diff --git a/src_viewer/viewer.rs b/src_viewer/viewer.rs index c07b8d54..a26a112a 100644 --- a/src_viewer/viewer.rs +++ b/src_viewer/viewer.rs @@ -468,12 +468,24 @@ impl NexusViewer { #[cfg(feature = "cuda")] fn init_cuda(&mut self) -> Option { match khal::backend::cuda::Cuda::new(0) { - Ok(cuda) => Some(KhalGpuBackend::Cuda(cuda)), + Ok(cuda) => { + // Also report on stderr: the UI banner is invisible to + // scripted/headless runs (`--cuda --run`). + match cuda.compute_capability() { + Ok((maj, min)) => { + eprintln!("[nexus] backend = native CUDA (sm_{maj}{min})") + } + Err(_) => eprintln!("[nexus] backend = native CUDA"), + } + Some(KhalGpuBackend::Cuda(cuda)) + } Err(e) => { - self.ui.gpu_init_error = Some(format!( + let msg = format!( "CUDA backend not available, initialization failed:\n\"{:?}\"\n", e - )); + ); + eprintln!("[nexus] {msg}"); + self.ui.gpu_init_error = Some(msg); None } } @@ -482,12 +494,17 @@ impl NexusViewer { #[cfg(feature = "metal")] fn init_metal(&mut self) -> Option { match khal::backend::metal::Metal::new() { - Ok(metal) => Some(KhalGpuBackend::Metal(metal)), + Ok(metal) => { + eprintln!("[nexus] backend = native Metal"); + Some(KhalGpuBackend::Metal(metal)) + } Err(e) => { - self.ui.gpu_init_error = Some(format!( + let msg = format!( "Metal backend not available, initialization failed:\n\"{:?}\"\n", e - )); + ); + eprintln!("[nexus] {msg}"); + self.ui.gpu_init_error = Some(msg); None } } From 69aeef2140a4154d37bfd27e3337c304a4393957 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?S=C3=A9bastien=20Crozet?= Date: Sun, 27 Sep 2026 14:37:14 +0200 Subject: [PATCH 2/4] feat: add support for compute graphs (cuda-graphs only for now) --- crates/examples2d/all_examples2.rs | 6 + crates/examples3d/all_examples3.rs | 6 + crates/nexus_python3d/src/nexus.rs | 6 + src/pipeline.rs | 156 +++++++++++++++++- src/state.rs | 17 ++ src_mpm/grid/grid.rs | 11 ++ src_mpm/pipeline.rs | 51 +++++- src_rbd/broad_phase/lbvh.rs | 11 ++ src_rbd/dynamics/multibody/env_reset.rs | 2 +- src_rbd/pipeline/insertion_removal.rs | 2 + src_rbd/pipeline/mod.rs | 2 +- src_rbd/pipeline/rbd_state.rs | 96 ++++++++++- src_rbd/pipeline/rbd_state_from_rapier.rs | 2 + src_rbd/pipeline/rbd_step.rs | 53 +++--- src_rbd/utils/compute_graph.rs | 86 ++++++++++ src_rbd/utils/mod.rs | 2 + .../dynamics/multibody/scatter_motor.rs | 2 +- src_viewer/ui.rs | 10 ++ src_viewer/viewer.rs | 13 ++ 19 files changed, 501 insertions(+), 33 deletions(-) create mode 100644 src_rbd/utils/compute_graph.rs diff --git a/crates/examples2d/all_examples2.rs b/crates/examples2d/all_examples2.rs index 325e619d..98a1b149 100644 --- a/crates/examples2d/all_examples2.rs +++ b/crates/examples2d/all_examples2.rs @@ -81,6 +81,7 @@ struct CliOptions { list: bool, cpu: bool, cuda: bool, + compute_graphs: bool, metal: bool, run: bool, } @@ -92,6 +93,7 @@ fn parse_command_line() -> CliOptions { list: false, cpu: false, cuda: false, + compute_graphs: false, metal: false, run: false, }; @@ -102,6 +104,7 @@ fn parse_command_line() -> CliOptions { "--list" => opts.list = true, "--cpu" => opts.cpu = true, "--cuda" => opts.cuda = true, + "--compute-graphs" => opts.compute_graphs = true, "--metal" => opts.metal = true, "--run" => opts.run = true, _ => {} @@ -146,6 +149,9 @@ pub async fn main() { if opts.cuda { viewer = viewer.with_backend(nexus_viewer2d::BackendType::Cuda); } + if opts.compute_graphs { + viewer = viewer.with_compute_graphs(true); + } #[cfg(feature = "metal")] if opts.metal { viewer = viewer.with_backend(nexus_viewer2d::BackendType::Metal); diff --git a/crates/examples3d/all_examples3.rs b/crates/examples3d/all_examples3.rs index b90582e6..5ca27e52 100644 --- a/crates/examples3d/all_examples3.rs +++ b/crates/examples3d/all_examples3.rs @@ -115,6 +115,7 @@ struct CliOptions { list: bool, cpu: bool, cuda: bool, + compute_graphs: bool, metal: bool, run: bool, } @@ -126,6 +127,7 @@ fn parse_command_line() -> CliOptions { list: false, cpu: false, cuda: false, + compute_graphs: false, metal: false, run: false, }; @@ -136,6 +138,7 @@ fn parse_command_line() -> CliOptions { "--list" => opts.list = true, "--cpu" => opts.cpu = true, "--cuda" => opts.cuda = true, + "--compute-graphs" => opts.compute_graphs = true, "--metal" => opts.metal = true, "--run" => opts.run = true, _ => {} @@ -182,6 +185,9 @@ pub async fn main() { if opts.cuda { viewer = viewer.with_backend(nexus_viewer3d::BackendType::Cuda); } + if opts.compute_graphs { + viewer = viewer.with_compute_graphs(true); + } #[cfg(feature = "metal")] if opts.metal { viewer = viewer.with_backend(nexus_viewer3d::BackendType::Metal); diff --git a/crates/nexus_python3d/src/nexus.rs b/crates/nexus_python3d/src/nexus.rs index b0bb9e60..caa55932 100644 --- a/crates/nexus_python3d/src/nexus.rs +++ b/crates/nexus_python3d/src/nexus.rs @@ -472,6 +472,12 @@ impl NexusState { self.0.set_rbd_steps_per_frame(steps); } + /// Replays each frame's physics through a compute graph on backends that + /// support graph capture (ignored elsewhere). + fn set_compute_graphs_enabled(&mut self, enabled: bool) { + self.0.set_compute_graphs_enabled(enabled); + } + fn set_rbd_gravity(&mut self, viewer: PyRef, gravity: Vec3) { self.0 .set_rbd_gravity(viewer.backend(), [gravity.0.x, gravity.0.y, gravity.0.z]); diff --git a/src/pipeline.rs b/src/pipeline.rs index e79b5424..a880f1c0 100644 --- a/src/pipeline.rs +++ b/src/pipeline.rs @@ -1,8 +1,9 @@ #[cfg(feature = "mpm")] use crate::mpm::pipeline::MpmPipeline; use crate::rbd::pipeline::RbdPipeline; +use crate::rbd::utils::ComputeGraphPlan; use crate::state::NexusState; -use khal::backend::{GpuBackend, GpuBackendError, GpuTimestamps}; +use khal::backend::{Backend, Encoder, GpuBackend, GpuBackendError, GpuGraph, GpuTimestamps}; bitflags::bitflags! { /// A bit mask identifying nexus pipelines. @@ -36,6 +37,22 @@ impl NexusPipeline { Ok(()) } + /// Launches a captured graph inside a timed pass so the profiler shows the + /// replayed frame as one entry. + fn launch_graph( + backend: &GpuBackend, + graph: &GpuGraph, + label: &str, + timestamps: Option<&mut GpuTimestamps>, + ) -> Result<(), GpuBackendError> { + let mut encoder = backend.begin_encoding(); + { + let _pass = encoder.begin_pass(label, timestamps); + graph.launch(backend)?; + } + backend.submit(encoder) + } + /// Advances the physics simulation by one GPU timestep. /// /// The compute pipelines are compiled lazily the first time their @@ -45,6 +62,11 @@ impl NexusPipeline { /// In addition, resources are loaded lazily on the GPU, so the first step /// after inserting/removing entities can be slower too. Call `Self::finalize` /// to pay that cost upfront. + /// + /// With [`NexusState::compute_graphs`] set, on a backend that supports + /// graph capture, the rigid-body steps and the MPM substeps of a frame are + /// recorded once into compute graphs and replayed on the following frames + /// until the scene's structure changes. pub async fn simulate( &mut self, backend: &GpuBackend, @@ -63,14 +85,72 @@ impl NexusPipeline { // apart instead of stalling the step on a blocking readback. let mut timestamps = timestamps.filter(|ts| ts.is_idle()); + let use_graphs = state.compute_graphs && backend.supports_graphs(); + // Rigid-bodies. `auto_resize_buffers` grows the collision-pair / coloring // buffers when the previous step overflowed them. if let Some(rbd) = state.rbd.as_mut() { self.preload_pipelines(backend, NexusPipelineMask::RBD)?; let pipeline = self.rbd_pipeline.as_mut().unwrap_or_else(|| unreachable!()); let steps = state.rbd_steps_per_frame.max(1); - for _ in 0..steps { - state.run_stats = pipeline.step(backend, rbd, timestamps.as_deref_mut())?; + let plan = if use_graphs { + let key = rbd.graph_key(steps); + rbd.compute_graph.plan(key) + } else { + rbd.compute_graph.invalidate(); + ComputeGraphPlan::Direct + }; + match plan { + ComputeGraphPlan::Replay => { + let graph = rbd.compute_graph.graph().unwrap_or_else(|| unreachable!()); + Self::launch_graph( + backend, + graph, + "[RBD] compute graph", + timestamps.as_deref_mut(), + )?; + } + ComputeGraphPlan::Capture => { + // Recording: nothing executes, and pass timestamps would + // record into the graph, so run without them. + backend.begin_capture()?; + let mut stats = Ok(state.run_stats.clone()); + for _ in 0..steps { + stats = pipeline.step(backend, rbd, None); + if stats.is_err() { + break; + } + } + match (stats, backend.end_capture()) { + (Ok(stats), Ok(graph)) => { + state.run_stats = stats; + Self::launch_graph( + backend, + &graph, + "[RBD] compute graph", + timestamps.as_deref_mut(), + )?; + rbd.compute_graph.captured(graph); + } + (Err(e), _) | (_, Err(e)) => { + eprintln!( + "[nexus] compute graph capture of the rigid-body step failed ({e}); \ + running it directly for this state" + ); + rbd.compute_graph.capture_failed(); + // The recorded (unexecuted) steps must still happen. + for _ in 0..steps { + state.run_stats = + pipeline.step(backend, rbd, timestamps.as_deref_mut())?; + } + } + } + } + ComputeGraphPlan::Direct => { + for _ in 0..steps { + state.run_stats = pipeline.step(backend, rbd, timestamps.as_deref_mut())?; + } + } } pipeline.auto_resize_buffers(backend, rbd)?; } @@ -85,8 +165,74 @@ impl NexusPipeline { // Upload the per-substep dt once, then run the substep loop. let substeps = state.mpm_substeps.max(1); let _ = mpm.write_substep_params(backend, substeps); - for _ in 0..substeps { - let _ = pipeline.step(backend, mpm, timestamps.as_deref_mut()); + let parity = mpm.grid.parity() as usize; + let plan = if use_graphs { + let key = mpm.graph_key(substeps); + mpm.compute_graphs[parity].plan(key) + } else { + for cache in &mut mpm.compute_graphs { + cache.invalidate(); + } + ComputeGraphPlan::Direct + }; + match plan { + ComputeGraphPlan::Replay => { + let graph = mpm.compute_graphs[parity] + .graph() + .unwrap_or_else(|| unreachable!()); + Self::launch_graph( + backend, + graph, + "[MPM] compute graph", + timestamps.as_deref_mut(), + )?; + // Mirror the host-side buffer swaps the recorded substeps did. + for _ in 0..substeps % 2 { + mpm.grid.swap_buffers(); + } + } + ComputeGraphPlan::Capture => { + backend.begin_capture()?; + let mut result = Ok(()); + for _ in 0..substeps { + result = pipeline.step(backend, mpm, None); + if result.is_err() { + break; + } + } + match (result, backend.end_capture()) { + (Ok(()), Ok(graph)) => { + Self::launch_graph( + backend, + &graph, + "[MPM] compute graph", + timestamps.as_deref_mut(), + )?; + mpm.compute_graphs[parity].captured(graph); + } + (Err(e), _) | (_, Err(e)) => { + eprintln!( + "[nexus] compute graph capture of the MPM substeps failed ({e}); \ + running them directly for this state" + ); + mpm.compute_graphs[parity].capture_failed(); + // The recorded substeps did not execute: undo the + // host-side buffer swaps they performed, then run + // them for real. + for _ in 0..substeps % 2 { + mpm.grid.swap_buffers(); + } + for _ in 0..substeps { + let _ = pipeline.step(backend, mpm, timestamps.as_deref_mut()); + } + } + } + } + ComputeGraphPlan::Direct => { + for _ in 0..substeps { + let _ = pipeline.step(backend, mpm, timestamps.as_deref_mut()); + } + } } } diff --git a/src/state.rs b/src/state.rs index 9a0976d2..ca57904d 100644 --- a/src/state.rs +++ b/src/state.rs @@ -180,6 +180,11 @@ pub struct NexusState { rbd_dirty: bool, /// Number of rigid-body solver steps advanced per [`NexusPipeline::simulate`](crate::pipeline::NexusPipeline::simulate) call. pub rbd_steps_per_frame: u32, + /// Replay each frame's physics through a compute graph instead of + /// re-encoding it, on backends that support graph capture (ignored + /// elsewhere). The graphs are re-recorded whenever the scene's structure + /// changes; see `RbdState::compute_graph` and `MpmState::compute_graphs`. + pub compute_graphs: bool, /// Per-environment GPU collider-slot reservation. When > 0, the GPU /// [`RbdState`] is built with this many slots (rather than exactly the /// current body count), leaving room for [`Self::add_rigid_body`] to append @@ -209,6 +214,7 @@ impl NexusState { rbd_sim_params: vec![RbdSimParams::tgs_soft()], rbd_dirty: false, rbd_steps_per_frame: 1, + compute_graphs: false, rbd_reserve_per_env: 0, rbd2gpu: vec![Coarena::new()], #[cfg(feature = "mpm")] @@ -367,6 +373,17 @@ impl NexusState { self.rbd_steps_per_frame } + /// Replays each frame's physics through a compute graph (see + /// [`Self::compute_graphs`]). + pub fn set_compute_graphs_enabled(&mut self, enabled: bool) { + self.compute_graphs = enabled; + } + + /// Whether each frame's physics is replayed through a compute graph. + pub fn compute_graphs_enabled(&self) -> bool { + self.compute_graphs + } + /// Current entity counts (rigid bodies, colliders, joints, multibody DOFs, /// particles) for display in the UI. Rigid-body /// counts are summed across all environments. diff --git a/src_mpm/grid/grid.rs b/src_mpm/grid/grid.rs index e5cad982..db88009f 100644 --- a/src_mpm/grid/grid.rs +++ b/src_mpm/grid/grid.rs @@ -433,6 +433,9 @@ pub struct GpuGrid { pub hmap_entries: Tensor, /// Pong buffer for hmap entries. pub prev_hmap_entries: Tensor, + /// Flipped by every [`Self::swap_buffers`]: which of the two meta / + /// hash-map buffer assignments is the current one. + parity: bool, /// Grid node data (momentum, mass, CDF). pub nodes: Tensor, /// Active block headers tracking particle ranges. @@ -525,6 +528,7 @@ impl GpuGrid { let debug = Tensor::vector(backend, [0u32, 0], BufferUsages::STORAGE)?; Ok(Self { + parity: false, cpu_meta, meta, prev_meta, @@ -543,5 +547,12 @@ impl GpuGrid { pub fn swap_buffers(&mut self) { std::mem::swap(&mut self.meta, &mut self.prev_meta); std::mem::swap(&mut self.prev_hmap_entries, &mut self.hmap_entries); + self.parity = !self.parity; + } + + /// Which of the two meta / hash-map buffer assignments is current; + /// flipped by every [`Self::swap_buffers`]. + pub fn parity(&self) -> bool { + self.parity } } diff --git a/src_mpm/pipeline.rs b/src_mpm/pipeline.rs index ba199b31..952d1dbe 100644 --- a/src_mpm/pipeline.rs +++ b/src_mpm/pipeline.rs @@ -15,7 +15,7 @@ use khal::backend::{Backend, Encoder, GpuBackend, GpuBackendError, GpuTimestamps use khal::{BufferUsages, Shader}; use nexus_rbd::dynamics::GpuBodySet; use nexus_rbd::math::{Pose, Vector}; -use nexus_rbd::utils::{GpuPrefixSum, PrefixSumWorkspace}; +use nexus_rbd::utils::{ComputeGraphCache, GpuPrefixSum, PrefixSumWorkspace}; use vortx::tensor::Tensor; use nexus_rbd::dynamics::body::{BodyCoupling, RapierBodyCouplingEntry}; @@ -94,9 +94,56 @@ pub struct MpmState { pub timestep_bounds_staging: Tensor, prefix_sum: PrefixSumWorkspace, coupling: Vec, + /// Cached compute graphs of this state's substep loop, keyed by + /// [`Self::graph_key`], one per grid double-buffer parity (see + /// [`GpuGrid::parity`]): each substep swaps the grid's current/previous + /// buffers on the host, so a recorded loop is only valid when the buffers + /// are in the assignment they had at capture time. With an odd substep + /// count the parity flips every frame and the two slots alternate. Driven + /// by `NexusPipeline::simulate`. + pub compute_graphs: [ComputeGraphCache; 2], +} + +/// Everything that shapes the GPU work recorded by one substep loop +/// ([`MpmPipeline::step`] run `substeps` times): the live particle / +/// rigid-particle / body counts (they size the dispatches), the grid +/// capacity and the substep count. A cached compute graph of the loop is +/// valid exactly as long as this key is unchanged. +#[derive(Clone, PartialEq, Eq, Hash, Debug)] +pub struct MpmGraphKey { + /// Number of substeps recorded per frame. + pub substeps: u32, + /// Whether CPIC rigid coupling is enabled. + pub use_cpic: bool, + /// Number of particles. + pub num_particles: usize, + /// Number of rigid particles sampled on the coupled colliders. + pub num_rigid_particles: u64, + /// Number of coupled rigid bodies. + pub num_coupled_bodies: usize, + /// Capacity of the grid's block hash map. + pub hmap_capacity: u32, + /// Length of the grid's hash-map entry buffer. + pub num_hmap_entries: u64, + /// Length of the coupled-body slot buffer. + pub num_rbd_body_slots: u64, } impl MpmState { + /// The [`MpmGraphKey`] of this state for `substeps` substeps per frame. + pub fn graph_key(&self, substeps: u32) -> MpmGraphKey { + MpmGraphKey { + substeps, + use_cpic: self.use_cpic, + num_particles: self.particles.len(), + num_rigid_particles: self.rigid_particles.len(), + num_coupled_bodies: self.coupling.len(), + hmap_capacity: self.grid.cpu_meta.hmap_capacity, + num_hmap_entries: self.grid.hmap_entries.len(), + num_rbd_body_slots: self.rbd_body_slots.len(), + } + } + /// Creates an empty MPM state with no particles and no coupled bodies. /// /// The grid is preallocated to hold `grid_capacity` cells. Physical @@ -159,6 +206,7 @@ impl MpmState { rbd_body_slots: Tensor::vector(backend, [], BufferUsages::STORAGE)?, timestep_bounds, timestep_bounds_staging, + compute_graphs: Default::default(), prefix_sum, coupling: Vec::new(), }) @@ -349,6 +397,7 @@ impl MpmState { // Standalone MPM: no rigid-body pipeline to write poses back to. rbd_body_slots: Tensor::vector(backend, [], BufferUsages::STORAGE)?, coupling, + compute_graphs: Default::default(), timestep_bounds, timestep_bounds_staging, base_dt: params.dt, diff --git a/src_rbd/broad_phase/lbvh.rs b/src_rbd/broad_phase/lbvh.rs index 80fb69d9..97eed52f 100644 --- a/src_rbd/broad_phase/lbvh.rs +++ b/src_rbd/broad_phase/lbvh.rs @@ -53,6 +53,8 @@ pub struct LbvhState { /// a resize re-seeds `n_sort` with the capacity). Avoids rewriting `n_sort` /// every frame when the live collider count hasn't changed. n_sort_active: Option<(u32, u32)>, + /// Bumped whenever a GPU buffer is (re)allocated; part of the compute-graph key. + pub(crate) generation: u64, unsorted_morton_keys: Tensor, sorted_morton_keys: Tensor, unsorted_colliders: Tensor, @@ -72,6 +74,12 @@ pub struct Lbvh { } impl LbvhState { + /// `(active colliders per batch, batches)` the radix sort was last told + /// about; part of the compute-graph key (a change re-uploads `n_sort`). + pub(crate) fn n_sort_active(&self) -> Option<(u32, u32)> { + self.n_sort_active + } + /// Creates a new LBVH state with default buffer usage flags. pub fn new(backend: &GpuBackend) -> Self { Self::with_usages(backend, BufferUsages::STORAGE) @@ -82,6 +90,7 @@ impl LbvhState { Self { n_sort: Tensor::scalar(backend, 0, usages).unwrap(), n_sort_active: None, + generation: 0, domain_aabb: Tensor::scalar_uninit(backend, usages).unwrap(), unsorted_morton_keys: Tensor::vector_uninit(backend, 0, usages).unwrap(), sorted_morton_keys: Tensor::vector_uninit(backend, 0, usages).unwrap(), @@ -104,12 +113,14 @@ impl LbvhState { fn resize_buffers(&mut self, backend: &GpuBackend, colliders_len: u32, num_batches: u32) { if (self.domain_aabb.len() as u32) < num_batches { + self.generation += 1; self.domain_aabb = Tensor::vector_uninit(backend, num_batches, self.buffer_usages).unwrap(); } // NOTE: colliders_len is the total colliders count, already taking all batches into account. if (self.tree.len() as u32) < 2 * colliders_len { + self.generation += 1; self.unsorted_morton_keys = Tensor::vector_uninit(backend, colliders_len, self.buffer_usages).unwrap(); self.sorted_morton_keys = diff --git a/src_rbd/dynamics/multibody/env_reset.rs b/src_rbd/dynamics/multibody/env_reset.rs index a2abef99..af0f7ebd 100644 --- a/src_rbd/dynamics/multibody/env_reset.rs +++ b/src_rbd/dynamics/multibody/env_reset.rs @@ -10,7 +10,7 @@ //! offset applied in-kernel. //! //! Reset loops should prefer the second: it is what keeps a rollout free of -//! per-step host writes, and therefore capturable into a CUDA graph. +//! per-step host writes, and therefore capturable into a compute graph. use super::multibody_set::GpuMultibodySet; use crate::math::Vector; diff --git a/src_rbd/pipeline/insertion_removal.rs b/src_rbd/pipeline/insertion_removal.rs index 098d6545..80b3d2b9 100644 --- a/src_rbd/pipeline/insertion_removal.rs +++ b/src_rbd/pipeline/insertion_removal.rs @@ -276,6 +276,8 @@ impl RbdState { contacts_capacity_cpu, collision_pairs_capacity_cpu, collision_pairs_len_cpu: 0, + graph_generation: 0, + compute_graph: Default::default(), #[cfg(feature = "dim3")] mb_cons_demand_cpu: 0, batch_indices, diff --git a/src_rbd/pipeline/mod.rs b/src_rbd/pipeline/mod.rs index cf08ce98..3aeb3406 100644 --- a/src_rbd/pipeline/mod.rs +++ b/src_rbd/pipeline/mod.rs @@ -16,5 +16,5 @@ mod test_batched_stacks; #[cfg(feature = "dim3")] pub use rbd_state::RbdSnapshot; -pub use rbd_state::{RbdCapacities, RbdResizePolicy, RbdState, RunStats}; +pub use rbd_state::{RbdCapacities, RbdGraphKey, RbdResizePolicy, RbdState, RunStats}; pub use rbd_step::{FORCE_FUSED_SWEEPS, RbdPipeline}; diff --git a/src_rbd/pipeline/rbd_state.rs b/src_rbd/pipeline/rbd_state.rs index e4bf3ab5..05cc8235 100644 --- a/src_rbd/pipeline/rbd_state.rs +++ b/src_rbd/pipeline/rbd_state.rs @@ -17,7 +17,7 @@ use crate::shaders::dynamics::{ }; use crate::shaders::shapes::Shape; use crate::shaders::utils::BatchIndices; -use crate::utils::PrefixSumWorkspace; +use crate::utils::{ComputeGraphCache, PrefixSumWorkspace}; use khal::BufferUsages; use khal::backend::{Backend, GpuBackend, GpuReadback}; @@ -191,6 +191,12 @@ pub struct RbdState { /// batches, harvested by the non-blocking readback in [`RbdPipeline::auto_resize_buffers`](crate::pipeline::RbdPipeline::auto_resize_buffers). /// Surfaced in the viewer UI; lags the GPU by a frame or two like the resize. pub(super) collision_pairs_len_cpu: u32, + /// Bumped on every GPU buffer (re)allocation or capacity change; part of + /// [`Self::graph_key`]. + pub(super) graph_generation: u64, + /// Cached compute graph of this state's step loop, keyed by + /// [`Self::graph_key`]. Driven by `NexusPipeline::simulate`. + pub compute_graph: ComputeGraphCache, /// CPU mirror of the multibody contact-constraint slot demand, refreshed /// by the same (asynchronous) readback as `collision_pairs_len_cpu`. #[cfg(feature = "dim3")] @@ -286,6 +292,7 @@ impl RbdState { /// `GpuMultibodySet::set_impulse_joints`). Call whenever a cap edit /// happens that any kernel reads via its `batch_ids` uniform. pub(super) fn rebuild_batch_indices(&mut self, backend: &GpuBackend) { + self.graph_generation += 1; #[allow(unused_mut)] // Only mutated with the dim3 (multibody) feature. let mut bi = BatchIndices { num_batches: self.num_batches, @@ -319,6 +326,9 @@ impl RbdState { /// Grows [`Self::color_uniforms`] so indices `0..n` are available. pub(super) fn ensure_color_uniforms(&mut self, backend: &GpuBackend, n: u32) { + if (self.color_uniforms.len() as u32) < n { + self.graph_generation += 1; + } for c in self.color_uniforms.len() as u32..n { self.color_uniforms .push(Tensor::scalar(backend, c, BufferUsages::UNIFORM).unwrap()); @@ -888,3 +898,87 @@ impl RbdState { self.reset_templates_bodies = Some(tpl); } } + +/// Everything that shapes the GPU work recorded by one +/// [`RbdPipeline::step`](crate::pipeline::RbdPipeline::step) run +/// `steps_per_frame` times: buffer identities (any reallocation bumps a +/// generation), the live counts and capacities that size dispatches or +/// host-side loops, and the solver path taken. A cached compute graph of the +/// step is valid exactly as long as this key is unchanged. +#[derive(Clone, PartialEq, Eq, Hash, Debug)] +pub struct RbdGraphKey { + /// Bumped by every (re)allocation or capacity change of the state's buffers. + pub generation: u64, + /// Bumped by every (re)allocation of the LBVH buffers. + pub lbvh_generation: u64, + /// `(active colliders per batch, batches)` the radix sort was last told about. + pub lbvh_n_sort_active: Option<(u32, u32)>, + /// Number of steps recorded per frame. + pub steps_per_frame: u32, + /// Number of simulation environments. + pub num_batches: u32, + /// Number of active colliders over all environments. + pub num_active_colliders: u32, + /// Collider slots per environment. + pub num_colliders_per_batch: u32, + /// Maximum number of graph colors the solver sweeps. + pub max_colors: u32, + /// Number of per-color uniform buffers. + pub num_color_uniforms: usize, + /// Capacity of the contact buffer. + pub contacts_capacity: u32, + /// Capacity of the collision-pair buffer. + pub collision_pairs_capacity: u32, + /// Solver iterations per step. + pub num_solver_iterations: u32, + /// Whether rigid-body contacts are skipped by the solver. + pub rb_contacts_inert: bool, + /// Number of graph colors of the impulse joints. + pub joints_num_colors: u32, + /// Whether there are no impulse joints. + pub joints_empty: bool, + /// Capacity of the multibody contact-constraint slabs. + #[cfg(feature = "dim3")] + pub mb_contact_constraints_capacity: u32, + /// Whether there are no multibodies. + #[cfg(feature = "dim3")] + pub multibodies_empty: bool, + /// Number of graph colors of the multibody impulse joints. + #[cfg(feature = "dim3")] + pub mb_imp_joint_num_colors: u32, + /// Whether the fused colored-sweep kernels are used. + pub fused_color_sweeps: bool, +} + +impl RbdState { + /// The [`RbdGraphKey`] of this state for `steps_per_frame` steps per frame. + pub fn graph_key(&self, steps_per_frame: u32) -> RbdGraphKey { + RbdGraphKey { + generation: self.graph_generation, + lbvh_generation: self.lbvh.generation, + lbvh_n_sort_active: self.lbvh.n_sort_active(), + steps_per_frame, + num_batches: self.num_batches(), + num_active_colliders: self.num_active_colliders(), + num_colliders_per_batch: self.num_colliders_per_batch(), + max_colors: self.max_colors, + num_color_uniforms: self.color_uniforms.len(), + contacts_capacity: self.contacts_capacity_cpu, + collision_pairs_capacity: self.collision_pairs_capacity_cpu, + num_solver_iterations: self.num_solver_iterations, + rb_contacts_inert: self.rb_contacts_inert(), + joints_num_colors: self.joints.num_colors(), + joints_empty: self.joints.is_empty(), + // NOTE: not `mb_cons_demand_cpu`: it is a readback mirror that + // changes as the scene evolves; the resize it may trigger bumps + // `graph_generation`, which is what the graph depends on. + #[cfg(feature = "dim3")] + mb_contact_constraints_capacity: self.mb_contact_constraints_capacity(), + #[cfg(feature = "dim3")] + multibodies_empty: self.multibodies.is_empty(), + #[cfg(feature = "dim3")] + mb_imp_joint_num_colors: self.multibodies.mb_imp_joint_num_colors(), + fused_color_sweeps: crate::pipeline::RbdPipeline::fused_color_sweeps(self), + } + } +} diff --git a/src_rbd/pipeline/rbd_state_from_rapier.rs b/src_rbd/pipeline/rbd_state_from_rapier.rs index a758c053..b5b03916 100644 --- a/src_rbd/pipeline/rbd_state_from_rapier.rs +++ b/src_rbd/pipeline/rbd_state_from_rapier.rs @@ -840,6 +840,8 @@ impl RbdState { contacts_capacity_cpu, collision_pairs_capacity_cpu, collision_pairs_len_cpu: 0, + graph_generation: 0, + compute_graph: Default::default(), #[cfg(feature = "dim3")] mb_cons_demand_cpu: 0, batch_indices, diff --git a/src_rbd/pipeline/rbd_step.rs b/src_rbd/pipeline/rbd_step.rs index 666623ae..49088822 100644 --- a/src_rbd/pipeline/rbd_step.rs +++ b/src_rbd/pipeline/rbd_step.rs @@ -96,6 +96,34 @@ impl RbdPipeline { self.step_impl(backend, state, timestamps, encoder, false) } + /// Whether the step uses the fused colored-sweep kernels (one workgroup per + /// batch walking every color) instead of one dispatch per color. + /// + /// Chosen from the expected pair count: small pair counts with many + /// environments benefit from the fused kernels. The fused path can also be + /// forced regardless of size (an A/B knob: an env var natively, + /// [`FORCE_FUSED_SWEEPS`] on wasm where env vars do not exist). It is + /// correct at any size, just serialized past ~64 lanes, which may still win + /// where per-dispatch latency rules, i.e. small batch counts in the browser. + /// + /// The estimate is per batch: the read-back counter (or the capacity when + /// the readback is disabled) is a total over the whole flat pair buffer. + pub fn fused_color_sweeps(state: &RbdState) -> bool { + let readback_enabled = state.capacities.solver_colors_resize_policy + != RbdResizePolicy::Fixed + || state.capacities.collisions_resize_policy != RbdResizePolicy::Fixed; + let est_pairs = if readback_enabled { + state.collision_pairs_len_cpu.div_ceil(state.num_batches) + } else { + state + .collision_pairs_capacity_cpu + .div_ceil(state.num_batches) + }; + est_pairs <= 128 + || FORCE_FUSED_SWEEPS.load(core::sync::atomic::Ordering::Relaxed) + || std::env::var("NEXUS_FUSED_SWEEPS").as_deref() == Ok("1") + } + fn step_impl( &self, backend: &GpuBackend, @@ -273,29 +301,7 @@ impl RbdPipeline { } } - let readback_enabled = state.capacities.solver_colors_resize_policy - != RbdResizePolicy::Fixed - || state.capacities.collisions_resize_policy != RbdResizePolicy::Fixed; - // Estimated pairs per batch: the counter (and the capacity fallback) - // are totals over the whole flat pair buffer. - let est_pairs = if readback_enabled { - state.collision_pairs_len_cpu.div_ceil(state.num_batches) - } else { - state - .collision_pairs_capacity_cpu - .div_ceil(state.num_batches) - }; - - // Choose the kernel depending on the expected pairs count: small pair - // counts with many environments benefit from the fused kernels. The - // fused path can also be forced regardless of size (an A/B knob: an env - // var natively, [`FORCE_FUSED_SWEEPS`] on wasm where env vars do not - // exist). It is correct at any size, just serialized past ~64 lanes, - // which may still win where per-dispatch latency rules, i.e. small - // batch counts in the browser. - let fused_color_sweeps = est_pairs <= 128 - || FORCE_FUSED_SWEEPS.load(core::sync::atomic::Ordering::Relaxed) - || std::env::var("NEXUS_FUSED_SWEEPS").as_deref() == Ok("1"); + let fused_color_sweeps = Self::fused_color_sweeps(state); // In small scenes, submit less frequently. In big scenes submit more // to overlap compute and encoding. @@ -679,6 +685,7 @@ impl RbdPipeline { if grow_colors || resize_pairs || resize_mb { backend.synchronize()?; + state.graph_generation += 1; } if grow_colors { diff --git a/src_rbd/utils/compute_graph.rs b/src_rbd/utils/compute_graph.rs new file mode 100644 index 00000000..6c759a30 --- /dev/null +++ b/src_rbd/utils/compute_graph.rs @@ -0,0 +1,86 @@ +//! Caching of captured compute graphs. + +use khal::backend::GpuGraph; + +/// What a frame should do with the work described by a key, as decided by +/// [`ComputeGraphCache::plan`]. +#[derive(Copy, Clone, PartialEq, Eq, Debug)] +pub enum ComputeGraphPlan { + /// Run the work directly: the key changed since the previous frame (or + /// its capture failed). + Direct, + /// Record the work into a new graph, then launch that graph. + Capture, + /// Launch the cached graph. + Replay, +} + +/// A cached compute graph of some recorded GPU work, keyed by a value +/// describing everything that shapes that work: buffer identities, dispatch +/// sizes, host-side loop counts, solver paths. +/// +/// A graph is captured only once the same key has been seen on two +/// consecutive frames: one-time work (buffer growth, uploads of changed +/// counts) happens on the first frame with a new key and must not be +/// recorded. A failed capture is remembered and not retried until the key +/// changes. +pub struct ComputeGraphCache { + graph: Option<(K, GpuGraph)>, + /// Key of the last [`Self::plan`] call. + current: Option, + /// Key whose capture failed. + failed: Option, +} + +impl Default for ComputeGraphCache { + fn default() -> Self { + Self { + graph: None, + current: None, + failed: None, + } + } +} + +impl ComputeGraphCache { + /// Drops the cached graph and the key history, e.g. when graph replay is + /// switched off. + pub fn invalidate(&mut self) { + self.graph = None; + self.current = None; + } + + /// Decides what this frame should do for the work described by `key`. + pub fn plan(&mut self, key: K) -> ComputeGraphPlan { + let plan = match &self.graph { + Some((cached, _)) if *cached == key => ComputeGraphPlan::Replay, + _ if self.current.as_ref() == Some(&key) && self.failed.as_ref() != Some(&key) => { + ComputeGraphPlan::Capture + } + _ => ComputeGraphPlan::Direct, + }; + if plan != ComputeGraphPlan::Replay { + self.graph = None; + } + self.current = Some(key); + plan + } + + /// The cached graph, if [`Self::plan`] just returned + /// [`ComputeGraphPlan::Replay`]. + pub fn graph(&self) -> Option<&GpuGraph> { + self.graph.as_ref().map(|(_, graph)| graph) + } + + /// Stores the graph captured for the key of the last [`Self::plan`] call. + pub fn captured(&mut self, graph: GpuGraph) { + let key = self.current.clone().unwrap_or_else(|| unreachable!()); + self.graph = Some((key, graph)); + } + + /// Records that capturing the work of the last [`Self::plan`] call failed, + /// so it runs directly until its key changes. + pub fn capture_failed(&mut self) { + self.failed = self.current.clone(); + } +} diff --git a/src_rbd/utils/mod.rs b/src_rbd/utils/mod.rs index 01d5e830..a1104b3b 100644 --- a/src_rbd/utils/mod.rs +++ b/src_rbd/utils/mod.rs @@ -3,8 +3,10 @@ //! This module provides general-purpose GPU algorithms that support the collision //! detection and physics simulation pipelines. +pub use compute_graph::{ComputeGraphCache, ComputeGraphPlan}; pub use prefix_sum::{GpuPrefixSum, PrefixSumWorkspace}; pub use radix_sort::{RadixSort, RadixSortWorkspace}; +mod compute_graph; mod prefix_sum; mod radix_sort; diff --git a/src_rbd_shaders/dynamics/multibody/scatter_motor.rs b/src_rbd_shaders/dynamics/multibody/scatter_motor.rs index 16ef0da6..3ec99f22 100644 --- a/src_rbd_shaders/dynamics/multibody/scatter_motor.rs +++ b/src_rbd_shaders/dynamics/multibody/scatter_motor.rs @@ -2,7 +2,7 @@ //! straight into `links_static` on the GPU, replacing a host-side //! `set_motors` + whole-mirror upload every step. This is what lets an RL //! policy drive the motors without a host round-trip, and therefore what makes -//! a rollout capturable into a CUDA graph (no per-step host writes). +//! a rollout capturable into a compute graph (no per-step host writes). //! //! `links_static` is batch-interleaved: link `l` of env `e` lives at //! `l ยท num_envs + e`. Targets are row-major `[num_actuated x num_envs]`, diff --git a/src_viewer/ui.rs b/src_viewer/ui.rs index 9731ecce..76c6d954 100644 --- a/src_viewer/ui.rs +++ b/src_viewer/ui.rs @@ -448,6 +448,16 @@ fn backend_selector(ui: &mut egui::Ui, state: &mut UiState, gpu_available: bool) { new_backend = Some(BackendType::Cuda); } + #[cfg(feature = "cuda")] + if state.backend_type == BackendType::Cuda { + ui.indent("compute-graphs", |ui| { + ui.checkbox(&mut state.compute_graphs, "Compute graphs") + .on_hover_text( + "Record each frame's physics dispatches once into a compute graph and replay \ + it with a single launch (re-recorded whenever the scene changes).", + ); + }); + } #[cfg(feature = "metal")] if ui diff --git a/src_viewer/viewer.rs b/src_viewer/viewer.rs index a26a112a..32a35ebf 100644 --- a/src_viewer/viewer.rs +++ b/src_viewer/viewer.rs @@ -103,6 +103,8 @@ pub struct UiState { pub sync_time: Duration, pub ui_sections: UiSections, pub backend_type: BackendType, + /// Replay each frame's GPU work through a compute graph (backends with graph support only). + pub compute_graphs: bool, pub gpu_init_error: Option, /// Names + kinds of all registered demos, used to populate the demo picker. pub demos: Vec<(String, DemoKind)>, @@ -343,6 +345,7 @@ impl NexusViewer { show_performance: true, }, backend_type: BackendType::Gpu, + compute_graphs: false, gpu_init_error: None, demos, selected_demo: 0, @@ -367,6 +370,13 @@ impl NexusViewer { self } + /// Replay each frame's physics through a compute graph (backends with graph support only; + /// no effect elsewhere). Also toggled from the backend panel. + pub fn with_compute_graphs(mut self, enabled: bool) -> Self { + self.ui.compute_graphs = enabled; + self + } + pub fn with_cpu(mut self) -> Self { self.ui.backend_type = BackendType::Cpu; self @@ -967,6 +977,9 @@ impl NexusViewer { state.set_mpm_gravity(s.mpm_gravity); state.set_rbd_steps_per_frame(s.rbd_steps_per_frame); } + // The compute-graph toggle is a testbed-wide choice, not a scene + // setting: it always flows from the backend panel into the scene. + state.set_compute_graphs_enabled(self.ui.compute_graphs); if self.direct_render_path() { self.sync_without_readback(state).await?; From 681bf5fa61290f4be81d582944919c2e15aa643d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?S=C3=A9bastien=20Crozet?= Date: Wed, 30 Sep 2026 19:42:50 +0200 Subject: [PATCH 3/4] Use khal and vortx main from GitHub instead of local path patches --- Cargo.toml | 17 +++++++++-------- 1 file changed, 9 insertions(+), 8 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 174998b8..ef0117eb 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -93,14 +93,15 @@ rust.unexpected_cfgs = { level = "warn", check-cfg = [ ] } [patch.crates-io] -# TODO: drop once khal 0.3.1 (WebGPU SPIR-V passthrough entry-point fix) and the -# khal/vortx `cuda-oxide` feature are published. -khal = { path = "../khal/crates/khal" } -khal-derive = { path = "../khal/crates/khal-derive" } -khal-std = { path = "../khal/crates/khal-std" } -khal-builder = { path = "../khal/crates/khal-builder" } -vortx = { path = "../vortx" } -vortx-shaders = { path = "../vortx/vortx-shaders" } +# khal and vortx from GitHub: the `cuda-oxide` feature, the WebGPU SPIR-V +# passthrough entry-point fix, and compute-graph capture are not on crates.io +# yet. Drop once khal 0.3.1 / vortx 0.4.1 are published. +khal = { git = "https://github.com/dimforge/khal", branch = "main" } +khal-derive = { git = "https://github.com/dimforge/khal", branch = "main" } +khal-std = { git = "https://github.com/dimforge/khal", branch = "main" } +khal-builder = { git = "https://github.com/dimforge/khal", branch = "main" } +vortx = { git = "https://github.com/dimforge/vortx", branch = "main" } +vortx-shaders = { git = "https://github.com/dimforge/vortx", branch = "main" } #rapier2d = { path = "../rapier/crates/rapier2d" } #rapier3d = { path = "../rapier/crates/rapier3d" } #rapier3d-mjcf = { path = "../rapier/crates/rapier3d-mjcf" } From dd1e37b05855316944735af0675031610e336562 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?S=C3=A9bastien=20Crozet?= Date: Thu, 1 Oct 2026 10:16:23 +0200 Subject: [PATCH 4/4] =?UTF-8?q?chore:=C2=A0remove=20the=20FORCE=5FFUSED=5F?= =?UTF-8?q?SWEEPS=20static=20and=20NEXUS=5FFUSED=5FSWEEPS=20env=20override?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src_rbd/pipeline/mod.rs | 2 +- src_rbd/pipeline/rbd_step.rs | 15 ++------------- 2 files changed, 3 insertions(+), 14 deletions(-) diff --git a/src_rbd/pipeline/mod.rs b/src_rbd/pipeline/mod.rs index 3aeb3406..ce56e6e9 100644 --- a/src_rbd/pipeline/mod.rs +++ b/src_rbd/pipeline/mod.rs @@ -17,4 +17,4 @@ mod test_batched_stacks; #[cfg(feature = "dim3")] pub use rbd_state::RbdSnapshot; pub use rbd_state::{RbdCapacities, RbdGraphKey, RbdResizePolicy, RbdState, RunStats}; -pub use rbd_step::{FORCE_FUSED_SWEEPS, RbdPipeline}; +pub use rbd_step::RbdPipeline; diff --git a/src_rbd/pipeline/rbd_step.rs b/src_rbd/pipeline/rbd_step.rs index 49088822..c15c50d7 100644 --- a/src_rbd/pipeline/rbd_step.rs +++ b/src_rbd/pipeline/rbd_step.rs @@ -17,12 +17,6 @@ use khal::BufferUsages; use khal::backend::{Backend, Encoder, GpuBackend, GpuBackendError, GpuTimestamps}; use vortx::tensor::Tensor; -/// Forces the fused colored-sweep kernels regardless of the estimated pair -/// count: the programmatic twin of `NEXUS_FUSED_SWEEPS=1`, for targets without -/// environment variables (wasm). -pub static FORCE_FUSED_SWEEPS: core::sync::atomic::AtomicBool = - core::sync::atomic::AtomicBool::new(false); - /// The main GPU physics pipeline coordinating all simulation stages. pub struct RbdPipeline { mprops_update: GpuMpropsUpdate, @@ -100,11 +94,8 @@ impl RbdPipeline { /// batch walking every color) instead of one dispatch per color. /// /// Chosen from the expected pair count: small pair counts with many - /// environments benefit from the fused kernels. The fused path can also be - /// forced regardless of size (an A/B knob: an env var natively, - /// [`FORCE_FUSED_SWEEPS`] on wasm where env vars do not exist). It is - /// correct at any size, just serialized past ~64 lanes, which may still win - /// where per-dispatch latency rules, i.e. small batch counts in the browser. + /// environments benefit from the fused kernels. The fused path is correct + /// at any size, just serialized past ~64 lanes. /// /// The estimate is per batch: the read-back counter (or the capacity when /// the readback is disabled) is a total over the whole flat pair buffer. @@ -120,8 +111,6 @@ impl RbdPipeline { .div_ceil(state.num_batches) }; est_pairs <= 128 - || FORCE_FUSED_SWEEPS.load(core::sync::atomic::Ordering::Relaxed) - || std::env::var("NEXUS_FUSED_SWEEPS").as_deref() == Ok("1") } fn step_impl(