use std::sync::Arc;
use cudarc::driver::PushKernelArg;
use super::error::DistError;
use super::fold::shard_plan;
use crate::mamba_ssm::gpu::buffers::GpuBuffer;
use crate::mamba_ssm::gpu::device::GpuDevice;
const DET_SUM_RANKS_SRC: &str = r#"
extern "C" __global__ void det_sum_ranks(
float* __restrict__ out,
const float* __restrict__ stacked,
int world,
int len
) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= len) return;
float acc = stacked[i];
for (int r = 1; r < world; ++r) {
acc += stacked[(size_t)r * (size_t)len + (size_t)i];
}
out[i] = acc;
}
"#;
struct DeviceScratch {
ptr: cudarc::driver::sys::CUdeviceptr,
_ctx: Arc<cudarc::driver::CudaContext>,
}
impl DeviceScratch {
fn alloc(ctx: &Arc<cudarc::driver::CudaContext>, elems: usize) -> Result<Self, DistError> {
let bytes = elems.max(1) * std::mem::size_of::<f32>();
let mut ptr: cudarc::driver::sys::CUdeviceptr = 0;
unsafe {
let r = cudarc::driver::sys::cuMemAlloc_v2(&mut ptr, bytes);
if r != cudarc::driver::sys::CUresult::CUDA_SUCCESS {
return Err(DistError::Transport(format!(
"reducer scratch alloc ({bytes} B): {r:?}"
)));
}
}
Ok(Self {
ptr,
_ctx: ctx.clone(),
})
}
}
impl Drop for DeviceScratch {
fn drop(&mut self) {
let _ = unsafe { cudarc::driver::sys::cuMemFree_v2(self.ptr) };
}
}
pub struct DetReduceKernel {
func: cudarc::driver::CudaFunction,
ctx: Arc<cudarc::driver::CudaContext>,
}
impl DetReduceKernel {
pub fn compile(ordinal: usize) -> Result<Self, DistError> {
let device = GpuDevice::new(ordinal)
.map_err(|e| DistError::Transport(format!("reducer device init: {e}")))?;
let arch = GpuDevice::nvrtc_arch(device.compute_capability);
let opts = cudarc::nvrtc::CompileOptions {
arch: Some(arch),
..Default::default()
};
let ptx = cudarc::nvrtc::compile_ptx_with_opts(DET_SUM_RANKS_SRC, opts)
.map_err(|e| DistError::Transport(format!("NVRTC det_sum_ranks: {e:?}")))?;
let ctx = device.context().clone();
let module = ctx
.load_module(ptx)
.map_err(|e| DistError::Transport(format!("det_sum_ranks module load: {e:?}")))?;
let func = module
.load_function("det_sum_ranks")
.map_err(|e| DistError::Transport(format!("det_sum_ranks lookup: {e:?}")))?;
Ok(Self { func, ctx })
}
pub fn launch(
&self,
out_ptr: cudarc::driver::sys::CUdeviceptr,
stacked_ptr: cudarc::driver::sys::CUdeviceptr,
world: usize,
len: usize,
stream: &cudarc::driver::CudaStream,
) -> Result<(), DistError> {
if len == 0 {
return Ok(());
}
if len > i32::MAX as usize || world > i32::MAX as usize {
return Err(DistError::Transport(format!(
"det_sum_ranks: len {len} / world {world} exceed the i32 kernel ABI"
)));
}
let world_i = world as i32;
let len_i = len as i32;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: ((len as u32).div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let mut b = stream.launch_builder(&self.func);
b.arg(&out_ptr);
b.arg(&stacked_ptr);
b.arg(&world_i);
b.arg(&len_i);
unsafe { b.launch(cfg) }
.map_err(|e| DistError::Transport(format!("det_sum_ranks launch: {e:?}")))?;
Ok(())
}
}
#[cfg(feature = "nccl")]
pub(super) struct FixedOrderState {
kernel: DetReduceKernel,
stacked: std::cell::RefCell<std::collections::HashMap<usize, DeviceScratch>>,
}
#[cfg(feature = "nccl")]
impl FixedOrderState {
pub(super) fn compile(ordinal: usize) -> Result<Self, DistError> {
Ok(Self {
kernel: DetReduceKernel::compile(ordinal)?,
stacked: std::cell::RefCell::new(std::collections::HashMap::new()),
})
}
fn stacked_ptr(
&self,
arena_len: usize,
world: usize,
my_len: usize,
) -> Result<cudarc::driver::sys::CUdeviceptr, DistError> {
let mut map = self.stacked.borrow_mut();
if let std::collections::hash_map::Entry::Vacant(slot) = map.entry(arena_len) {
slot.insert(DeviceScratch::alloc(&self.kernel.ctx, world * my_len)?);
}
let Some(scratch) = map.get(&arena_len) else {
return Err(DistError::Transport(
"reducer stacked scratch missing after allocation".into(),
));
};
Ok(scratch.ptr)
}
}
#[cfg(feature = "nccl")]
pub(super) fn reduce_sum_fixed_order_nccl(
comm: &super::comm::MambaComm,
state: &FixedOrderState,
arena: &mut GpuBuffer,
stream: &cudarc::driver::CudaStream,
) -> Result<(), DistError> {
let world = comm.world();
let me = comm.rank();
if world <= 1 {
return Ok(());
}
let n = arena.len();
let plan = shard_plan(n, world);
let my = plan[me];
let f32_size = std::mem::size_of::<f32>() as u64;
let arena_base = arena.cached_ptr();
let stacked_base = state.stacked_ptr(n, world, my.len)?;
super::comm::MambaComm::group(|| {
for (p, shard) in plan.iter().enumerate() {
if p == me {
continue;
}
if shard.len > 0 {
comm.send_f32(
arena_base + shard.start as u64 * f32_size,
shard.len,
p,
stream,
)?;
}
if my.len > 0 {
comm.recv_f32(
stacked_base + (p * my.len) as u64 * f32_size,
my.len,
p,
stream,
)?;
}
}
Ok(())
})?;
if my.len > 0 {
copy_d2d(
stacked_base + (me * my.len) as u64 * f32_size,
arena_base + my.start as u64 * f32_size,
my.len,
stream,
)?;
state.kernel.launch(
arena_base + my.start as u64 * f32_size,
stacked_base,
world,
my.len,
stream,
)?;
}
super::comm::MambaComm::group(|| {
for (o, shard) in plan.iter().enumerate() {
if shard.len > 0 {
comm.broadcast_f32(
arena_base + shard.start as u64 * f32_size,
shard.len,
o,
stream,
)?;
}
}
Ok(())
})
}
fn copy_d2d(
dst: cudarc::driver::sys::CUdeviceptr,
src: cudarc::driver::sys::CUdeviceptr,
len: usize,
stream: &cudarc::driver::CudaStream,
) -> Result<(), DistError> {
if len == 0 {
return Ok(());
}
let bytes = len * std::mem::size_of::<f32>();
let r =
unsafe { cudarc::driver::sys::cuMemcpyDtoDAsync_v2(dst, src, bytes, stream.cu_stream()) };
if r != cudarc::driver::sys::CUresult::CUDA_SUCCESS {
return Err(DistError::Transport(format!("reducer D2D copy: {r:?}")));
}
Ok(())
}
pub struct LoopbackWorld {
arenas: Vec<GpuBuffer>,
pub reverse_delivery: bool,
}
impl LoopbackWorld {
pub fn new(arenas: Vec<GpuBuffer>) -> Self {
Self {
arenas,
reverse_delivery: false,
}
}
pub fn arenas(&self) -> &[GpuBuffer] {
&self.arenas
}
pub fn run_round(
&mut self,
kernel: &DetReduceKernel,
stream: &cudarc::driver::CudaStream,
) -> Result<(), DistError> {
let world = self.arenas.len();
if world <= 1 {
return Ok(());
}
let n = self.arenas[0].len();
for a in &self.arenas {
if a.len() != n {
return Err(DistError::Transport(
"loopback arenas must share one length".into(),
));
}
}
let plan = shard_plan(n, world);
let f32_size = std::mem::size_of::<f32>() as u64;
for (me, my) in plan.iter().enumerate() {
if my.len == 0 {
continue;
}
let stacked = DeviceScratch::alloc(&kernel.ctx, world * my.len)?;
let order: Vec<usize> = if self.reverse_delivery {
(0..world).rev().collect()
} else {
(0..world).collect()
};
for src in order {
copy_d2d(
stacked.ptr + (src * my.len) as u64 * f32_size,
self.arenas[src].cached_ptr() + my.start as u64 * f32_size,
my.len,
stream,
)?;
}
kernel.launch(
self.arenas[me].cached_ptr() + my.start as u64 * f32_size,
stacked.ptr,
world,
my.len,
stream,
)?;
stream
.synchronize()
.map_err(|e| DistError::Transport(format!("loopback sync: {e:?}")))?;
}
for (o, shard) in plan.iter().enumerate() {
if shard.len == 0 {
continue;
}
for r in 0..world {
if r == o {
continue;
}
copy_d2d(
self.arenas[r].cached_ptr() + shard.start as u64 * f32_size,
self.arenas[o].cached_ptr() + shard.start as u64 * f32_size,
shard.len,
stream,
)?;
}
}
stream
.synchronize()
.map_err(|e| DistError::Transport(format!("loopback sync: {e:?}")))?;
Ok(())
}
}
pub fn reduce_sum_reference(addends: &[&[f32]], out: &mut [f32]) {
let world = addends.len();
assert!(world > 0, "need at least one addend");
for a in addends {
assert_eq!(a.len(), out.len(), "addend length mismatch");
}
for (i, o) in out.iter_mut().enumerate() {
let mut acc = addends[0][i];
for a in &addends[1..] {
acc += a[i];
}
*o = acc;
}
}