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
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}