Files
rustytorch/crates/training/rtx-distributed/src/dtensor/tests.rs
T
2026-03-04 00:08:42 +00:00

214 lines
7.0 KiB
Rust

//! 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)
));
}
}