Skip to main content

torsh_optim/distributed/
utils.rs

1//! Utility functions for creating and managing distributed optimizers
2//!
3//! This module provides convenient factory functions and utilities for setting up
4//! distributed training with various optimizer types and configurations.
5
6use super::core::{DistributedConfig, DistributedOptimizer};
7use crate::{Adam, OptimizerResult, SGD};
8use parking_lot::RwLock;
9use std::sync::Arc;
10use torsh_core::error::Result;
11use torsh_tensor::Tensor;
12
13/// Create a distributed SGD optimizer with standard configuration
14///
15/// This is a convenience function for creating a distributed SGD optimizer
16/// with commonly used settings for distributed training.
17///
18/// # Arguments
19///
20/// * `params` - Parameters to optimize
21/// * `lr` - Learning rate
22/// * `world_size` - Total number of processes in distributed training
23/// * `rank` - Current process rank (0 to world_size-1)
24/// * `momentum` - Optional momentum factor
25/// * `weight_decay` - Optional weight decay (L2 penalty)
26///
27/// # Returns
28///
29/// A distributed SGD optimizer ready for training
30///
31/// # Example
32///
33/// ```rust
34/// # use torsh_tensor::creation::randn;
35/// # use torsh_core::error::Result;
36/// # use parking_lot::RwLock;
37/// # use std::sync::Arc;
38/// # fn main() -> Result<()> {
39/// use torsh_optim::distributed::utils::distributed_sgd;
40///
41/// // Create some parameters
42/// let param1 = Arc::new(RwLock::new(randn::<f32>(&[10, 20])?));
43/// let params = vec![param1];
44///
45/// let optimizer = distributed_sgd(
46///     params,
47///     0.1,        // learning rate
48///     4,          // world size (4 workers)
49///     0,          // rank (worker 0)
50///     Some(0.9),  // momentum
51///     Some(1e-4)  // weight decay
52/// )?;
53/// # Ok(())
54/// # }
55/// ```
56pub fn distributed_sgd(
57    params: Vec<Arc<RwLock<Tensor>>>,
58    lr: f32,
59    world_size: usize,
60    rank: usize,
61    momentum: Option<f32>,
62    weight_decay: Option<f32>,
63) -> OptimizerResult<DistributedOptimizer<SGD>> {
64    let sgd = SGD::new(params, lr, momentum, None, weight_decay, false);
65    let config = DistributedConfig {
66        world_size,
67        rank,
68        ..Default::default()
69    };
70    DistributedOptimizer::new(sgd, config)
71}
72
73/// Create a distributed Adam optimizer with standard configuration
74///
75/// This is a convenience function for creating a distributed Adam optimizer
76/// with commonly used settings for distributed training.
77///
78/// # Arguments
79///
80/// * `params` - Parameters to optimize
81/// * `lr` - Learning rate
82/// * `world_size` - Total number of processes in distributed training
83/// * `rank` - Current process rank (0 to world_size-1)
84/// * `betas` - Optional Adam beta parameters (momentum and RMS decay)
85/// * `eps` - Optional epsilon for numerical stability
86/// * `weight_decay` - Optional weight decay (L2 penalty)
87///
88/// # Returns
89///
90/// A distributed Adam optimizer ready for training
91///
92/// # Example
93///
94/// ```rust
95/// # use torsh_tensor::creation::randn;
96/// # use torsh_core::error::Result;
97/// # use parking_lot::RwLock;
98/// # use std::sync::Arc;
99/// # fn main() -> Result<()> {
100/// use torsh_optim::distributed::utils::distributed_adam;
101///
102/// // Create some parameters
103/// let param1 = Arc::new(RwLock::new(randn::<f32>(&[10, 20])?));
104/// let params = vec![param1];
105///
106/// let optimizer = distributed_adam(
107///     params,
108///     1e-3,                    // learning rate
109///     8,                       // world size (8 workers)
110///     3,                       // rank (worker 3)
111///     Some((0.9, 0.999)),     // betas
112///     Some(1e-8),             // epsilon
113///     Some(0.01)              // weight decay
114/// )?;
115/// # Ok(())
116/// # }
117/// ```
118pub fn distributed_adam(
119    params: Vec<Arc<RwLock<Tensor>>>,
120    lr: f32,
121    world_size: usize,
122    rank: usize,
123    betas: Option<(f32, f32)>,
124    eps: Option<f32>,
125    weight_decay: Option<f32>,
126) -> OptimizerResult<DistributedOptimizer<Adam>> {
127    let adam = Adam::new(params, Some(lr), betas, eps, weight_decay, false);
128    let config = DistributedConfig {
129        world_size,
130        rank,
131        ..Default::default()
132    };
133    DistributedOptimizer::new(adam, config)
134}
135
136/// Create a distributed optimizer with custom configuration
137///
138/// This function allows for full customization of the distributed training setup
139/// by accepting a custom DistributedConfig.
140///
141/// # Arguments
142///
143/// * `optimizer` - Base optimizer to wrap with distributed functionality
144/// * `config` - Distributed training configuration
145///
146/// # Returns
147///
148/// A distributed optimizer with the specified configuration
149///
150/// # Example
151///
152/// ```rust
153/// # use torsh_tensor::creation::randn;
154/// # use torsh_core::error::Result;
155/// # use parking_lot::RwLock;
156/// # use std::sync::Arc;
157/// # fn main() -> Result<()> {
158/// use torsh_optim::distributed::{utils::distributed_optimizer, core::*};
159/// use torsh_optim::AdamW;
160///
161/// // Create some parameters
162/// let param1 = Arc::new(RwLock::new(randn::<f32>(&[10, 20])?));
163/// let params = vec![param1];
164///
165/// let config = DistributedConfig {
166///     backend: DistributedBackend::NCCL,
167///     sync_strategy: SyncStrategy::AllReduce,
168///     world_size: 16,
169///     rank: 0,
170///     gradient_compression: true,
171///     bucket_size_mb: 50.0,
172///     overlap_communication: true,
173///     ..Default::default()
174/// };
175///
176/// let base_optimizer = AdamW::new(params, Some(1e-4), None, None, Some(0.01), false);
177/// let distributed_opt = distributed_optimizer(base_optimizer, config)?;
178/// # Ok(())
179/// # }
180/// ```
181pub fn distributed_optimizer<O: crate::Optimizer>(
182    optimizer: O,
183    config: DistributedConfig,
184) -> OptimizerResult<DistributedOptimizer<O>> {
185    DistributedOptimizer::new(optimizer, config)
186}
187
188/// Create distributed optimizer configurations for common scenarios
189pub mod configs {
190    use super::super::core::{DistributedBackend, DistributedConfig, SyncStrategy};
191
192    /// Configuration for CPU-based distributed training with MPI
193    pub fn cpu_mpi_config(world_size: usize, rank: usize) -> DistributedConfig {
194        DistributedConfig {
195            backend: DistributedBackend::MPI,
196            sync_strategy: SyncStrategy::AllReduce,
197            world_size,
198            rank,
199            gradient_compression: false,  // Less beneficial on CPU
200            bucket_size_mb: 10.0,         // Smaller buckets for CPU
201            overlap_communication: false, // Less effective on CPU
202            ..Default::default()
203        }
204    }
205
206    /// Configuration for GPU-based distributed training with NCCL
207    pub fn gpu_nccl_config(world_size: usize, rank: usize) -> DistributedConfig {
208        DistributedConfig {
209            backend: DistributedBackend::NCCL,
210            sync_strategy: SyncStrategy::AllReduce,
211            world_size,
212            rank,
213            gradient_compression: world_size >= 8, // Enable for large clusters
214            bucket_size_mb: 25.0,                  // Standard bucket size
215            overlap_communication: true,           // Enable overlap for GPUs
216            ..Default::default()
217        }
218    }
219
220    /// Configuration for mixed CPU/GPU training with Gloo
221    pub fn mixed_gloo_config(world_size: usize, rank: usize) -> DistributedConfig {
222        DistributedConfig {
223            backend: DistributedBackend::Gloo,
224            sync_strategy: SyncStrategy::AllReduce,
225            world_size,
226            rank,
227            gradient_compression: world_size > 4,
228            bucket_size_mb: 15.0,
229            overlap_communication: true,
230            ..Default::default()
231        }
232    }
233
234    /// Configuration optimized for large-scale training (many workers)
235    pub fn large_scale_config(world_size: usize, rank: usize) -> DistributedConfig {
236        DistributedConfig {
237            backend: DistributedBackend::NCCL,
238            sync_strategy: SyncStrategy::ReduceScatter, // More efficient for large scales
239            world_size,
240            rank,
241            gradient_compression: true,  // Essential for large scale
242            bucket_size_mb: 50.0,        // Larger buckets for efficiency
243            overlap_communication: true, // Critical for performance
244            ..Default::default()
245        }
246    }
247
248    /// Configuration for bandwidth-limited environments
249    pub fn low_bandwidth_config(world_size: usize, rank: usize) -> DistributedConfig {
250        DistributedConfig {
251            backend: DistributedBackend::Gloo,
252            sync_strategy: SyncStrategy::AllReduce,
253            world_size,
254            rank,
255            gradient_compression: true, // Always enable for low bandwidth
256            bucket_size_mb: 5.0,        // Small buckets for frequent communication
257            overlap_communication: true,
258            ..Default::default()
259        }
260    }
261}
262
263/// Utilities for monitoring and debugging distributed training
264pub mod monitoring {
265    use super::super::core::{CommunicationStats, DistributedOptimizer};
266    use crate::Optimizer;
267    use std::collections::HashMap;
268
269    /// Collect communication statistics from distributed optimizers
270    pub fn collect_communication_stats<O: Optimizer>(
271        optimizers: &[DistributedOptimizer<O>],
272    ) -> HashMap<usize, CommunicationStats> {
273        optimizers
274            .iter()
275            .enumerate()
276            .map(|(i, opt)| (i, opt.get_communication_stats()))
277            .collect()
278    }
279
280    /// Calculate aggregate statistics across all workers
281    pub fn aggregate_communication_stats(
282        stats: &HashMap<usize, CommunicationStats>,
283    ) -> CommunicationStats {
284        if stats.is_empty() {
285            return CommunicationStats::default();
286        }
287
288        let total_communications: u64 = stats.values().map(|s| s.total_communications).sum();
289        let total_bytes: u64 = stats.values().map(|s| s.total_bytes_transferred).sum();
290        let avg_time: f32 = stats
291            .values()
292            .map(|s| s.average_communication_time_ms)
293            .sum::<f32>()
294            / stats.len() as f32;
295        let avg_compression: f32 = stats
296            .values()
297            .map(|s| s.gradient_compression_ratio)
298            .sum::<f32>()
299            / stats.len() as f32;
300
301        CommunicationStats {
302            total_communications,
303            total_bytes_transferred: total_bytes,
304            average_communication_time_ms: avg_time,
305            gradient_compression_ratio: avg_compression,
306        }
307    }
308
309    /// Print a summary of distributed training performance
310    pub fn print_performance_summary<O: Optimizer>(optimizers: &[DistributedOptimizer<O>]) {
311        let stats = collect_communication_stats(optimizers);
312        let aggregate = aggregate_communication_stats(&stats);
313
314        println!("=== Distributed Training Performance Summary ===");
315        println!("Number of workers: {}", optimizers.len());
316        println!("Total communications: {}", aggregate.total_communications);
317        println!(
318            "Total bytes transferred: {:.2} MB",
319            aggregate.total_bytes_transferred as f64 / 1024.0 / 1024.0
320        );
321        println!(
322            "Average communication time: {:.2} ms",
323            aggregate.average_communication_time_ms
324        );
325        println!(
326            "Average compression ratio: {:.2}x",
327            aggregate.gradient_compression_ratio
328        );
329
330        // Per-worker breakdown
331        println!("\n--- Per-Worker Breakdown ---");
332        for (worker_id, stat) in &stats {
333            println!(
334                "Worker {}: {} communications, {:.2} MB, {:.2} ms avg",
335                worker_id,
336                stat.total_communications,
337                stat.total_bytes_transferred as f64 / 1024.0 / 1024.0,
338                stat.average_communication_time_ms
339            );
340        }
341    }
342}
343
344#[cfg(test)]
345mod tests {
346    use super::super::core::*;
347    use super::*;
348
349    #[test]
350    fn test_distributed_sgd_creation() {
351        // This would normally require actual tensors, but for testing we'll just check
352        // that the function signature is correct
353        // In a real test, you'd create actual tensors and verify the optimizer works
354    }
355
356    #[test]
357    fn test_config_creation() {
358        let config = configs::cpu_mpi_config(4, 0);
359        assert_eq!(config.world_size, 4);
360        assert_eq!(config.rank, 0);
361        assert!(matches!(config.backend, DistributedBackend::MPI));
362
363        let gpu_config = configs::gpu_nccl_config(8, 3);
364        assert_eq!(gpu_config.world_size, 8);
365        assert_eq!(gpu_config.rank, 3);
366        assert!(matches!(gpu_config.backend, DistributedBackend::NCCL));
367        assert!(gpu_config.gradient_compression); // Should be enabled for 8 workers
368    }
369
370    #[test]
371    fn test_large_scale_config() {
372        let config = configs::large_scale_config(64, 0);
373        assert!(matches!(config.sync_strategy, SyncStrategy::ReduceScatter));
374        assert!(config.gradient_compression);
375        assert_eq!(config.bucket_size_mb, 50.0);
376    }
377}