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::{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);
73        module.visit(&mut visitor);
74    }
75
76    /// Return the accumulated gradients and reset the accumulator state.
77    pub fn grads(&mut self) -> GradientsParams {
78        let mut grads = GradientsParams::new();
79        core::mem::swap(&mut self.grads, &mut grads);
80
81        grads
82    }
83}
84
85#[derive(new)]
86struct ModuleGradsAccumulator<'a, M> {
87    grads: &'a mut GradientsParams,
88    grads_new: GradientsParams,
89    phantom: PhantomData<M>,
90}
91
92impl<B: AutodiffBackend, M: AutodiffModule<B>> ModuleVisitor<B> for ModuleGradsAccumulator<'_, M> {
93    fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<B, D>>) {
94        let grad_updated = match self.grads_new.remove::<B::InnerBackend, D>(param.id) {
95            Some(new) => match self.grads.remove::<B::InnerBackend, D>(param.id) {
96                Some(grad) => grad.add(new),
97                None => new,
98            },
99            None => match self.grads.remove::<B::InnerBackend, D>(param.id) {
100                Some(grad) => grad,
101                None => return,
102            },
103        };
104
105        self.grads
106            .register::<B::InnerBackend, D>(param.id, grad_updated);
107    }
108}
109
110#[cfg(test)]
111mod tests {
112    use super::*;
113    use crate::TestAutodiffBackend;
114    use ruda_model::tensor::{Distribution, backend::Backend};
115    use ruda_nn::{Linear, LinearConfig};
116
117    #[test]
118    fn test_accumulate_gradients_one_step() {
119        let device = Default::default();
120        let mut accumulator = GradientsAccumulator::new();
121        let layer = layer::<TestAutodiffBackend>(&device);
122        let loss = layer.forward(random_tensor::<TestAutodiffBackend>(&device));
123        let grads = GradientsParams::from_grads(loss.backward(), &layer);
124
125        accumulator.accumulate(&layer, grads);
126
127        let grads = accumulator.grads();
128        assert!(!grads.is_empty())
129    }
130
131    #[test]
132    fn test_accumulate_gradients_two_steps() {
133        let device = Default::default();
134        let mut accumulator = GradientsAccumulator::new();
135        let layer = layer::<TestAutodiffBackend>(&device);
136        let loss_1 = layer.forward(random_tensor(&device));
137        let loss_2 = layer.forward(random_tensor(&device));
138        let grads_1 = GradientsParams::from_grads(loss_1.backward(), &layer);
139        let grads_2 = GradientsParams::from_grads(loss_2.backward(), &layer);
140
141        accumulator.accumulate(&layer, grads_1);
142        accumulator.accumulate(&layer, grads_2);
143
144        let grads = accumulator.grads();
145        assert_eq!(grads.len(), 2)
146    }
147
148    fn layer<B: Backend>(device: &B::Device) -> Linear<B> {
149        LinearConfig::new(20, 20).init(device)
150    }
151
152    fn random_tensor<B: Backend>(device: &B::Device) -> Tensor<B, 2> {
153        Tensor::<B, 2>::random([2, 20], Distribution::Default, device)
154    }
155}