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

512 lines
17 KiB
Rust

#![cfg(feature = "disabled_tests")]
use rtx_preprocessing::{
DistributedDataLoader, FaultToleranceConfig, PreprocessingError, RebalancingStrategy, Result,
ShardingConfig, ShardingStrategy, WorkerInfo,
rtx_stubs::rtx_tensor::{Device, Tensor},
};
use std::fs::{File, create_dir_all};
use std::io::Write;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use tempfile::TempDir;
/// Helper function to create multiple test data files
fn create_test_dataset(dir: &Path, file_count: usize, samples_per_file: usize) -> Vec<PathBuf> {
let mut file_paths = Vec::new();
for file_idx in 0..file_count {
let file_name = format!("data_{:03}.bin", file_idx);
let file_path = dir.join(&file_name);
let mut file = File::create(&file_path).unwrap();
// Write sequential data for testing
let start_value = file_idx * samples_per_file;
for i in 0..samples_per_file {
let value = (start_value + i) as f32;
file.write_all(&value.to_le_bytes()).unwrap();
}
file.flush().unwrap();
file_paths.push(file_path);
}
file_paths
}
#[test]
fn test_distributed_data_loader_creation() {
let config = ShardingConfig::default();
let workers = vec![
WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()),
WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()),
];
let loader = DistributedDataLoader::new(config.clone(), workers.clone());
assert_eq!(loader.config().num_shards, config.num_shards);
assert_eq!(loader.workers().len(), 2);
assert!(!loader.is_initialized());
}
#[test]
fn test_distributed_data_loader_with_round_robin_sharding() {
let config = ShardingConfig {
num_shards: 4,
sharding_strategy: ShardingStrategy::RoundRobin,
rebalancing_strategy: RebalancingStrategy::Static,
fault_tolerance: FaultToleranceConfig::default(),
};
let workers = vec![
WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()),
WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()),
];
let loader = DistributedDataLoader::new(config, workers);
assert_eq!(
loader.config().sharding_strategy,
ShardingStrategy::RoundRobin
);
assert_eq!(loader.config().num_shards, 4);
}
#[test]
fn test_distributed_data_loader_initialization() {
let temp_dir = TempDir::new().unwrap();
let file_paths = create_test_dataset(temp_dir.path(), 8, 256); // 8 files, 256 samples each
let config = ShardingConfig {
num_shards: 4,
sharding_strategy: ShardingStrategy::RoundRobin,
rebalancing_strategy: RebalancingStrategy::Static,
fault_tolerance: FaultToleranceConfig::default(),
};
let workers = vec![
WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()),
WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()),
];
let mut loader = DistributedDataLoader::new(config, workers);
// Initialize with data files
let result = loader.initialize_shards(&file_paths);
assert!(result.is_ok());
assert!(loader.is_initialized());
// Check that shards were created
let shard_info = loader.get_shard_info();
assert_eq!(shard_info.len(), 4);
// Check that all files were assigned to shards
let total_files: usize = shard_info.iter().map(|s| s.file_paths.len()).sum();
assert_eq!(total_files, 8);
}
#[test]
fn test_hash_based_sharding() {
let temp_dir = TempDir::new().unwrap();
let file_paths = create_test_dataset(temp_dir.path(), 12, 100);
let config = ShardingConfig {
num_shards: 3,
sharding_strategy: ShardingStrategy::HashBased,
rebalancing_strategy: RebalancingStrategy::Static,
fault_tolerance: FaultToleranceConfig::default(),
};
let workers = vec![
WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()),
WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()),
WorkerInfo::new(2, "worker2".to_string(), "127.0.0.1:8002".to_string()),
];
let mut loader = DistributedDataLoader::new(config, workers);
loader.initialize_shards(&file_paths).unwrap();
let shard_info = loader.get_shard_info();
// Hash-based sharding should distribute files somewhat evenly
for shard in &shard_info {
assert!(shard.file_paths.len() > 0);
assert!(shard.file_paths.len() <= 6); // At most half the files in one shard
}
// Verify same file always goes to same shard (deterministic)
let first_file_shard = shard_info
.iter()
.find(|s| s.file_paths.contains(&file_paths[0]))
.unwrap()
.shard_id;
// Reinitialize and check consistency
loader.reset();
loader.initialize_shards(&file_paths).unwrap();
let new_shard_info = loader.get_shard_info();
let new_first_file_shard = new_shard_info
.iter()
.find(|s| s.file_paths.contains(&file_paths[0]))
.unwrap()
.shard_id;
assert_eq!(first_file_shard, new_first_file_shard);
}
#[test]
fn test_size_aware_sharding() {
let temp_dir = TempDir::new().unwrap();
// Create files with different sizes
let mut file_paths = Vec::new();
let sizes = vec![100, 200, 50, 300, 150]; // Different sample counts
for (idx, &size) in sizes.iter().enumerate() {
let file_name = format!("data_{}.bin", idx);
let file_path = temp_dir.path().join(&file_name);
let mut file = File::create(&file_path).unwrap();
for i in 0..size {
let value = i as f32;
file.write_all(&value.to_le_bytes()).unwrap();
}
file.flush().unwrap();
file_paths.push(file_path);
}
let config = ShardingConfig {
num_shards: 2,
sharding_strategy: ShardingStrategy::SizeAware,
rebalancing_strategy: RebalancingStrategy::Static,
fault_tolerance: FaultToleranceConfig::default(),
};
let workers = vec![
WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()),
WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()),
];
let mut loader = DistributedDataLoader::new(config, workers);
loader.initialize_shards(&file_paths).unwrap();
let shard_info = loader.get_shard_info();
// Calculate total size for each shard
let total_sizes: Vec<usize> = shard_info
.iter()
.map(|shard| {
shard
.file_paths
.iter()
.map(|path| {
let file_idx = path
.file_stem()
.unwrap()
.to_str()
.unwrap()
.split('_')
.last()
.unwrap()
.parse::<usize>()
.unwrap();
sizes[file_idx] * 4 // f32 = 4 bytes
})
.sum()
})
.collect();
// Size-aware sharding should balance shard sizes
let size_diff = (total_sizes[0] as i64 - total_sizes[1] as i64).abs();
let avg_size = (total_sizes[0] + total_sizes[1]) / 2;
// Sizes should be within 50% of each other for good balancing
assert!(size_diff as usize <= avg_size / 2);
}
#[test]
fn test_shard_assignment_retrieval() {
let temp_dir = TempDir::new().unwrap();
let file_paths = create_test_dataset(temp_dir.path(), 6, 100);
let config = ShardingConfig::default();
let workers = vec![
WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()),
WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()),
];
let mut loader = DistributedDataLoader::new(config, workers.clone());
loader.initialize_shards(&file_paths).unwrap();
// Test worker assignment
for worker in &workers {
let assigned_shards = loader.get_shards_for_worker(worker.worker_id);
assert!(!assigned_shards.is_empty());
for shard_id in assigned_shards {
let shard = loader.get_shard(shard_id).unwrap();
assert!(shard.assigned_worker_id == worker.worker_id);
}
}
// Test file to shard mapping
for file_path in &file_paths {
let shard_id = loader.get_shard_for_file(file_path).unwrap();
let shard = loader.get_shard(shard_id).unwrap();
assert!(shard.file_paths.contains(file_path));
}
}
#[test]
fn test_dynamic_rebalancing() {
let temp_dir = TempDir::new().unwrap();
let file_paths = create_test_dataset(temp_dir.path(), 8, 100);
let config = ShardingConfig {
num_shards: 4,
sharding_strategy: ShardingStrategy::RoundRobin,
rebalancing_strategy: RebalancingStrategy::Dynamic {
imbalance_threshold: 0.3,
rebalance_interval_secs: 10,
},
fault_tolerance: FaultToleranceConfig::default(),
};
let workers = vec![
WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()),
WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()),
];
let mut loader = DistributedDataLoader::new(config, workers);
loader.initialize_shards(&file_paths).unwrap();
// Simulate load imbalance by marking one worker as slow
loader.report_worker_performance(0, 1000); // Worker 0 slow (1000ms)
loader.report_worker_performance(1, 100); // Worker 1 fast (100ms)
// Check if rebalancing is triggered
let needs_rebalancing = loader.needs_rebalancing();
assert!(needs_rebalancing);
// Trigger rebalancing
let result = loader.rebalance();
assert!(result.is_ok());
// Verify that shards were redistributed
let shard_info = loader.get_shard_info();
let worker1_shards = shard_info
.iter()
.filter(|s| s.assigned_worker_id == 1)
.count();
let worker0_shards = shard_info
.iter()
.filter(|s| s.assigned_worker_id == 0)
.count();
// Worker 1 (faster) should get more shards
assert!(worker1_shards >= worker0_shards);
}
#[test]
fn test_fault_tolerance_detection() {
let temp_dir = TempDir::new().unwrap();
let file_paths = create_test_dataset(temp_dir.path(), 4, 100);
let config = ShardingConfig {
num_shards: 4,
sharding_strategy: ShardingStrategy::RoundRobin,
rebalancing_strategy: RebalancingStrategy::Static,
fault_tolerance: FaultToleranceConfig {
enable_health_checks: true,
health_check_interval_secs: 5,
max_failures: 3,
recovery_timeout_secs: 30,
},
};
let workers = vec![
WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()),
WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()),
];
let mut loader = DistributedDataLoader::new(config, workers);
loader.initialize_shards(&file_paths).unwrap();
// All workers should initially be healthy
assert!(loader.is_worker_healthy(0));
assert!(loader.is_worker_healthy(1));
// Simulate worker failure
loader.report_worker_failure(0);
loader.report_worker_failure(0);
loader.report_worker_failure(0); // 3 failures = unhealthy
assert!(!loader.is_worker_healthy(0));
assert!(loader.is_worker_healthy(1));
// Check that failed worker's shards are reassigned
let failed_worker_shards = loader.get_shards_for_worker(0);
assert!(failed_worker_shards.is_empty());
// All shards should now be assigned to healthy worker
let healthy_worker_shards = loader.get_shards_for_worker(1);
assert_eq!(healthy_worker_shards.len(), 4);
}
#[test]
fn test_fault_tolerance_recovery() {
let temp_dir = TempDir::new().unwrap();
let file_paths = create_test_dataset(temp_dir.path(), 6, 100);
let config = ShardingConfig {
num_shards: 2,
sharding_strategy: ShardingStrategy::RoundRobin,
rebalancing_strategy: RebalancingStrategy::Static,
fault_tolerance: FaultToleranceConfig {
enable_health_checks: true,
health_check_interval_secs: 1,
max_failures: 2,
recovery_timeout_secs: 5,
},
};
let workers = vec![
WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()),
WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()),
];
let mut loader = DistributedDataLoader::new(config, workers);
loader.initialize_shards(&file_paths).unwrap();
// Simulate worker failure and recovery
loader.report_worker_failure(0);
loader.report_worker_failure(0); // Worker 0 becomes unhealthy
assert!(!loader.is_worker_healthy(0));
// Simulate worker recovery
loader.report_worker_recovery(0);
assert!(loader.is_worker_healthy(0));
// Verify shards are redistributed after recovery
let result = loader.redistribute_after_recovery(0);
assert!(result.is_ok());
let worker0_shards = loader.get_shards_for_worker(0);
assert!(!worker0_shards.is_empty());
}
#[test]
fn test_concurrent_access() {
use std::sync::Arc;
use std::thread;
let temp_dir = TempDir::new().unwrap();
let file_paths = create_test_dataset(temp_dir.path(), 16, 50);
let config = ShardingConfig {
num_shards: 8,
sharding_strategy: ShardingStrategy::RoundRobin,
rebalancing_strategy: RebalancingStrategy::Static,
fault_tolerance: FaultToleranceConfig::default(),
};
let workers = vec![
WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()),
WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()),
WorkerInfo::new(2, "worker2".to_string(), "127.0.0.1:8002".to_string()),
WorkerInfo::new(3, "worker3".to_string(), "127.0.0.1:8003".to_string()),
];
let mut loader = DistributedDataLoader::new(config, workers);
loader.initialize_shards(&file_paths).unwrap();
let loader = Arc::new(loader);
let mut handles = vec![];
// Spawn multiple threads doing concurrent operations
for worker_id in 0..4 {
let loader_clone = Arc::clone(&loader);
let handle = thread::spawn(move || {
// Each thread queries different aspects
match worker_id {
0 => {
// Thread 0: Query shard assignments
let shards = loader_clone.get_shards_for_worker(worker_id);
assert!(!shards.is_empty());
}
1 => {
// Thread 1: Check worker health
let is_healthy = loader_clone.is_worker_healthy(worker_id);
assert!(is_healthy);
}
2 => {
// Thread 2: Get shard info
let shard_info = loader_clone.get_shard_info();
assert_eq!(shard_info.len(), 8);
}
3 => {
// Thread 3: Report performance metrics
loader_clone.report_worker_performance(worker_id, 150);
}
_ => unreachable!(),
}
});
handles.push(handle);
}
// Wait for all threads to complete
for handle in handles {
handle.join().unwrap();
}
// Verify system is still consistent
let shard_info = loader.get_shard_info();
assert_eq!(shard_info.len(), 8);
let total_files: usize = shard_info.iter().map(|s| s.file_paths.len()).sum();
assert_eq!(total_files, 16);
}
#[test]
fn test_distributed_loader_statistics() {
let temp_dir = TempDir::new().unwrap();
let file_paths = create_test_dataset(temp_dir.path(), 10, 100);
let config = ShardingConfig::default();
let workers = vec![
WorkerInfo::new(0, "worker0".to_string(), "127.0.0.1:8000".to_string()),
WorkerInfo::new(1, "worker1".to_string(), "127.0.0.1:8001".to_string()),
];
let mut loader = DistributedDataLoader::new(config, workers);
loader.initialize_shards(&file_paths).unwrap();
// Get initial statistics
let stats = loader.get_statistics();
assert_eq!(stats.total_shards, 2);
assert_eq!(stats.total_files, 10);
assert_eq!(stats.active_workers, 2);
assert_eq!(stats.failed_workers, 0);
// Simulate some operations
loader.report_worker_performance(0, 200);
loader.report_worker_performance(1, 150);
// Simulate worker failure
loader.report_worker_failure(0);
loader.report_worker_failure(0);
loader.report_worker_failure(0);
let updated_stats = loader.get_statistics();
assert_eq!(updated_stats.active_workers, 1);
assert_eq!(updated_stats.failed_workers, 1);
// Performance metrics should be recorded
assert!(updated_stats.average_worker_performance > 0.0);
}