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 pending(&self) -> &GradientsParams { &self.grads }
35
36 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 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 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 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 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 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}