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    #[cfg(feature = "collective")]
114    pub(crate) fn primitive<B: Backend>(
115        &self,
116        id: ParamId,
117    ) -> Option<ruda_model::tensor::TensorPrimitive<B>> {
118        self.container.get(&id)
119    }
120
121    /// Change the device of each tensor gradients registered for the given [module](AutodiffModule).
122    pub fn to_device<B: AutodiffBackend, M: AutodiffModule<B>>(
123        mut self,
124        device: &B::Device,
125        module: &M,
126    ) -> Self {
127        let mut visitor = GradientsParamsChangeDevice::<M, B>::new(device, &mut self);
128        module.visit(&mut visitor);
129        self
130    }
131
132    /// Syncs the gradient params with the other peers in the collective.
133    #[cfg(feature = "collective")]
134    pub fn all_reduce<B: Backend>(
135        mut self,
136        peer_id: PeerId,
137        op: ReduceOperation,
138    ) -> Result<Self, CollectiveError> {
139        let mut ids = self
140            .container
141            .ids()
142            .into_iter()
143            .copied()
144            .collect::<Vec<ParamId>>();
145        // This is crucial, since the all-reduce operations need to happen in the same order for the same parameters on all nodes!
146        ids.sort();
147
148        for id in ids {
149            let Some(grad) = self.container.remove::<B>(&id) else {
150                todo!()
151            };
152
153            let grad = match grad {
154                ruda_model::tensor::TensorPrimitive::Float(grad) => {
155                    let grad = all_reduce::<B>(peer_id, grad, op)?;
156                    ruda_model::tensor::TensorPrimitive::Float(grad)
157                }
158                ruda_model::tensor::TensorPrimitive::QFloat(_grad) => {
159                    unimplemented!("quantized all-reduce unimplemented")
160                }
161            };
162
163            self.container.register::<B>(id, grad);
164        }
165
166        Ok(self)
167    }
168}
169
170#[cfg(test)]
171mod tests {
172    use super::*;
173    use crate::TestAutodiffBackend;
174    use ruda_model::module::{Module, list_param_ids};
175    use ruda_model::tensor::{Distribution, backend::Backend};
176    use ruda_nn::{Linear, LinearConfig};
177
178    #[test]
179    fn test_convert_grads() {
180        let device = Default::default();
181        let layer_1 = layer::<TestAutodiffBackend>(&device);
182        let mut layer_2 = layer_1.clone();
183        layer_2 = layer_2.fork(&device);
184        let loss_1 = layer_1.forward(random_tensor(&device));
185        let loss_2 = layer_2.forward(random_tensor(&device));
186        let grads_1 = GradientsParams::from_grads(loss_1.backward(), &layer_1);
187        let grads_2 = GradientsParams::from_grads(loss_2.backward(), &layer_2);
188
189        let param_ids_1 = list_param_ids(&layer_1);
190        let param_ids_2 = list_param_ids(&layer_2);
191
192        assert_eq!(param_ids_1, param_ids_2);
193        assert_eq!(grads_1.len(), param_ids_1.len());
194        assert_eq!(grads_2.len(), param_ids_2.len());
195    }
196
197    fn layer<B: Backend>(device: &B::Device) -> Linear<B> {
198        LinearConfig::new(20, 20).init(device)
199    }
200
201    fn random_tensor<B: Backend>(device: &B::Device) -> Tensor<B, 2> {
202        Tensor::<B, 2>::random([2, 20], Distribution::Default, device)
203    }
204}