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
76#[cfg(feature="collective")]
77impl<B:Backend,O:crate::data_parallel::zero2::ElementwiseShardOptimizer<B>>
78    crate::data_parallel::zero2::ElementwiseShardOptimizer<B> for Fp32MasterOptimizer<O> {
79    fn validate_element_sharding(&self)->Result<(),&'static str> {
80        if self.grad_clipping.is_some() {
81            return Err("tensor-wide clipping must run before ZeRO-2 gradient sharding");
82        }
83        self.optimizer.validate_element_sharding()
84    }
85    fn shard_gradient_dtype(&self,_storage:DType)->DType {DType::F32}
86}
87
88impl<B: Backend, O: SimpleOptimizer<B>> SimpleOptimizer<B> for Fp32MasterOptimizer<O> {
89    type State<const D: usize> = Fp32MasterState<B, D, O::State<D>>;
90
91    fn step<const D: usize>(
92        &self,
93        lr: LearningRate,
94        tensor: Tensor<B, D>,
95        grad: Tensor<B, D>,
96        state: Option<Self::State<D>>,
97    ) -> (Tensor<B, D>, Option<Self::State<D>>) {
98        let storage_dtype = tensor.dtype();
99        assert!(
100            matches!(storage_dtype, DType::F32 | DType::F16 | DType::BF16),
101            "FP32 master optimizer requires FP32/FP16/BF16 parameter storage"
102        );
103        assert_eq!(
104            tensor.dims(),
105            grad.dims(),
106            "master parameter/gradient shape mismatch"
107        );
108        let (master, inner) = match state {
109            Some(state) => {
110                assert_eq!(state.master.dtype(), DType::F32, "master must remain FP32");
111                assert_eq!(
112                    state.master.dims(),
113                    tensor.dims(),
114                    "master state shape mismatch"
115                );
116                (state.master, state.inner)
117            }
118            None => (tensor.cast(DType::F32), None),
119        };
120        let grad = grad.cast(DType::F32);
121        let grad = if self.gradient_scale == 1. {
122            grad
123        } else {
124            grad / self.gradient_scale
125        };
126        let grad = match &self.grad_clipping {
127            Some(clipping) => clipping.clip_gradient(grad),
128            None => grad,
129        };
130        let (master, inner) = self.optimizer.step(lr, master, grad, inner);
131        assert_eq!(
132            master.dtype(),
133            DType::F32,
134            "wrapped optimizer must return FP32 parameters"
135        );
136        let tensor = master.clone().cast(storage_dtype);
137        (tensor, Some(Fp32MasterState { master, inner }))
138    }
139
140    fn to_device<const D: usize>(state: Self::State<D>, device: &B::Device) -> Self::State<D> {
141        Fp32MasterState {
142            master: state.master.to_device(device),
143            inner: state.inner.map(|inner| O::to_device(inner, device)),
144        }
145    }
146}
147
148impl<B: Backend, const D: usize, S: Record<B> + Clone> Record<B> for Fp32MasterState<B, D, S> {
149    type Item<P: PrecisionSettings> = <(Tensor<B, D>, Option<S>) as Record<B>>::Item<P>;
150
151    fn into_item<P: PrecisionSettings>(self) -> Self::Item<P> {
152        (self.master, self.inner).into_item::<P>()
153    }
154
155    fn from_item<P: PrecisionSettings>(item: Self::Item<P>, device: &B::Device) -> Self {
156        let (master, inner): (Tensor<B, D>, Option<S>) = Record::from_item::<P>(item, device);
157        Self {
158            master: master.cast(DType::F32),
159            inner,
160        }
161    }
162}
163
164#[cfg(test)]
165mod tests {
166    use super::*;
167    use crate::{
168        AdamWConfig, GradientsParams, Optimizer, SgdConfig, TestAutodiffBackend, TestBackend,
169    };
170    use ruda_model::{
171        module::{Initializer, Module},
172        record::{BinBytesRecorder, FullPrecisionSettings, Recorder},
173    };
174    use ruda_nn::LinearConfig;
175
176    #[test]
177    fn master_retains_updates_smaller_than_half_storage_resolution() {
178        let device = Default::default();
179        for dtype in [DType::F16, DType::BF16] {
180            let original = SgdConfig::new().build::<TestBackend>();
181            let optimizer = Fp32MasterOptimizer::new(original.clone());
182            let mut parameter = Tensor::<TestBackend, 1>::from_floats([1.], &device).cast(dtype);
183            let gradient = Tensor::<TestBackend, 1>::ones([1], &device).cast(dtype);
184            let mut expected = parameter.clone().cast(DType::F32);
185            let mut state = None;
186            let mut expected_state = None;
187            for _ in 0..64 {
188                (expected, expected_state) = original.step(
189                    0.0001,
190                    expected,
191                    gradient.clone().cast(DType::F32),
192                    expected_state,
193                );
194                (parameter, state) = optimizer.step(0.0001, parameter, gradient.clone(), state);
195                assert_eq!(parameter.dtype(), dtype);
196                let master = &state.as_ref().unwrap().master;
197                assert_eq!(master.dtype(), DType::F32);
198                master.to_data().assert_eq(&expected.to_data(), false);
199                parameter
200                    .to_data()
201                    .assert_eq(&expected.clone().cast(dtype).to_data(), false);
202            }
203            assert!(parameter.cast(DType::F32).into_scalar() < 1.);
204        }
205    }
206
207    #[test]
208    fn adamw_master_and_original_moments_continue_after_record_roundtrip() {
209        let device = Default::default();
210        for dtype in [DType::F16, DType::BF16] {
211            let original = AdamWConfig::new()
212                .with_amsgrad(true)
213                .with_cautious_weight_decay(true)
214                .build();
215            let optimizer = Fp32MasterOptimizer::new(original.clone());
216            let mut parameter =
217                Tensor::<TestBackend, 1>::from_floats([1., -2.], &device).cast(dtype);
218            let gradient = Tensor::<TestBackend, 1>::from_floats([0.25, -0.5], &device).cast(dtype);
219            let mut expected = parameter.clone().cast(DType::F32);
220            let mut state = None;
221            let mut expected_state = None;
222            for _ in 0..3 {
223                (expected, expected_state) = original.step(
224                    0.001,
225                    expected,
226                    gradient.clone().cast(DType::F32),
227                    expected_state,
228                );
229                (parameter, state) = optimizer.step(0.001, parameter, gradient.clone(), state);
230                let saved = state.take().unwrap();
231                saved.master.to_data().assert_eq(&expected.to_data(), false);
232                let inner = saved.inner.as_ref().unwrap();
233                let reference = expected_state.as_ref().unwrap();
234                assert_eq!(inner.momentum.time, reference.momentum.time);
235                for (actual, expected) in [
236                    (&inner.momentum.moment_1, &reference.momentum.moment_1),
237                    (&inner.momentum.moment_2, &reference.momentum.moment_2),
238                    (
239                        inner.momentum.max_moment_2.as_ref().unwrap(),
240                        reference.momentum.max_moment_2.as_ref().unwrap(),
241                    ),
242                ] {
243                    assert_eq!(actual.dtype(), DType::F32);
244                    actual.to_data().assert_eq(&expected.to_data(), false);
245                }
246                let recorder = BinBytesRecorder::<FullPrecisionSettings>::default();
247                let bytes =
248                    <BinBytesRecorder<FullPrecisionSettings> as Recorder<TestBackend>>::record(
249                        &recorder,
250                        saved,
251                        (),
252                    )
253                    .unwrap();
254                state = Some(
255                    <BinBytesRecorder<FullPrecisionSettings> as Recorder<TestBackend>>::load(
256                        &recorder, bytes, &device,
257                    )
258                    .unwrap(),
259                );
260                parameter
261                    .to_data()
262                    .assert_eq(&expected.clone().cast(dtype).to_data(), false);
263            }
264        }
265    }
266
267    #[test]
268    fn explicit_clipping_uses_original_algorithm_after_fp32_conversion() {
269        let device = Default::default();
270        let original = SgdConfig::new().build::<TestBackend>();
271        let gradient =
272            Tensor::<TestBackend, 1>::from_floats([1000., -2000.], &device).cast(DType::F16);
273        for clipping in [GradientClipping::Value(1.), GradientClipping::Norm(1.)] {
274            let parameter = Tensor::<TestBackend, 1>::ones([2], &device).cast(DType::F16);
275            let fp32_gradient = clipping.clip_gradient(gradient.clone().cast(DType::F32));
276            let (expected, _) = original.step(
277                0.01,
278                parameter.clone().cast(DType::F32),
279                fp32_gradient,
280                None,
281            );
282            for scale in [1., 16.] {
283                let optimizer = Fp32MasterOptimizer::new(original.clone())
284                    .with_gradient_scale(scale)
285                    .with_grad_clipping(clipping.clone());
286                let (actual, state) =
287                    optimizer.step(0.01, parameter.clone(), gradient.clone() * scale, None);
288                state
289                    .unwrap()
290                    .master
291                    .to_data()
292                    .assert_eq(&expected.to_data(), false);
293                actual
294                    .to_data()
295                    .assert_eq(&expected.clone().cast(DType::F16).to_data(), false);
296            }
297        }
298    }
299
300    #[test]
301    fn muon_master_reuses_original_fp32_matrix_update_and_momentum() {
302        let device = Default::default();
303        for dtype in [DType::F16, DType::BF16] {
304            let original = crate::MuonConfig::new()
305                .with_stable_normalization(true)
306                .build::<TestBackend>();
307            let optimizer = Fp32MasterOptimizer::new(original.clone());
308            let mut parameter = Tensor::<TestBackend, 2>::ones([2, 3], &device).cast(dtype);
309            let gradient =
310                Tensor::<TestBackend, 2>::from_floats([[1., 2., 3.], [-2., 1., 4.]], &device)
311                    .cast(dtype);
312            let mut expected = parameter.clone().cast(DType::F32);
313            let mut state = None;
314            let mut expected_state = None;
315            for _ in 0..3 {
316                (expected, expected_state) = original.step(
317                    0.01,
318                    expected,
319                    gradient.clone().cast(DType::F32),
320                    expected_state,
321                );
322                (parameter, state) = optimizer.step(0.01, parameter, gradient.clone(), state);
323                let saved = state.as_ref().unwrap();
324                saved.master.to_data().assert_eq(&expected.to_data(), false);
325                let velocity = saved.inner.as_ref().unwrap().momentum.velocity();
326                assert_eq!(velocity.dtype(), DType::F32);
327                velocity.to_data().assert_eq(
328                    &expected_state
329                        .as_ref()
330                        .unwrap()
331                        .momentum
332                        .velocity()
333                        .to_data(),
334                    false,
335                );
336                parameter
337                    .to_data()
338                    .assert_eq(&expected.clone().cast(dtype).to_data(), false);
339            }
340        }
341    }
342
343    #[test]
344    fn original_module_adaptor_preserves_param_id_and_recorded_master() {
345        let device = Default::default();
346        let mut model = LinearConfig::new(2, 1)
347            .with_bias(false)
348            .with_initializer(Initializer::Ones)
349            .init::<TestAutodiffBackend>(&device)
350            .to_dtype(ruda_model::tensor::FloatDType::BF16);
351        let id = model.weight.id;
352        let mut optimizer = Fp32MasterOptimizer::new(AdamWConfig::new().build()).init();
353        for _ in 0..3 {
354            let input = Tensor::<TestAutodiffBackend, 2>::from_floats([[1., 2.]], &device)
355                .cast(DType::BF16);
356            let gradients = model.forward(input).sum().backward();
357            let gradients = GradientsParams::from_grads(gradients, &model);
358            model = optimizer.step(0.001, model, gradients);
359            assert_eq!(model.weight.id, id);
360            assert_eq!(model.weight.val().dtype(), DType::BF16);
361            let record = optimizer.to_record();
362            assert_eq!(record.len(), 1);
363            let state: Fp32MasterState<TestBackend, 2, crate::AdamWState<TestBackend, 2>> =
364                record[&id].clone().into_state();
365            assert_eq!(state.master.dtype(), DType::F32);
366            assert_eq!(state.master.dims(), [2, 1]);
367        }
368    }
369}