//! NCCL-based communication primitives implementation //! //! This module provides NCCL-native implementations of all collective communication //! operations, replacing the custom ring-allreduce algorithms with optimal NCCL //! collective operations for superior multi-GPU performance. use crate::comm::{AllGatherOutput, CommHandle, CommunicationPrimitive, ReduceOp}; use crate::error::{DistributedError, Result}; use crate::nccl::NcclBackend; use crate::{Tensor, TensorShape}; use async_trait::async_trait; use std::sync::Arc; /// NCCL-based communication primitive implementation pub struct NcclCommunicationPrimitive { /// NCCL backend reference backend: Arc, /// World size world_size: usize, /// Rank rank: usize, } impl std::fmt::Debug for NcclCommunicationPrimitive { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("NcclCommunicationPrimitive") .field("world_size", &self.world_size) .field("rank", &self.rank) .finish() } } impl NcclCommunicationPrimitive { /// Create a new NCCL communication primitive pub fn new(backend: Arc, world_size: usize, rank: usize) -> Self { Self { backend, world_size, rank, } } /// Get world communicator from backend fn get_world_comm(&self) -> Result> { let comm_guard = self.backend.world_comm(); let guard = comm_guard.read(); guard .as_ref() .ok_or_else(|| { DistributedError::communication("nccl", "world communicator not initialized") }) .cloned() } /// Validate tensor for communication operations fn validate_tensor(&self, tensor: &Tensor) -> Result<()> { if tensor.shape().dims().iter().product::() == 0 { return Err(DistributedError::tensor("cannot communicate empty tensor")); } // Check if tensor is on a supported device match tensor.device() { crate::Device::Cuda(_) => Ok(()), _ => { // Non-CUDA tensors need to be moved to GPU for NCCL operations // This includes CPU tensors if the cpu feature is enabled Ok(()) } } } /// Convert tensor to appropriate device for NCCL operation async fn prepare_tensor_for_nccl(&self, tensor: &Tensor) -> Result { match tensor.device() { crate::Device::Cuda(_) => Ok(tensor.clone()), _ => { // Move non-CUDA tensor to GPU let gpu_device = crate::Device::Cuda(self.rank); // Use rank as device ID Ok(tensor.to_device(&gpu_device)?) } } } } #[async_trait] impl CommunicationPrimitive for NcclCommunicationPrimitive { async fn allreduce(&self, tensor: &mut Tensor, op: ReduceOp) -> Result<()> { self.validate_tensor(tensor)?; // Get the world communicator let comm = self.get_world_comm()?; tracing::debug!( "NCCL AllReduce operation {} on tensor shape {:?} for rank {}/{}", op, tensor.shape(), self.rank, self.world_size ); // Use NCCL's native allreduce implementation comm.allreduce(tensor, op).await?; Ok(()) } async fn broadcast(&self, tensor: &mut Tensor, root: i32) -> Result<()> { if root < 0 || root >= self.world_size as i32 { return Err(DistributedError::communication( "nccl", format!("invalid root rank: {root}"), )); } self.validate_tensor(tensor)?; // Get the world communicator let comm = self.get_world_comm()?; tracing::debug!( "NCCL Broadcast from rank {} to rank {}/{} with tensor shape {:?}", root, self.rank, self.world_size, tensor.shape() ); // Use NCCL's native broadcast implementation comm.broadcast(tensor, root).await?; Ok(()) } async fn allgather(&self, tensor: &Tensor) -> Result { self.validate_tensor(tensor)?; // Get the world communicator let comm = self.get_world_comm()?; tracing::debug!( "NCCL AllGather on tensor shape {:?} for rank {}/{}", tensor.shape(), self.rank, self.world_size ); // Use NCCL's native allgather implementation let result_tensor = comm.allgather(tensor).await?; Ok(AllGatherOutput::Tensor(result_tensor)) } async fn reduce_scatter(&self, tensor: &Tensor, op: ReduceOp) -> Result { self.validate_tensor(tensor)?; let input_count = tensor.numel(); if !input_count.is_multiple_of(self.world_size) { return Err(DistributedError::tensor( "tensor size must be divisible by world size for reduce_scatter", )); } // Get the world communicator let comm = self.get_world_comm()?; tracing::debug!( "NCCL ReduceScatter operation {} on tensor shape {:?} for rank {}/{}", op, tensor.shape(), self.rank, self.world_size ); // Use NCCL's native reduce_scatter implementation let result = comm.reduce_scatter(tensor, op).await?; Ok(result) } async fn reduce(&self, tensor: &mut Tensor, op: ReduceOp, root: i32) -> Result<()> { if root < 0 || root >= self.world_size as i32 { return Err(DistributedError::communication( "nccl", format!("invalid root rank: {root}"), )); } self.validate_tensor(tensor)?; tracing::debug!( "NCCL Reduce operation {} to root {} for rank {}/{}", op, root, self.rank, self.world_size ); // NCCL doesn't have a native reduce operation, so we implement it // using allreduce followed by broadcast of zeros to non-root ranks if self.rank == root as usize { // Root performs allreduce to get the result let mut temp_tensor = tensor.clone(); self.allreduce(&mut temp_tensor, op).await?; *tensor = temp_tensor; } else { // Non-root ranks participate in allreduce but don't keep the result let mut temp_tensor = tensor.clone(); self.allreduce(&mut temp_tensor, op).await?; // Clear the tensor for non-root ranks *tensor = Tensor::zeros(tensor.shape().clone(), tensor.device())?; } Ok(()) } async fn gather(&self, tensor: &Tensor, root: i32) -> Result>> { if root < 0 || root >= self.world_size as i32 { return Err(DistributedError::communication( "nccl", format!("invalid root rank: {root}"), )); } self.validate_tensor(tensor)?; tracing::debug!( "NCCL Gather to root {} for rank {}/{}", root, self.rank, self.world_size ); // Implement gather using allgather followed by filtering let all_gathered = self.allgather(tensor).await?; if self.rank == root as usize { // Root gets all tensors match all_gathered { AllGatherOutput::Tensor(gathered_tensor) => { // Split the gathered tensor into individual tensors let mut result = Vec::new(); let input_size = tensor.numel(); let total_size = gathered_tensor.numel(); if total_size != input_size * self.world_size { return Err(DistributedError::tensor("gathered tensor size mismatch")); } // Create views/slices of the gathered tensor for each rank for rank in 0..self.world_size { let start_idx = rank * input_size; let end_idx = start_idx + input_size; // For now, create a copy. In a real implementation, // we would create tensor views/slices let rank_data = gathered_tensor.data()?; let rank_slice = &rank_data[start_idx..end_idx]; let rank_tensor = Tensor::from_data( rank_slice.to_vec(), tensor.shape().dims().to_vec(), tensor.device(), )?; result.push(rank_tensor); } Ok(Some(result)) } AllGatherOutput::TensorList(tensors) => Ok(Some(tensors)), } } else { // Non-root ranks don't get the result Ok(None) } } async fn scatter(&self, input: &Tensor, output: &mut Tensor, root: i32) -> Result<()> { if root < 0 || root >= self.world_size as i32 { return Err(DistributedError::communication( "nccl", format!("invalid root rank: {root}"), )); } tracing::debug!( "NCCL Scatter from root {} for rank {}/{}", root, self.rank, self.world_size ); // NCCL doesn't have native scatter, so we implement it using broadcast // and then each rank takes its portion if self.rank == root as usize { // Root broadcasts the input tensor let mut broadcast_tensor = input.clone(); self.broadcast(&mut broadcast_tensor, root).await?; // Root takes its portion let input_data = broadcast_tensor.data()?; let chunk_size = input_data.len() / self.world_size; let my_chunk_start = self.rank * chunk_size; let my_chunk_end = my_chunk_start + chunk_size; let my_data = input_data[my_chunk_start..my_chunk_end].to_vec(); *output = Tensor::from_data(my_data, vec![chunk_size], output.device())?; } else { // Non-root ranks receive the broadcast and take their portion let input_size = input.numel(); let mut broadcast_tensor = Tensor::zeros(TensorShape::new(vec![input_size])?, input.device())?; self.broadcast(&mut broadcast_tensor, root).await?; // Take our portion let input_data = broadcast_tensor.data()?; let chunk_size = input_data.len() / self.world_size; let my_chunk_start = self.rank * chunk_size; let my_chunk_end = my_chunk_start + chunk_size; let my_data = input_data[my_chunk_start..my_chunk_end].to_vec(); *output = Tensor::from_data(my_data, vec![chunk_size], output.device())?; } Ok(()) } async fn send(&self, tensor: &Tensor, dst: i32) -> Result<()> { if dst < 0 || dst >= self.world_size as i32 || dst == self.rank as i32 { return Err(DistributedError::communication( "nccl", format!("invalid destination rank: {dst}"), )); } self.validate_tensor(tensor)?; // Get the world communicator let comm = self.get_world_comm()?; tracing::debug!( "NCCL Send tensor shape {:?} from rank {} to rank {}", tensor.shape(), self.rank, dst ); // Use NCCL's native send implementation comm.send(tensor, dst).await?; Ok(()) } async fn recv(&self, tensor: &mut Tensor, src: i32) -> Result<()> { if src < 0 || src >= self.world_size as i32 || src == self.rank as i32 { return Err(DistributedError::communication( "nccl", format!("invalid source rank: {src}"), )); } self.validate_tensor(tensor)?; // Get the world communicator let comm = self.get_world_comm()?; tracing::debug!( "NCCL Recv tensor shape {:?} at rank {} from rank {}", tensor.shape(), self.rank, src ); // Use NCCL's native recv implementation comm.recv(tensor, src).await?; Ok(()) } async fn isend(&self, tensor: &Tensor, dst: i32) -> Result { if dst < 0 || dst >= self.world_size as i32 || dst == self.rank as i32 { return Err(DistributedError::communication( "nccl", format!("invalid destination rank: {dst}"), )); } tracing::debug!( "NCCL Non-blocking send tensor shape {:?} from rank {} to rank {}", tensor.shape(), self.rank, dst ); // Get communicator and perform blocking send for now // In a real implementation, this would use NCCL's stream-based operations let comm = self.get_world_comm()?; // Perform the send operation (guard is dropped here) comm.send(tensor, dst).await?; // Return completed handle Ok(CommHandle::new()) } async fn irecv(&self, tensor: &mut Tensor, src: i32) -> Result { if src < 0 || src >= self.world_size as i32 || src == self.rank as i32 { return Err(DistributedError::communication( "nccl", format!("invalid source rank: {src}"), )); } tracing::debug!( "NCCL Non-blocking recv tensor shape {:?} at rank {} from rank {}", tensor.shape(), self.rank, src ); // Get communicator and perform blocking recv for now // In a real implementation, this would use NCCL's stream-based operations let comm = self.get_world_comm()?; // Perform the recv operation (guard is dropped here) comm.recv(tensor, src).await?; // Return completed handle Ok(CommHandle::new()) } } /// Validate that a tensor is suitable for NCCL operations /// /// This function checks that: /// - The tensor is not empty /// - The tensor is on a CUDA device (for NCCL operations) pub fn validate_tensor_for_nccl(tensor: &Tensor) -> Result<()> { // Check that tensor is not empty if tensor.numel() == 0 { return Err(DistributedError::tensor( "cannot use NCCL with empty tensor", )); } // Check that tensor is on CUDA device match tensor.device() { crate::Device::Cuda(_) => Ok(()), _ => { // For now, allow CPU tensors as they may be transferred to GPU // In production, this should be stricter Ok(()) } } } #[cfg(test)] mod tests { use super::*; use crate::nccl::{NcclBackend, NcclConfig}; #[tokio::test] async fn test_nccl_communication_primitive_creation() { let config = NcclConfig::default(); let backend = Arc::new(NcclBackend::new(config)); let comm = NcclCommunicationPrimitive::new(backend, 2, 0); assert_eq!(comm.world_size, 2); assert_eq!(comm.rank, 0); } #[tokio::test] async fn test_tensor_validation() { let config = NcclConfig::default(); let backend = Arc::new(NcclBackend::new(config)); let comm = NcclCommunicationPrimitive::new(backend, 2, 0); // Test empty tensor validation let empty_tensor = Tensor::zeros( TensorShape::new(vec![0]).unwrap(), &crate::Device::default(), ) .unwrap(); assert!(comm.validate_tensor(&empty_tensor).is_err()); // Test valid tensor validation let valid_tensor = Tensor::ones(&vec![3, 3], &crate::Device::default()).unwrap(); assert!(comm.validate_tensor(&valid_tensor).is_ok()); } }