Skip to main content

burn_optim/optim/
base.rs

1use burn_core::Tensor;
2
3use burn_core::module::ParamId;
4use burn_core::tensor::Device;
5
6use super::GradientsParams;
7use alloc::vec::Vec;
8
9#[derive(Default)]
10/// Exposes multiple gradients for each parameter.
11pub struct MultiGradientsParams {
12    /// Each [GradientsParams] has its associated [Device].
13    pub grads: Vec<(GradientsParams, Device)>,
14}
15
16impl MultiGradientsParams {
17    /// Removes the gradients for the given [parameter id](ParamId).
18    ///
19    /// Potentially accumulates the gradients from multiple sources using a device associated with
20    /// a parameter id. The same parameter will be accumulated using the same device during
21    /// all training.
22    pub fn remove<const D: usize>(&mut self, id: ParamId) -> Option<(Tensor<D>, Device)> {
23        let (mut tensor, device, index) = self.select(id)?;
24
25        for (i, (grads, _)) in self.grads.iter_mut().enumerate() {
26            if i == index {
27                continue;
28            }
29
30            if let Some(grad) = grads.remove::<D>(id) {
31                tensor = tensor + grad.to_device(&device);
32            }
33        }
34
35        Some((tensor, device))
36    }
37
38    fn select<const D: usize>(&mut self, id: ParamId) -> Option<(Tensor<D>, Device, usize)> {
39        let id_val = id.val() as usize;
40        for i in 0..self.grads.len() {
41            let selected_device_index = (id_val + i) % self.grads.len();
42
43            if let Some(acc) = self.grads[selected_device_index].0.remove::<D>(id) {
44                let device = &self.grads[selected_device_index].1;
45                return Some((acc.to_device(device), device.clone(), selected_device_index));
46            }
47        }
48
49        None
50    }
51}