Skip to main content

distributed_optimizer

Function distributed_optimizer 

Source
pub fn distributed_optimizer<O: Optimizer>(
    optimizer: O,
    config: DistributedConfig,
) -> OptimizerResult<DistributedOptimizer<O>>
Expand description

Create a distributed optimizer with custom configuration

This function allows for full customization of the distributed training setup by accepting a custom DistributedConfig.

§Arguments

  • optimizer - Base optimizer to wrap with distributed functionality
  • config - Distributed training configuration

§Returns

A distributed optimizer with the specified configuration

§Example

use torsh_optim::distributed::{utils::distributed_optimizer, core::*};
use torsh_optim::AdamW;

// Create some parameters
let param1 = Arc::new(RwLock::new(randn::<f32>(&[10, 20])?));
let params = vec![param1];

let config = DistributedConfig {
    backend: DistributedBackend::NCCL,
    sync_strategy: SyncStrategy::AllReduce,
    world_size: 16,
    rank: 0,
    gradient_compression: true,
    bucket_size_mb: 50.0,
    overlap_communication: true,
    ..Default::default()
};

let base_optimizer = AdamW::new(params, Some(1e-4), None, None, Some(0.01), false);
let distributed_opt = distributed_optimizer(base_optimizer, config)?;