Skip to main content

ruda_optim/optim/
grad_accum.rs

1
2use core::marker::PhantomData;
3
4use ruda_model::module::{AutodiffModule, ModuleVisitor, Param};
5use ruda_model::record::RecorderError;
6use ruda_model::tensor::{FloatDType, Tensor, backend::AutodiffBackend};
7
8use super::{GradientsParams, GradientsParamsRecord};
9
10/// Accumulate gradients into a single [Gradients](AutodiffBackend::Gradients) object.
11pub struct GradientsAccumulator<M> {
12    grads: GradientsParams,
13    phantom: PhantomData<M>,
14}
15
16impl<M> Default for GradientsAccumulator<M> {
17    fn default() -> Self {
18        Self::new()
19    }
20}
21
22impl<M> GradientsAccumulator<M> {
23    /// Create a new gradients accumulator.
24    pub fn new() -> Self {
25        Self {
26            grads: GradientsParams::new(),
27            phantom: PhantomData,
28        }
29    }
30}
31
32impl<M> GradientsAccumulator<M> {
33    /// Borrow pending gradients without clearing an accumulation window.
34    pub fn pending(&self) -> &GradientsParams { &self.grads }
35
36    /// Snapshot pending gradients without resetting the accumulation window.
37    pub fn try_to_record<B: AutodiffBackend>(&self) -> Result<GradientsParamsRecord, RecorderError>
38    where
39        M: AutodiffModule<B>,
40    {
41        self.grads.try_to_record::<B::InnerBackend>()
42    }
43
44    /// Asynchronously snapshot pending gradients without resetting the accumulator.
45    pub async fn to_record_async<B: AutodiffBackend>(
46        &self,
47    ) -> Result<GradientsParamsRecord, RecorderError>
48    where
49        M: AutodiffModule<B>,
50    {
51        self.grads.to_record_async::<B::InnerBackend>().await
52    }
53
54    /// Replace pending gradients with checkpoint state on the given device.
55    ///
56    /// The model IDs and the caller's accumulation count must be restored from
57    /// the same checkpoint. A rejected record leaves the accumulator unchanged.
58    pub fn load_record<B: AutodiffBackend>(
59        &mut self,
60        record: GradientsParamsRecord,
61        device: &B::Device,
62    ) -> Result<(), RecorderError>
63    where
64        M: AutodiffModule<B>,
65    {
66        self.grads = GradientsParams::from_record::<B::InnerBackend>(record, device)?;
67        Ok(())
68    }
69
70    /// Accumulate the given gradients for each parameter in the given module.
71    pub fn accumulate<B: AutodiffBackend>(&mut self, module: &M, grads: GradientsParams)
72    where
73        M: AutodiffModule<B>,
74    {
75        let mut visitor = ModuleGradsAccumulator::<M>::new(&mut self.grads, grads, None);
76        module.visit(&mut visitor);
77    }
78
79    /// Accumulate after converting incoming and pending gradients to the given dtype.
80    ///
81    /// FP32 accumulation can retain small additions to half-precision gradients.
82    /// Parameter storage, loss normalization and accumulation counts are unchanged.
83    /// Checkpoints preserve the pending gradient dtype; select the same dtype on
84    /// subsequent calls after restoring. No accumulator reset is performed.
85    pub fn accumulate_with_dtype<B: AutodiffBackend>(
86        &mut self,
87        module: &M,
88        grads: GradientsParams,
89        dtype: FloatDType,
90    ) where
91        M: AutodiffModule<B>,
92    {
93        let mut visitor = ModuleGradsAccumulator::<M>::new(&mut self.grads, grads, Some(dtype));
94        module.visit(&mut visitor);
95    }
96
97    /// Return the accumulated gradients and reset the accumulator state.
98    pub fn grads(&mut self) -> GradientsParams {
99        let mut grads = GradientsParams::new();
100        core::mem::swap(&mut self.grads, &mut grads);
101
102        grads
103    }
104}
105
106#[derive(new)]
107struct ModuleGradsAccumulator<'a, M> {
108    grads: &'a mut GradientsParams,
109    grads_new: GradientsParams,
110    dtype: Option<FloatDType>,
111    phantom: PhantomData<M>,
112}
113
114impl<B: AutodiffBackend, M: AutodiffModule<B>> ModuleVisitor<B> for ModuleGradsAccumulator<'_, M> {
115    fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<B, D>>) {
116        let dtype = self.dtype;
117        let convert = |tensor: Tensor<B::InnerBackend, D>| match dtype {
118            Some(dtype) => tensor.cast(dtype),
119            None => tensor,
120        };
121        let grad_updated = match self
122            .grads_new
123            .remove::<B::InnerBackend, D>(param.id)
124            .map(convert)
125        {
126            Some(new) => match self
127                .grads
128                .remove::<B::InnerBackend, D>(param.id)
129                .map(convert)
130            {
131                Some(grad) => grad.add(new),
132                None => new,
133            },
134            None => match self
135                .grads
136                .remove::<B::InnerBackend, D>(param.id)
137                .map(convert)
138            {
139                Some(grad) => grad,
140                None => return,
141            },
142        };
143
144        self.grads
145            .register::<B::InnerBackend, D>(param.id, grad_updated);
146    }
147}
148
149#[cfg(test)]
150mod tests {
151    use super::*;
152    use crate::{TestAutodiffBackend, TestBackend};
153    use ruda_model::module::Module;
154    use ruda_model::tensor::{Distribution, backend::Backend};
155    use ruda_nn::{Linear, LinearConfig};
156
157    #[test]
158    fn test_accumulate_gradients_one_step() {
159        let device = Default::default();
160        let mut accumulator = GradientsAccumulator::new();
161        let layer = layer::<TestAutodiffBackend>(&device);
162        let loss = layer.forward(random_tensor::<TestAutodiffBackend>(&device));
163        let grads = GradientsParams::from_grads(loss.backward(), &layer);
164
165        accumulator.accumulate(&layer, grads);
166
167        let grads = accumulator.grads();
168        assert!(!grads.is_empty())
169    }
170
171    #[test]
172    fn test_accumulate_gradients_two_steps() {
173        let device = Default::default();
174        let mut accumulator = GradientsAccumulator::new();
175        let layer = layer::<TestAutodiffBackend>(&device);
176        let loss_1 = layer.forward(random_tensor(&device));
177        let loss_2 = layer.forward(random_tensor(&device));
178        let grads_1 = GradientsParams::from_grads(loss_1.backward(), &layer);
179        let grads_2 = GradientsParams::from_grads(loss_2.backward(), &layer);
180
181        accumulator.accumulate(&layer, grads_1);
182        accumulator.accumulate(&layer, grads_2);
183
184        let grads = accumulator.grads();
185        assert_eq!(grads.len(), 2)
186    }
187
188    #[test]
189    fn fp32_accumulation_retains_small_half_gradients_after_pending_record_restore() {
190        let device = Default::default();
191        for dtype in [FloatDType::F16, FloatDType::BF16] {
192            let model = LinearConfig::new(1, 1)
193                .with_bias(false)
194                .init::<TestAutodiffBackend>(&device)
195                .to_dtype(dtype);
196            let id = model.weight.id;
197            let gradient = |value| {
198                let mut grads = GradientsParams::new();
199                grads.register(
200                    id,
201                    Tensor::<TestBackend, 2>::full([1, 1], value, &device).cast(dtype),
202                );
203                grads
204            };
205            let mut original = GradientsAccumulator::new();
206            let mut fp32 = GradientsAccumulator::new();
207            original.accumulate(&model, gradient(1024.));
208            fp32.accumulate_with_dtype(&model, gradient(1024.), FloatDType::F32);
209            for step in 0..8 {
210                original.accumulate(&model, gradient(0.125));
211                fp32.accumulate_with_dtype(&model, gradient(0.125), FloatDType::F32);
212                if step == 3 {
213                    let record = fp32.try_to_record::<TestAutodiffBackend>().unwrap();
214                    fp32 = GradientsAccumulator::new();
215                    fp32.load_record::<TestAutodiffBackend>(record, &device)
216                        .unwrap();
217                }
218            }
219            let result = fp32.grads().remove::<TestBackend, 2>(id).unwrap();
220            assert_eq!(result.dtype(), FloatDType::F32.into());
221            assert_eq!(result.into_scalar(), 1025.);
222            let unchanged = original.grads().remove::<TestBackend, 2>(id).unwrap();
223            assert_eq!(unchanged.dtype(), dtype.into());
224            assert_eq!(unchanged.cast(FloatDType::F32).into_scalar(), 1024.);
225            assert!(fp32.grads().is_empty());
226        }
227    }
228
229    fn layer<B: Backend>(device: &B::Device) -> Linear<B> {
230        LinearConfig::new(20, 20).init(device)
231    }
232
233    fn random_tensor<B: Backend>(device: &B::Device) -> Tensor<B, 2> {
234        Tensor::<B, 2>::random([2, 20], Distribution::Default, device)
235    }
236}