//! Communication primitives for distributed training //! //! This module provides implementations of collective communication operations //! including AllReduce, Broadcast, AllGather, ReduceScatter, and point-to-point //! operations like Send/Recv. use crate::error::{DistributedError, Result}; use crate::group::ProcessGroup; use crate::tensor_ext::TensorExt; use rtx_tensor::Tensor; use serde::{Deserialize, Serialize}; use std::fmt; #[cfg(feature = "hpc-channels")] use std::sync::atomic::{AtomicU64, Ordering}; /// Global step counter for gradient sync events #[cfg(feature = "hpc-channels")] static GRADIENT_SYNC_STEP: AtomicU64 = AtomicU64::new(0); /// Reduction operations for collective communication #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] pub enum ReduceOp { /// Sum all values Sum, /// Take maximum value Max, /// Take minimum value Min, /// Multiply all values Product, /// Logical AND (for boolean tensors) And, /// Logical OR (for boolean tensors) Or, /// Bitwise AND (for integer tensors) BitAnd, /// Bitwise OR (for integer tensors) BitOr, /// Bitwise XOR (for integer tensors) BitXor, } impl fmt::Display for ReduceOp { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { match self { Self::Sum => write!(f, "sum"), Self::Max => write!(f, "max"), Self::Min => write!(f, "min"), Self::Product => write!(f, "product"), Self::And => write!(f, "and"), Self::Or => write!(f, "or"), Self::BitAnd => write!(f, "bitand"), Self::BitOr => write!(f, "bitor"), Self::BitXor => write!(f, "bitxor"), } } } /// Output type for AllGather operations #[derive(Debug, Clone)] pub enum AllGatherOutput { /// Single concatenated tensor containing all gathered data Tensor(Tensor), /// List of tensors, one from each process TensorList(Vec), } /// AllReduce operation trait #[derive(Debug, Clone, Copy)] pub struct AllReduceOp; impl AllReduceOp { /// Apply reduction operation to tensor data pub fn apply(op: ReduceOp, data: &mut [f32], world_size: i32) -> Result<()> { match op { ReduceOp::Sum => { // For simulation, multiply by world_size to simulate sum from all ranks for value in data.iter_mut() { *value *= world_size as f32; } } ReduceOp::Max => { // Max operation keeps values unchanged in simulation // In real implementation, would take max across all ranks } ReduceOp::Min => { // Min operation keeps values unchanged in simulation // In real implementation, would take min across all ranks } ReduceOp::Product => { // Product operation for value in data.iter_mut() { *value = value.powf(world_size as f32); } } _ => { return Err(DistributedError::communication( "allreduce", format!("unsupported reduce operation: {op}"), )); } } Ok(()) } } /// Communication primitive trait for process groups #[async_trait::async_trait] pub trait CommunicationPrimitive { /// AllReduce collective operation /// /// Reduces tensors across all processes and distributes the result async fn allreduce(&self, tensor: &mut Tensor, op: ReduceOp) -> Result<()>; /// Broadcast operation /// /// Broadcasts a tensor from the root process to all other processes async fn broadcast(&self, tensor: &mut Tensor, root: i32) -> Result<()>; /// AllGather collective operation /// /// Gathers tensors from all processes and makes them available to all async fn allgather(&self, tensor: &Tensor) -> Result; /// ReduceScatter collective operation /// /// Reduces tensors across all processes and scatters the results async fn reduce_scatter(&self, tensor: &Tensor, op: ReduceOp) -> Result; /// Reduce operation /// /// Reduces tensors from all processes to the root process async fn reduce(&self, tensor: &mut Tensor, op: ReduceOp, root: i32) -> Result<()>; /// Gather operation /// /// Gathers tensors from all processes to the root process async fn gather(&self, tensor: &Tensor, root: i32) -> Result>>; /// Scatter operation /// /// Scatters tensor from root process to all other processes async fn scatter(&self, input: &Tensor, output: &mut Tensor, root: i32) -> Result<()>; /// Point-to-point send operation async fn send(&self, tensor: &Tensor, dst: i32) -> Result<()>; /// Point-to-point receive operation async fn recv(&self, tensor: &mut Tensor, src: i32) -> Result<()>; /// Non-blocking send operation async fn isend(&self, tensor: &Tensor, dst: i32) -> Result; /// Non-blocking receive operation async fn irecv(&self, tensor: &mut Tensor, src: i32) -> Result; } /// Handle for non-blocking communication operations #[derive(Debug)] pub struct CommHandle { /// Unique identifier for the operation pub id: uuid::Uuid, /// Whether the operation has completed pub completed: bool, /// Optional result of the operation pub result: Option>, } impl CommHandle { /// Create a new communication handle pub fn new() -> Self { Self { id: uuid::Uuid::new_v4(), completed: false, result: None, } } /// Wait for the operation to complete pub async fn wait(&mut self) -> Result<()> { // Simulate async completion if !self.completed { tokio::time::sleep(std::time::Duration::from_millis(1)).await; self.completed = true; self.result = Some(Ok(())); } self.result.take().unwrap_or(Ok(())) } /// Check if the operation has completed without waiting pub fn test(&self) -> bool { self.completed } } impl Default for CommHandle { fn default() -> Self { Self::new() } } /// Implementation of communication primitives for ProcessGroup #[async_trait::async_trait] impl CommunicationPrimitive for ProcessGroup { async fn allreduce(&self, tensor: &mut Tensor, op: ReduceOp) -> Result<()> { // Validate tensor if tensor.shape().dims().iter().product::() == 0 { return Err(DistributedError::tensor("cannot allreduce empty tensor")); } #[cfg(feature = "hpc-channels")] let start_time = std::time::Instant::now(); tracing::debug!( "AllReduce operation {} on tensor shape {:?} for rank {}/{}", op, tensor.shape(), self.rank(), self.world_size() ); // Use RNCCL backend if available #[cfg(feature = "rnccl")] { if let Some(rnccl_backend) = self.rnccl_backend() { return rnccl_backend.allreduce(tensor, op).await; } } // Get tensor data and apply operation (modify in-place via new tensor) let mut data = tensor.data()?; #[cfg(feature = "hpc-channels")] let tensor_bytes = data.len() * std::mem::size_of::(); // Apply the reduction operation (simulated for testing) AllReduceOp::apply(op, &mut data, self.world_size() as i32)?; // Create new tensor with modified data *tensor = Tensor::from_data(data, tensor.shape().dims().to_vec(), tensor.device())?; // Simulate communication latency let latency = std::time::Duration::from_micros(10 * self.world_size() as u64); tokio::time::sleep(latency).await; // Publish gradient sync event to HPC channels #[cfg(feature = "hpc-channels")] { let step = GRADIENT_SYNC_STEP.fetch_add(1, Ordering::Relaxed); let duration_ms = start_time.elapsed().as_secs_f64() * 1000.0; // Get or create the gradient sync channel and publish let tx = hpc_channels::broadcast::( hpc_channels::channels::TORCH_GRADIENT_SYNC, 256, ); let _ = tx.send(crate::hpc_bridge::GradientSyncEvent { step, total_bytes: tensor_bytes as u64, duration_ms, algorithm: "ring".to_string(), world_size: self.world_size() as u32, timestamp_ms: std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) .map(|d| d.as_millis() as u64) .unwrap_or(0), }); } Ok(()) } async fn broadcast(&self, tensor: &mut Tensor, root: i32) -> Result<()> { if root < 0 || root >= self.world_size() as i32 { return Err(DistributedError::communication( "broadcast", format!("invalid root rank: {root}"), )); } tracing::debug!( "Broadcast from rank {} to rank {}/{} with tensor shape {:?}", root, self.rank(), self.world_size(), tensor.shape() ); // Use RNCCL backend if available #[cfg(feature = "rnccl")] { if let Some(rnccl_backend) = self.rnccl_backend() { return rnccl_backend.broadcast(tensor, root).await; } } // If this is not the root, simulate receiving data if self.rank() != root as usize { // In a real implementation, this would receive data from root // For simulation, we'll modify the tensor to indicate broadcast let mut data = tensor.data()?; for value in &mut data { *value = root as f32; // Set to root rank value for testing } *tensor = Tensor::from_data(data, tensor.shape().dims().to_vec(), tensor.device())?; } // Simulate communication latency tokio::time::sleep(std::time::Duration::from_millis(1)).await; Ok(()) } async fn allgather(&self, tensor: &Tensor) -> Result { tracing::debug!( "AllGather on tensor shape {:?} for rank {}/{}", tensor.shape(), self.rank(), self.world_size() ); // Use RNCCL backend if available #[cfg(feature = "rnccl")] { if let Some(rnccl_backend) = self.rnccl_backend() { return rnccl_backend.allgather(tensor).await; } } let input_shape = tensor.shape(); let world_size = self.world_size(); // Create output tensor with concatenated shape let mut output_dims = input_shape.dims().to_vec(); if !output_dims.is_empty() { output_dims[0] *= world_size; } else { output_dims.push(world_size); } let output_shape = crate::TensorShape::new(output_dims)?; // Simulate gathering data from all ranks let output_tensor = { let input_data = tensor.data()?; let mut output_data = vec![0.0f32; output_shape.numel()]; for rank in 0..world_size { let offset = rank * input_data.len(); for (i, &value) in input_data.iter().enumerate() { if offset + i < output_data.len() { // Simulate data from different ranks output_data[offset + i] = value + rank as f32; } } } Tensor::from_data( output_data, output_shape.dims().to_vec(), &crate::Device::default(), )? }; // Simulate communication latency tokio::time::sleep(std::time::Duration::from_millis(2)).await; Ok(AllGatherOutput::Tensor(output_tensor)) } async fn reduce_scatter(&self, tensor: &Tensor, op: ReduceOp) -> Result { // Use RNCCL backend if available #[cfg(feature = "rnccl")] { if let Some(rnccl_backend) = self.rnccl_backend() { return rnccl_backend.reduce_scatter(tensor, op).await; } } let world_size = self.world_size(); let input_shape = tensor.shape(); // Output shape should be input_shape[0] / world_size if input_shape.dims().is_empty() || !input_shape.dims()[0].is_multiple_of(world_size) { return Err(DistributedError::tensor( "tensor first dimension must be divisible by world size", )); } let mut output_dims = input_shape.dims().to_vec(); output_dims[0] /= world_size; let output_shape = crate::TensorShape::new(output_dims)?; // Simulate reduce-scatter operation let input_data = tensor.data()?; let mut output_data = vec![0.0f32; output_shape.numel()]; let chunk_size = input_data.len() / world_size; let my_chunk_start = self.rank() * chunk_size; for i in 0..chunk_size.min(output_data.len()) { if my_chunk_start + i < input_data.len() { output_data[i] = match op { ReduceOp::Sum => input_data[my_chunk_start + i] * world_size as f32, ReduceOp::Max | ReduceOp::Min => input_data[my_chunk_start + i], _ => { return Err(DistributedError::communication( "reduce_scatter", format!("unsupported reduce operation: {op}"), )); } }; } } let output_tensor = Tensor::from_data( output_data, output_shape.dims().to_vec(), &crate::Device::default(), )?; tokio::time::sleep(std::time::Duration::from_millis(1)).await; Ok(output_tensor) } 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( "reduce", format!("invalid root rank: {root}"), )); } // Use RNCCL backend if available #[cfg(feature = "rnccl")] { if let Some(rnccl_backend) = self.rnccl_backend() { return rnccl_backend.reduce(tensor, op, root).await; } } // Only the root rank gets the reduced result if self.rank() == root as usize { let mut data = tensor.data()?; AllReduceOp::apply(op, &mut data, self.world_size() as i32)?; *tensor = Tensor::from_data( data, tensor.shape().dims().to_vec(), &crate::Device::default(), )?; } tokio::time::sleep(std::time::Duration::from_millis(1)).await; Ok(()) } async fn gather(&self, tensor: &Tensor, root: i32) -> Result>> { if root < 0 || root >= self.world_size() as i32 { return Err(DistributedError::communication( "gather", format!("invalid root rank: {root}"), )); } // Use RNCCL backend if available #[cfg(feature = "rnccl")] { if let Some(rnccl_backend) = self.rnccl_backend() { return rnccl_backend.gather(tensor, root).await; } } if self.rank() == root as usize { // Root collects all tensors let mut gathered = Vec::with_capacity(self.world_size()); for rank in 0..self.world_size() { // Simulate receiving from each rank let rank_tensor = { let source_data = tensor.data()?; let mut rank_data = vec![0.0f32; source_data.len()]; for (i, &value) in source_data.iter().enumerate() { if i < rank_data.len() { rank_data[i] = value + rank as f32; } } Tensor::from_data( rank_data, tensor.shape().dims().to_vec(), &crate::Device::default(), )? }; gathered.push(rank_tensor); } tokio::time::sleep(std::time::Duration::from_millis(1)).await; Ok(Some(gathered)) } else { // Non-root ranks don't receive anything tokio::time::sleep(std::time::Duration::from_millis(1)).await; 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( "scatter", format!("invalid root rank: {root}"), )); } // Use RNCCL backend if available #[cfg(feature = "rnccl")] { if let Some(rnccl_backend) = self.rnccl_backend() { return rnccl_backend.scatter(input, output, root).await; } } let world_size = self.world_size(); let input_data = input.data()?; let chunk_size = input_data.len() / world_size; let my_chunk_start = self.rank() * chunk_size; // Copy the appropriate chunk to output let mut output_data = vec![0.0f32; chunk_size]; for i in 0..chunk_size { if my_chunk_start + i < input_data.len() { output_data[i] = input_data[my_chunk_start + i]; } } *output = Tensor::from_data(output_data, vec![chunk_size], &crate::Device::default())?; tokio::time::sleep(std::time::Duration::from_millis(1)).await; 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( "send", format!("invalid destination rank: {dst}"), )); } // Use RNCCL backend if available #[cfg(feature = "rnccl")] { if let Some(rnccl_backend) = self.rnccl_backend() { return rnccl_backend.send(tensor, dst).await; } } tracing::debug!( "Send tensor shape {:?} from rank {} to rank {}", tensor.shape(), self.rank(), dst ); // Simulate send latency tokio::time::sleep(std::time::Duration::from_micros(100)).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( "recv", format!("invalid source rank: {src}"), )); } // Use RNCCL backend if available #[cfg(feature = "rnccl")] { if let Some(rnccl_backend) = self.rnccl_backend() { return rnccl_backend.recv(tensor, src).await; } } tracing::debug!( "Recv tensor shape {:?} at rank {} from rank {}", tensor.shape(), self.rank(), src ); // Simulate receiving data { let data = tensor.data_mut(); for (i, value) in data.iter_mut().enumerate() { *value = src as f32 + i as f32 * 0.1; // Simulate received data } } // Simulate recv latency tokio::time::sleep(std::time::Duration::from_micros(100)).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( "isend", format!("invalid destination rank: {dst}"), )); } // Use RNCCL backend if available #[cfg(feature = "rnccl")] { if let Some(rnccl_backend) = self.rnccl_backend() { return rnccl_backend.isend(tensor, dst).await; } } tracing::debug!( "Non-blocking send tensor shape {:?} from rank {} to rank {}", tensor.shape(), self.rank(), dst ); let handle = CommHandle::new(); // In a real implementation, would start async send operation // For simulation, we return a handle that completes quickly Ok(handle) } 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( "irecv", format!("invalid source rank: {src}"), )); } // Use RNCCL backend if available #[cfg(feature = "rnccl")] { if let Some(rnccl_backend) = self.rnccl_backend() { return rnccl_backend.irecv(tensor, src).await; } } tracing::debug!( "Non-blocking recv tensor shape {:?} at rank {} from rank {}", tensor.shape(), self.rank(), src ); let handle = CommHandle::new(); // In a real implementation, would start async recv operation // For simulation, we return a handle that completes quickly Ok(handle) } } #[cfg(test)] mod tests { use super::*; #[allow(unused_imports)] use crate::{Backend, BackendConfig, ProcessGroup}; #[test] fn test_reduce_op_display() { assert_eq!(format!("{}", ReduceOp::Sum), "sum"); assert_eq!(format!("{}", ReduceOp::Max), "max"); assert_eq!(format!("{}", ReduceOp::BitXor), "bitxor"); } #[test] fn test_allreduce_op_apply() { let mut data = vec![1.0, 2.0, 3.0]; AllReduceOp::apply(ReduceOp::Sum, &mut data, 2).unwrap(); assert_eq!(data, vec![2.0, 4.0, 6.0]); } #[test] fn test_comm_handle() { let mut handle = CommHandle::new(); assert!(!handle.completed); assert!(!handle.test()); } #[tokio::test] async fn test_comm_handle_wait() { let mut handle = CommHandle::new(); let result = handle.wait().await; assert!(result.is_ok()); assert!(handle.completed); } #[tokio::test] async fn test_communication_primitive() { use crate::WorldInfo; let world_info = WorldInfo::new(2, 0, Backend::Cpu); let pg = ProcessGroup::new(Backend::Cpu, world_info).unwrap(); // Test AllReduce let shape_dims = vec![3]; let mut tensor = Tensor::ones(&shape_dims, &crate::Device::default()).unwrap(); assert!(pg.allreduce(&mut tensor, ReduceOp::Sum).await.is_ok()); // Test Broadcast let mut tensor = Tensor::ones(&shape_dims, &crate::Device::default()).unwrap(); assert!(pg.broadcast(&mut tensor, 0).await.is_ok()); } }