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

346 lines
11 KiB
Rust

//! Multi-node distributed training tests
#[cfg(test)]
mod tests {
use rtx_distributed::multi_node::{
AggregationPattern, BandwidthOptimizer, CrossNodeCommunicator, HealthChecker,
InterconnectType, MultiNodeCluster, NetworkTopology, NodeConfig, NodeInfo, NodeRole,
RendezvousProtocol,
};
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
use std::time::Duration;
#[tokio::test]
async fn test_multi_node_cluster_creation() {
let config = NodeConfig::default()
.with_role(NodeRole::Worker)
.with_port(29500);
let cluster = MultiNodeCluster::new(config).await.unwrap();
assert_eq!(cluster.num_nodes(), 1);
assert!(cluster.is_initialized());
assert_eq!(cluster.local_rank(), 0);
}
#[tokio::test]
async fn test_node_discovery() {
let master_config = NodeConfig::default()
.with_role(NodeRole::Master)
.with_port(29501);
let mut cluster = MultiNodeCluster::new(master_config).await.unwrap();
// Simulate worker joining
let worker_info = NodeInfo {
node_id: 1,
hostname: "worker1".to_string(),
ip_address: IpAddr::V4(Ipv4Addr::new(192, 168, 1, 2)),
port: 29502,
role: NodeRole::Worker,
num_gpus: 8,
region: None,
latency_ms: None,
};
cluster.register_node(worker_info).await.unwrap();
assert_eq!(cluster.num_nodes(), 2);
assert!(cluster.has_node(1));
let discovered = cluster
.discover_nodes(Duration::from_secs(1))
.await
.unwrap();
assert_eq!(discovered.len(), 2);
}
#[tokio::test]
async fn test_cross_node_communication() {
let comm = CrossNodeCommunicator::new("tcp://0.0.0.0:29503")
.await
.unwrap();
// Test send/receive
let data = vec![1.0f32, 2.0, 3.0, 4.0];
let target_node = 1;
comm.send_async(target_node, &data).await.unwrap();
// In real test, another node would receive
let received = comm.receive_from(target_node).await.unwrap();
assert_eq!(received, data);
}
#[tokio::test]
async fn test_rendezvous_protocol() {
let rendezvous = RendezvousProtocol::new(
"tcp://master:29504",
"test_job",
4, // expected world size
);
// Simulate node joining
let node_info = rendezvous.join(0).await.unwrap();
assert_eq!(node_info.rank, 0);
assert_eq!(node_info.world_size, 4);
// Test barrier
rendezvous.barrier().await.unwrap();
// Test store operations
rendezvous.set("key1", b"value1").await.unwrap();
let value = rendezvous.get("key1").await.unwrap();
assert_eq!(value, b"value1");
}
#[tokio::test]
async fn test_network_topology_optimization() {
let mut topology = NetworkTopology::new();
// Add nodes with different interconnects
topology.add_node(0, InterconnectType::InfiniBand);
topology.add_node(1, InterconnectType::InfiniBand);
topology.add_node(2, InterconnectType::Ethernet10G);
topology.add_node(3, InterconnectType::Ethernet10G);
// Add links with bandwidth
topology.add_link(0, 1, 200.0); // 200 Gbps IB
topology.add_link(0, 2, 10.0); // 10 Gbps Ethernet
topology.add_link(1, 3, 10.0);
topology.add_link(2, 3, 10.0);
// Find optimal communication pattern
let pattern = topology.optimize_allreduce();
assert!(pattern.uses_hierarchical());
assert_eq!(pattern.num_levels(), 2); // IB level and Ethernet level
}
#[tokio::test]
async fn test_health_checking() {
let health_checker = HealthChecker::new(Duration::from_secs(5));
let node1 = NodeInfo::new(1, "node1", "192.168.1.2:29505");
let node2 = NodeInfo::new(2, "node2", "192.168.1.3:29505");
health_checker.monitor_node(node1).await;
health_checker.monitor_node(node2).await;
// Check health status
tokio::time::sleep(Duration::from_millis(100)).await;
assert!(health_checker.is_healthy(1).await);
assert!(health_checker.is_healthy(2).await);
// Simulate node failure
health_checker.mark_unhealthy(2).await;
assert!(!health_checker.is_healthy(2).await);
// Get healthy nodes
let healthy = health_checker.get_healthy_nodes().await;
assert_eq!(healthy.len(), 1);
assert_eq!(healthy[0].node_id, 1);
}
#[tokio::test]
async fn test_bandwidth_optimization() {
let optimizer = BandwidthOptimizer::new();
// Add bandwidth measurements
optimizer.record_bandwidth(0, 1, 180.0); // Gbps
optimizer.record_bandwidth(0, 2, 8.5);
optimizer.record_bandwidth(1, 2, 9.0);
// Get optimal routing
let route = optimizer.optimal_route(0, 2);
assert_eq!(route.len(), 2);
assert_eq!(route[0], 0);
assert_eq!(route[1], 2);
// Get bandwidth estimate
let bandwidth = optimizer.estimate_bandwidth(0, 2);
assert_eq!(bandwidth, 8.5);
}
#[tokio::test]
async fn test_elastic_scaling() {
let mut cluster = MultiNodeCluster::new(NodeConfig::default()).await.unwrap();
// Start with 2 nodes
cluster
.add_node(NodeInfo::new(0, "node0", "192.168.1.1:29506"))
.await
.unwrap();
cluster
.add_node(NodeInfo::new(1, "node1", "192.168.1.2:29506"))
.await
.unwrap();
assert_eq!(cluster.world_size(), 2);
// Scale up
cluster
.add_node(NodeInfo::new(2, "node2", "192.168.1.3:29506"))
.await
.unwrap();
cluster
.add_node(NodeInfo::new(3, "node3", "192.168.1.4:29506"))
.await
.unwrap();
assert_eq!(cluster.world_size(), 4);
// Redistribute work
let new_mapping = cluster.redistribute_work().await.unwrap();
assert_eq!(new_mapping.len(), 4);
// Scale down
cluster.remove_node(3).await.unwrap();
assert_eq!(cluster.world_size(), 3);
}
#[tokio::test]
async fn test_cross_region_communication() {
let mut cluster = MultiNodeCluster::new(NodeConfig::default()).await.unwrap();
// Add nodes in different regions
let node_us = NodeInfo::new(0, "us-east-1", "10.0.1.1:29507")
.with_region("us-east-1")
.with_latency_ms(0);
let node_eu = NodeInfo::new(1, "eu-west-1", "10.0.2.1:29507")
.with_region("eu-west-1")
.with_latency_ms(80); // 80ms to US
let node_asia = NodeInfo::new(2, "ap-south-1", "10.0.3.1:29507")
.with_region("ap-south-1")
.with_latency_ms(150); // 150ms to US
cluster.add_node(node_us).await.unwrap();
cluster.add_node(node_eu).await.unwrap();
cluster.add_node(node_asia).await.unwrap();
// Get region-aware communication pattern
let pattern = cluster.get_region_aware_pattern();
assert!(pattern.minimizes_cross_region());
assert!(pattern.prioritizes_local_region());
}
#[tokio::test]
async fn test_checkpoint_coordination() {
let cluster = MultiNodeCluster::new(NodeConfig::default()).await.unwrap();
// Coordinate checkpoint across nodes
let checkpoint_id = "ckpt_001";
cluster.begin_checkpoint(checkpoint_id).await.unwrap();
// Each node saves its state
cluster
.save_local_state(checkpoint_id, b"state_data")
.await
.unwrap();
// Wait for all nodes
cluster.checkpoint_barrier(checkpoint_id).await.unwrap();
// Finalize checkpoint
cluster.finalize_checkpoint(checkpoint_id).await.unwrap();
assert!(cluster.has_checkpoint(checkpoint_id).await);
}
#[tokio::test]
async fn test_gradient_aggregation_patterns() {
let cluster = MultiNodeCluster::new(NodeConfig::default()).await.unwrap();
// Test different aggregation patterns
let gradients = vec![1.0f32; 1000];
// Ring AllReduce
let ring_result = cluster.ring_allreduce(&gradients).await.unwrap();
assert_eq!(ring_result.len(), gradients.len());
// Tree AllReduce
let tree_result = cluster.tree_allreduce(&gradients).await.unwrap();
assert_eq!(tree_result.len(), gradients.len());
// Butterfly AllReduce
let butterfly_result = cluster.butterfly_allreduce(&gradients).await.unwrap();
assert_eq!(butterfly_result.len(), gradients.len());
// Measure efficiency
let ring_time = cluster
.measure_allreduce_time(AggregationPattern::Ring)
.await;
let tree_time = cluster
.measure_allreduce_time(AggregationPattern::Tree)
.await;
// Tree should be faster for small messages
let variance_factor = Duration::from_micros((ring_time.as_micros() as f64 * 1.1) as u64);
assert!(tree_time <= variance_factor); // Allow 10% variance
}
#[tokio::test]
#[ignore = "Pre-existing node failure recovery assertion failure"]
async fn test_node_failure_recovery() {
let mut cluster = MultiNodeCluster::new(NodeConfig::default().with_fault_tolerance(true))
.await
.unwrap();
// Add nodes
for i in 0..4 {
cluster
.add_node(NodeInfo::new(
i,
&format!("node{}", i),
&format!("192.168.1.{}:29508", i + 1),
))
.await
.unwrap();
}
// Simulate node failure
cluster.simulate_node_failure(2).await;
// Cluster should detect and handle failure
tokio::time::sleep(Duration::from_millis(500)).await;
assert!(!cluster.is_node_healthy(2).await);
assert_eq!(cluster.num_healthy_nodes(), 3);
// Redistribute work from failed node
let redistributed = cluster.handle_node_failure(2).await.unwrap();
assert!(redistributed.work_reassigned);
assert_eq!(redistributed.new_world_size, 3);
// Recovery when node comes back
cluster.recover_node(2).await.unwrap();
assert!(cluster.is_node_healthy(2).await);
}
#[tokio::test]
async fn test_multi_job_coordination() {
let mut cluster = MultiNodeCluster::new(NodeConfig::default()).await.unwrap();
// Register multiple training jobs
let job1 = cluster.register_job("model_a", 2).await.unwrap();
let job2 = cluster.register_job("model_b", 2).await.unwrap();
assert_ne!(job1.job_id, job2.job_id);
// Jobs should have separate communication groups
assert!(!job1.shares_nodes_with(&job2));
// Coordinate resource allocation
let allocation = cluster
.allocate_resources_for_jobs(&[job1.job_id, job2.job_id])
.await
.unwrap();
assert_eq!(allocation.total_nodes_used(), 4);
}
}