Skip to main content

burn_optim/optim/module/
module_optimizer.rs

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
20/// Scalar key (per parameter) under which the parameter's state rank is persisted.
21///
22/// Reserved: a custom [`Optimizer::State`](crate::Optimizer::State) must not have a top-level
23/// scalar field named `__rank`, as it would collide with this key in the record.
24const 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/// Keep a reference to the optimizer to avoid matching every step.
40#[derive(Clone)]
41struct OptimizationContext {
42    optim: Arc<dyn DynOptimizer>,
43    grad_clipping: Option<GradientClipping>,
44    path: Option<String>,
45    state: DynState,
46}
47
48/// Optimizes a whole module by applying a per-parameter [`Optimizer`] to each of its parameters.
49///
50/// It is non-generic over the module and optimizer: any `O: Optimizer` is type-erased behind a
51/// dynamic optimizer, and per-parameter states are kept as type-erased states keyed by
52/// [`ParamId`](burn::module::ParamId). Build one with `optimizer.into()` or
53/// `OptimizerConfig::init()`.
54///
55/// It is possible to use different optimizers for different parameters. To do so, use the
56/// [ModuleOptimizer::with_group] function to add an optimizer for all parameters matching the
57/// provided group.
58#[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    /// Check if the optimizer has gradient clipping.
82    /// If there are multiple optimizers, checks if any group has gradient clipping.
83    pub fn has_gradient_clipping(&self) -> bool {
84        self.optimizers.iter().any(|g| g.grad_clipping.is_some())
85    }
86
87    /// Access the gradient clipping.
88    /// If there are multiple optimizers, returns the first optimizer's [GradientClipping].
89    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    /// Sets the gradient clipping.
98    /// If there are multiple optimizers, assigns it to the first one.
99    ///
100    /// # Arguments
101    ///
102    /// * `gradient_clipping` - The gradient clipping.
103    ///
104    /// # Returns
105    ///
106    /// The optimizer.
107    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    /// Adds an optimizer specific to a parameter group.
130    ///
131    /// Parameters matching this group will be optimized using the provided optimizer
132    /// and gradient clipping configuration.
133    ///
134    /// ### Matching Rules
135    /// * **Precedence:** If a parameter matches multiple groups, the *last* group added takes precedence.
136    /// * **Fallback:** The first optimizer added must match all parameters to act as a global fallback.
137    ///
138    /// ### Side Effects
139    /// * **State Reset:** Adding a new group will reset any existing optimizer states for parameters
140    ///   that match the new group.
141    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    /// Update the `module` parameters with the given `gradients`, advancing the optimizer state.
163    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    /// Like [`step`](Self::step), but accumulating gradients sourced from multiple devices.
173    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    /// Decompose the optimizer state into a serializable [`OptimizerRecord`].
199    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, &param_state.state, &mut sink);
210
211            // Persist the parameter rank explicitly so the state can be reconstructed even when it
212            // carries no tensors, and without inferring the rank from tensor shapes.
213            scalars.insert(
214                join_path(&prefix, RANK_KEY),
215                burn_pack::Scalar::from(param_state.state.rank()),
216            );
217            // Save parameter path to be able to match to the right group when loading.
218            if let Some(path) = &param_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    /// Load the optimizer state from an [`OptimizerRecord`].
244    ///
245    /// State tensors are materialized on the default device; no device argument is needed because
246    /// each parameter's state is migrated to that parameter's (gradient's) device on the next
247    /// [`step`](ModuleOptimizer::step) — see the `to_device` call in the step path. The load device
248    /// is therefore irrelevant to correctness.
249    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        // Recover each parameter's rank from its persisted `__rank` scalar (authoritative). Keys
255        // are `"{param_id}.__rank"`, so strip the dotted suffix to recover the id.
256        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            // Fall back to inferring rank from a tensor shape if no `__rank` scalar was present.
281            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            // Skip parameters whose state can't be reconstructed (truncated/foreign record); they
292            // are re-initialized lazily on the next step rather than aborting the load.
293            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    /// Serialize the optimizer state to an in-memory burnpack byte buffer.
311    pub fn into_bytes(&self) -> Result<Bytes, RecordError> {
312        self.to_record().into_bytes()
313    }
314
315    /// Load the optimizer state from an in-memory burnpack byte buffer.
316    pub fn from_bytes(self, bytes: Bytes) -> Result<Self, RecordError> {
317        Ok(self.load_record(OptimizerRecord::from_bytes(bytes)?))
318    }
319
320    /// Save the optimizer state to a burnpack file on disk.
321    #[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    /// Load the optimizer state from a burnpack file on disk.
327    #[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
333/// Wrapper to unify the `remove` method for [GradientsParams] and [MultiGradientsParams].
334pub enum GradAdaptor {
335    /// Wrapper for [`GradientsParams`].
336    Single(GradientsParams),
337
338    /// Wrapper for [`MultiGradientsParams`].
339    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    /// Remove a gradient parameter by ID.
356    ///
357    /// # Returns
358    /// Maybe the (tensor, device) pair.
359    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    /// to_record / load_record must fully preserve a stateful optimizer's internal state so that
548    /// a step on the restored optimizer is numerically identical to one taken on the original.
549    #[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    /// The paths saved in OptimizerRecord enable load_record to route each parameter to the
581    /// correct group optimizer. A step on the restored optimizer must be numerically identical
582    /// to one on the original — for both the group (Adam) and the default (SGD) optimizer.
583    #[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    /// Adding a group after training must clear accumulated state for matching params
632    /// while leaving state for non-matching params untouched.
633    #[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        // Switch layer_a to SGD. Its Adam state must be cleared.
645        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        // Large inputs produce gradients that clearly exceed the clipping threshold.
679        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    /// An OptimizerRecord with empty paths map must load cleanly.
709    /// Parameters default to the default optimizer.
710    #[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        // Simulate a record with empty paths.
720        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}