use cudarc::driver::LaunchConfig;
const BLOCK_1D: u32 = 256;
pub fn grid_1d(n: usize) -> LaunchConfig {
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_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 {
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,
}
}