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,
}
}
pub fn grid_parallel_scan_bwd(batch: usize, d_inner: usize, bytes_per_act: usize) -> LaunchConfig {
debug_assert!(bytes_per_act == 2 || bytes_per_act == 4);
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_fixed_floats = 2 * NWARPS + 2 * MAX_DSTATE + 2 * NTHREADS as usize;
let bwd_extra_floats = 2 * NWARPS + 2 * MAX_DSTATE + 2 * NTHREADS as usize + MAX_DSTATE;
let fixed_bytes = (fwd_fixed_floats + bwd_extra_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,
}
}
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);
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,
}
}