rtx-cfd embedded3 item 7: e3_step.cu + step::device::{DeviceStep, StepTimers} on the shared runtime (FMA off); gate 7 HELD: device = host ≤ 7e-12 under tight tolerances on MMS/Beltrami/Poiseuille, equal CG counts on every step, periodic planes within 2e-16; the default-tolerance differences are the projection's inner stop
CI / Format Check (push) Failing after 4s
CI / Build (ubuntu-latest) (push) Failing after 3s
CI / Clippy Check (push) Failing after 4s
CI / Build CPU-Only (Explicit) (push) Failing after 3s
Documentation / Build API Documentation (push) Failing after 2s
Documentation / Build User Guide (push) Successful in 4s
Performance Benchmarks / Run Benchmarks (push) Successful in 3m4s
CI / Build (macos-latest) (push) Canceled after 0s
CI / Test (macos-latest) (push) Canceled after 0s
CI / Test (ubuntu-latest) (push) Canceled after 0s
CI / Python Bindings (maturin) (macos-latest) (push) Canceled after 0s
CI / Python Bindings (maturin) (ubuntu-latest) (push) Canceled after 0s
CI / Distributed Training Tests (push) Canceled after 0s
CI / CI Success (push) Canceled after 0s
CI / WASM Build + Size Check (push) Canceled after 0s

