ruda_optim/optim/
grads.rs1
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#[derive(Default, Debug)]
24pub struct GradientsParams {
25 container: TensorContainer<ParamId>,
26}
27
28impl GradientsParams {
29 pub fn new() -> Self {
31 Self::default()
32 }
33
34 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 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 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 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 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 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 pub fn len(&self) -> usize {
105 self.container.len()
106 }
107
108 pub fn is_empty(&self) -> bool {
110 self.len() == 0
111 }
112
113 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 #[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 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}