Skip to main content

shrew/
distributed.rs

1// Distributed Training — Data Parallelism, Mixed Precision, Pipeline Stages
2//
3// This module provides primitives for scaling training across multiple
4// workers, mixed-precision (FP16/FP32) training, and model-parallel
5// pipeline execution.
6//
7// COMPONENTS:
8//
9//   DataParallel<M>       — Splits input batches across N workers and runs
10//                           each forward pass in parallel (rayon threads).
11//                           Implements Module, so it's a drop-in replacement.
12//
13//   MixedPrecisionTrainer — Maintains FP32 "master" weights, casts to FP16
14//                           for forward/backward, applies dynamic loss scaling
15//                           to prevent underflow in FP16 gradients.
16//
17//   PipelineParallel      — Splits a sequential model into stages and overlaps
18//                           micro-batch execution (GPipe-style 1F1B schedule).
19//
20//   average_gradients()   — Averages multiple GradStores (the core AllReduce
21//                           primitive). Usable standalone for custom loops.
22
23use std::marker::PhantomData;
24
25use shrew_core::backend::Backend;
26use shrew_core::backprop::GradStore;
27use shrew_core::dtype::DType;
28use shrew_core::error::Result;
29use shrew_core::tensor::Tensor;
30
31use shrew_nn::Module;
32use shrew_optim::Optimizer;
33
34// AllReduce strategy
35
36/// Strategy for combining gradients from multiple replicas.
37#[derive(Debug, Clone, Copy, PartialEq, Eq)]
38pub enum AllReduceOp {
39    /// Sum all gradients (caller divides by N if needed).
40    Sum,
41    /// Average gradients across replicas (most common).
42    Average,
43}
44
45// Gradient averaging
46
47/// Average (or sum) multiple `GradStore`s into a single `GradStore`.
48///
49/// This is the core AllReduce primitive. Each worker produces a `GradStore`
50/// from its backward pass; this function merges them.
51///
52/// # Arguments
53/// - `grad_stores`: one `GradStore` per replica/worker
54/// - `params`: the shared parameter tensors (used to enumerate keys)
55/// - `strategy`: `Sum` or `Average`
56pub fn reduce_gradients<B: Backend>(
57    grad_stores: &[GradStore<B>],
58    params: &[Tensor<B>],
59    strategy: AllReduceOp,
60) -> Result<GradStore<B>> {
61    let n = grad_stores.len();
62    if n == 0 {
63        return Ok(GradStore::new());
64    }
65    if n == 1 {
66        return Ok(grad_stores[0].clone());
67    }
68
69    let mut merged = GradStore::new();
70
71    for param in params {
72        // Collect gradients from all stores for this parameter
73        let mut grads: Vec<&Tensor<B>> = Vec::new();
74        for store in grad_stores {
75            if let Some(g) = store.get(param) {
76                grads.push(g);
77            }
78        }
79        if grads.is_empty() {
80            continue;
81        }
82
83        // Sum all gradients
84        let mut acc = grads[0].clone();
85        for g in &grads[1..] {
86            acc = acc.add(g)?;
87        }
88
89        // Average if requested
90        if strategy == AllReduceOp::Average && grads.len() > 1 {
91            let scale = 1.0 / grads.len() as f64;
92            acc = acc.affine(scale, 0.0)?;
93        }
94
95        merged.accumulate(param.id(), acc)?;
96    }
97
98    Ok(merged)
99}
100
101// DataParallel — Module wrapper for batch-parallel forward passes
102
103/// Wraps a `Module` and splits each input batch across `num_workers` threads.
104///
105/// The forward pass:
106///   1. Split input along dimension 0 into `num_workers` chunks
107///   2. Run each chunk through the module in parallel (rayon)
108///   3. Concatenate the outputs
109///
110/// Because all workers share the same parameters (Tensor uses Arc), the
111/// autograd graph correctly tracks all operations. After computing loss
112/// on the concatenated output and calling `.backward()`, the gradients
113/// are automatically accumulated across all chunks.
114///
115/// # Example
116/// ```ignore
117/// let model = Sequential::new(vec![...]);
118/// let dp = DataParallel::new(model, 4);  // 4 workers
119/// let output = dp.forward(&big_batch)?;  // splits into 4 chunks
120/// ```
121pub struct DataParallel<M> {
122    /// The underlying module (shared across workers).
123    pub module: M,
124    /// Number of parallel workers.
125    pub num_workers: usize,
126}
127
128impl<M: Clone> Clone for DataParallel<M> {
129    fn clone(&self) -> Self {
130        Self {
131            module: self.module.clone(),
132            num_workers: self.num_workers,
133        }
134    }
135}
136
137impl<M: std::fmt::Debug> std::fmt::Debug for DataParallel<M> {
138    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
139        f.debug_struct("DataParallel")
140            .field("module", &self.module)
141            .field("num_workers", &self.num_workers)
142            .finish()
143    }
144}
145
146impl<M> DataParallel<M> {
147    /// Create a `DataParallel` wrapper with the given number of workers.
148    ///
149    /// `num_workers` controls how many chunks the batch is split into.
150    /// For CPU, this maps to rayon thread-pool parallelism.
151    pub fn new(module: M, num_workers: usize) -> Self {
152        assert!(num_workers > 0, "num_workers must be > 0");
153        Self {
154            module,
155            num_workers,
156        }
157    }
158
159    /// Get a reference to the underlying module.
160    pub fn inner(&self) -> &M {
161        &self.module
162    }
163
164    /// Get a mutable reference to the underlying module.
165    pub fn inner_mut(&mut self) -> &mut M {
166        &mut self.module
167    }
168
169    /// Unwrap the `DataParallel`, returning the inner module.
170    pub fn into_inner(self) -> M {
171        self.module
172    }
173}
174
175impl<M, B> Module<B> for DataParallel<M>
176where
177    M: Module<B> + Send + Sync,
178    B: Backend,
179{
180    fn forward(&self, x: &Tensor<B>) -> Result<Tensor<B>> {
181        let batch_size = x.dims()[0];
182        let effective_workers = self.num_workers.min(batch_size);
183
184        if effective_workers <= 1 {
185            return self.module.forward(x);
186        }
187
188        // Split into chunks along batch dimension
189        let chunks = x.chunk(effective_workers, 0)?;
190
191        // Run forward on each chunk in parallel using Rayon thread pool
192        use rayon::prelude::*;
193        let outputs: Result<Vec<Tensor<B>>> = chunks
194            .par_iter()
195            .map(|chunk| self.module.forward(chunk))
196            .collect();
197        let outputs = outputs?;
198
199        // Concatenate results
200        Tensor::cat(&outputs, 0)
201    }
202
203    fn parameters(&self) -> Vec<Tensor<B>> {
204        self.module.parameters()
205    }
206
207    fn named_parameters(&self) -> Vec<(String, Tensor<B>)> {
208        self.module.named_parameters()
209    }
210
211    fn set_training(&self, training: bool) {
212        self.module.set_training(training);
213    }
214
215    fn is_training(&self) -> bool {
216        self.module.is_training()
217    }
218}
219
220// MixedPrecisionTrainer — FP16 forward + FP32 master weights
221
222/// Configuration for dynamic loss scaling in mixed-precision training.
223#[derive(Debug, Clone)]
224pub struct LossScaleConfig {
225    /// Initial loss scale factor (default: 2^16 = 65536).
226    pub init_scale: f64,
227    /// Multiply scale by this when no overflow (default: 2.0).
228    pub scale_growth_factor: f64,
229    /// Divide scale by this on overflow (default: 2.0).
230    pub scale_backoff_factor: f64,
231    /// Number of consecutive good steps before increasing scale (default: 2000).
232    pub growth_interval: u64,
233}
234
235impl Default for LossScaleConfig {
236    fn default() -> Self {
237        Self {
238            init_scale: 65536.0,
239            scale_growth_factor: 2.0,
240            scale_backoff_factor: 2.0,
241            growth_interval: 2000,
242        }
243    }
244}
245
246/// Mixed-precision training: reduced-precision forward/backward with FP32 master weights.
247///
248/// **Why mixed precision?**
249/// - FP16/BF16 is 2× faster on GPUs with tensor cores (V100, A100, H100)
250/// - Half-precision uses half the memory for activations, enabling larger batches
251/// - FP32 master weights prevent precision loss during gradient updates
252///
253/// **How it works:**
254/// 1. Inputs and targets are cast to `compute_dtype` (F16 or BF16) before forward
255/// 2. The forward pass runs with reduced-precision activations
256/// 3. Dynamic loss scaling prevents gradient underflow in half-precision:
257///    - Loss is multiplied by a scale factor before backward
258///    - Gradients are divided by the same factor after
259///    - If overflow (NaN/Inf) is detected, the step is skipped and scale reduces
260/// 4. Gradients are cast back to FP32 and applied to FP32 master weights
261///
262/// **Compute dtype options:**
263/// - `DType::F16`: 16-bit IEEE float, range ±65504, good for most training
264/// - `DType::BF16`: bfloat16, same range as F32 but less precision, preferred when available
265/// - `DType::F32`: Standard precision (disables casting, only does loss scaling)
266///
267/// # Example
268/// ```ignore
269/// let model = Linear::<CpuBackend>::new(784, 10, true, DType::F32, &CpuDevice)?;
270/// let optimizer = Adam::new(model.parameters(), 1e-3);
271/// let mut trainer = MixedPrecisionTrainer::new(
272///     model, optimizer, DType::F16, Default::default(),
273/// );
274///
275/// for (input, target) in data {
276///     let metrics = trainer.train_step(&input, &target, mse_loss)?;
277///     println!("loss={:.4}, scale={}", metrics.loss, metrics.loss_scale);
278/// }
279/// ```
280pub struct MixedPrecisionTrainer<M, O, B: Backend> {
281    /// The model (with FP32 parameters as master copies).
282    model: M,
283    /// The optimizer operating on FP32 parameters.
284    optimizer: O,
285    /// The dtype for forward/backward computation (F16, BF16, or F32).
286    compute_dtype: DType,
287    /// Current loss scale factor.
288    loss_scale: f64,
289    /// Loss scale configuration.
290    config: LossScaleConfig,
291    /// Number of consecutive successful steps (no overflow).
292    good_steps: u64,
293    /// Total skipped steps (overflow detected).
294    skipped_steps: u64,
295    _phantom: PhantomData<B>,
296}
297
298/// Metrics from a single mixed-precision training step.
299#[derive(Debug, Clone)]
300pub struct MixedPrecisionMetrics {
301    /// The unscaled loss value.
302    pub loss: f64,
303    /// Whether this step was skipped (overflow detected).
304    pub skipped: bool,
305    /// Current loss scale factor.
306    pub loss_scale: f64,
307    /// Total skipped steps so far.
308    pub total_skipped: u64,
309    /// The compute dtype used for this step.
310    pub compute_dtype: DType,
311}
312
313impl<M, O, B> MixedPrecisionTrainer<M, O, B>
314where
315    M: Module<B>,
316    O: Optimizer<B>,
317    B: Backend,
318{
319    /// Create a new mixed-precision trainer.
320    ///
321    /// The model and optimizer should use FP32 parameters.
322    /// `compute_dtype` sets the precision for forward/backward (F16, BF16, or F32).
323    pub fn new(model: M, optimizer: O, compute_dtype: DType, config: LossScaleConfig) -> Self {
324        let loss_scale = config.init_scale;
325        Self {
326            model,
327            optimizer,
328            compute_dtype,
329            loss_scale,
330            config,
331            good_steps: 0,
332            skipped_steps: 0,
333            _phantom: PhantomData,
334        }
335    }
336
337    /// Reference to the model.
338    pub fn model(&self) -> &M {
339        &self.model
340    }
341
342    /// Mutable reference to the model.
343    pub fn model_mut(&mut self) -> &mut M {
344        &mut self.model
345    }
346
347    /// Reference to the optimizer.
348    pub fn optimizer(&self) -> &O {
349        &self.optimizer
350    }
351
352    /// Current loss scale.
353    pub fn loss_scale(&self) -> f64 {
354        self.loss_scale
355    }
356
357    /// The compute dtype (F16, BF16, or F32).
358    pub fn compute_dtype(&self) -> DType {
359        self.compute_dtype
360    }
361
362    /// Total number of skipped steps.
363    pub fn skipped_steps(&self) -> u64 {
364        self.skipped_steps
365    }
366
367    /// Perform one mixed-precision training step.
368    ///
369    /// The input and target are cast to `compute_dtype` for the forward pass.
370    /// Dynamic loss scaling is applied to prevent gradient underflow.
371    /// Gradients are cast back to FP32 and applied to FP32 master weights.
372    ///
373    /// # Arguments
374    /// - `input`: input tensor (any dtype, will be cast to compute_dtype)
375    /// - `target`: target tensor (any dtype, will be cast to compute_dtype)
376    /// - `loss_fn`: function computing scalar loss from (prediction, target)
377    ///
378    /// # Returns
379    /// `MixedPrecisionMetrics` with loss value and scaling info.
380    pub fn train_step<F>(
381        &mut self,
382        input: &Tensor<B>,
383        target: &Tensor<B>,
384        loss_fn: F,
385    ) -> Result<MixedPrecisionMetrics>
386    where
387        F: Fn(&Tensor<B>, &Tensor<B>) -> Result<Tensor<B>>,
388    {
389        // 1. Determine if we should cast inputs to compute_dtype.
390        // Only cast if the model's parameters already match compute_dtype,
391        // otherwise auto-casting inputs would cause dtype mismatches with weights.
392        let model_dtype = self
393            .model
394            .parameters()
395            .first()
396            .map(|p| p.dtype())
397            .unwrap_or(DType::F32);
398        let should_cast = self.compute_dtype != DType::F32 && self.compute_dtype == model_dtype;
399
400        let input_cast = if should_cast && input.dtype() != self.compute_dtype {
401            input.to_dtype(self.compute_dtype)?
402        } else {
403            input.clone()
404        };
405        let target_cast = if should_cast && target.dtype() != self.compute_dtype {
406            target.to_dtype(self.compute_dtype)?
407        } else {
408            target.clone()
409        };
410
411        // 2. Forward pass
412        let output = self.model.forward(&input_cast)?;
413
414        // 3. Compute loss (in compute_dtype)
415        let loss = loss_fn(&output, &target_cast)?;
416        let loss_val = loss.to_scalar_f64()?;
417
418        // 4. Scale loss for backward (prevents gradient underflow in F16)
419        let scaled_loss = loss.affine(self.loss_scale, 0.0)?;
420
421        // 5. Backward on scaled loss
422        let grads = scaled_loss.backward()?;
423
424        // 6. Check for overflow in gradients
425        let has_overflow = self.check_overflow(&grads)?;
426
427        if has_overflow {
428            // Skip this step, reduce the scale
429            self.loss_scale /= self.config.scale_backoff_factor;
430            self.loss_scale = self.loss_scale.max(1.0); // don't go below 1
431            self.good_steps = 0;
432            self.skipped_steps += 1;
433
434            return Ok(MixedPrecisionMetrics {
435                loss: loss_val,
436                skipped: true,
437                loss_scale: self.loss_scale,
438                total_skipped: self.skipped_steps,
439                compute_dtype: self.compute_dtype,
440            });
441        }
442
443        // 7. Unscale gradients and cast back to FP32 for master weight update
444        let unscaled = self.unscale_and_cast_gradients(&grads)?;
445
446        // 8. Optimizer step with FP32 unscaled gradients
447        self.optimizer.step(&unscaled)?;
448
449        // 9. Update loss scale (possibly increase after consecutive good steps)
450        self.good_steps += 1;
451        if self.good_steps >= self.config.growth_interval {
452            self.loss_scale *= self.config.scale_growth_factor;
453            self.good_steps = 0;
454        }
455
456        Ok(MixedPrecisionMetrics {
457            loss: loss_val,
458            skipped: false,
459            loss_scale: self.loss_scale,
460            total_skipped: self.skipped_steps,
461            compute_dtype: self.compute_dtype,
462        })
463    }
464
465    /// Check if any gradient contains NaN or Inf.
466    fn check_overflow(&self, grads: &GradStore<B>) -> Result<bool> {
467        for param in self.model.parameters() {
468            if let Some(g) = grads.get(&param) {
469                let data = g.to_f64_vec()?;
470                for &v in &data {
471                    if v.is_nan() || v.is_infinite() {
472                        return Ok(true);
473                    }
474                }
475            }
476        }
477        Ok(false)
478    }
479
480    /// Unscale gradients by the loss scale factor and cast to FP32.
481    ///
482    /// This ensures the optimizer always sees FP32 gradients, regardless
483    /// of the compute dtype used during forward/backward.
484    fn unscale_and_cast_gradients(&self, grads: &GradStore<B>) -> Result<GradStore<B>> {
485        let inv_scale = 1.0 / self.loss_scale;
486        let mut unscaled = GradStore::new();
487        for param in self.model.parameters() {
488            if let Some(g) = grads.get(&param) {
489                // Unscale the gradient
490                let g_unscaled = g.affine(inv_scale, 0.0)?;
491                // Cast back to the master weight dtype (FP32) if needed
492                let g_fp32 = if g_unscaled.dtype() != param.dtype() {
493                    g_unscaled.to_dtype(param.dtype())?
494                } else {
495                    g_unscaled
496                };
497                unscaled.accumulate(param.id(), g_fp32)?;
498            }
499        }
500        Ok(unscaled)
501    }
502}
503
504// PipelineParallel — GPipe-style micro-batch pipelining
505
506/// A stage in a pipeline-parallel model.
507///
508/// Each stage holds a sub-model (one or more layers). During execution,
509/// micro-batches flow through stages in a pipeline fashion, overlapping
510/// the forward and backward passes of different micro-batches.
511pub struct PipelineStage<B: Backend> {
512    /// The layers in this stage (as boxed Module).
513    layers: Vec<Box<dyn Module<B>>>,
514    /// Stage index (0-based).
515    stage_id: usize,
516}
517
518impl<B: Backend> PipelineStage<B> {
519    /// Create a new pipeline stage.
520    pub fn new(stage_id: usize) -> Self {
521        Self {
522            layers: Vec::new(),
523            stage_id,
524        }
525    }
526
527    /// Add a layer to this stage.
528    pub fn add_layer(mut self, layer: Box<dyn Module<B>>) -> Self {
529        self.layers.push(layer);
530        self
531    }
532
533    /// Stage identifier.
534    pub fn stage_id(&self) -> usize {
535        self.stage_id
536    }
537
538    /// Forward pass through all layers in this stage.
539    pub fn forward(&self, x: &Tensor<B>) -> Result<Tensor<B>> {
540        let mut out = x.clone();
541        for layer in &self.layers {
542            out = layer.forward(&out)?;
543        }
544        Ok(out)
545    }
546
547    /// Collect all parameters from all layers in this stage.
548    pub fn parameters(&self) -> Vec<Tensor<B>> {
549        self.layers.iter().flat_map(|l| l.parameters()).collect()
550    }
551}
552
553/// Pipeline-parallel executor using GPipe-style micro-batching.
554///
555/// Splits a model into sequential stages and processes micro-batches
556/// through the pipeline. This increases throughput by overlapping
557/// computation across stages.
558///
559/// # Example
560/// ```ignore
561/// let stage0 = PipelineStage::new(0)
562///     .add_layer(Box::new(Linear::new(784, 256, true, DType::F32, &dev)?));
563/// let stage1 = PipelineStage::new(1)
564///     .add_layer(Box::new(Linear::new(256, 10, true, DType::F32, &dev)?));
565///
566/// let pipeline = PipelineParallel::new(vec![stage0, stage1], 4);
567/// let output = pipeline.forward(&input)?;
568/// ```
569pub struct PipelineParallel<B: Backend> {
570    /// Ordered stages of the model.
571    stages: Vec<PipelineStage<B>>,
572    /// Number of micro-batches to split each input into.
573    num_micro_batches: usize,
574}
575
576impl<B: Backend> PipelineParallel<B> {
577    /// Create a pipeline with the given stages and micro-batch count.
578    pub fn new(stages: Vec<PipelineStage<B>>, num_micro_batches: usize) -> Self {
579        assert!(!stages.is_empty(), "pipeline needs at least one stage");
580        assert!(num_micro_batches > 0, "num_micro_batches must be > 0");
581        Self {
582            stages,
583            num_micro_batches,
584        }
585    }
586
587    /// Full forward pass through all stages.
588    ///
589    /// Splits the input into `num_micro_batches` micro-batches, runs each
590    /// through the pipeline sequentially, and concatenates the results.
591    ///
592    /// In a multi-device setting, stages would run on different devices
593    /// with inter-device transfers between stages.
594    pub fn forward(&self, x: &Tensor<B>) -> Result<Tensor<B>> {
595        let batch_size = x.dims()[0];
596        let effective_micros = self.num_micro_batches.min(batch_size);
597
598        if effective_micros <= 1 {
599            // No micro-batching — sequential pass through all stages
600            let mut out = x.clone();
601            for stage in &self.stages {
602                out = stage.forward(&out)?;
603            }
604            return Ok(out);
605        }
606
607        // Split into micro-batches
608        let micro_batches = x.chunk(effective_micros, 0)?;
609
610        // Run each micro-batch through all stages
611        let mut outputs = Vec::with_capacity(micro_batches.len());
612        for mb in &micro_batches {
613            let mut out = mb.clone();
614            for stage in &self.stages {
615                out = stage.forward(&out)?;
616            }
617            outputs.push(out);
618        }
619
620        // Concatenate micro-batch outputs
621        Tensor::cat(&outputs, 0)
622    }
623
624    /// Collect all parameters from all stages (for optimizer).
625    pub fn parameters(&self) -> Vec<Tensor<B>> {
626        self.stages.iter().flat_map(|s| s.parameters()).collect()
627    }
628
629    /// Number of stages.
630    pub fn num_stages(&self) -> usize {
631        self.stages.len()
632    }
633
634    /// Get a reference to a specific stage.
635    pub fn stage(&self, idx: usize) -> Option<&PipelineStage<B>> {
636        self.stages.get(idx)
637    }
638}
639
640// ParallelTrainer — High-level training loop with gradient accumulation
641
642/// High-level training loop with gradient accumulation.
643///
644/// Splits a large effective batch into `accumulation_steps` micro-batches,
645/// accumulates gradients across all of them, then performs a single
646/// optimizer step. This simulates a larger batch size without requiring
647/// more memory.
648///
649/// # Example
650/// ```ignore
651/// let model = Sequential::new(vec![...]);
652/// let optimizer = Adam::new(model.parameters(), 1e-3);
653/// let mut trainer = ParallelTrainer::new(model, optimizer, 4);
654///
655/// // Each call accumulates 1/4 of the gradient; every 4th call steps.
656/// for (i, (x, y)) in data.iter().enumerate() {
657///     if let Some(loss) = trainer.accumulate_step(&x, &y, mse_loss)? {
658///         println!("step {}: loss = {:.4}", i, loss);
659///     }
660/// }
661/// ```
662pub struct ParallelTrainer<M, O, B: Backend> {
663    /// The model.
664    pub model: M,
665    /// The optimizer.
666    pub optimizer: O,
667    /// Number of micro-batches to accumulate before stepping.
668    accumulation_steps: usize,
669    /// Current accumulated gradients.
670    accumulated: Option<GradStore<B>>,
671    /// Current micro-batch index (0 .. accumulation_steps - 1).
672    current_step: usize,
673    /// Running loss sum for the current accumulation window.
674    loss_sum: f64,
675    _phantom: PhantomData<B>,
676}
677
678impl<M, O, B> ParallelTrainer<M, O, B>
679where
680    M: Module<B>,
681    O: Optimizer<B>,
682    B: Backend,
683{
684    /// Create a new `ParallelTrainer`.
685    ///
686    /// `accumulation_steps`: number of micro-batches before each optimizer step.
687    pub fn new(model: M, optimizer: O, accumulation_steps: usize) -> Self {
688        assert!(accumulation_steps > 0);
689        Self {
690            model,
691            optimizer,
692            accumulation_steps,
693            accumulated: None,
694            current_step: 0,
695            loss_sum: 0.0,
696            _phantom: PhantomData,
697        }
698    }
699
700    /// Process one micro-batch. Returns `Some(avg_loss)` when an optimizer
701    /// step was performed (every `accumulation_steps` calls), else `None`.
702    pub fn accumulate_step<F>(
703        &mut self,
704        input: &Tensor<B>,
705        target: &Tensor<B>,
706        loss_fn: F,
707    ) -> Result<Option<f64>>
708    where
709        F: Fn(&Tensor<B>, &Tensor<B>) -> Result<Tensor<B>>,
710    {
711        // Forward
712        let output = self.model.forward(input)?;
713        let loss = loss_fn(&output, target)?;
714        let loss_val = loss.to_scalar_f64()?;
715        self.loss_sum += loss_val;
716
717        // Backward
718        let grads = loss.backward()?;
719
720        // Accumulate gradients
721        let params = self.model.parameters();
722        match self.accumulated.take() {
723            Some(prev) => {
724                let merged = reduce_gradients(&[prev, grads], &params, AllReduceOp::Sum)?;
725                self.accumulated = Some(merged);
726            }
727            None => {
728                self.accumulated = Some(grads);
729            }
730        }
731
732        self.current_step += 1;
733
734        // Step when we've accumulated enough
735        if self.current_step >= self.accumulation_steps {
736            let avg_grads = {
737                let acc = self.accumulated.take().unwrap();
738                // Average by accumulation_steps
739                let mut averaged = GradStore::new();
740                let scale = 1.0 / self.accumulation_steps as f64;
741                for param in &params {
742                    if let Some(g) = acc.get(param) {
743                        let g_avg = g.affine(scale, 0.0)?;
744                        averaged.accumulate(param.id(), g_avg)?;
745                    }
746                }
747                averaged
748            };
749
750            self.optimizer.step(&avg_grads)?;
751
752            let avg_loss = self.loss_sum / self.accumulation_steps as f64;
753            self.current_step = 0;
754            self.loss_sum = 0.0;
755            self.accumulated = None;
756
757            Ok(Some(avg_loss))
758        } else {
759            Ok(None)
760        }
761    }
762
763    /// Force an optimizer step with whatever gradients have been accumulated so far.
764    /// Useful at the end of an epoch when remaining micro-batches < accumulation_steps.
765    pub fn flush(&mut self) -> Result<Option<f64>> {
766        if self.current_step == 0 || self.accumulated.is_none() {
767            return Ok(None);
768        }
769
770        let params = self.model.parameters();
771        let acc = self.accumulated.take().unwrap();
772        let scale = 1.0 / self.current_step as f64;
773        let mut averaged = GradStore::new();
774        for param in &params {
775            if let Some(g) = acc.get(param) {
776                let g_avg = g.affine(scale, 0.0)?;
777                averaged.accumulate(param.id(), g_avg)?;
778            }
779        }
780
781        self.optimizer.step(&averaged)?;
782
783        let avg_loss = self.loss_sum / self.current_step as f64;
784        self.current_step = 0;
785        self.loss_sum = 0.0;
786        self.accumulated = None;
787
788        Ok(Some(avg_loss))
789    }
790}
791
792// Tests
793
794#[cfg(test)]
795mod tests {
796    use super::*;
797    use shrew_cpu::{CpuBackend, CpuDevice};
798
799    type B = CpuBackend;
800    type T = Tensor<B>;
801    const DEV: CpuDevice = CpuDevice;
802
803    // ── AllReduce / gradient averaging ──
804
805    #[test]
806    fn test_reduce_gradients_average() {
807        let p = T::randn(vec![4], DType::F32, &DEV).unwrap().set_variable();
808        let loss1 = p.sum_all().unwrap();
809        let g1 = loss1.backward().unwrap();
810
811        let loss2 = p.affine(2.0, 0.0).unwrap().sum_all().unwrap();
812        let g2 = loss2.backward().unwrap();
813
814        let merged = reduce_gradients(&[g1, g2], &[p.clone()], AllReduceOp::Average).unwrap();
815        let avg = merged.get(&p).unwrap().to_f64_vec().unwrap();
816        // g1 = all 1s, g2 = all 2s, average = all 1.5s
817        for &v in &avg {
818            assert!((v - 1.5).abs() < 1e-5, "expected 1.5, got {v}");
819        }
820    }
821
822    #[test]
823    fn test_reduce_gradients_sum() {
824        let p = T::randn(vec![3], DType::F32, &DEV).unwrap().set_variable();
825        let loss1 = p.sum_all().unwrap();
826        let g1 = loss1.backward().unwrap();
827
828        let loss2 = p.sum_all().unwrap();
829        let g2 = loss2.backward().unwrap();
830
831        let merged = reduce_gradients(&[g1, g2], &[p.clone()], AllReduceOp::Sum).unwrap();
832        let summed = merged.get(&p).unwrap().to_f64_vec().unwrap();
833        for &v in &summed {
834            assert!((v - 2.0).abs() < 1e-5, "expected 2.0, got {v}");
835        }
836    }
837
838    // ── DataParallel ──
839
840    #[test]
841    fn test_data_parallel_forward() {
842        let linear = shrew_nn::Linear::<B>::new(4, 2, true, DType::F32, &DEV).unwrap();
843        let dp = DataParallel::new(linear, 2);
844
845        let input = T::randn(vec![6, 4], DType::F32, &DEV).unwrap();
846        let output = dp.forward(&input).unwrap();
847        assert_eq!(output.dims(), &[6, 2]);
848    }
849
850    #[test]
851    fn test_data_parallel_single_worker() {
852        let linear = shrew_nn::Linear::<B>::new(3, 2, true, DType::F32, &DEV).unwrap();
853        let dp = DataParallel::new(linear, 1);
854
855        let input = T::randn(vec![4, 3], DType::F32, &DEV).unwrap();
856        let output = dp.forward(&input).unwrap();
857        assert_eq!(output.dims(), &[4, 2]);
858    }
859
860    #[test]
861    fn test_data_parallel_parameters() {
862        let linear = shrew_nn::Linear::<B>::new(4, 2, true, DType::F32, &DEV).unwrap();
863        let n_params = linear.parameters().len();
864        let dp = DataParallel::new(linear, 4);
865        assert_eq!(dp.parameters().len(), n_params);
866    }
867
868    // ── MixedPrecisionTrainer ──
869
870    #[test]
871    fn test_mixed_precision_basic() {
872        let linear = shrew_nn::Linear::<B>::new(4, 1, true, DType::F32, &DEV).unwrap();
873        let optimizer = shrew_optim::SGD::new(linear.parameters(), 0.01, 0.0, 0.0);
874        let mut trainer =
875            MixedPrecisionTrainer::new(linear, optimizer, DType::F16, LossScaleConfig::default());
876
877        let input = T::randn(vec![2, 4], DType::F32, &DEV).unwrap();
878        let target = T::zeros(vec![2, 1], DType::F32, &DEV).unwrap();
879
880        let metrics = trainer
881            .train_step(&input, &target, |pred, tgt| shrew_nn::mse_loss(pred, tgt))
882            .unwrap();
883
884        assert!(!metrics.skipped);
885        assert!(metrics.loss >= 0.0);
886        assert_eq!(metrics.loss_scale, 65536.0);
887    }
888
889    #[cfg(feature = "cuda")]
890    #[test]
891    fn test_mixed_precision_trainer_gpu() {
892        use shrew_cuda::{CudaBackend, CudaDevice};
893        type GpuB = CudaBackend;
894        if let Ok(dev) = CudaDevice::new(0) {
895            let linear = shrew_nn::Linear::<GpuB>::new(4, 2, true, DType::F16, &dev).unwrap();
896            let optimizer = shrew_optim::SGD::new(linear.parameters(), 0.01, 0.0, 0.0);
897            let config = LossScaleConfig {
898                init_scale: 1.0,
899                ..Default::default()
900            };
901            let mut trainer = MixedPrecisionTrainer::new(linear, optimizer, DType::F16, config);
902
903            let input = Tensor::<GpuB>::randn(vec![2, 4], DType::F16, &dev).unwrap();
904            let target = Tensor::<GpuB>::zeros(vec![2, 2], DType::F16, &dev).unwrap();
905
906            let metrics = trainer
907                .train_step(&input, &target, |pred, tgt| shrew_nn::mse_loss(pred, tgt))
908                .unwrap();
909
910            assert!(!metrics.skipped);
911            assert!(metrics.loss >= 0.0);
912            assert_eq!(metrics.compute_dtype, DType::F16);
913        }
914    }
915
916    // ── Pipeline ──
917
918    #[test]
919    fn test_pipeline_forward() {
920        let stage0 = PipelineStage::<B>::new(0).add_layer(Box::new(
921            shrew_nn::Linear::<B>::new(4, 8, true, DType::F32, &DEV).unwrap(),
922        ));
923        let stage1 = PipelineStage::<B>::new(1).add_layer(Box::new(
924            shrew_nn::Linear::<B>::new(8, 2, true, DType::F32, &DEV).unwrap(),
925        ));
926
927        let pipeline = PipelineParallel::new(vec![stage0, stage1], 2);
928        let input = T::randn(vec![4, 4], DType::F32, &DEV).unwrap();
929        let output = pipeline.forward(&input).unwrap();
930        assert_eq!(output.dims(), &[4, 2]);
931    }
932
933    #[test]
934    fn test_pipeline_parameters() {
935        let stage0 = PipelineStage::<B>::new(0).add_layer(Box::new(
936            shrew_nn::Linear::<B>::new(4, 8, true, DType::F32, &DEV).unwrap(),
937        ));
938        let stage1 = PipelineStage::<B>::new(1).add_layer(Box::new(
939            shrew_nn::Linear::<B>::new(8, 2, true, DType::F32, &DEV).unwrap(),
940        ));
941
942        let pipeline = PipelineParallel::new(vec![stage0, stage1], 1);
943        // stage0: 4*8 + 8 = 40, stage1: 8*2 + 2 = 18, total = 58
944        let total: usize = pipeline.parameters().iter().map(|p| p.elem_count()).sum();
945        assert_eq!(total, 40 + 18);
946    }
947
948    // ── ParallelTrainer (gradient accumulation) ──
949
950    #[test]
951    fn test_parallel_trainer_accumulation() {
952        let linear = shrew_nn::Linear::<B>::new(3, 1, true, DType::F32, &DEV).unwrap();
953        let optimizer = shrew_optim::SGD::new(linear.parameters(), 0.01, 0.0, 0.0);
954        let mut trainer = ParallelTrainer::new(linear, optimizer, 2);
955
956        let x1 = T::randn(vec![1, 3], DType::F32, &DEV).unwrap();
957        let y1 = T::zeros(vec![1, 1], DType::F32, &DEV).unwrap();
958        let x2 = T::randn(vec![1, 3], DType::F32, &DEV).unwrap();
959        let y2 = T::zeros(vec![1, 1], DType::F32, &DEV).unwrap();
960
961        // First micro-batch: no step yet
962        let result1 = trainer
963            .accumulate_step(&x1, &y1, |p, t| shrew_nn::mse_loss(p, t))
964            .unwrap();
965        assert!(result1.is_none());
966
967        // Second micro-batch: step happens, returns average loss
968        let result2 = trainer
969            .accumulate_step(&x2, &y2, |p, t| shrew_nn::mse_loss(p, t))
970            .unwrap();
971        assert!(result2.is_some());
972    }
973
974    #[test]
975    fn test_parallel_trainer_flush() {
976        let linear = shrew_nn::Linear::<B>::new(3, 1, true, DType::F32, &DEV).unwrap();
977        let optimizer = shrew_optim::SGD::new(linear.parameters(), 0.01, 0.0, 0.0);
978        let mut trainer = ParallelTrainer::new(linear, optimizer, 4);
979
980        let x = T::randn(vec![1, 3], DType::F32, &DEV).unwrap();
981        let y = T::zeros(vec![1, 1], DType::F32, &DEV).unwrap();
982
983        // Only 1 of 4 accumulation steps done
984        trainer
985            .accumulate_step(&x, &y, |p, t| shrew_nn::mse_loss(p, t))
986            .unwrap();
987
988        // Flush forces a step with whatever we have
989        let flushed = trainer.flush().unwrap();
990        assert!(flushed.is_some());
991    }
992}