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
10pub 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 pub fn new() -> Self {
25 Self {
26 grads: GradientsParams::new(),
27 phantom: PhantomData,
28 }
29 }
30}
31
32impl<M> GradientsAccumulator<M> {
33 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 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 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 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 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 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}