Co-Authored-By: Claude Fable 5.1 <[email protected]>
This commit is contained in:
Omar Sobh
2026-09-17 14:56:46 -05:00
co-authored by Claude Fable 5.1
parent 8821e18520
commit 54911b4db3
4 changed files with 1555 additions and 0 deletions
@@ -0,0 +1,559 @@
//! embedded3 item 7: the PISO step device-resident (`e3_step.cu` + the device CG). The
//! fields live on the device; the host holds the solver's parameters and
//! functions, evaluates the boundary tables per step (kB) and the steady
//! momentum source once, and reads back scalars. `download` mirrors the
//! fields into a `Field` at instants. Compiled with FMA contraction
//! off so the predictors are the host's arithmetic to the bit.
use super::{Side, Solver, StepResult};
use crate::solvers::incompressible::embedded3::Grid;
use crate::solvers::incompressible::embedded3::field::Field;
use crate::solvers::incompressible::embedded3::poisson::device::{cfg, load_module, runtime};
use crate::solvers::incompressible::embedded3::poisson::device_cg::DeviceCg;
use crate::solvers::incompressible::poisson::MultigridParameters;
use crate::solvers::incompressible::simple::ConvectionScheme;
use cudarc::driver::{
CudaFunction, CudaModule, CudaSlice, DevicePtr, DeviceRepr, LaunchConfig, PushKernelArg,
ValidAsZeroBits,
};
use std::sync::{Arc, OnceLock};
use std::time::Instant;
const KERNELS: &str = include_str!("../../../../kernels/cuda/e3_step.cu");
struct Kernels {
_module: Arc<CudaModule>,
predict_u: CudaFunction,
predict_v: CudaFunction,
predict_w: CudaFunction,
sides_x: CudaFunction,
sides_y: CudaFunction,
sides_z: CudaFunction,
divergence: CudaFunction,
correct_u: CudaFunction,
correct_v: CudaFunction,
correct_w: CudaFunction,
add_p: CudaFunction,
reduce: CudaFunction,
}
static KERNELS_ONCE: OnceLock<Kernels> = OnceLock::new();
fn kernels() -> &'static Kernels {
KERNELS_ONCE.get_or_init(|| {
let module = load_module(KERNELS, "e3_step.cu", true);
let f = |name: &str| module.load_function(name).expect(name);
Kernels {
predict_u: f("e3_step_predict_u"),
predict_v: f("e3_step_predict_v"),
predict_w: f("e3_step_predict_w"),
sides_x: f("e3_step_sides_x"),
sides_y: f("e3_step_sides_y"),
sides_z: f("e3_step_sides_z"),
divergence: f("e3_step_divergence"),
correct_u: f("e3_step_correct_u"),
correct_v: f("e3_step_correct_v"),
correct_w: f("e3_step_correct_w"),
add_p: f("e3_step_add_p_and_imbalance"),
reduce: f("e3_step_reduce"),
_module: module,
}
})
}
/// `struct E3Params` in e3_step.cu.
#[repr(C)]
#[derive(Clone, Copy)]
struct E3Params {
nx: i32,
ny: i32,
nz: i32,
periodic_z: i32,
bx0: i32,
bx1: i32,
by0: i32,
by1: i32,
bz0: i32,
bz1: i32,
scheme: i32,
pad: i32,
dx: f64,
dy: f64,
dz: f64,
dt: f64,
rho: f64,
nu: f64,
}
unsafe impl DeviceRepr for E3Params {}
unsafe impl ValidAsZeroBits for E3Params {}
/// `struct E3Ptrs` in e3_step.cu: 33 device pointers.
#[repr(C)]
#[derive(Clone, Copy)]
struct E3Ptrs {
ptrs: [u64; 33],
}
unsafe impl DeviceRepr for E3Ptrs {}
unsafe impl ValidAsZeroBits for E3Ptrs {}
fn side_code(s: Side) -> i32 {
match s {
Side::Velocity => 0,
Side::SlipWall => 1,
Side::PressureOutlet => 2,
Side::Periodic => 3,
}
}
fn scheme_code(s: ConvectionScheme) -> i32 {
match s {
ConvectionScheme::Upwind => 0,
ConvectionScheme::TvdVanAlbada => 1,
ConvectionScheme::TvdVanLeer => 2,
}
}
/// Step timers (`RTX_PROFILE`): nanoseconds per phase and the step count.
#[derive(Debug, Clone, Copy, Default)]
pub struct StepTimers {
pub predictor_ns: u64,
pub poisson_ns: u64,
pub apply_ns: u64,
pub transfer_ns: u64,
pub steps: u64,
pub cg_iterations: u64,
}
pub struct DeviceStep {
pub solver: Solver,
grid: Grid,
u: CudaSlice<f64>,
v: CudaSlice<f64>,
w: CudaSlice<f64>,
p: CudaSlice<f64>,
u_old: CudaSlice<f64>,
v_old: CudaSlice<f64>,
w_old: CudaSlice<f64>,
u_star: CudaSlice<f64>,
v_star: CudaSlice<f64>,
w_star: CudaSlice<f64>,
p_prime: CudaSlice<f64>,
sp: CudaSlice<f64>,
su: CudaSlice<f64>,
sv: CudaSlice<f64>,
sw: CudaSlice<f64>,
tables: Vec<CudaSlice<f64>>,
partial: CudaSlice<f64>,
scalar: CudaSlice<f64>,
n_blocks: usize,
cg: Option<DeviceCg>,
cg_dt: f64,
timers: Option<StepTimers>,
initialized: bool,
}
impl DeviceStep {
/// Allocates the device fields for `grid`; the momentum source is
/// tabulated at `t = 0` (steady sources only in Stage 1).
pub fn new(solver: Solver, grid: Grid) -> Self {
let rt = runtime();
let nu = (grid.nx + 1) * grid.ny * grid.nz;
let nv = grid.nx * (grid.ny + 1) * grid.nz;
let nw = grid.nx * grid.ny * (grid.nz + 1);
let nc = grid.cells();
let zeros = |n: usize| rt.stream.alloc_zeros::<f64>(n).expect("alloc");
let (su, sv, sw) = solver.source_tables(grid, 0.0);
let up = |v: &Vec<f64>| -> CudaSlice<f64> {
rt.stream
.memcpy_stod(if v.is_empty() { &[0.0f64][..] } else { v })
.expect("upload")
};
let tables: Vec<CudaSlice<f64>> = solver
.boundary_tables(grid, 0.0)
.iter()
.flat_map(|side| side.iter().map(up).collect::<Vec<_>>())
.collect();
let n_blocks = nc.div_ceil(256).max(1);
let timers = std::env::var("RTX_PROFILE")
.is_ok()
.then(StepTimers::default);
Self {
solver,
grid,
u: zeros(nu),
v: zeros(nv),
w: zeros(nw),
p: zeros(nc),
u_old: zeros(nu),
v_old: zeros(nv),
w_old: zeros(nw),
u_star: zeros(nu),
v_star: zeros(nv),
w_star: zeros(nw),
p_prime: zeros(nc),
sp: zeros(nc),
su: up(&su),
sv: up(&sv),
sw: up(&sw),
tables,
partial: zeros(n_blocks),
scalar: zeros(1),
n_blocks,
cg: None,
cg_dt: 0.0,
timers,
initialized: false,
}
}
pub fn timers(&self) -> Option<StepTimers> {
self.timers
}
pub fn grid(&self) -> Grid {
self.grid
}
/// The host field onto the device (u, v, w, p; p' too for a warm start).
pub fn upload(&mut self, field: &Field) {
let rt = runtime();
assert_eq!(field.grid, self.grid);
rt.stream.memcpy_htod(&field.u, &mut self.u).expect("u");
rt.stream.memcpy_htod(&field.v, &mut self.v).expect("v");
rt.stream.memcpy_htod(&field.w, &mut self.w).expect("w");
rt.stream.memcpy_htod(&field.p, &mut self.p).expect("p");
rt.stream
.memcpy_htod(&field.p_prime, &mut self.p_prime)
.expect("p'");
rt.stream.synchronize().expect("sync");
}
/// The device field into the host mirror.
pub fn download(&self, field: &mut Field) {
let rt = runtime();
assert_eq!(field.grid, self.grid);
rt.stream.memcpy_dtoh(&self.u, &mut field.u).expect("u");
rt.stream.memcpy_dtoh(&self.v, &mut field.v).expect("v");
rt.stream.memcpy_dtoh(&self.w, &mut field.w).expect("w");
rt.stream.memcpy_dtoh(&self.p, &mut field.p).expect("p");
rt.stream
.memcpy_dtoh(&self.p_prime, &mut field.p_prime)
.expect("p'");
rt.stream.memcpy_dtoh(&self.sp, &mut field.sp).expect("sp");
rt.stream.synchronize().expect("sync");
}
fn params(&self, dt: f64) -> E3Params {
let g = self.grid;
let b = self.solver.params.boundaries;
E3Params {
nx: g.nx as i32,
ny: g.ny as i32,
nz: g.nz as i32,
periodic_z: i32::from(b.z0 == Side::Periodic),
bx0: side_code(b.x0),
bx1: side_code(b.x1),
by0: side_code(b.y0),
by1: side_code(b.y1),
bz0: side_code(b.z0),
bz1: side_code(b.z1),
scheme: scheme_code(self.solver.params.convection_scheme),
pad: 0,
dx: g.dx,
dy: g.dy,
dz: g.dz,
dt,
rho: self.solver.fluid.density,
nu: self.solver.fluid.viscosity / self.solver.fluid.density,
}
}
fn ptrs(&self) -> E3Ptrs {
let rt = runtime();
let s = &rt.stream;
let p = |x: &CudaSlice<f64>| x.device_ptr(s).0;
let mut ptrs = [0u64; 33];
let base = [
&self.u,
&self.v,
&self.w,
&self.p,
&self.u_old,
&self.v_old,
&self.w_old,
&self.u_star,
&self.v_star,
&self.w_star,
&self.p_prime,
&self.sp,
&self.su,
&self.sv,
&self.sw,
];
for (k, b) in base.iter().enumerate() {
ptrs[k] = p(b);
}
for (k, t) in self.tables.iter().enumerate() {
ptrs[15 + k] = p(t);
}
E3Ptrs { ptrs }
}
fn upload_tables(&mut self, t: f64) {
let rt = runtime();
let host = self.solver.boundary_tables(self.grid, t);
let mut k = 0;
for side in &host {
for comp in side {
if !comp.is_empty() {
rt.stream
.memcpy_htod(comp, &mut self.tables[k])
.expect("table");
}
k += 1;
}
}
}
fn reduce(&mut self) -> f64 {
let rt = runtime();
let k = kernels();
let nb = self.n_blocks as i32;
unsafe {
rt.stream
.launch_builder(&k.reduce)
.arg(&nb)
.arg(&self.partial)
.arg(&mut self.scalar)
.launch(LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
})
.expect("e3_step_reduce");
}
let mut one = vec![0.0f64];
rt.stream
.memcpy_dtoh(&self.scalar, &mut one)
.expect("scalar");
rt.stream.synchronize().expect("sync");
one[0]
}
fn launch_sides(&self, prm: E3Params, ptrs: &E3Ptrs, stamp: i32) {
let rt = runtime();
let k = kernels();
let g = self.grid;
unsafe {
rt.stream
.launch_builder(&k.sides_x)
.arg(&prm)
.arg(ptrs)
.arg(&stamp)
.launch(cfg(g.ny * g.nz))
.expect("sides_x");
rt.stream
.launch_builder(&k.sides_y)
.arg(&prm)
.arg(ptrs)
.arg(&stamp)
.launch(cfg(g.nx * g.nz))
.expect("sides_y");
rt.stream
.launch_builder(&k.sides_z)
.arg(&prm)
.arg(ptrs)
.arg(&stamp)
.launch(cfg(g.nx * g.ny))
.expect("sides_z");
}
}
/// Stamp the `t = time` boundary data (lazily by the first step).
pub fn initialize(&mut self) {
let t = self.solver.time();
self.upload_tables(t);
let prm = self.params(1.0);
let ptrs = self.ptrs();
self.launch_sides(prm, &ptrs, 1);
self.initialized = true;
}
/// One step of `dt` on the device.
pub fn advance(&mut self, dt: f64) -> StepResult {
assert!(dt > 0.0 && dt.is_finite());
if !self.initialized {
self.initialize();
}
let rt = runtime();
let k = kernels();
let g = self.grid;
let t_old = self.solver.time();
let t_new = t_old + dt;
let t0 = Instant::now();
// History shift.
rt.stream
.memcpy_dtod(&self.u, &mut self.u_old)
.expect("u_old");
rt.stream
.memcpy_dtod(&self.v, &mut self.v_old)
.expect("v_old");
rt.stream
.memcpy_dtod(&self.w, &mut self.w_old)
.expect("w_old");
// Predictor with the t_old tables.
self.upload_tables(t_old);
let prm = self.params(dt);
let ptrs = self.ptrs();
let nu = (g.nx + 1) * g.ny * g.nz;
let nv = g.nx * (g.ny + 1) * g.nz;
let nw = g.nx * g.ny * (g.nz + 1);
unsafe {
rt.stream
.launch_builder(&k.predict_u)
.arg(&prm)
.arg(&ptrs)
.launch(cfg(nu))
.expect("predict_u");
rt.stream
.launch_builder(&k.predict_v)
.arg(&prm)
.arg(&ptrs)
.launch(cfg(nv))
.expect("predict_v");
rt.stream
.launch_builder(&k.predict_w)
.arg(&prm)
.arg(&ptrs)
.launch(cfg(nw))
.expect("predict_w");
}
// Outlet zero-gradient + periodic copy (no stamping), then u* = u.
self.launch_sides(prm, &ptrs, 0);
// The new interval's boundary data.
self.upload_tables(t_new);
self.launch_sides(prm, &ptrs, 1);
rt.stream
.memcpy_dtod(&self.u, &mut self.u_star)
.expect("u*");
rt.stream
.memcpy_dtod(&self.v, &mut self.v_star)
.expect("v*");
rt.stream
.memcpy_dtod(&self.w, &mut self.w_star)
.expect("w*");
rt.stream.synchronize().expect("sync");
let t_pred = t0.elapsed();
// The operator (a cache keyed on dt; no body: constant otherwise).
let t1 = Instant::now();
if self.cg.is_none() || self.cg_dt != dt {
let problem = self.solver.poisson_operator(g, dt);
let params = MultigridParameters {
precision: self.solver.params.poisson_precision,
smoother: self.solver.params.poisson_smoother,
..MultigridParameters::default()
};
self.cg = Some(DeviceCg::new(&problem, &params));
self.cg_dt = dt;
}
let anchor = self.solver.anchor_cell(g);
let mut total = 0;
let mut final_residual = f64::INFINITY;
let mut cg_iterations = 0;
let mut t_poisson = std::time::Duration::ZERO;
let mut t_apply = std::time::Duration::ZERO;
for corrector in 0..self.solver.params.corrector_steps.max(1) {
let tp = Instant::now();
unsafe {
rt.stream
.launch_builder(&k.divergence)
.arg(&prm)
.arg(&ptrs)
.arg(&mut self.partial)
.launch(cfg(g.cells()))
.expect("divergence");
}
let source_scale = self.reduce();
let inner_stop = self.solver.inner_stop(g, source_scale);
if corrector > 0 {
rt.stream.memset_zeros(&mut self.p_prime).expect("p' = 0");
}
let sol = {
let cg = self.cg.as_mut().expect("cg");
cg.solve_device(&self.sp, &mut self.p_prime, inner_stop, anchor, 0)
};
cg_iterations += sol.iterations;
t_poisson += tp.elapsed();
let ta = Instant::now();
unsafe {
rt.stream
.launch_builder(&k.correct_u)
.arg(&prm)
.arg(&ptrs)
.launch(cfg(nu))
.expect("correct_u");
rt.stream
.launch_builder(&k.correct_v)
.arg(&prm)
.arg(&ptrs)
.launch(cfg(nv))
.expect("correct_v");
rt.stream
.launch_builder(&k.correct_w)
.arg(&prm)
.arg(&ptrs)
.launch(cfg(nw))
.expect("correct_w");
}
if prm.periodic_z != 0 {
self.launch_sides(prm, &ptrs, 0);
}
unsafe {
rt.stream
.launch_builder(&k.add_p)
.arg(&prm)
.arg(&ptrs)
.arg(&mut self.partial)
.launch(cfg(g.cells()))
.expect("add_p");
}
let imbalance = self.reduce();
let reference_flux = self.solver.reference_flux(g);
let mass_residual = if reference_flux > 0.0 {
imbalance / reference_flux
} else {
imbalance
};
final_residual = mass_residual;
total += 1;
t_apply += ta.elapsed();
if mass_residual < self.solver.params.tolerance {
break;
}
rt.stream
.memcpy_dtod(&self.u, &mut self.u_star)
.expect("u*");
rt.stream
.memcpy_dtod(&self.v, &mut self.v_star)
.expect("v*");
rt.stream
.memcpy_dtod(&self.w, &mut self.w_star)
.expect("w*");
}
let _ = t1;
self.solver.set_time(t_new);
if let Some(tm) = self.timers.as_mut() {
tm.predictor_ns += t_pred.as_nanos() as u64;
tm.poisson_ns += t_poisson.as_nanos() as u64;
tm.apply_ns += t_apply.as_nanos() as u64;
tm.steps += 1;
tm.cg_iterations += cg_iterations as u64;
}
StepResult {
converged: final_residual < self.solver.params.tolerance,
corrector_steps_performed: total,
final_residual,
poisson_iterations: cg_iterations,
}
}
}
@@ -4,6 +4,8 @@
//! `dz = 1`) every number is the 2D solver's. The fluid predicates are the
//! wall's hooks (item 9).
#[cfg(feature = "cuda")]
pub mod device;
mod predictor;
mod projection;