Skip to main content

ruda_optim/optim/
grads.rs

1
2use ruda_model::{
3    Tensor,
4    tensor::{
5        backend::{AutodiffBackend, Backend},
6        container::TensorContainer,
7    },
8};
9#[cfg(feature = "collective")]
10use ruccl::{CollectiveError, PeerId, ReduceOperation, all_reduce};
11
12use ruda_model::module::{AutodiffModule, ParamId};
13
14use super::visitor::{GradientsParamsChangeDevice, GradientsParamsConverter};
15
16mod record;
17pub use record::GradientsParamsRecord;
18
19#[cfg(feature = "collective")]
20mod collective;
21
22/// Data type that contains gradients for parameters.
23#[derive(Default, Debug)]
24pub struct GradientsParams {
25    container: TensorContainer<ParamId>,
26}
27
28impl GradientsParams {
29    /// Creates a new [GradientsParams](GradientsParams).
30    pub fn new() -> Self {
31        Self::default()
32    }
33
34    /// Extract each tensor gradients for the given [module](AutodiffModule).
35    ///
36    /// Note: This consumes the gradients. See ['from_module'] to extract gradients only for
37    ///  a specific module.
38    pub fn from_grads<B: AutodiffBackend, M: AutodiffModule<B>>(
39        grads: B::Gradients,
40        module: &M,
41    ) -> Self {
42        let mut grads = grads;
43        Self::from_module(&mut grads, module)
44    }
45
46    /// Extract each tensor gradients for the given [module](AutodiffModule).
47    pub fn from_module<B: AutodiffBackend, M: AutodiffModule<B>>(
48        grads: &mut B::Gradients,
49        module: &M,
50    ) -> Self {
51        let mut grads_params = GradientsParams::new();
52        let mut visitor = GradientsParamsConverter::<M, B>::new(grads, &mut grads_params, None);
53        module.visit(&mut visitor);
54        grads_params
55    }
56
57    /// Extract tensor gradients for the given [module](AutodiffModule) and given parameters.
58    pub fn from_params<B: AutodiffBackend, M: AutodiffModule<B>>(
59        grads: &mut B::Gradients,
60        module: &M,
61        params: &[ParamId],
62    ) -> Self {
63        let mut grads_params = GradientsParams::new();
64        let mut visitor =
65            GradientsParamsConverter::<M, B>::new(grads, &mut grads_params, Some(params.to_vec()));
66        module.visit(&mut visitor);
67        grads_params
68    }
69
70    /// Get the gradients for the given [parameter id](ParamId).
71    ///
72    /// # Notes
73    ///
74    /// You should use [remove](GradientsParams::remove) if you want to get the gradients
75    /// only one time.
76    pub fn get<B, const D: usize>(&self, id: ParamId) -> Option<Tensor<B, D>>
77    where
78        B: Backend,
79    {
80        self.container.get(&id).map(Tensor::from_primitive)
81    }
82
83    /// Remove the gradients for the given [parameter id](ParamId).
84    pub fn remove<B, const D: usize>(&mut self, id: ParamId) -> Option<Tensor<B, D>>
85    where
86        B: Backend,
87    {
88        self.container.remove(&id).map(Tensor::from_primitive)
89    }
90
91    /// Register a gradients tensor for the given [parameter id](ParamId).
92    ///
93    /// # Notes
94    ///
95    /// If a tensor is already registered for the given [parameter id](ParamId), it will be replaced.
96    pub fn register<B, const D: usize>(&mut self, id: ParamId, value: Tensor<B, D>)
97    where
98        B: Backend,
99    {
100        self.container.register(id, value.into_primitive())
101    }
102
103    /// The number of gradients tensors registered.
104    pub fn len(&self) -> usize {
105        self.container.len()
106    }
107
108    /// If any tensor is contained.
109    pub fn is_empty(&self) -> bool {
110        self.len() == 0
111    }
112
113    /// Change the device of each tensor gradients registered for the given [module](AutodiffModule).
114    pub fn to_device<B: AutodiffBackend, M: AutodiffModule<B>>(
115        mut self,
116        device: &B::Device,
117        module: &M,
118    ) -> Self {
119        let mut visitor = GradientsParamsChangeDevice::<M, B>::new(device, &mut self);
120        module.visit(&mut visitor);
121        self
122    }
123
124    /// Syncs the gradient params with the other peers in the collective.
125    #[cfg(feature = "collective")]
126    pub fn all_reduce<B: Backend>(
127        mut self,
128        peer_id: PeerId,
129        op: ReduceOperation,
130    ) -> Result<Self, CollectiveError> {
131        let mut ids = self
132            .container
133            .ids()
134            .into_iter()
135            .copied()
136            .collect::<Vec<ParamId>>();
137        // This is crucial, since the all-reduce operations need to happen in the same order for the same parameters on all nodes!
138        ids.sort();
139
140        for id in ids {
141            let Some(grad) = self.container.remove::<B>(&id) else {
142                todo!()
143            };
144
145            let grad = match grad {
146                ruda_model::tensor::TensorPrimitive::Float(grad) => {
147                    let grad = all_reduce::<B>(peer_id, grad, op)?;
148                    ruda_model::tensor::TensorPrimitive::Float(grad)
149                }
150                ruda_model::tensor::TensorPrimitive::QFloat(_grad) => {
151                    unimplemented!("quantized all-reduce unimplemented")
152                }
153            };
154
155            self.container.register::<B>(id, grad);
156        }
157
158        Ok(self)
159    }
160}
161
162#[cfg(test)]
163mod tests {
164    use super::*;
165    use crate::TestAutodiffBackend;
166    use ruda_model::module::{Module, list_param_ids};
167    use ruda_model::tensor::{Distribution, backend::Backend};
168    use ruda_nn::{Linear, LinearConfig};
169
170    #[test]
171    fn test_convert_grads() {
172        let device = Default::default();
173        let layer_1 = layer::<TestAutodiffBackend>(&device);
174        let mut layer_2 = layer_1.clone();
175        layer_2 = layer_2.fork(&device);
176        let loss_1 = layer_1.forward(random_tensor(&device));
177        let loss_2 = layer_2.forward(random_tensor(&device));
178        let grads_1 = GradientsParams::from_grads(loss_1.backward(), &layer_1);
179        let grads_2 = GradientsParams::from_grads(loss_2.backward(), &layer_2);
180
181        let param_ids_1 = list_param_ids(&layer_1);
182        let param_ids_2 = list_param_ids(&layer_2);
183
184        assert_eq!(param_ids_1, param_ids_2);
185        assert_eq!(grads_1.len(), param_ids_1.len());
186        assert_eq!(grads_2.len(), param_ids_2.len());
187    }
188
189    fn layer<B: Backend>(device: &B::Device) -> Linear<B> {
190        LinearConfig::new(20, 20).init(device)
191    }
192
193    fn random_tensor<B: Backend>(device: &B::Device) -> Tensor<B, 2> {
194        Tensor::<B, 2>::random([2, 20], Distribution::Default, device)
195    }
196}