214 lines
7.0 KiB
Rust
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)
|
|
));
|
|
}
|
|
}
|