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> {
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 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 NTHREADS: u32 = 128;
const NWARPS: usize = NTHREADS as usize / 32;
const MAX_DSTATE: usize = 256;
const CHUNK_SIZE: usize = NTHREADS as usize * 8; let smem_floats = 2 * NWARPS + 2 * MAX_DSTATE + 2 * NTHREADS as usize + CHUNK_SIZE;
LaunchConfig {
grid_dim: (batch as u32, d_inner as u32, 1),
block_dim: (NTHREADS, 1, 1),
shared_mem_bytes: (smem_floats * std::mem::size_of::<f32>()) 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 NTHREADS: u32 = 128;
const NWARPS: usize = NTHREADS as usize / 32;
const MAX_DSTATE: usize = 256;
const CHUNK_SIZE: usize = NTHREADS as usize * 8;
let fwd_total_floats = 2 * NWARPS + 2 * MAX_DSTATE + 2 * NTHREADS as usize + CHUNK_SIZE;
let bwd_extra_floats = 2 * NWARPS + 3 * MAX_DSTATE + 2 * NTHREADS as usize;
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: (NTHREADS, 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 NTHREADS: u32 = 128;
const NWARPS: usize = NTHREADS as usize / 32;
const MAX_DSTATE: usize = 256;
const CHUNK_SIZE: usize = NTHREADS as usize * 8;
let fixed_floats = 2 * NWARPS + 2 * MAX_DSTATE + 2 * NTHREADS as usize;
let fixed_bytes = fixed_floats * std::mem::size_of::<f32>();
let stage_bytes = CHUNK_SIZE * bytes_per_act;
LaunchConfig {
grid_dim: (batch as u32, d_inner as u32, 1),
block_dim: (NTHREADS, 1, 1),
shared_mem_bytes: (fixed_bytes + stage_bytes) as u32,
}
}