Skip to main content

ruda_optim/optim/simple/
adaptor.rs

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/// Wrapper struct that adapts any [simple optimizer](SimpleOptimizer) into
18/// an [optimizer](Optimizer).
19#[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    /// Access the wrapped [`SimpleOptimizer`].
55    pub fn optim(&self) -> &O {
56        &self.optim
57    }
58
59    /// Check if the optimizer has gradient clipping.
60    pub fn has_gradient_clipping(&self) -> bool {
61        self.grad_clipping.is_some()
62    }
63
64    /// Access the gradient clipping.
65    pub fn grad_clipping(&self) -> Option<&GradientClipping> {
66        self.grad_clipping.as_ref()
67    }
68
69    /// Sets the gradient clipping.
70    ///
71    /// # Arguments
72    ///
73    /// * `gradient_clipping` - The gradient clipping.
74    ///
75    /// # Returns
76    ///
77    /// The optimizer.
78    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
120/// Wrapper to unify the `remove` method for [GradientsParams] and [MultiGradientsParams].
121pub enum GradAdaptor {
122    /// Wrapper for [`GradientsParams`].
123    Single(GradientsParams),
124
125    /// Wrapper for [`MultiGradientsParams`].
126    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    /// Remove a gradient parameter by ID.
143    ///
144    /// # Returns
145    /// Maybe the (tensor, device) pair.
146    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}