1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
38pub enum AllReduceOp {
39 Sum,
41 Average,
43}
44
45pub 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 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 let mut acc = grads[0].clone();
85 for g in &grads[1..] {
86 acc = acc.add(g)?;
87 }
88
89 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
101pub struct DataParallel<M> {
122 pub module: M,
124 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 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 pub fn inner(&self) -> &M {
161 &self.module
162 }
163
164 pub fn inner_mut(&mut self) -> &mut M {
166 &mut self.module
167 }
168
169 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 let chunks = x.chunk(effective_workers, 0)?;
190
191 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 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#[derive(Debug, Clone)]
224pub struct LossScaleConfig {
225 pub init_scale: f64,
227 pub scale_growth_factor: f64,
229 pub scale_backoff_factor: f64,
231 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
246pub struct MixedPrecisionTrainer<M, O, B: Backend> {
281 model: M,
283 optimizer: O,
285 compute_dtype: DType,
287 loss_scale: f64,
289 config: LossScaleConfig,
291 good_steps: u64,
293 skipped_steps: u64,
295 _phantom: PhantomData<B>,
296}
297
298#[derive(Debug, Clone)]
300pub struct MixedPrecisionMetrics {
301 pub loss: f64,
303 pub skipped: bool,
305 pub loss_scale: f64,
307 pub total_skipped: u64,
309 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 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 pub fn model(&self) -> &M {
339 &self.model
340 }
341
342 pub fn model_mut(&mut self) -> &mut M {
344 &mut self.model
345 }
346
347 pub fn optimizer(&self) -> &O {
349 &self.optimizer
350 }
351
352 pub fn loss_scale(&self) -> f64 {
354 self.loss_scale
355 }
356
357 pub fn compute_dtype(&self) -> DType {
359 self.compute_dtype
360 }
361
362 pub fn skipped_steps(&self) -> u64 {
364 self.skipped_steps
365 }
366
367 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 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 let output = self.model.forward(&input_cast)?;
413
414 let loss = loss_fn(&output, &target_cast)?;
416 let loss_val = loss.to_scalar_f64()?;
417
418 let scaled_loss = loss.affine(self.loss_scale, 0.0)?;
420
421 let grads = scaled_loss.backward()?;
423
424 let has_overflow = self.check_overflow(&grads)?;
426
427 if has_overflow {
428 self.loss_scale /= self.config.scale_backoff_factor;
430 self.loss_scale = self.loss_scale.max(1.0); 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 let unscaled = self.unscale_and_cast_gradients(&grads)?;
445
446 self.optimizer.step(&unscaled)?;
448
449 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 fn check_overflow(&self, grads: &GradStore<B>) -> Result<bool> {
467 for param in self.model.parameters() {
468 if let Some(g) = grads.get(¶m) {
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 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(¶m) {
489 let g_unscaled = g.affine(inv_scale, 0.0)?;
491 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
504pub struct PipelineStage<B: Backend> {
512 layers: Vec<Box<dyn Module<B>>>,
514 stage_id: usize,
516}
517
518impl<B: Backend> PipelineStage<B> {
519 pub fn new(stage_id: usize) -> Self {
521 Self {
522 layers: Vec::new(),
523 stage_id,
524 }
525 }
526
527 pub fn add_layer(mut self, layer: Box<dyn Module<B>>) -> Self {
529 self.layers.push(layer);
530 self
531 }
532
533 pub fn stage_id(&self) -> usize {
535 self.stage_id
536 }
537
538 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 pub fn parameters(&self) -> Vec<Tensor<B>> {
549 self.layers.iter().flat_map(|l| l.parameters()).collect()
550 }
551}
552
553pub struct PipelineParallel<B: Backend> {
570 stages: Vec<PipelineStage<B>>,
572 num_micro_batches: usize,
574}
575
576impl<B: Backend> PipelineParallel<B> {
577 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 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 let mut out = x.clone();
601 for stage in &self.stages {
602 out = stage.forward(&out)?;
603 }
604 return Ok(out);
605 }
606
607 let micro_batches = x.chunk(effective_micros, 0)?;
609
610 let mut outputs = Vec::with_capacity(micro_batches.len());
612 for mb in µ_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 Tensor::cat(&outputs, 0)
622 }
623
624 pub fn parameters(&self) -> Vec<Tensor<B>> {
626 self.stages.iter().flat_map(|s| s.parameters()).collect()
627 }
628
629 pub fn num_stages(&self) -> usize {
631 self.stages.len()
632 }
633
634 pub fn stage(&self, idx: usize) -> Option<&PipelineStage<B>> {
636 self.stages.get(idx)
637 }
638}
639
640pub struct ParallelTrainer<M, O, B: Backend> {
663 pub model: M,
665 pub optimizer: O,
667 accumulation_steps: usize,
669 accumulated: Option<GradStore<B>>,
671 current_step: usize,
673 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 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 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 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 let grads = loss.backward()?;
719
720 let params = self.model.parameters();
722 match self.accumulated.take() {
723 Some(prev) => {
724 let merged = reduce_gradients(&[prev, grads], ¶ms, 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 if self.current_step >= self.accumulation_steps {
736 let avg_grads = {
737 let acc = self.accumulated.take().unwrap();
738 let mut averaged = GradStore::new();
740 let scale = 1.0 / self.accumulation_steps as f64;
741 for param in ¶ms {
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 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 ¶ms {
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#[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 #[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 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 #[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 #[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 #[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 let total: usize = pipeline.parameters().iter().map(|p| p.elem_count()).sum();
945 assert_eq!(total, 40 + 18);
946 }
947
948 #[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 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 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 trainer
985 .accumulate_step(&x, &y, |p, t| shrew_nn::mse_loss(p, t))
986 .unwrap();
987
988 let flushed = trainer.flush().unwrap();
990 assert!(flushed.is_some());
991 }
992}