1#[cfg(feature = "distributed")]
2use ruda_model::tensor::backend::distributed::DistributedParamId;
3use ruda_model::{prelude::Backend, tensor::Device};
4
5use super::{SimpleOptimizer, record::AdaptorRecord};
6use crate::{
7 LearningRate, MultiGradientsParams,
8 grad_clipping::GradientClipping,
9 optim::{GradientsParams, Optimizer},
10};
11
12use ruda_model::module::{AutodiffModule, ModuleMapper, Param, ParamId};
13use ruda_model::tensor::{Bool, Int, Tensor, backend::AutodiffBackend, container::TensorContainer};
14use core::marker::PhantomData;
15use hashbrown::HashMap;
16
17#[derive(Clone)]
20pub struct OptimizerAdaptor<O, M, B>
21where
22 O: SimpleOptimizer<B::InnerBackend>,
23 M: AutodiffModule<B>,
24 B: AutodiffBackend,
25{
26 optim: O,
27 records: HashMap<ParamId, AdaptorRecord<O, B>>,
28 module: PhantomData<M>,
29 grad_clipping: Option<GradientClipping>,
30}
31
32impl<O, B, M> From<O> for OptimizerAdaptor<O, M, B>
33where
34 B: AutodiffBackend,
35 M: AutodiffModule<B>,
36 O: SimpleOptimizer<B::InnerBackend>,
37{
38 fn from(optim: O) -> Self {
39 Self {
40 optim,
41 records: HashMap::new(),
42 module: PhantomData,
43 grad_clipping: None,
44 }
45 }
46}
47
48impl<O, M, B> OptimizerAdaptor<O, M, B>
49where
50 O: SimpleOptimizer<B::InnerBackend>,
51 M: AutodiffModule<B>,
52 B: AutodiffBackend,
53{
54 pub fn optim(&self) -> &O {
56 &self.optim
57 }
58
59 pub fn has_gradient_clipping(&self) -> bool {
61 self.grad_clipping.is_some()
62 }
63
64 pub fn grad_clipping(&self) -> Option<&GradientClipping> {
66 self.grad_clipping.as_ref()
67 }
68
69 pub fn with_grad_clipping(mut self, gradient_clipping: GradientClipping) -> Self {
79 self.grad_clipping = Some(gradient_clipping);
80 self
81 }
82
83 fn step_common(&mut self, lr: LearningRate, module: M, mut grads: GradAdaptor) -> M {
84 module.map(&mut SimpleOptimizerMapper::<B, O>::new(
85 &self.optim,
86 &mut self.records,
87 &mut grads,
88 lr,
89 self.grad_clipping.as_ref(),
90 ))
91 }
92}
93
94impl<O, B, M> Optimizer<M, B> for OptimizerAdaptor<O, M, B>
95where
96 B: AutodiffBackend,
97 M: AutodiffModule<B>,
98 O: SimpleOptimizer<B::InnerBackend>,
99{
100 type Record = HashMap<ParamId, AdaptorRecord<O, B>>;
101
102 fn step(&mut self, lr: LearningRate, module: M, grads: GradientsParams) -> M {
103 self.step_common(lr, module, grads.into())
104 }
105
106 fn step_multi(&mut self, lr: LearningRate, module: M, grads: MultiGradientsParams) -> M {
107 self.step_common(lr, module, grads.into())
108 }
109
110 fn to_record(&self) -> Self::Record {
111 self.records.clone()
112 }
113
114 fn load_record(mut self, record: Self::Record) -> Self {
115 self.records = record;
116 self
117 }
118}
119
120pub enum GradAdaptor {
122 Single(GradientsParams),
124
125 Multi(MultiGradientsParams),
127}
128
129impl From<GradientsParams> for GradAdaptor {
130 fn from(grads: GradientsParams) -> Self {
131 Self::Single(grads)
132 }
133}
134
135impl From<MultiGradientsParams> for GradAdaptor {
136 fn from(grads: MultiGradientsParams) -> Self {
137 Self::Multi(grads)
138 }
139}
140
141impl GradAdaptor {
142 pub fn remove<B: Backend, const D: usize>(
147 &mut self,
148 id: ParamId,
149 ) -> Option<(Tensor<B, D>, Device<B>)> {
150 match self {
151 GradAdaptor::Single(grads) => grads.remove(id).map(|t| {
152 let device = t.device();
153 (t, device)
154 }),
155 GradAdaptor::Multi(grads) => grads.remove(id),
156 }
157 }
158}
159
160#[derive(new)]
161struct SimpleOptimizerMapper<'a, B, O>
162where
163 B: AutodiffBackend,
164 O: SimpleOptimizer<B::InnerBackend>,
165{
166 optimizer: &'a O,
167 records: &'a mut HashMap<ParamId, AdaptorRecord<O, B>>,
168 grads: &'a mut GradAdaptor,
169 lr: LearningRate,
170 grad_clipping: Option<&'a GradientClipping>,
171 #[new(default)]
172 updated: TensorContainer<ParamId>,
173}
174
175impl<B, O> ModuleMapper<B> for SimpleOptimizerMapper<'_, B, O>
176where
177 B: AutodiffBackend,
178 O: SimpleOptimizer<B::InnerBackend>,
179{
180 fn map_float<const D: usize>(&mut self, param: Param<Tensor<B, D>>) -> Param<Tensor<B, D>> {
181 if !param.is_require_grad() {
182 return param;
183 }
184 let (id, tensor, mapper) = param.consume();
185 if let Some(updated) = self.updated.get::<B>(&id) {
186 return Param::from_mapped_value(id, Tensor::from_primitive(updated), mapper);
187 }
188 let grad = self.grads.remove(id);
189
190 let tensor = if let Some((grad, device)) = grad {
191 let is_require_grad = tensor.is_require_grad();
192 #[cfg(feature = "distributed")]
193 let is_distributed = tensor.is_distributed();
194
195 let (key, record) = self.records.remove_entry(&id).unzip();
196 let tensor = if tensor.device() != device {
197 tensor.to_device(&device)
198 } else {
199 tensor
200 };
201
202 debug_assert_eq!(
203 grad.device(),
204 device,
205 "The gradient is on the provided device"
206 );
207 let clipped_grad = if let Some(g_clipping) = self.grad_clipping {
208 g_clipping.clip_gradient(grad)
209 } else {
210 grad
211 };
212
213 debug_assert_eq!(
214 tensor.device(),
215 device,
216 "Tensor and gradients are on the same device."
217 );
218
219 let (tensor, state) = self.optimizer.step(
220 self.lr,
221 tensor.inner(),
222 clipped_grad,
223 record.map(|record| O::to_device(record.into_state(), &device)),
224 );
225
226 if let Some(state) = state {
227 self.records
228 .insert(key.unwrap_or(id), AdaptorRecord::from_state(state));
229 }
230
231 let mut tensor = Tensor::from_inner(tensor);
232 if is_require_grad {
233 tensor = tensor.require_grad();
234 }
235 #[cfg(feature = "distributed")]
236 if is_distributed {
237 tensor = tensor.set_distributed(DistributedParamId::from(id.val()))
238 }
239
240 self.updated.register::<B>(id, tensor.clone().into_primitive());
241 tensor
242 } else {
243 tensor
244 };
245
246 Param::from_mapped_value(id, tensor, mapper)
247 }
248
249 fn map_int<const D: usize>(&mut self, param: Param<Tensor<B, D, Int>>) -> Param<Tensor<B, D, Int>> {
250 param
251 }
252
253 fn map_bool<const D: usize>(
254 &mut self,
255 param: Param<Tensor<B, D, Bool>>,
256 ) -> Param<Tensor<B, D, Bool>> {
257 param
258 }
259}