//! Tests for distributed tensors. #[cfg(test)] mod tests { use crate::device_mesh::DeviceMesh; use crate::dtensor::*; use rtx_tensor::Device; use std::sync::Arc; #[test] fn test_tensor_spec_local_shape() { // Mock mesh: 2 devices let mesh = Arc::new(DeviceMesh::new_simple(2, "dp")); // Global shape [1024, 4096], sharded on first mesh dim along tensor dim 1 let spec = TensorSpec::new(vec![1024, 4096]).with_placement(0, Placement::Shard { tensor_dim: 1 }); let local_shape = spec.local_shape(&mesh).unwrap(); assert_eq!(local_shape, vec![1024, 2048]); // 4096 / 2 = 2048 } #[test] fn test_tensor_spec_replicated() { let spec = TensorSpec::new(vec![1024, 1024]).with_placement(0, Placement::Replicate); assert!(spec.is_replicated()); assert!(!spec.is_sharded()); } #[test] fn test_tensor_spec_sharded() { let spec = TensorSpec::new(vec![1024, 1024]).with_placement(0, Placement::Shard { tensor_dim: 0 }); assert!(!spec.is_replicated()); assert!(spec.is_sharded()); } #[test] fn test_placement_partial() { let spec = TensorSpec::new(vec![1024, 1024]).with_placement( 0, Placement::Partial { reduce_op: PartialReduceOp::Sum, }, ); assert!(spec.is_partial()); } #[test] fn test_concatenate_tensors() { use rtx_tensor::Tensor; let device = Device::cpu(); let t1 = Tensor::from_data(vec![1.0, 2.0], vec![2], &device).unwrap(); let t2 = Tensor::from_data(vec![3.0, 4.0], vec![2], &device).unwrap(); let result = DTensor::concatenate_tensors(&[t1, t2], 0).unwrap(); let data = result.to_vec().unwrap(); assert_eq!(data, vec![1.0, 2.0, 3.0, 4.0]); } #[test] fn test_slice_tensor_data() { // Shape [4, 2], slice dim 0 from index 1 to 3 let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0]; let shape = vec![4, 2]; let result = DTensor::slice_tensor_data(&data, &shape, 0, 1, 3).unwrap(); // Should get rows 1 and 2: [3, 4, 5, 6] assert_eq!(result, vec![3.0, 4.0, 5.0, 6.0]); } #[test] fn test_dtensor_all_gather() { use crate::group::ProcessGroup; use crate::{Backend, WorldInfo}; use rtx_tensor::Tensor; // Create a simple mesh with 2 devices let mesh = Arc::new(DeviceMesh::new_simple(2, "dp")); let device = Device::cpu(); // Create a sharded DTensor (each rank has half the data) let local_data = vec![1.0, 2.0, 3.0, 4.0]; // 4 elements on this rank let local_tensor = Tensor::from_data(local_data, vec![4], &device).unwrap(); let spec = TensorSpec::new(vec![8]) // Global shape is 8 .with_placement(0, Placement::Shard { tensor_dim: 0 }); let dtensor = DTensor::from_local(local_tensor, spec, mesh.clone()).unwrap(); // Create process group let world_info = WorldInfo::new(2, 0, Backend::Cpu); let pg = ProcessGroup::new(Backend::Cpu, world_info).unwrap(); // All-gather should collect from all ranks let gathered = dtensor.all_gather(0, 0, &pg).unwrap(); // Check that placement is now Replicate assert!(matches!( gathered.spec.placements.get(0), Some(Placement::Replicate) )); } #[test] fn test_dtensor_all_reduce() { use crate::comm::ReduceOp; use crate::group::ProcessGroup; use crate::{Backend, WorldInfo}; use rtx_tensor::Tensor; let mesh = Arc::new(DeviceMesh::new_simple(2, "dp")); let device = Device::cpu(); // Create a partial DTensor let local_data = vec![1.0, 2.0, 3.0, 4.0]; let local_tensor = Tensor::from_data(local_data, vec![4], &device).unwrap(); let spec = TensorSpec::new(vec![4]).with_placement( 0, Placement::Partial { reduce_op: PartialReduceOp::Sum, }, ); let dtensor = DTensor::from_local(local_tensor, spec, mesh.clone()).unwrap(); // Create process group let world_info = WorldInfo::new(2, 0, Backend::Cpu); let pg = ProcessGroup::new(Backend::Cpu, world_info).unwrap(); // All-reduce with sum let reduced = dtensor.all_reduce(0, ReduceOp::Sum, &pg).unwrap(); // Check that placement is now Replicate assert!(matches!( reduced.spec.placements.get(0), Some(Placement::Replicate) )); // Check that values were summed (multiplied by world_size in simulation) let reduced_data = reduced.local_shard().to_vec().unwrap(); assert_eq!(reduced_data, vec![2.0, 4.0, 6.0, 8.0]); // Each value * 2 } #[test] fn test_dtensor_reduce_scatter() { use crate::comm::ReduceOp; use crate::group::ProcessGroup; use crate::{Backend, WorldInfo}; use rtx_tensor::Tensor; let mesh = Arc::new(DeviceMesh::new_simple(2, "dp")); let device = Device::cpu(); // Create a replicated DTensor let local_data = vec![1.0, 2.0, 3.0, 4.0]; let local_tensor = Tensor::from_data(local_data, vec![4], &device).unwrap(); let spec = TensorSpec::new(vec![4]).with_placement(0, Placement::Replicate); let dtensor = DTensor::from_local(local_tensor, spec, mesh.clone()).unwrap(); // Create process group let world_info = WorldInfo::new(2, 0, Backend::Cpu); let pg = ProcessGroup::new(Backend::Cpu, world_info).unwrap(); // Reduce-scatter let scattered = dtensor.reduce_scatter(0, 0, ReduceOp::Sum, &pg).unwrap(); // Check that placement is now Shard assert!(matches!( scattered.spec.placements.get(0), Some(Placement::Shard { tensor_dim: 0 }) )); } #[test] fn test_dtensor_redistribute() { use crate::group::ProcessGroup; use crate::{Backend, WorldInfo}; use rtx_tensor::Tensor; let mesh = Arc::new(DeviceMesh::new_simple(2, "dp")); let device = Device::cpu(); // Create a sharded DTensor let local_data = vec![1.0, 2.0]; let local_tensor = Tensor::from_data(local_data, vec![2], &device).unwrap(); let spec = TensorSpec::new(vec![4]).with_placement(0, Placement::Shard { tensor_dim: 0 }); let dtensor = DTensor::from_local(local_tensor, spec, mesh.clone()).unwrap(); // Create process group let world_info = WorldInfo::new(2, 0, Backend::Cpu); let pg = ProcessGroup::new(Backend::Cpu, world_info).unwrap(); // Target: Replicate let target_spec = TensorSpec::new(vec![4]).with_placement(0, Placement::Replicate); let redistributed = dtensor.redistribute(target_spec, &pg).unwrap(); // Should now be replicated assert!(matches!( redistributed.spec.placements.get(0), Some(Placement::Replicate) )); } }