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