//! Resharding operations for DTensor redistribution. //! //! This module provides efficient implementations of tensor resharding operations //! needed for distributed training, including: //! - AllGather: Shard -> Replicate //! - ReduceScatter: Partial -> Shard //! - AllToAll: Shard(dim_a) -> Shard(dim_b) //! - AllReduce: Partial -> Replicate //! //! These operations are fundamental to FSDP2's per-parameter sharding approach. use crate::comm::{CommunicationPrimitive, ReduceOp}; use crate::device_mesh::DeviceMesh; use crate::dtensor::{DTensor, PartialReduceOp, Placement, TensorSpec}; use crate::error::{DistributedError, Result}; use crate::group::ProcessGroup; use rtx_tensor::{Device, Tensor}; use std::sync::Arc; // ============================================================================= // Resharding Operations // ============================================================================= /// Reshard a DTensor to a new placement specification. /// /// This function determines the required communication pattern and /// executes the resharding operation. pub async fn reshard_dtensor( dtensor: &DTensor, target_placements: &[Placement], pg: &ProcessGroup, ) -> Result { let current_placements = &dtensor.spec().placements; // No resharding needed if placements match if current_placements == target_placements { return Ok(dtensor.clone()); } // Handle each mesh dimension's resharding let mut result = dtensor.clone(); for (mesh_dim, (src, dst)) in current_placements .iter() .zip(target_placements.iter()) .enumerate() { if src == dst { continue; } result = match (src, dst) { // Shard -> Replicate: AllGather (Placement::Shard { tensor_dim }, Placement::Replicate) => { all_gather_reshard(&result, mesh_dim, *tensor_dim, pg).await? } // Replicate -> Shard: Local slice (no communication) (Placement::Replicate, Placement::Shard { tensor_dim }) => { replicate_to_shard(&result, mesh_dim, *tensor_dim)? } // Partial -> Replicate: AllReduce (Placement::Partial { reduce_op }, Placement::Replicate) => { all_reduce_reshard(&result, mesh_dim, *reduce_op, pg).await? } // Partial -> Shard: ReduceScatter (Placement::Partial { reduce_op }, Placement::Shard { tensor_dim }) => { reduce_scatter_reshard(&result, mesh_dim, *tensor_dim, *reduce_op, pg).await? } // Shard -> Shard (different dims): AllToAll ( Placement::Shard { tensor_dim: src_dim, }, Placement::Shard { tensor_dim: dst_dim, }, ) if src_dim != dst_dim => { all_to_all_reshard(&result, mesh_dim, *src_dim, *dst_dim, pg).await? } // Same placement type - no action needed (src_p, dst_p) if std::mem::discriminant(src_p) == std::mem::discriminant(dst_p) => { result } // Unsupported transition _ => { return Err(DistributedError::configuration(format!( "Unsupported resharding: {:?} -> {:?}", src, dst ))); } }; } Ok(result) } // ============================================================================= // AllGather: Shard -> Replicate // ============================================================================= /// Perform AllGather to convert sharded tensor to replicated. async fn all_gather_reshard( dtensor: &DTensor, _mesh_dim: usize, tensor_dim: usize, pg: &ProcessGroup, ) -> Result { let local_shard = dtensor.local_shard().clone(); let gathered = pg.allgather(&local_shard).await?; // Concatenate gathered tensors along the sharded dimension let full_tensor = match gathered { crate::comm::AllGatherOutput::Tensor(t) => t, crate::comm::AllGatherOutput::TensorList(tensors) => { concat_along_dim(&tensors, tensor_dim)? } }; // Create new DTensor with replicated placement let new_spec = TensorSpec::new(dtensor.global_shape().to_vec()) .with_placement(0, Placement::Replicate) .with_requires_grad(dtensor.requires_grad()); DTensor::from_local(full_tensor, new_spec, dtensor.mesh().clone()) } /// Concatenate tensors along a specific dimension. fn concat_along_dim(tensors: &[Tensor], dim: usize) -> Result { if tensors.is_empty() { return Err(DistributedError::tensor( "Cannot concatenate empty tensor list", )); } let first_shape = tensors[0].shape().dims().to_vec(); if dim >= first_shape.len() { return Err(DistributedError::tensor(format!( "Concatenation dim {} out of range for {}D tensor", dim, first_shape.len() ))); } // Calculate output shape let total_dim_size: usize = tensors .iter() .map(|t| t.shape().dims().get(dim).copied().unwrap_or(0)) .sum(); let mut output_shape = first_shape.clone(); output_shape[dim] = total_dim_size; // Simple case: dim 0 concatenation if dim == 0 { let mut combined_data = Vec::new(); for tensor in tensors { let data = tensor .data() .map_err(|e| DistributedError::tensor(e.to_string()))?; combined_data.extend(data); } return Tensor::from_data(combined_data, output_shape, &Device::cpu()) .map_err(|e| DistributedError::tensor(e.to_string())); } // General case: need to interleave data let output_numel: usize = output_shape.iter().product(); let mut output_data = vec![0.0f32; output_numel]; // Calculate strides for output let mut output_strides: Vec = vec![1; output_shape.len()]; for i in (0..output_shape.len() - 1).rev() { output_strides[i] = output_strides[i + 1] * output_shape[i + 1]; } // Copy each tensor's data to the correct position let mut offset_in_dim = 0; for tensor in tensors { let tensor_shape = tensor.shape().dims().to_vec(); let tensor_data = tensor .data() .map_err(|e| DistributedError::tensor(e.to_string()))?; let dim_size = tensor_shape[dim]; // Calculate tensor strides let mut tensor_strides: Vec = vec![1; tensor_shape.len()]; for i in (0..tensor_shape.len() - 1).rev() { tensor_strides[i] = tensor_strides[i + 1] * tensor_shape[i + 1]; } // Copy data let tensor_numel: usize = tensor_shape.iter().product(); for flat_idx in 0..tensor_numel { // Convert to coordinates let mut coords = vec![0; tensor_shape.len()]; let mut remaining = flat_idx; for d in 0..tensor_shape.len() { coords[d] = remaining / tensor_strides[d]; remaining %= tensor_strides[d]; } // Adjust coordinate for concatenation dimension coords[dim] += offset_in_dim; // Calculate output flat index let output_flat_idx: usize = coords .iter() .zip(output_strides.iter()) .map(|(c, s)| c * s) .sum(); output_data[output_flat_idx] = tensor_data[flat_idx]; } offset_in_dim += dim_size; } Tensor::from_data(output_data, output_shape, &Device::cpu()) .map_err(|e| DistributedError::tensor(e.to_string())) } // ============================================================================= // Replicate -> Shard: Local Slice // ============================================================================= /// Extract local shard from replicated tensor (no communication). fn replicate_to_shard(dtensor: &DTensor, mesh_dim: usize, tensor_dim: usize) -> Result { let mesh = dtensor.mesh(); let mesh_size = mesh.shape().get(mesh_dim).copied().unwrap_or(1); let mesh_coord = dtensor.local_coord().get(mesh_dim).copied().unwrap_or(0); let global_shape = dtensor.global_shape(); let global_dim_size = global_shape.get(tensor_dim).copied().unwrap_or(0); let shard_size = (global_dim_size + mesh_size - 1) / mesh_size; let start_idx = mesh_coord * shard_size; let end_idx = ((mesh_coord + 1) * shard_size).min(global_dim_size); // Extract the local shard let local_shard = slice_tensor(&dtensor.local_shard(), tensor_dim, start_idx, end_idx)?; // Create new spec with sharded placement let new_spec = TensorSpec::new(global_shape.to_vec()) .with_placement(mesh_dim, Placement::Shard { tensor_dim }) .with_requires_grad(dtensor.requires_grad()); DTensor::from_local(local_shard, new_spec, mesh.clone()) } /// Slice a tensor along a dimension. fn slice_tensor(tensor: &Tensor, dim: usize, start: usize, end: usize) -> Result { let shape = tensor.shape().dims().to_vec(); if dim >= shape.len() { return Err(DistributedError::tensor(format!( "Slice dim {} out of range for {}D tensor", dim, shape.len() ))); } let mut output_shape = shape.clone(); output_shape[dim] = end - start; // Calculate strides let mut strides: Vec = vec![1; shape.len()]; for i in (0..shape.len() - 1).rev() { strides[i] = strides[i + 1] * shape[i + 1]; } let data = tensor .data() .map_err(|e| DistributedError::tensor(e.to_string()))?; let output_numel: usize = output_shape.iter().product(); let mut output_data = vec![0.0f32; output_numel]; // Copy sliced data let tensor_numel: usize = shape.iter().product(); let mut out_idx = 0; for flat_idx in 0..tensor_numel { // Convert to coordinates let mut coords = vec![0; shape.len()]; let mut remaining = flat_idx; for d in 0..shape.len() { coords[d] = remaining / strides[d]; remaining %= strides[d]; } // Check if this element is in the slice if coords[dim] >= start && coords[dim] < end { if out_idx < output_data.len() { output_data[out_idx] = data[flat_idx]; out_idx += 1; } } } Tensor::from_data(output_data, output_shape, &Device::cpu()) .map_err(|e| DistributedError::tensor(e.to_string())) } // ============================================================================= // AllReduce: Partial -> Replicate // ============================================================================= /// Perform AllReduce to reduce partial results across devices. async fn all_reduce_reshard( dtensor: &DTensor, _mesh_dim: usize, reduce_op: PartialReduceOp, pg: &ProcessGroup, ) -> Result { let mut local_tensor = dtensor.local_shard().clone(); let comm_reduce_op = match reduce_op { PartialReduceOp::Sum => ReduceOp::Sum, PartialReduceOp::Mean => ReduceOp::Sum, // Sum then divide PartialReduceOp::Max => ReduceOp::Max, PartialReduceOp::Min => ReduceOp::Min, }; pg.allreduce(&mut local_tensor, comm_reduce_op).await?; // For mean, divide by world size if reduce_op == PartialReduceOp::Mean { let world_size = pg.world_size() as f32; local_tensor = local_tensor .div_scalar(world_size) .map_err(|e| DistributedError::tensor(e.to_string()))?; } // Create new DTensor with replicated placement let new_spec = TensorSpec::new(dtensor.global_shape().to_vec()) .with_placement(0, Placement::Replicate) .with_requires_grad(dtensor.requires_grad()); DTensor::from_local(local_tensor, new_spec, dtensor.mesh().clone()) } // ============================================================================= // ReduceScatter: Partial -> Shard // ============================================================================= /// Perform ReduceScatter to reduce and shard in one operation. async fn reduce_scatter_reshard( dtensor: &DTensor, mesh_dim: usize, tensor_dim: usize, reduce_op: PartialReduceOp, pg: &ProcessGroup, ) -> Result { let local_tensor = dtensor.local_shard().clone(); let comm_reduce_op = match reduce_op { PartialReduceOp::Sum => ReduceOp::Sum, PartialReduceOp::Mean => ReduceOp::Sum, PartialReduceOp::Max => ReduceOp::Max, PartialReduceOp::Min => ReduceOp::Min, }; let scattered: Tensor = pg.reduce_scatter(&local_tensor, comm_reduce_op)?; // For mean, divide by world size let result = if reduce_op == PartialReduceOp::Mean { let world_size = pg.world_size() as f32; scattered .div_scalar(world_size) .map_err(|e| DistributedError::tensor(e.to_string()))? } else { scattered }; // Create new DTensor with sharded placement let new_spec = TensorSpec::new(dtensor.global_shape().to_vec()) .with_placement(mesh_dim, Placement::Shard { tensor_dim }) .with_requires_grad(dtensor.requires_grad()); DTensor::from_local(result, new_spec, dtensor.mesh().clone()) } // ============================================================================= // AllToAll: Shard(dim_a) -> Shard(dim_b) // ============================================================================= /// Perform AllToAll to change sharding dimension. async fn all_to_all_reshard( dtensor: &DTensor, mesh_dim: usize, _src_dim: usize, dst_dim: usize, pg: &ProcessGroup, ) -> Result { // AllToAll is more complex - for now, we use AllGather + local slice // A more efficient implementation would use actual AllToAll communication // First, gather to replicated let gathered_spec = TensorSpec::new(dtensor.global_shape().to_vec()) .with_placement(0, Placement::Replicate) .with_requires_grad(dtensor.requires_grad()); let local_shard = dtensor.local_shard().clone(); let gathered = pg.allgather(&local_shard).await?; let full_tensor = match gathered { crate::comm::AllGatherOutput::Tensor(t) => t, crate::comm::AllGatherOutput::TensorList(tensors) => concat_along_dim(&tensors, 0)?, }; let gathered_dtensor = DTensor::from_local(full_tensor, gathered_spec, dtensor.mesh().clone())?; // Then slice for new sharding dimension replicate_to_shard(&gathered_dtensor, mesh_dim, dst_dim) } // ============================================================================= // Utility Functions // ============================================================================= /// Check if two DTensors have compatible placements for element-wise operations. pub fn placements_compatible(a: &DTensor, b: &DTensor) -> bool { a.spec().placements == b.spec().placements } /// Redistribute DTensor to match another DTensor's placement. pub async fn redistribute_like( dtensor: &DTensor, target: &DTensor, pg: &ProcessGroup, ) -> Result { reshard_dtensor(dtensor, &target.spec().placements, pg).await } /// Create a replicated DTensor from a local tensor. pub fn replicate_tensor(tensor: Tensor, mesh: Arc) -> Result { let shape: Vec = tensor.shape().dims().to_vec(); let spec = TensorSpec::new(shape).with_placement(0, Placement::Replicate); DTensor::from_local(tensor, spec, mesh) } /// Create a sharded DTensor from a local shard. pub fn shard_tensor( local_shard: Tensor, global_shape: Vec, mesh_dim: usize, tensor_dim: usize, mesh: Arc, ) -> Result { let spec = TensorSpec::new(global_shape).with_placement(mesh_dim, Placement::Shard { tensor_dim }); DTensor::from_local(local_shard, spec, mesh) } // ============================================================================= // Tests // ============================================================================= #[cfg(test)] mod tests { use super::*; #[test] fn test_concat_along_dim0() { let t1 = Tensor::from_data(vec![1.0, 2.0], vec![2], &Device::cpu()).unwrap(); let t2 = Tensor::from_data(vec![3.0, 4.0], vec![2], &Device::cpu()).unwrap(); let result = concat_along_dim(&[t1, t2], 0).unwrap(); let data = result.data().unwrap(); assert_eq!(data, vec![1.0, 2.0, 3.0, 4.0]); assert_eq!(result.shape().dims(), &[4]); } #[test] fn test_concat_along_dim1() { let t1 = Tensor::from_data(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], &Device::cpu()).unwrap(); let t2 = Tensor::from_data(vec![5.0, 6.0, 7.0, 8.0], vec![2, 2], &Device::cpu()).unwrap(); let result = concat_along_dim(&[t1, t2], 1).unwrap(); let data = result.data().unwrap(); // [1, 2] concat [5, 6] along dim 1 = [1, 2, 5, 6] // [3, 4] concat [7, 8] along dim 1 = [3, 4, 7, 8] assert_eq!(data, vec![1.0, 2.0, 5.0, 6.0, 3.0, 4.0, 7.0, 8.0]); assert_eq!(result.shape().dims(), &[2, 4]); } #[test] fn test_slice_tensor_dim0() { let tensor = Tensor::from_data(vec![1.0, 2.0, 3.0, 4.0], vec![4], &Device::cpu()).unwrap(); let sliced = slice_tensor(&tensor, 0, 1, 3).unwrap(); assert_eq!(sliced.data().unwrap(), vec![2.0, 3.0]); assert_eq!(sliced.shape().dims(), &[2]); } #[test] fn test_slice_tensor_dim1() { let tensor = Tensor::from_data( vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3], &Device::cpu(), ) .unwrap(); let sliced = slice_tensor(&tensor, 1, 0, 2).unwrap(); assert_eq!(sliced.data().unwrap(), vec![1.0, 2.0, 4.0, 5.0]); assert_eq!(sliced.shape().dims(), &[2, 2]); } #[test] fn test_replicate_tensor() { let mesh = Arc::new(DeviceMesh::new_simple(2, "dp")); let tensor = Tensor::from_data(vec![1.0, 2.0, 3.0], vec![3], &Device::cpu()).unwrap(); let dtensor = replicate_tensor(tensor, mesh).unwrap(); assert!(!dtensor.spec().is_sharded()); assert_eq!(dtensor.global_shape(), &[3]); } #[test] fn test_shard_tensor() { let mesh = Arc::new(DeviceMesh::new_simple(2, "dp")); let local_shard = Tensor::from_data(vec![1.0, 2.0], vec![2], &Device::cpu()).unwrap(); let dtensor = shard_tensor(local_shard, vec![4], 0, 0, mesh).unwrap(); assert!(dtensor.spec().is_sharded()); assert_eq!(dtensor.global_shape(), &[4]); } #[test] fn test_placements_compatible() { let mesh = Arc::new(DeviceMesh::new_simple(2, "dp")); let t1 = Tensor::from_data(vec![1.0, 2.0], vec![2], &Device::cpu()).unwrap(); let t2 = Tensor::from_data(vec![3.0, 4.0], vec![2], &Device::cpu()).unwrap(); let d1 = replicate_tensor(t1, mesh.clone()).unwrap(); let d2 = replicate_tensor(t2, mesh).unwrap(); assert!(placements_compatible(&d1, &d2)); } }