Skip to main content

ruda_optim/optim/
fp32_master.rs

1use super::{SimpleOptimizer, adaptor::OptimizerAdaptor};
2use crate::LearningRate;
3use crate::grad_clipping::GradientClipping;
4use ruda_model::{
5    module::AutodiffModule,
6    record::{PrecisionSettings, Record},
7    tensor::{
8        DType, Tensor,
9        backend::{AutodiffBackend, Backend},
10    },
11};
12
13/// Opt-in FP32 master parameters for an existing simple optimizer.
14///
15/// The wrapped optimizer receives FP32 parameters and gradients; its algorithm,
16/// hyperparameters and per-parameter state are unchanged. Returned parameters
17/// retain their incoming storage dtype. Existing optimizers do not opt in.
18#[derive(Clone)]
19pub struct Fp32MasterOptimizer<O> {
20    optimizer: O,
21    grad_clipping: Option<GradientClipping>,
22    gradient_scale: f32,
23}
24
25/// Master parameter and the original optimizer's state, recorded together.
26/// Use full-precision recorder settings to preserve the master for continuation.
27#[derive(Clone)]
28pub struct Fp32MasterState<B: Backend, const D: usize, S: Record<B> + Clone> {
29    /// Authoritative FP32 parameter after the most recent update.
30    pub master: Tensor<B, D>,
31    /// State returned by the wrapped optimizer.
32    pub inner: Option<S>,
33}
34
35impl<O> Fp32MasterOptimizer<O> {
36    /// Wrap an existing optimizer without modifying its configuration.
37    pub fn new(optimizer: O) -> Self {
38        Self {
39            optimizer,
40            grad_clipping: None,
41            gradient_scale: 1.,
42        }
43    }
44
45    /// Explicitly apply RUDA's existing per-parameter clipping after FP32 conversion.
46    pub fn with_grad_clipping(mut self, clipping: GradientClipping) -> Self {
47        self.grad_clipping = Some(clipping);
48        self
49    }
50
51    /// Explicitly divide gradients by this loss scale in FP32 before clipping.
52    ///
53    /// The caller scales the loss and uses one scale for the accumulation window.
54    /// No dynamic scaling, non-finite policy or optimizer-step skipping is enabled.
55    pub fn with_gradient_scale(mut self, scale: f32) -> Self {
56        assert!(
57            scale.is_finite() && scale > 0. && scale.recip().is_finite(),
58            "gradient scale must be finite, positive and have a finite reciprocal"
59        );
60        self.gradient_scale = scale;
61        self
62    }
63
64    /// Build RUDA's original module adaptor, preserving parameter IDs and records.
65    /// Clipping is disabled unless explicitly configured on the wrapper or adaptor.
66    pub fn init<B, M>(self) -> OptimizerAdaptor<Self, M, B>
67    where
68        B: AutodiffBackend,
69        M: AutodiffModule<B>,
70        O: SimpleOptimizer<B::InnerBackend>,
71    {
72        OptimizerAdaptor::from(self)
73    }
74}
75
76impl<B: Backend, O: SimpleOptimizer<B>> SimpleOptimizer<B> for Fp32MasterOptimizer<O> {
77    type State<const D: usize> = Fp32MasterState<B, D, O::State<D>>;
78
79    fn step<const D: usize>(
80        &self,
81        lr: LearningRate,
82        tensor: Tensor<B, D>,
83        grad: Tensor<B, D>,
84        state: Option<Self::State<D>>,
85    ) -> (Tensor<B, D>, Option<Self::State<D>>) {
86        let storage_dtype = tensor.dtype();
87        assert!(
88            matches!(storage_dtype, DType::F32 | DType::F16 | DType::BF16),
89            "FP32 master optimizer requires FP32/FP16/BF16 parameter storage"
90        );
91        assert_eq!(
92            tensor.dims(),
93            grad.dims(),
94            "master parameter/gradient shape mismatch"
95        );
96        let (master, inner) = match state {
97            Some(state) => {
98                assert_eq!(state.master.dtype(), DType::F32, "master must remain FP32");
99                assert_eq!(
100                    state.master.dims(),
101                    tensor.dims(),
102                    "master state shape mismatch"
103                );
104                (state.master, state.inner)
105            }
106            None => (tensor.cast(DType::F32), None),
107        };
108        let grad = grad.cast(DType::F32);
109        let grad = if self.gradient_scale == 1. {
110            grad
111        } else {
112            grad / self.gradient_scale
113        };
114        let grad = match &self.grad_clipping {
115            Some(clipping) => clipping.clip_gradient(grad),
116            None => grad,
117        };
118        let (master, inner) = self.optimizer.step(lr, master, grad, inner);
119        assert_eq!(
120            master.dtype(),
121            DType::F32,
122            "wrapped optimizer must return FP32 parameters"
123        );
124        let tensor = master.clone().cast(storage_dtype);
125        (tensor, Some(Fp32MasterState { master, inner }))
126    }
127
128    fn to_device<const D: usize>(state: Self::State<D>, device: &B::Device) -> Self::State<D> {
129        Fp32MasterState {
130            master: state.master.to_device(device),
131            inner: state.inner.map(|inner| O::to_device(inner, device)),
132        }
133    }
134}
135
136impl<B: Backend, const D: usize, S: Record<B> + Clone> Record<B> for Fp32MasterState<B, D, S> {
137    type Item<P: PrecisionSettings> = <(Tensor<B, D>, Option<S>) as Record<B>>::Item<P>;
138
139    fn into_item<P: PrecisionSettings>(self) -> Self::Item<P> {
140        (self.master, self.inner).into_item::<P>()
141    }
142
143    fn from_item<P: PrecisionSettings>(item: Self::Item<P>, device: &B::Device) -> Self {
144        let (master, inner): (Tensor<B, D>, Option<S>) = Record::from_item::<P>(item, device);
145        Self {
146            master: master.cast(DType::F32),
147            inner,
148        }
149    }
150}
151
152#[cfg(test)]
153mod tests {
154    use super::*;
155    use crate::{
156        AdamWConfig, GradientsParams, Optimizer, SgdConfig, TestAutodiffBackend, TestBackend,
157    };
158    use ruda_model::{
159        module::{Initializer, Module},
160        record::{BinBytesRecorder, FullPrecisionSettings, Recorder},
161    };
162    use ruda_nn::LinearConfig;
163
164    #[test]
165    fn master_retains_updates_smaller_than_half_storage_resolution() {
166        let device = Default::default();
167        for dtype in [DType::F16, DType::BF16] {
168            let original = SgdConfig::new().build::<TestBackend>();
169            let optimizer = Fp32MasterOptimizer::new(original.clone());
170            let mut parameter = Tensor::<TestBackend, 1>::from_floats([1.], &device).cast(dtype);
171            let gradient = Tensor::<TestBackend, 1>::ones([1], &device).cast(dtype);
172            let mut expected = parameter.clone().cast(DType::F32);
173            let mut state = None;
174            let mut expected_state = None;
175            for _ in 0..64 {
176                (expected, expected_state) = original.step(
177                    0.0001,
178                    expected,
179                    gradient.clone().cast(DType::F32),
180                    expected_state,
181                );
182                (parameter, state) = optimizer.step(0.0001, parameter, gradient.clone(), state);
183                assert_eq!(parameter.dtype(), dtype);
184                let master = &state.as_ref().unwrap().master;
185                assert_eq!(master.dtype(), DType::F32);
186                master.to_data().assert_eq(&expected.to_data(), false);
187                parameter
188                    .to_data()
189                    .assert_eq(&expected.clone().cast(dtype).to_data(), false);
190            }
191            assert!(parameter.cast(DType::F32).into_scalar() < 1.);
192        }
193    }
194
195    #[test]
196    fn adamw_master_and_original_moments_continue_after_record_roundtrip() {
197        let device = Default::default();
198        for dtype in [DType::F16, DType::BF16] {
199            let original = AdamWConfig::new()
200                .with_amsgrad(true)
201                .with_cautious_weight_decay(true)
202                .build();
203            let optimizer = Fp32MasterOptimizer::new(original.clone());
204            let mut parameter =
205                Tensor::<TestBackend, 1>::from_floats([1., -2.], &device).cast(dtype);
206            let gradient = Tensor::<TestBackend, 1>::from_floats([0.25, -0.5], &device).cast(dtype);
207            let mut expected = parameter.clone().cast(DType::F32);
208            let mut state = None;
209            let mut expected_state = None;
210            for _ in 0..3 {
211                (expected, expected_state) = original.step(
212                    0.001,
213                    expected,
214                    gradient.clone().cast(DType::F32),
215                    expected_state,
216                );
217                (parameter, state) = optimizer.step(0.001, parameter, gradient.clone(), state);
218                let saved = state.take().unwrap();
219                saved.master.to_data().assert_eq(&expected.to_data(), false);
220                let inner = saved.inner.as_ref().unwrap();
221                let reference = expected_state.as_ref().unwrap();
222                assert_eq!(inner.momentum.time, reference.momentum.time);
223                for (actual, expected) in [
224                    (&inner.momentum.moment_1, &reference.momentum.moment_1),
225                    (&inner.momentum.moment_2, &reference.momentum.moment_2),
226                    (
227                        inner.momentum.max_moment_2.as_ref().unwrap(),
228                        reference.momentum.max_moment_2.as_ref().unwrap(),
229                    ),
230                ] {
231                    assert_eq!(actual.dtype(), DType::F32);
232                    actual.to_data().assert_eq(&expected.to_data(), false);
233                }
234                let recorder = BinBytesRecorder::<FullPrecisionSettings>::default();
235                let bytes =
236                    <BinBytesRecorder<FullPrecisionSettings> as Recorder<TestBackend>>::record(
237                        &recorder,
238                        saved,
239                        (),
240                    )
241                    .unwrap();
242                state = Some(
243                    <BinBytesRecorder<FullPrecisionSettings> as Recorder<TestBackend>>::load(
244                        &recorder, bytes, &device,
245                    )
246                    .unwrap(),
247                );
248                parameter
249                    .to_data()
250                    .assert_eq(&expected.clone().cast(dtype).to_data(), false);
251            }
252        }
253    }
254
255    #[test]
256    fn explicit_clipping_uses_original_algorithm_after_fp32_conversion() {
257        let device = Default::default();
258        let original = SgdConfig::new().build::<TestBackend>();
259        let gradient =
260            Tensor::<TestBackend, 1>::from_floats([1000., -2000.], &device).cast(DType::F16);
261        for clipping in [GradientClipping::Value(1.), GradientClipping::Norm(1.)] {
262            let parameter = Tensor::<TestBackend, 1>::ones([2], &device).cast(DType::F16);
263            let fp32_gradient = clipping.clip_gradient(gradient.clone().cast(DType::F32));
264            let (expected, _) = original.step(
265                0.01,
266                parameter.clone().cast(DType::F32),
267                fp32_gradient,
268                None,
269            );
270            for scale in [1., 16.] {
271                let optimizer = Fp32MasterOptimizer::new(original.clone())
272                    .with_gradient_scale(scale)
273                    .with_grad_clipping(clipping.clone());
274                let (actual, state) =
275                    optimizer.step(0.01, parameter.clone(), gradient.clone() * scale, None);
276                state
277                    .unwrap()
278                    .master
279                    .to_data()
280                    .assert_eq(&expected.to_data(), false);
281                actual
282                    .to_data()
283                    .assert_eq(&expected.clone().cast(DType::F16).to_data(), false);
284            }
285        }
286    }
287
288    #[test]
289    fn muon_master_reuses_original_fp32_matrix_update_and_momentum() {
290        let device = Default::default();
291        for dtype in [DType::F16, DType::BF16] {
292            let original = crate::MuonConfig::new()
293                .with_stable_normalization(true)
294                .build::<TestBackend>();
295            let optimizer = Fp32MasterOptimizer::new(original.clone());
296            let mut parameter = Tensor::<TestBackend, 2>::ones([2, 3], &device).cast(dtype);
297            let gradient =
298                Tensor::<TestBackend, 2>::from_floats([[1., 2., 3.], [-2., 1., 4.]], &device)
299                    .cast(dtype);
300            let mut expected = parameter.clone().cast(DType::F32);
301            let mut state = None;
302            let mut expected_state = None;
303            for _ in 0..3 {
304                (expected, expected_state) = original.step(
305                    0.01,
306                    expected,
307                    gradient.clone().cast(DType::F32),
308                    expected_state,
309                );
310                (parameter, state) = optimizer.step(0.01, parameter, gradient.clone(), state);
311                let saved = state.as_ref().unwrap();
312                saved.master.to_data().assert_eq(&expected.to_data(), false);
313                let velocity = saved.inner.as_ref().unwrap().momentum.velocity();
314                assert_eq!(velocity.dtype(), DType::F32);
315                velocity.to_data().assert_eq(
316                    &expected_state
317                        .as_ref()
318                        .unwrap()
319                        .momentum
320                        .velocity()
321                        .to_data(),
322                    false,
323                );
324                parameter
325                    .to_data()
326                    .assert_eq(&expected.clone().cast(dtype).to_data(), false);
327            }
328        }
329    }
330
331    #[test]
332    fn original_module_adaptor_preserves_param_id_and_recorded_master() {
333        let device = Default::default();
334        let mut model = LinearConfig::new(2, 1)
335            .with_bias(false)
336            .with_initializer(Initializer::Ones)
337            .init::<TestAutodiffBackend>(&device)
338            .to_dtype(ruda_model::tensor::FloatDType::BF16);
339        let id = model.weight.id;
340        let mut optimizer = Fp32MasterOptimizer::new(AdamWConfig::new().build()).init();
341        for _ in 0..3 {
342            let input = Tensor::<TestAutodiffBackend, 2>::from_floats([[1., 2.]], &device)
343                .cast(DType::BF16);
344            let gradients = model.forward(input).sum().backward();
345            let gradients = GradientsParams::from_grads(gradients, &model);
346            model = optimizer.step(0.001, model, gradients);
347            assert_eq!(model.weight.id, id);
348            assert_eq!(model.weight.val().dtype(), DType::BF16);
349            let record = optimizer.to_record();
350            assert_eq!(record.len(), 1);
351            let state: Fp32MasterState<TestBackend, 2, crate::AdamWState<TestBackend, 2>> =
352                record[&id].clone().into_state();
353            assert_eq!(state.master.dtype(), DType::F32);
354            assert_eq!(state.master.dims(), [2, 1]);
355        }
356    }
357}