ruda_optim/optim/
grad_accum.rs1
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
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);
73 module.visit(&mut visitor);
74 }
75
76 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}