use std::ops::Range;
use std::sync::Arc;
use anyhow::{Context, Result};
use cudarc::driver::sys::CUstream;
use cudarc::driver::{CudaContext, CudaEvent, CudaStream};
use cudarc::nccl::sys::{
ncclBcast, ncclComm_t, ncclCommDestroy, ncclDataType_t, ncclGroupEnd, ncclGroupStart,
};
use velo::EventManager;
use crate::BlockId;
use kvbm_common::LogicalLayoutHandle;
use kvbm_physical::layout::PhysicalLayout;
use kvbm_physical::transfer::TransferCompleteNotification;
use super::CollectiveOps;
use super::bootstrap::{NcclBootstrap, check_nccl_result};
pub trait LayoutResolver: Send + Sync {
fn resolve_layout(&self, logical: LogicalLayoutHandle) -> Result<PhysicalLayout>;
}
pub trait CudaEventRegistrar: Send + Sync {
fn register_cuda_event(&self, event: CudaEvent) -> TransferCompleteNotification;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum CommOwnership {
Owned,
Borrowed,
}
enum NcclStream {
Owned(Arc<CudaStream>),
Borrowed(CUstream),
}
impl NcclStream {
fn raw(&self) -> CUstream {
match self {
NcclStream::Owned(stream) => stream.cu_stream(),
NcclStream::Borrowed(ptr) => *ptr,
}
}
fn as_owned(&self) -> Option<&Arc<CudaStream>> {
match self {
NcclStream::Owned(stream) => Some(stream),
NcclStream::Borrowed(_) => None,
}
}
}
pub struct NcclCollectives {
comm: ncclComm_t,
ownership: CommOwnership,
rank: usize,
world_size: usize,
nccl_stream: NcclStream,
#[allow(dead_code)]
cuda_context: Arc<CudaContext>,
event_system: EventManager,
event_registrar: Arc<dyn CudaEventRegistrar>,
layout_resolver: Arc<dyn LayoutResolver>,
}
impl NcclCollectives {
pub fn from_bootstrap(
bootstrap: &NcclBootstrap,
rank: usize,
cuda_context: Arc<CudaContext>,
event_system: EventManager,
event_registrar: Arc<dyn CudaEventRegistrar>,
layout_resolver: Arc<dyn LayoutResolver>,
) -> Result<Self> {
let nccl_stream = cuda_context
.new_stream()
.context("Failed to create NCCL stream")?;
let comm = bootstrap
.init_communicator(rank, nccl_stream.cu_stream())
.context("Failed to initialize NCCL communicator")?;
Ok(Self {
comm,
ownership: CommOwnership::Owned,
rank,
world_size: bootstrap.world_size(),
nccl_stream: NcclStream::Owned(nccl_stream),
cuda_context,
event_system,
event_registrar,
layout_resolver,
})
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn from_borrowed(
comm_ptr: usize,
stream_ptr: usize,
rank: usize,
world_size: usize,
cuda_context: Arc<CudaContext>,
event_system: EventManager,
event_registrar: Arc<dyn CudaEventRegistrar>,
layout_resolver: Arc<dyn LayoutResolver>,
) -> Self {
Self {
comm: comm_ptr as ncclComm_t,
ownership: CommOwnership::Borrowed,
rank,
world_size,
nccl_stream: NcclStream::Borrowed(stream_ptr as CUstream),
cuda_context,
event_system,
event_registrar,
layout_resolver,
}
}
fn broadcast_regions(&self, regions: &[(usize, usize)], root: i32) -> Result<()> {
if regions.is_empty() {
return Ok(());
}
let stream = self.nccl_stream.raw();
let result = unsafe { ncclGroupStart() };
check_nccl_result(result).context("ncclGroupStart failed")?;
for (ptr, size) in regions {
let result = unsafe {
ncclBcast(
*ptr as *mut std::ffi::c_void,
*size,
ncclDataType_t::ncclChar, root,
self.comm,
stream.cast(),
)
};
check_nccl_result(result).context("ncclBcast failed")?;
}
let result = unsafe { ncclGroupEnd() };
check_nccl_result(result).context("ncclGroupEnd failed")?;
Ok(())
}
fn collect_regions(
&self,
layout: &PhysicalLayout,
block_ids: &[BlockId],
layer_range: Option<Range<usize>>,
) -> Result<Vec<(usize, usize)>> {
let num_layers = layout.layout().num_layers();
let outer_dim = layout.layout().outer_dim();
let layer_range = layer_range.unwrap_or(0..num_layers);
let mut regions =
Vec::with_capacity(block_ids.len() * (layer_range.end - layer_range.start) * outer_dim);
for &block_id in block_ids {
for layer_id in layer_range.clone() {
for outer_id in 0..outer_dim {
let region = layout.memory_region(block_id, layer_id, outer_id)?;
regions.push((region.addr, region.size));
}
}
}
Ok(regions)
}
fn create_completion_notification(&self) -> Result<TransferCompleteNotification> {
if let Some(stream) = self.nccl_stream.as_owned() {
let cuda_event = stream
.record_event(None)
.context("Failed to record CUDA event")?;
Ok(self.event_registrar.register_cuda_event(cuda_event))
} else {
tracing::warn!(
"Using borrowed stream - returning immediate completion. \
Caller must ensure stream synchronization."
);
let nova_event = self.event_system.new_event()?;
let handle = nova_event.handle();
nova_event.trigger()?;
let awaiter = self.event_system.awaiter(handle)?;
Ok(TransferCompleteNotification::from_awaiter(awaiter))
}
}
}
impl CollectiveOps for NcclCollectives {
fn broadcast(
&self,
src: LogicalLayoutHandle,
dst: LogicalLayoutHandle,
src_block_ids: &[BlockId],
dst_block_ids: &[BlockId],
layer_range: Option<Range<usize>>,
) -> Result<TransferCompleteNotification> {
let src_layout = self.layout_resolver.resolve_layout(src)?;
let dst_layout = self.layout_resolver.resolve_layout(dst)?;
let layout = if self.rank == 0 {
&src_layout
} else {
&dst_layout
};
let block_ids = if self.rank == 0 {
src_block_ids
} else {
dst_block_ids
};
let regions = self.collect_regions(layout, block_ids, layer_range)?;
tracing::debug!(
rank = self.rank,
world_size = self.world_size,
num_regions = regions.len(),
total_bytes = regions.iter().map(|(_, size)| size).sum::<usize>(),
"Starting NCCL broadcast"
);
self.broadcast_regions(®ions, 0)?;
self.create_completion_notification()
}
fn rank(&self) -> usize {
self.rank
}
fn world_size(&self) -> usize {
self.world_size
}
}
impl Drop for NcclCollectives {
fn drop(&mut self) {
if self.ownership == CommOwnership::Owned {
let result = unsafe { ncclCommDestroy(self.comm) };
if let Err(e) = check_nccl_result(result) {
tracing::warn!("Failed to destroy NCCL communicator: {:?}", e);
}
}
}
}
unsafe impl Send for NcclCollectives {}
unsafe impl Sync for NcclCollectives {}
#[cfg(test)]
mod tests {
use super::*;
use cudarc::driver::{CudaContext, CudaSlice, DevicePtr};
use cudarc::nccl::sys::{ncclCommDestroy, ncclCommInitAll};
use std::ffi::c_int;
use std::sync::{Arc, Barrier};
use std::thread;
fn cuda_device_count() -> usize {
CudaContext::device_count().unwrap_or(0) as usize
}
unsafe fn init_all_comms(num_devices: usize) -> Result<Vec<usize>> {
let mut comms: Vec<ncclComm_t> = vec![std::ptr::null_mut(); num_devices];
let devices: Vec<c_int> = (0..num_devices as c_int).collect();
let result =
unsafe { ncclCommInitAll(comms.as_mut_ptr(), num_devices as c_int, devices.as_ptr()) };
check_nccl_result(result).context("ncclCommInitAll failed")?;
Ok(comms.into_iter().map(|c| c as usize).collect())
}
unsafe fn destroy_comms(comms: &[usize]) {
for &comm in comms {
unsafe {
let _ = ncclCommDestroy(comm as ncclComm_t);
}
}
}
fn get_device_ptr(slice: &CudaSlice<u8>, stream: &CudaStream) -> usize {
let (ptr, _guard) = slice.device_ptr(stream);
ptr as usize
}
#[test]
#[cfg(feature = "testing-nccl")]
fn test_nccl_broadcast_multi_gpu_raw() {
let num_devices = cuda_device_count();
if num_devices < 2 {
println!(
"Skipping test: {} GPUs available, need at least 2",
num_devices
);
return;
}
let world_size = 2;
println!("Testing NCCL broadcast with {} GPUs", world_size);
let comms = unsafe { init_all_comms(world_size) }.expect("Failed to init NCCL comms");
let contexts: Vec<Arc<CudaContext>> = (0..world_size)
.map(|i| CudaContext::new(i).expect("Failed to create CUDA context"))
.collect();
let streams: Vec<Arc<CudaStream>> = contexts
.iter()
.map(|ctx| ctx.new_stream().expect("Failed to create stream"))
.collect();
let test_size = 1024 * 1024; let test_pattern: u8 = 0xAB;
let buffers: Vec<CudaSlice<u8>> = streams
.iter()
.map(|stream| {
let zeros = vec![0u8; test_size];
stream
.clone_htod(&zeros)
.expect("Failed to allocate buffer")
})
.collect();
{
let host_data = vec![test_pattern; test_size];
let buffer = streams[0]
.clone_htod(&host_data)
.expect("Failed to copy to device 0");
let src_ptr = get_device_ptr(&buffer, &streams[0]);
let dst_ptr = get_device_ptr(&buffers[0], &streams[0]);
unsafe {
cudarc::driver::result::memcpy_dtod_async(
dst_ptr as u64,
src_ptr as u64,
test_size,
streams[0].cu_stream(),
)
.expect("dtod copy failed");
}
streams[0].synchronize().expect("sync failed");
}
let buffer_ptrs: Vec<usize> = buffers
.iter()
.zip(streams.iter())
.map(|(buf, stream)| get_device_ptr(buf, stream))
.collect();
let barrier = Arc::new(Barrier::new(world_size));
let handles: Vec<_> = (0..world_size)
.map(|rank| {
let comm = comms[rank]; let stream = streams[rank].clone();
let buffer_ptr = buffer_ptrs[rank];
let barrier = barrier.clone();
thread::spawn(move || {
barrier.wait();
let result = unsafe {
ncclBcast(
buffer_ptr as *mut std::ffi::c_void,
test_size,
ncclDataType_t::ncclChar,
0, comm as ncclComm_t, stream.cu_stream().cast(),
)
};
check_nccl_result(result).expect("ncclBcast failed");
stream.synchronize().expect("Stream sync failed");
println!("Rank {} completed broadcast", rank);
})
})
.collect();
for handle in handles {
handle.join().expect("Thread panicked");
}
for (rank, (stream, buffer)) in streams.iter().zip(buffers.iter()).enumerate() {
let host_data = stream
.clone_dtoh(buffer)
.expect("Failed to copy from device");
assert_eq!(
host_data[0], test_pattern,
"Rank {} first byte mismatch",
rank
);
assert_eq!(
host_data[test_size - 1],
test_pattern,
"Rank {} last byte mismatch",
rank
);
assert_eq!(
host_data[test_size / 2],
test_pattern,
"Rank {} middle byte mismatch",
rank
);
let mismatch_count = host_data.iter().filter(|&&b| b != test_pattern).count();
assert_eq!(
mismatch_count, 0,
"Rank {} has {} mismatched bytes",
rank, mismatch_count
);
println!("Rank {} verified: all {} bytes correct", rank, test_size);
}
unsafe { destroy_comms(&comms) };
println!("Test passed!");
}
#[test]
#[cfg(feature = "testing-nccl")]
fn test_nccl_grouped_broadcast_multi_gpu() {
let num_devices = cuda_device_count();
if num_devices < 2 {
println!(
"Skipping test: {} GPUs available, need at least 2",
num_devices
);
return;
}
let world_size = 2;
println!("Testing NCCL grouped broadcast with {} GPUs", world_size);
let comms = unsafe { init_all_comms(world_size) }.expect("Failed to init NCCL comms");
let contexts: Vec<Arc<CudaContext>> = (0..world_size)
.map(|i| CudaContext::new(i).expect("Failed to create CUDA context"))
.collect();
let streams: Vec<Arc<CudaStream>> = contexts
.iter()
.map(|ctx| ctx.new_stream().expect("Failed to create stream"))
.collect();
let num_regions = 4;
let region_size = 256 * 1024;
let buffers: Vec<Vec<CudaSlice<u8>>> = streams
.iter()
.map(|stream| {
(0..num_regions)
.map(|_| {
let zeros = vec![0u8; region_size];
stream.clone_htod(&zeros).expect("Failed to allocate")
})
.collect()
})
.collect();
for (region_idx, buffer) in buffers[0].iter().enumerate() {
let pattern = (region_idx + 1) as u8 * 0x11; let host_data = vec![pattern; region_size];
let src_buffer = streams[0]
.clone_htod(&host_data)
.expect("Failed to allocate src");
let src_ptr = get_device_ptr(&src_buffer, &streams[0]);
let dst_ptr = get_device_ptr(buffer, &streams[0]);
unsafe {
cudarc::driver::result::memcpy_dtod_async(
dst_ptr as u64,
src_ptr as u64,
region_size,
streams[0].cu_stream(),
)
.expect("dtod copy failed");
}
}
streams[0].synchronize().expect("sync failed");
let barrier = Arc::new(Barrier::new(world_size));
let buffer_ptrs: Vec<Vec<usize>> = buffers
.iter()
.zip(streams.iter())
.map(|(rank_buffers, stream)| {
rank_buffers
.iter()
.map(|b| get_device_ptr(b, stream))
.collect()
})
.collect();
let handles: Vec<_> = (0..world_size)
.map(|rank| {
let comm = comms[rank]; let stream = streams[rank].clone();
let ptrs = buffer_ptrs[rank].clone();
let barrier = barrier.clone();
thread::spawn(move || {
barrier.wait();
unsafe {
check_nccl_result(ncclGroupStart()).expect("ncclGroupStart failed");
for ptr in &ptrs {
let result = ncclBcast(
*ptr as *mut std::ffi::c_void,
region_size,
ncclDataType_t::ncclChar,
0,
comm as ncclComm_t,
stream.cu_stream().cast(),
);
check_nccl_result(result).expect("ncclBcast failed");
}
check_nccl_result(ncclGroupEnd()).expect("ncclGroupEnd failed");
}
stream.synchronize().expect("Stream sync failed");
println!("Rank {} completed grouped broadcast", rank);
})
})
.collect();
for handle in handles {
handle.join().expect("Thread panicked");
}
for (rank, (stream, rank_buffers)) in streams.iter().zip(buffers.iter()).enumerate() {
for (region_idx, buffer) in rank_buffers.iter().enumerate() {
let expected_pattern = (region_idx + 1) as u8 * 0x11;
let host_data = stream
.clone_dtoh(buffer)
.expect("Failed to copy from device");
let mismatch_count = host_data.iter().filter(|&&b| b != expected_pattern).count();
assert_eq!(
mismatch_count, 0,
"Rank {} region {} has {} mismatched bytes (expected 0x{:02x})",
rank, region_idx, mismatch_count, expected_pattern
);
}
println!(
"Rank {} verified: all {} regions correct",
rank, num_regions
);
}
unsafe { destroy_comms(&comms) };
println!("Grouped broadcast test passed!");
}
#[test]
#[cfg(feature = "testing-nccl")]
fn test_nccl_broadcast_large_transfer() {
let num_devices = cuda_device_count();
if num_devices < 2 {
println!(
"Skipping test: {} GPUs available, need at least 2",
num_devices
);
return;
}
let world_size = 2;
println!("Testing NCCL large broadcast with {} GPUs", world_size);
let comms = unsafe { init_all_comms(world_size) }.expect("Failed to init NCCL comms");
let contexts: Vec<Arc<CudaContext>> = (0..world_size)
.map(|i| CudaContext::new(i).expect("Failed to create CUDA context"))
.collect();
let streams: Vec<Arc<CudaStream>> = contexts
.iter()
.map(|ctx| ctx.new_stream().expect("Failed to create stream"))
.collect();
let test_size = 64 * 1024 * 1024;
println!("Transfer size: {} MB", test_size / (1024 * 1024));
let buffers: Vec<CudaSlice<u8>> = streams
.iter()
.map(|stream| {
let zeros = vec![0u8; test_size];
stream.clone_htod(&zeros).expect("Failed to allocate")
})
.collect();
{
let host_data: Vec<u8> = (0..test_size).map(|i| (i % 256) as u8).collect();
let src_buffer = streams[0]
.clone_htod(&host_data)
.expect("Failed to copy to device 0");
let src_ptr = get_device_ptr(&src_buffer, &streams[0]);
let dst_ptr = get_device_ptr(&buffers[0], &streams[0]);
unsafe {
cudarc::driver::result::memcpy_dtod_async(
dst_ptr as u64,
src_ptr as u64,
test_size,
streams[0].cu_stream(),
)
.expect("dtod copy failed");
}
streams[0].synchronize().expect("sync failed");
}
let buffer_ptrs: Vec<usize> = buffers
.iter()
.zip(streams.iter())
.map(|(buf, stream)| get_device_ptr(buf, stream))
.collect();
let barrier = Arc::new(Barrier::new(world_size));
let start = std::time::Instant::now();
let handles: Vec<_> = (0..world_size)
.map(|rank| {
let comm = comms[rank]; let stream = streams[rank].clone();
let buffer_ptr = buffer_ptrs[rank];
let barrier = barrier.clone();
thread::spawn(move || {
barrier.wait();
let result = unsafe {
ncclBcast(
buffer_ptr as *mut std::ffi::c_void,
test_size,
ncclDataType_t::ncclChar,
0,
comm as ncclComm_t,
stream.cu_stream().cast(),
)
};
check_nccl_result(result).expect("ncclBcast failed");
stream.synchronize().expect("Stream sync failed");
})
})
.collect();
for handle in handles {
handle.join().expect("Thread panicked");
}
let elapsed = start.elapsed();
let throughput_gbps =
(test_size as f64 / (1024.0 * 1024.0 * 1024.0)) / elapsed.as_secs_f64();
println!(
"Transfer completed in {:?} ({:.2} GB/s)",
elapsed, throughput_gbps
);
{
let host_data = streams[1]
.clone_dtoh(&buffers[1])
.expect("Failed to copy from device 1");
let samples = [
0,
test_size / 4,
test_size / 2,
test_size * 3 / 4,
test_size - 1,
];
for &idx in &samples {
let expected = (idx % 256) as u8;
assert_eq!(
host_data[idx], expected,
"Mismatch at index {}: expected {}, got {}",
idx, expected, host_data[idx]
);
}
let mismatch_count = host_data
.iter()
.enumerate()
.filter(|(i, b)| **b != (*i % 256) as u8)
.count();
assert_eq!(
mismatch_count, 0,
"Found {} mismatched bytes",
mismatch_count
);
}
unsafe { destroy_comms(&comms) };
println!("Large transfer test passed!");
}
}