1use burn_core as burn;
2use burn_core::module::ParamGroup;
3
4use super::Optimizer;
5use crate::lr_scheduler::module_lr_scheduler::ModuleLearningRate;
6use crate::{
7 DynOptimizer, DynState, MultiGradientsParams, OptimizerRecord, StateSink, StateSource,
8 grad_clipping::GradientClipping, optim::GradientsParams, optim::state::join_path,
9};
10
11use alloc::collections::BTreeMap;
12use alloc::string::ToString;
13use alloc::sync::Arc;
14use alloc::vec::Vec;
15use burn::module::{AutodiffModule, ModuleMapper, Param, ParamId};
16use burn::store::RecordError;
17use burn::tensor::{Bytes, Device, Tensor, TensorData};
18use hashbrown::HashMap;
19
20const RANK_KEY: &str = "__rank";
25
26#[derive(Clone)]
27struct OptimizerGroup {
28 group: ParamGroup,
29 optim: Arc<dyn DynOptimizer>,
30 grad_clipping: Option<GradientClipping>,
31}
32
33impl OptimizerGroup {
34 pub(crate) fn set_gradient_clipping(&mut self, gradient_clipping: GradientClipping) {
35 self.grad_clipping = Some(gradient_clipping)
36 }
37}
38
39#[derive(Clone)]
41struct OptimizationContext {
42 optim: Arc<dyn DynOptimizer>,
43 grad_clipping: Option<GradientClipping>,
44 path: Option<String>,
45 state: DynState,
46}
47
48#[derive(Clone)]
59pub struct ModuleOptimizer {
60 optimizers: Vec<OptimizerGroup>,
61 param_context: HashMap<ParamId, OptimizationContext>,
62}
63
64impl<O> From<O> for ModuleOptimizer
65where
66 O: Optimizer,
67{
68 fn from(optim: O) -> Self {
69 Self {
70 param_context: HashMap::new(),
71 optimizers: vec![OptimizerGroup {
72 group: ParamGroup::all(),
73 optim: Arc::new(optim),
74 grad_clipping: None,
75 }],
76 }
77 }
78}
79
80impl ModuleOptimizer {
81 pub fn has_gradient_clipping(&self) -> bool {
84 self.optimizers.iter().any(|g| g.grad_clipping.is_some())
85 }
86
87 pub fn grad_clipping(&self) -> Option<&GradientClipping> {
90 self.optimizers
91 .first()
92 .expect("Should have at least one optimizer")
93 .grad_clipping
94 .as_ref()
95 }
96
97 pub fn with_grad_clipping(mut self, gradient_clipping: GradientClipping) -> Self {
108 self.optimizers
109 .first_mut()
110 .expect("Should have at least one optimizer")
111 .set_gradient_clipping(gradient_clipping);
112 self
113 }
114
115 fn step_common<M: AutodiffModule>(
116 &mut self,
117 lr_policy: ModuleLearningRate,
118 module: M,
119 mut grads: GradAdaptor,
120 ) -> M {
121 module.map(&mut ModuleOptimizerMapper::new(
122 self.optimizers.iter().collect(),
123 &mut self.param_context,
124 &mut grads,
125 lr_policy,
126 ))
127 }
128
129 pub fn with_group<O>(
142 mut self,
143 group: ParamGroup,
144 optim: O,
145 grad_clipping: Option<GradientClipping>,
146 ) -> Self
147 where
148 O: DynOptimizer + 'static,
149 {
150 self.optimizers.push(OptimizerGroup {
151 group: group.clone(),
152 optim: Arc::new(optim),
153 grad_clipping,
154 });
155 self.param_context
156 .retain(|id, param_state| !group.matches(id, param_state.path.as_deref()));
157 self
158 }
159}
160
161impl ModuleOptimizer {
162 pub fn step<M: AutodiffModule>(
164 &mut self,
165 lr_module: impl Into<ModuleLearningRate>,
166 module: M,
167 grads: GradientsParams,
168 ) -> M {
169 self.step_common(lr_module.into(), module, grads.into())
170 }
171
172 pub fn step_multi<M: AutodiffModule>(
174 &mut self,
175 lr_module: impl Into<ModuleLearningRate>,
176 module: M,
177 grads: MultiGradientsParams,
178 ) -> M {
179 self.step_common(lr_module.into(), module, grads.into())
180 }
181
182 fn optim_from_param(
183 &self,
184 id: ParamId,
185 path: Option<&str>,
186 ) -> (&'_ Arc<dyn DynOptimizer>, Option<GradientClipping>) {
187 self.optimizers
188 .iter()
189 .filter_map(|val| {
190 val.group
191 .matches(&id, path)
192 .then_some((&val.optim, val.grad_clipping.clone()))
193 })
194 .next_back()
195 .expect("Should match at least one parameter group.")
196 }
197
198 pub fn to_record(&self) -> OptimizerRecord {
200 let mut tensors = Vec::new();
201 let mut scalars = BTreeMap::new();
202 let mut paths = BTreeMap::new();
203
204 for (id, param_state) in self.param_context.iter() {
205 let prefix = id.val().to_string();
206 let mut sink = StateSink::default();
207 param_state
208 .optim
209 .state_flatten(&prefix, ¶m_state.state, &mut sink);
210
211 scalars.insert(
214 join_path(&prefix, RANK_KEY),
215 burn_pack::Scalar::from(param_state.state.rank()),
216 );
217 if let Some(path) = ¶m_state.path {
219 paths.insert(prefix, path.clone());
220 }
221
222 for (name, data) in sink.tensors {
223 tensors.push(burn_pack::Tensor::new(
224 name,
225 data.dtype,
226 data.shape,
227 Some(id.val()),
228 data.bytes,
229 ));
230 }
231 for (name, value) in sink.scalars {
232 scalars.insert(name, value);
233 }
234 }
235
236 OptimizerRecord {
237 tensors,
238 scalars,
239 paths,
240 }
241 }
242
243 pub fn load_record(mut self, record: OptimizerRecord) -> Self {
250 let device = Device::default();
251 let mut ranks: BTreeMap<u64, usize> = BTreeMap::new();
252 let mut paths: BTreeMap<u64, String> = BTreeMap::new();
253
254 let suffix = alloc::format!(".{RANK_KEY}");
257 for (name, value) in record.scalars.iter() {
258 if let Some(id_str) = name.strip_suffix(&suffix)
259 && let (Ok(id), Ok(rank)) = (id_str.parse::<u64>(), usize::try_from(*value))
260 {
261 ranks.insert(id, rank);
262 }
263 }
264
265 for (name, path) in record.paths.iter() {
266 if let Ok(id) = name.parse::<u64>() {
267 paths.insert(id, path.to_string());
268 }
269 }
270
271 let mut source = StateSource::new(record.scalars);
272
273 for tensor in record.tensors {
274 let id = tensor
275 .param_id
276 .expect("Optimizer record tensors should carry a parameter id.");
277 let (name, dtype, shape, _, bytes) =
278 tensor.into_parts().expect("record tensors are resident");
279 let data = TensorData::from_bytes(bytes, shape, dtype);
280 ranks.entry(id).or_insert(data.shape.len());
282 source.insert_tensor(name, data);
283 }
284
285 let mut states = HashMap::new();
286 for (id, rank) in ranks {
287 let prefix = id.to_string();
288 let path = paths.get(&id);
289 let (optim, grad_clipping) =
290 self.optim_from_param(id.into(), path.map(|path| path.as_str()));
291 if let Some(state) = optim.state_unflatten(rank, &prefix, &mut source, &device) {
294 states.insert(
295 ParamId::from(id),
296 OptimizationContext {
297 optim: optim.clone(),
298 path: path.cloned(),
299 state,
300 grad_clipping,
301 },
302 );
303 }
304 }
305
306 self.param_context = states;
307 self
308 }
309
310 pub fn into_bytes(&self) -> Result<Bytes, RecordError> {
312 self.to_record().into_bytes()
313 }
314
315 pub fn from_bytes(self, bytes: Bytes) -> Result<Self, RecordError> {
317 Ok(self.load_record(OptimizerRecord::from_bytes(bytes)?))
318 }
319
320 #[cfg(feature = "std")]
322 pub fn save<P: AsRef<std::path::Path>>(&self, path: P) -> Result<(), RecordError> {
323 self.to_record().save(path)
324 }
325
326 #[cfg(feature = "std")]
328 pub fn load<P: AsRef<std::path::Path>>(self, path: P) -> Result<Self, RecordError> {
329 Ok(self.load_record(OptimizerRecord::load(path)?))
330 }
331}
332
333pub enum GradAdaptor {
335 Single(GradientsParams),
337
338 Multi(MultiGradientsParams),
340}
341
342impl From<GradientsParams> for GradAdaptor {
343 fn from(grads: GradientsParams) -> Self {
344 Self::Single(grads)
345 }
346}
347
348impl From<MultiGradientsParams> for GradAdaptor {
349 fn from(grads: MultiGradientsParams) -> Self {
350 Self::Multi(grads)
351 }
352}
353
354impl GradAdaptor {
355 pub fn remove<const D: usize>(&mut self, id: ParamId) -> Option<(Tensor<D>, Device)> {
360 match self {
361 GradAdaptor::Single(grads) => grads.remove(id).map(|t| {
362 let device = t.device();
363 (t, device)
364 }),
365 GradAdaptor::Multi(grads) => grads.remove(id),
366 }
367 }
368}
369
370struct ModuleOptimizerMapper<'a> {
371 path: Vec<String>,
372 optimizer_groups: Vec<&'a OptimizerGroup>,
373 states: &'a mut HashMap<ParamId, OptimizationContext>,
374 grads: &'a mut GradAdaptor,
375 lr_module: ModuleLearningRate,
376}
377
378impl<'a> ModuleOptimizerMapper<'a> {
379 pub(crate) fn new(
380 optimizer_groups: Vec<&'a OptimizerGroup>,
381 states: &'a mut HashMap<ParamId, OptimizationContext>,
382 grads: &'a mut GradAdaptor,
383 lr_module: ModuleLearningRate,
384 ) -> Self {
385 Self {
386 path: vec![],
387 optimizer_groups,
388 states,
389 grads,
390 lr_module,
391 }
392 }
393
394 fn optimizer_from_param(
395 &self,
396 id: ParamId,
397 path: Option<&str>,
398 ) -> (Arc<dyn DynOptimizer>, Option<GradientClipping>) {
399 self.optimizer_groups
400 .iter()
401 .filter_map(|val| {
402 val.group
403 .matches(&id, path)
404 .then_some((val.optim.clone(), val.grad_clipping.clone()))
405 })
406 .next_back()
407 .expect("Should match at least one parameter group.")
408 }
409}
410
411impl ModuleMapper for ModuleOptimizerMapper<'_> {
412 fn enter_module(&mut self, name: &str, _container_type: &str) {
413 self.path.push(name.to_string());
414 }
415
416 fn exit_module(&mut self, _name: &str, _container_type: &str) {
417 self.path.pop();
418 }
419
420 fn map_float<const D: usize>(&mut self, param: Param<Tensor<D>>) -> Param<Tensor<D>> {
421 let (id, tensor, mapper) = param.consume();
422 let grad = self.grads.remove(id);
423
424 let tensor = if let Some((grad, device)) = grad {
425 let is_require_grad = tensor.is_require_grad();
426 #[cfg(feature = "std")]
427 let is_distributed = tensor.is_distributed();
428
429 let entry = self.states.remove_entry(&id);
430 let key = entry.as_ref().map(|(k, _)| *k);
431 let tensor = if tensor.device() != device {
432 tensor.to_device(&device)
433 } else {
434 tensor
435 };
436
437 let path = self.path.join(".");
438 let (optim, grad_clipping, existing_dyn_state) = match entry.map(|(_, s)| s) {
439 Some(OptimizationContext {
440 optim,
441 grad_clipping,
442 state,
443 ..
444 }) => (optim, grad_clipping, Some(state)),
445 None => {
446 let (optim, grad_clipping) =
447 self.optimizer_from_param(id, Some(path.as_str())).clone();
448 (optim, grad_clipping, None)
449 }
450 };
451
452 debug_assert_eq!(
453 grad.device(),
454 device,
455 "The gradient is on the provided device"
456 );
457 let clipped_grad: Tensor<D> = if let Some(g_clipping) = grad_clipping.as_ref() {
458 g_clipping.clip_gradient(grad)
459 } else {
460 grad
461 };
462
463 debug_assert_eq!(
464 tensor.device(),
465 device,
466 "Tensor and gradients are on the same device."
467 );
468
469 let lr = self.lr_module.lr_from_param(id, Some(path.as_str()));
470 let (tensor, state) = optim.step_dyn(
471 D,
472 lr,
473 tensor.inner().into_bridge(),
474 clipped_grad.into_bridge(),
475 existing_dyn_state.map(|s| optim.to_device_dyn(s, &device)),
476 );
477
478 if let Some(state) = state {
479 self.states.insert(
480 key.unwrap_or(id),
481 OptimizationContext {
482 optim,
483 path: Some(path),
484 state,
485 grad_clipping,
486 },
487 );
488 }
489
490 let mut tensor = Tensor::from_inner(Tensor::from_bridge(tensor));
491
492 if is_require_grad {
493 tensor = tensor.require_grad();
494 }
495 #[cfg(feature = "std")]
496 if is_distributed {
497 tensor = tensor.set_distributed(id)
498 }
499
500 tensor
501 } else {
502 tensor
503 };
504
505 Param::from_mapped_value(id, tensor, mapper)
506 }
507}
508
509#[cfg(test)]
510mod tests {
511 use super::*;
512 use crate::{
513 AdamConfig, GradientsParams, SgdConfig,
514 lr_scheduler::module_lr_scheduler::ModuleLearningRate,
515 };
516 use burn::module::ParamGroup;
517 use burn::tensor::{Distribution, Tensor, Tolerance};
518 use burn_derive::Module;
519 use burn_nn::{Linear, LinearConfig};
520
521 #[derive(Module, Debug)]
522 struct TwoLayerModel {
523 layer_a: Linear,
524 layer_b: Linear,
525 }
526
527 fn make_model(device: &Device) -> TwoLayerModel {
528 TwoLayerModel {
529 layer_a: LinearConfig::new(4, 4).init(device),
530 layer_b: LinearConfig::new(4, 4).init(device),
531 }
532 }
533
534 fn make_grads(model: &TwoLayerModel, x: Tensor<2>) -> GradientsParams {
535 let out = model.layer_a.forward(x.clone()) + model.layer_b.forward(x);
536 GradientsParams::from_grads(out.mean().backward(), model)
537 }
538
539 fn lr() -> ModuleLearningRate {
540 ModuleLearningRate::from(0.01_f64)
541 }
542
543 fn sgd() -> ModuleOptimizer {
544 SgdConfig::new().init()
545 }
546
547 #[test]
550 fn default_optimizer_state_survives_round_trip() {
551 let device = Device::default().autodiff();
552 let mut model = make_model(&device);
553 let mut optim: ModuleOptimizer = AdamConfig::new().init();
554
555 for _ in 0..3 {
556 let x = Tensor::<2>::random([2, 4], Distribution::Default, &device);
557 model = optim.step(lr(), model.clone(), make_grads(&model, x));
558 }
559
560 let record = optim.to_record();
561 let mut reloaded: ModuleOptimizer = AdamConfig::new().init().load_record(record);
562
563 let x = Tensor::<2>::random([2, 4], Distribution::Default, &device);
564 let grads_a = make_grads(&model, x.clone());
565 let grads_b = make_grads(&model, x);
566 let from_orig = optim.step(lr(), model.clone(), grads_a);
567 let from_reload = reloaded.step(lr(), model, grads_b);
568
569 from_orig
570 .layer_a
571 .weight
572 .val()
573 .into_data()
574 .assert_approx_eq::<f32>(
575 &from_reload.layer_a.weight.val().into_data(),
576 Tolerance::absolute(1e-6),
577 );
578 }
579
580 #[test]
584 fn group_optimizer_routes_correctly_after_record_round_trip() {
585 let device = Device::default().autodiff();
586 let mut model = make_model(&device);
587
588 let make_optim = || {
589 sgd().with_group(
590 ParamGroup::from_predicate("layer_a"),
591 AdamConfig::new().build(),
592 None,
593 )
594 };
595 let mut optim = make_optim();
596
597 for _ in 0..3 {
598 let x = Tensor::<2>::random([2, 4], Distribution::Default, &device);
599 model = optim.step(lr(), model.clone(), make_grads(&model, x));
600 }
601
602 let record = optim.to_record();
603 let mut reloaded = make_optim().load_record(record);
604
605 let x = Tensor::<2>::random([2, 4], Distribution::Default, &device);
606 let grads_a = make_grads(&model, x.clone());
607 let grads_b = make_grads(&model, x);
608 let from_orig = optim.step(lr(), model.clone(), grads_a);
609 let from_reload = reloaded.step(lr(), model, grads_b);
610
611 from_orig
612 .layer_a
613 .weight
614 .val()
615 .into_data()
616 .assert_approx_eq::<f32>(
617 &from_reload.layer_a.weight.val().into_data(),
618 Tolerance::absolute(1e-6),
619 );
620 from_orig
621 .layer_b
622 .weight
623 .val()
624 .into_data()
625 .assert_approx_eq::<f32>(
626 &from_reload.layer_b.weight.val().into_data(),
627 Tolerance::absolute(1e-6),
628 );
629 }
630
631 #[test]
634 fn with_group_clears_state_for_matching_params() {
635 let device = Device::default().autodiff();
636 let mut model = make_model(&device);
637 let mut optim: ModuleOptimizer = AdamConfig::new().init();
638
639 for _ in 0..2 {
640 let x = Tensor::<2>::random([2, 4], Distribution::Default, &device);
641 model = optim.step(lr(), model.clone(), make_grads(&model, x));
642 }
643
644 optim = optim.with_group(
646 ParamGroup::from_predicate("layer_a"),
647 SgdConfig::new().build(),
648 None,
649 );
650
651 let x = Tensor::<2>::random([2, 4], Distribution::Default, &device);
652 _ = optim.step(lr(), model.clone(), make_grads(&model, x));
653
654 let record = optim.to_record();
655
656 let time_key_count = record.scalars.keys().filter(|k| k.contains("time")).count();
657 assert_eq!(
658 time_key_count, 2,
659 "only layer_b params should carry Adam's time scalar after group switch"
660 );
661 }
662
663 #[test]
664 fn group_grad_clipping_applies_only_to_matching_params() {
665 let device = Device::default().autodiff();
666 let mut model = make_model(&device);
667
668 let lr_value = 0.01;
669 let threshold = 0.01_f32;
670 let bound = (lr_value * threshold as f64) as f32;
671
672 let mut optim = sgd().with_group(
673 ParamGroup::from_predicate("layer_a"),
674 SgdConfig::new().build(),
675 Some(GradientClipping::Value(threshold)),
676 );
677
678 let x = Tensor::<2>::random([2, 4], Distribution::Uniform(100.0, 200.0), &device);
680 let weight_a_before = model.layer_a.weight.val();
681 let weight_b_before = model.layer_b.weight.val();
682
683 model = optim.step(
684 ModuleLearningRate::from(lr_value),
685 model.clone(),
686 make_grads(&model, x),
687 );
688
689 let diff_a = weight_a_before - model.layer_a.weight.val();
690 let diff_b = weight_b_before - model.layer_b.weight.val();
691
692 let eps = 1e-6_f32;
693 for value in diff_a.into_data().iter::<f32>() {
694 assert!(
695 value.abs() <= bound + eps,
696 "layer_a's group grad_clipping should keep every update within lr * threshold"
697 );
698 }
699 assert!(
700 diff_b
701 .into_data()
702 .iter::<f32>()
703 .any(|value| value.abs() > bound + eps),
704 "layer_b uses the default optimizer and should not be clipped"
705 );
706 }
707
708 #[test]
711 fn record_without_paths_loads_without_panic() {
712 let device = Device::default().autodiff();
713 let mut model = make_model(&device);
714 let mut optim: ModuleOptimizer = AdamConfig::new().init();
715
716 let x = Tensor::<2>::random([2, 4], Distribution::Default, &device);
717 model = optim.step(lr(), model.clone(), make_grads(&model, x));
718
719 let mut record = optim.to_record();
721 record.paths.clear();
722
723 let mut optim_loaded: ModuleOptimizer = AdamConfig::new().init().load_record(record);
724
725 let x = Tensor::<2>::random([2, 4], Distribution::Default, &device);
726 let grads_a = make_grads(&model, x.clone());
727 let grads_b = make_grads(&model, x);
728 let from_orig = optim.step(lr(), model.clone(), grads_a);
729 let from_reload = optim_loaded.step(lr(), model, grads_b);
730
731 from_orig
732 .layer_a
733 .weight
734 .val()
735 .into_data()
736 .assert_approx_eq::<f32>(
737 &from_reload.layer_a.weight.val().into_data(),
738 Tolerance::absolute(1e-6),
739 );
740 from_orig
741 .layer_b
742 .weight
743 .val()
744 .into_data()
745 .assert_approx_eq::<f32>(
746 &from_reload.layer_b.weight.val().into_data(),
747 Tolerance::absolute(1e-6),
748 );
749 }
750}