use cudarc::driver::LaunchConfig;
const BLOCK_1D: u32 = 256;
pub fn validate_kernel_arg_capacity(
batch: usize,
seq_len: usize,
d_inner: usize,
d_state: usize,
) -> Result<(), String> {
if seq_len == 0 {
return Err("seq_len must be > 0".into());
}
if batch == 0 {
return Err("batch must be > 0".into());
}
let elems = batch
.checked_mul(seq_len + 1)
.and_then(|v| v.checked_mul(d_inner))
.and_then(|v| v.checked_mul(d_state))
.ok_or("batch * (seq_len+1) * d_inner * d_state overflows usize")?;
if elems > i32::MAX as usize {
return Err(format!(
"batch({batch}) * (seq_len({seq_len})+1) * d_inner({d_inner}) * d_state({d_state}) \
= {elems} elements exceeds i32::MAX; CUDA kernels take element counts as 32-bit ints"
));
}
Ok(())
}
pub fn grid_1d(n: usize) -> LaunchConfig {
assert!(
n <= i32::MAX as usize,
"grid_1d: element count {n} exceeds i32::MAX"
);
let num_blocks = (n as u32).div_ceil(BLOCK_1D);
LaunchConfig {
grid_dim: (num_blocks, 1, 1),
block_dim: (BLOCK_1D, 1, 1),
shared_mem_bytes: 0,
}
}
pub fn grid_ssm(batch: usize, d_inner: usize) -> LaunchConfig {
grid_1d(batch * d_inner)
}
pub fn grid_reduce(total_elements: usize) -> LaunchConfig {
grid_1d(total_elements)
}
pub fn grid_col_tree_reduce(n_cols: usize) -> LaunchConfig {
LaunchConfig {
grid_dim: (n_cols as u32, 1, 1),
block_dim: (BLOCK_1D, 1, 1),
shared_mem_bytes: (BLOCK_1D as usize * std::mem::size_of::<f32>()) as u32,
}
}
pub fn grid_norm(batch: usize, dim: usize) -> LaunchConfig {
let block = (dim as u32).min(1024);
let block = block.next_power_of_two();
LaunchConfig {
grid_dim: (batch as u32, 1, 1),
block_dim: (block, 1, 1),
shared_mem_bytes: (block as usize * std::mem::size_of::<f32>()) as u32,
}
}
pub const SCAN_NTHREADS: usize = 128;
pub const SCAN_NITEMS: usize = 8;
pub const SCAN_CHUNK: usize = SCAN_NTHREADS * SCAN_NITEMS;
pub fn scan_tape_slim() -> bool {
static SLIM: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*SLIM.get_or_init(|| {
std::env::var("MAMBA_RS_SCAN_TAPE")
.map(|v| v != "full")
.unwrap_or(true)
})
}
pub fn scan_tape_len(batch: usize, seq_len: usize, d_inner: usize, d_state: usize) -> usize {
let n_chunks = seq_len.div_ceil(SCAN_CHUNK);
batch * d_inner * d_state * 3 * n_chunks
}
pub fn grid_parallel_scan(batch: usize, d_inner: usize) -> LaunchConfig {
assert!(
d_inner <= 65535,
"grid_parallel_scan: d_inner {d_inner} exceeds CUDA grid.y limit 65535"
);
const NWARPS: usize = SCAN_NTHREADS / 32;
const MAX_DSTATE: usize = 256;
let smem_floats = 2 * NWARPS + 2 * MAX_DSTATE + 2 * SCAN_NTHREADS + SCAN_CHUNK;
LaunchConfig {
grid_dim: (batch as u32, d_inner as u32, 1),
block_dim: (SCAN_NTHREADS as u32, 1, 1),
shared_mem_bytes: (smem_floats * std::mem::size_of::<f32>()) as u32,
}
}
pub const SCAN_BWD_DGROUP: usize = 4;
pub fn grid_parallel_scan_bwd_fold(
batch: usize,
d_inner: usize,
bytes_per_act: usize,
) -> LaunchConfig {
const NWARPS: usize = SCAN_NTHREADS / 32;
const MAX_DSTATE: usize = 256;
let g = SCAN_BWD_DGROUP;
let f32_floats = 4 * NWARPS + 4 * SCAN_NTHREADS + 3 * g * MAX_DSTATE + 3 * SCAN_NTHREADS;
let stage_bytes = 3 * g * SCAN_CHUNK * bytes_per_act;
LaunchConfig {
grid_dim: (batch as u32, (d_inner / g) as u32, 1),
block_dim: (SCAN_NTHREADS as u32, 1, 1),
shared_mem_bytes: (f32_floats * std::mem::size_of::<f32>() + stage_bytes) as u32,
}
}
pub fn grid_parallel_scan_bwd(batch: usize, d_inner: usize) -> LaunchConfig {
assert!(
d_inner <= 65535,
"grid_parallel_scan_bwd: d_inner {d_inner} exceeds CUDA grid.y limit 65535"
);
const NWARPS: usize = SCAN_NTHREADS / 32;
const MAX_DSTATE: usize = 256;
let fwd_total_floats = 2 * NWARPS + 2 * MAX_DSTATE + 2 * SCAN_NTHREADS + SCAN_CHUNK;
let bwd_extra_floats = 2 * NWARPS + 3 * MAX_DSTATE + 2 * SCAN_NTHREADS;
let total_bytes = (fwd_total_floats + bwd_extra_floats) * std::mem::size_of::<f32>();
LaunchConfig {
grid_dim: (batch as u32, d_inner as u32, 1),
block_dim: (SCAN_NTHREADS as u32, 1, 1),
shared_mem_bytes: total_bytes as u32,
}
}
pub fn grid_parallel_scan_typed(
batch: usize,
d_inner: usize,
bytes_per_act: usize,
) -> LaunchConfig {
debug_assert!(bytes_per_act == 2 || bytes_per_act == 4);
assert!(
d_inner <= 65535,
"grid_parallel_scan_typed: d_inner {d_inner} exceeds CUDA grid.y limit 65535"
);
const NWARPS: usize = SCAN_NTHREADS / 32;
const MAX_DSTATE: usize = 256;
let fixed_floats = 2 * NWARPS + 2 * MAX_DSTATE + 2 * SCAN_NTHREADS;
let fixed_bytes = fixed_floats * std::mem::size_of::<f32>();
let stage_bytes = SCAN_CHUNK * bytes_per_act;
LaunchConfig {
grid_dim: (batch as u32, d_inner as u32, 1),
block_dim: (SCAN_NTHREADS as u32, 1, 1),
shared_mem_bytes: (fixed_bytes + stage_bytes) as u32,
}
}
pub fn grid_conv_tiled(batch: usize, d_inner: usize, t: usize) -> LaunchConfig {
const TILE_T: usize = 128;
LaunchConfig {
grid_dim: (
((batch * d_inner) as u32).div_ceil(256),
(t as u32).div_ceil(TILE_T as u32),
1,
),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
}
}