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)]
10pub struct MultiGradientsParams {
12 pub grads: Vec<(GradientsParams, Device)>,
14}
15
16impl MultiGradientsParams {
17 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}