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 #[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 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 #[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 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}