use super::core::{DistributedConfig, DistributedOptimizer};
use crate::{Adam, OptimizerResult, SGD};
use parking_lot::RwLock;
use std::sync::Arc;
use torsh_core::error::Result;
use torsh_tensor::Tensor;
pub fn distributed_sgd(
params: Vec<Arc<RwLock<Tensor>>>,
lr: f32,
world_size: usize,
rank: usize,
momentum: Option<f32>,
weight_decay: Option<f32>,
) -> OptimizerResult<DistributedOptimizer<SGD>> {
let sgd = SGD::new(params, lr, momentum, None, weight_decay, false);
let config = DistributedConfig {
world_size,
rank,
..Default::default()
};
DistributedOptimizer::new(sgd, config)
}
pub fn distributed_adam(
params: Vec<Arc<RwLock<Tensor>>>,
lr: f32,
world_size: usize,
rank: usize,
betas: Option<(f32, f32)>,
eps: Option<f32>,
weight_decay: Option<f32>,
) -> OptimizerResult<DistributedOptimizer<Adam>> {
let adam = Adam::new(params, Some(lr), betas, eps, weight_decay, false);
let config = DistributedConfig {
world_size,
rank,
..Default::default()
};
DistributedOptimizer::new(adam, config)
}
pub fn distributed_optimizer<O: crate::Optimizer>(
optimizer: O,
config: DistributedConfig,
) -> OptimizerResult<DistributedOptimizer<O>> {
DistributedOptimizer::new(optimizer, config)
}
pub mod configs {
use super::super::core::{DistributedBackend, DistributedConfig, SyncStrategy};
pub fn cpu_mpi_config(world_size: usize, rank: usize) -> DistributedConfig {
DistributedConfig {
backend: DistributedBackend::MPI,
sync_strategy: SyncStrategy::AllReduce,
world_size,
rank,
gradient_compression: false, bucket_size_mb: 10.0, overlap_communication: false, ..Default::default()
}
}
pub fn gpu_nccl_config(world_size: usize, rank: usize) -> DistributedConfig {
DistributedConfig {
backend: DistributedBackend::NCCL,
sync_strategy: SyncStrategy::AllReduce,
world_size,
rank,
gradient_compression: world_size >= 8, bucket_size_mb: 25.0, overlap_communication: true, ..Default::default()
}
}
pub fn mixed_gloo_config(world_size: usize, rank: usize) -> DistributedConfig {
DistributedConfig {
backend: DistributedBackend::Gloo,
sync_strategy: SyncStrategy::AllReduce,
world_size,
rank,
gradient_compression: world_size > 4,
bucket_size_mb: 15.0,
overlap_communication: true,
..Default::default()
}
}
pub fn large_scale_config(world_size: usize, rank: usize) -> DistributedConfig {
DistributedConfig {
backend: DistributedBackend::NCCL,
sync_strategy: SyncStrategy::ReduceScatter, world_size,
rank,
gradient_compression: true, bucket_size_mb: 50.0, overlap_communication: true, ..Default::default()
}
}
pub fn low_bandwidth_config(world_size: usize, rank: usize) -> DistributedConfig {
DistributedConfig {
backend: DistributedBackend::Gloo,
sync_strategy: SyncStrategy::AllReduce,
world_size,
rank,
gradient_compression: true, bucket_size_mb: 5.0, overlap_communication: true,
..Default::default()
}
}
}
pub mod monitoring {
use super::super::core::{CommunicationStats, DistributedOptimizer};
use crate::Optimizer;
use std::collections::HashMap;
pub fn collect_communication_stats<O: Optimizer>(
optimizers: &[DistributedOptimizer<O>],
) -> HashMap<usize, CommunicationStats> {
optimizers
.iter()
.enumerate()
.map(|(i, opt)| (i, opt.get_communication_stats()))
.collect()
}
pub fn aggregate_communication_stats(
stats: &HashMap<usize, CommunicationStats>,
) -> CommunicationStats {
if stats.is_empty() {
return CommunicationStats::default();
}
let total_communications: u64 = stats.values().map(|s| s.total_communications).sum();
let total_bytes: u64 = stats.values().map(|s| s.total_bytes_transferred).sum();
let avg_time: f32 = stats
.values()
.map(|s| s.average_communication_time_ms)
.sum::<f32>()
/ stats.len() as f32;
let avg_compression: f32 = stats
.values()
.map(|s| s.gradient_compression_ratio)
.sum::<f32>()
/ stats.len() as f32;
CommunicationStats {
total_communications,
total_bytes_transferred: total_bytes,
average_communication_time_ms: avg_time,
gradient_compression_ratio: avg_compression,
}
}
pub fn print_performance_summary<O: Optimizer>(optimizers: &[DistributedOptimizer<O>]) {
let stats = collect_communication_stats(optimizers);
let aggregate = aggregate_communication_stats(&stats);
println!("=== Distributed Training Performance Summary ===");
println!("Number of workers: {}", optimizers.len());
println!("Total communications: {}", aggregate.total_communications);
println!(
"Total bytes transferred: {:.2} MB",
aggregate.total_bytes_transferred as f64 / 1024.0 / 1024.0
);
println!(
"Average communication time: {:.2} ms",
aggregate.average_communication_time_ms
);
println!(
"Average compression ratio: {:.2}x",
aggregate.gradient_compression_ratio
);
println!("\n--- Per-Worker Breakdown ---");
for (worker_id, stat) in &stats {
println!(
"Worker {}: {} communications, {:.2} MB, {:.2} ms avg",
worker_id,
stat.total_communications,
stat.total_bytes_transferred as f64 / 1024.0 / 1024.0,
stat.average_communication_time_ms
);
}
}
}
#[cfg(test)]
mod tests {
use super::super::core::*;
use super::*;
#[test]
fn test_distributed_sgd_creation() {
}
#[test]
fn test_config_creation() {
let config = configs::cpu_mpi_config(4, 0);
assert_eq!(config.world_size, 4);
assert_eq!(config.rank, 0);
assert!(matches!(config.backend, DistributedBackend::MPI));
let gpu_config = configs::gpu_nccl_config(8, 3);
assert_eq!(gpu_config.world_size, 8);
assert_eq!(gpu_config.rank, 3);
assert!(matches!(gpu_config.backend, DistributedBackend::NCCL));
assert!(gpu_config.gradient_compression); }
#[test]
fn test_large_scale_config() {
let config = configs::large_scale_config(64, 0);
assert!(matches!(config.sync_strategy, SyncStrategy::ReduceScatter));
assert!(config.gradient_compression);
assert_eq!(config.bucket_size_mb, 50.0);
}
}