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#[derive(Clone)]
19pub struct Fp32MasterOptimizer<O> {
20 optimizer: O,
21 grad_clipping: Option<GradientClipping>,
22 gradient_scale: f32,
23}
24
25#[derive(Clone)]
28pub struct Fp32MasterState<B: Backend, const D: usize, S: Record<B> + Clone> {
29 pub master: Tensor<B, D>,
31 pub inner: Option<S>,
33}
34
35impl<O> Fp32MasterOptimizer<O> {
36 pub fn new(optimizer: O) -> Self {
38 Self {
39 optimizer,
40 grad_clipping: None,
41 gradient_scale: 1.,
42 }
43 }
44
45 pub fn with_grad_clipping(mut self, clipping: GradientClipping) -> Self {
47 self.grad_clipping = Some(clipping);
48 self
49 }
50
51 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 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}