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}