Skip to main content

ruda_optim/optim/lbfgs/
mod.rs

1#![allow(clippy::excessive_precision)]
2
3
4use super::GradientsParams;
5use crate::LearningRate;
6use ruda_model::config::Config;
7use ruda_model::module::{AutodiffModule, Module, ModuleMapper, ModuleVisitor, Param, ParamId};
8use ruda_model::prelude::ToElement;
9use ruda_model::record::Record;
10use ruda_model::tensor::backend::Backend;
11use ruda_model::tensor::{Tensor, backend::AutodiffBackend, container::TensorContainer};
12use hashbrown::HashSet;
13use serde::{Deserialize, Serialize};
14
15use alloc::vec;
16use alloc::vec::Vec;
17#[cfg(not(feature = "std"))]
18#[allow(unused_imports)]
19use num_traits::Float as _;
20
21mod line_search;
22use line_search::strong_wolfe;
23#[cfg(test)]
24use line_search::cubic_interpolate;
25
26/// Strategy for the line search optimization phase
27#[derive(Clone, Default, Debug, Copy, PartialEq, Eq, Serialize, Deserialize)]
28pub enum LineSearchFn {
29    /// No line search performed
30    #[default]
31    None,
32    /// strong wolfe conditions
33    ///
34    /// See: <https://en.wikipedia.org/wiki/Wolfe_conditions>
35    StrongWolfe,
36}
37
38/// LBFGS Configuration.
39#[derive(Config, Debug)]
40pub struct LBFGSConfig {
41    /// Maximal number of iterations per optimization step (default: 20)
42    #[config(default = 20)]
43    pub max_iter: usize,
44    /// Update history size (default: 100).
45    #[config(default = 100)]
46    pub history_size: usize,
47    /// Termination tolerance on first order optimality (default: 1e-7).
48    #[config(default = 1e-7)]
49    pub tolerance_grad: f64,
50    /// Termination tolerance on function value/parameter changes (default: 1e-9).
51    #[config(default = 1e-9)]
52    pub tolerance_change: f64,
53    /// Maximal number of function evaluations per optimization step (default: max_iter * 1.25).
54    #[config(default = "None")]
55    pub max_eval: Option<usize>,
56    /// Either ‘strong_wolfe’ or None (default: None).
57    #[config(default = "LineSearchFn::None")]
58    pub line_search_fn: LineSearchFn,
59}
60
61impl LBFGSConfig {
62    /// Initialize AdamW optimizer
63    ///
64    /// # Returns
65    ///
66    /// Returns an optimizer that can be used to optimize a module
67    pub fn init<B: AutodiffBackend>(&self) -> LBFGS<B> {
68        // by default max_eval = max_iter * 5/4
69        let max_eval = self.max_eval.unwrap_or(self.max_iter * 5 / 4);
70        LBFGS {
71            config: LBFGSConfig {
72                max_iter: self.max_iter,
73                history_size: self.history_size,
74                tolerance_grad: self.tolerance_grad,
75                tolerance_change: self.tolerance_change,
76                max_eval: Some(max_eval),
77                line_search_fn: self.line_search_fn,
78            },
79            state: Default::default(),
80        }
81    }
82}
83
84/// Collects gradients in module visit order.
85struct FlattenGradsVisitorInner<'a, B: AutodiffBackend> {
86    grads: &'a GradientsParams,
87    tensors: &'a mut Vec<Tensor<B::InnerBackend, 1>>,
88    seen: HashSet<ParamId>,
89}
90
91impl<B: AutodiffBackend> ModuleVisitor<B> for FlattenGradsVisitorInner<'_, B> {
92    fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<B, D>>) {
93        let tensor = param.val();
94        if !tensor.is_require_grad() || !self.seen.insert(param.id) {
95            return;
96        }
97        let grad = self.grads.get::<B::InnerBackend, D>(param.id)
98            .unwrap_or_else(|| tensor.inner().zeros_like());
99        let numel = grad.shape().num_elements();
100        self.tensors.push(grad.reshape([numel]));
101    }
102}
103
104/// Flatten params to inner backend 1D tensor.
105fn flatten_params_inner<B: AutodiffBackend, M: Module<B>>(
106    module: &M,
107) -> Option<Tensor<B::InnerBackend, 1>> {
108    let mut tensors = Vec::new();
109    let mut visitor = FlattenParamsVisitorInner::<B> {
110        tensors: &mut tensors,
111        seen: HashSet::new(),
112    };
113    module.visit(&mut visitor);
114    if tensors.is_empty() {
115        return None;
116    }
117    Some(Tensor::cat(tensors, 0))
118}
119
120struct FlattenParamsVisitorInner<'a, B: AutodiffBackend> {
121    tensors: &'a mut Vec<Tensor<B::InnerBackend, 1>>,
122    seen: HashSet<ParamId>,
123}
124
125impl<B: AutodiffBackend> ModuleVisitor<B> for FlattenParamsVisitorInner<'_, B> {
126    fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<B, D>>) {
127        let tensor = param.val();
128        if !tensor.is_require_grad() || !self.seen.insert(param.id) {
129            return;
130        }
131        let t = tensor.inner();
132        let numel = t.shape().num_elements();
133        self.tensors.push(t.reshape([numel]));
134    }
135}
136
137/// Flatten gradients for a module.
138fn flatten_grads_inner<B: AutodiffBackend, M: Module<B>>(
139    module: &M,
140    grads: &GradientsParams,
141) -> Tensor<B::InnerBackend, 1> {
142    let mut tensors = Vec::new();
143    let mut visitor = FlattenGradsVisitorInner {
144        grads,
145        tensors: &mut tensors,
146        seen: HashSet::new(),
147    };
148    module.visit(&mut visitor);
149    if tensors.is_empty() {
150        return Tensor::empty([0], &module.devices()[0]);
151    }
152    Tensor::cat(tensors, 0)
153}
154
155/// Mapper that assigns each float param from a flat inner-backend 1D tensor.
156struct ParamsFromFlatMapperInner<'a, B: AutodiffBackend> {
157    flat: &'a Tensor<B::InnerBackend, 1>,
158    offset: &'a mut usize,
159    updated: TensorContainer<ParamId>,
160}
161
162impl<B: AutodiffBackend> ParamsFromFlatMapperInner<'_, B> {
163    fn take_slice(&mut self, numel: usize) -> Tensor<B::InnerBackend, 1> {
164        let start = *self.offset;
165        *self.offset += numel;
166        self.flat.clone().slice(start..*self.offset)
167    }
168}
169
170impl<B: AutodiffBackend> ModuleMapper<B> for ParamsFromFlatMapperInner<'_, B> {
171    fn map_float<const D: usize>(&mut self, param: Param<Tensor<B, D>>) -> Param<Tensor<B, D>> {
172        let (id, tensor, mapper) = param.consume();
173        if !tensor.is_require_grad() {
174            return Param::from_mapped_value(id, tensor, mapper);
175        }
176        if let Some(updated) = self.updated.get::<B>(&id) {
177            return Param::from_mapped_value(id, Tensor::from_primitive(updated), mapper);
178        }
179        let numel = tensor.shape().num_elements();
180        let slice_1d = self.take_slice(numel);
181        let new_inner = slice_1d.reshape(tensor.shape());
182        let new_tensor = Tensor::from_inner(new_inner).require_grad();
183        self.updated.register::<B>(id, new_tensor.clone().into_primitive());
184        Param::from_mapped_value(id, new_tensor, mapper)
185    }
186}
187
188/// Overwrite module parameters from a flat inner-backend 1D tensor
189fn set_params_from_flat_inner<B: AutodiffBackend, M: Module<B>>(
190    module: M,
191    flat: Tensor<B::InnerBackend, 1>,
192) -> M {
193    let mut offset = 0;
194    let mut mapper = ParamsFromFlatMapperInner {
195        flat: &flat,
196        offset: &mut offset,
197        updated: TensorContainer::new(),
198    };
199    module.map(&mut mapper)
200}
201
202/// L-BFGS optimizer state
203#[derive(Clone, Record)]
204pub struct LBFGSState<B: Backend> {
205    /// Historical displacement vectors
206    pub history_s: Vec<Tensor<B, 1>>,
207    /// Historical gradient difference vectors
208    pub history_y: Vec<Tensor<B, 1>>,
209    /// Search direction
210    pub d: Option<Tensor<B, 1>>,
211    /// Step size from the previous iteration
212    pub t: Option<f64>,
213    /// Flattened gradient from the previous iteration
214    pub prev_flat_grad: Option<Tensor<B, 1>>,
215    /// Loss value from the previous iteration
216    pub prev_loss: Option<f64>,
217    /// Global iteration count
218    pub g_iter: usize,
219}
220
221impl<B: Backend> LBFGSState<B> {
222    /// Moves all historical tensors to the target device.
223    pub fn to_device(self, device: &B::Device) -> Self {
224        Self {
225            history_s: self
226                .history_s
227                .into_iter()
228                .map(|t| t.to_device(device))
229                .collect(),
230            history_y: self
231                .history_y
232                .into_iter()
233                .map(|t| t.to_device(device))
234                .collect(),
235            d: self.d.map(|t| t.to_device(device)),
236            t: self.t,
237            prev_flat_grad: self.prev_flat_grad.map(|t| t.to_device(device)),
238            prev_loss: self.prev_loss,
239            g_iter: self.g_iter,
240        }
241    }
242}
243impl<B: Backend> Default for LBFGSState<B> {
244    fn default() -> Self {
245        Self {
246            history_s: Vec::new(),
247            history_y: Vec::new(),
248            d: None,
249            t: Some(1.0),
250            prev_flat_grad: None,
251            prev_loss: None,
252            g_iter: 0,
253        }
254    }
255}
256
257/// L-BFGS optimizer.
258///
259/// Ported from [pytorch](https://github.com/pytorch/pytorch/torch/optim/lbfgs.py). Heavily inspired by [miniFunc](https://www.cs.ubc.ca/~schmidtm/Software/minFunc.html)
260///
261/// See also:
262/// - [L-BFGS](https://en.wikipedia.org/wiki/Limited-memory_BFGS)
263///
264/// # Note
265/// This optimizer is memory intensive
266#[derive(Clone)]
267pub struct LBFGS<B: Backend + AutodiffBackend> {
268    config: LBFGSConfig,
269    state: LBFGSState<B::InnerBackend>,
270}
271
272impl<B: Backend + AutodiffBackend> LBFGS<B> {
273    /// Export the optimizer state for checkpointing.
274    pub fn to_record(&self) -> LBFGSState<B::InnerBackend> {
275        self.state.clone()
276    }
277
278    /// Restore optimizer state for the same configuration and ordered parameter layout.
279    pub fn load_record(mut self, record: LBFGSState<B::InnerBackend>) -> Self {
280        self.state = record;
281        self
282    }
283
284    /// A single optimization step for any tensor that represents the parameters of a model.
285    pub fn step<M, F>(&mut self, lr: LearningRate, mut module: M, mut closure: F) -> (M, f64)
286    where
287        M: AutodiffModule<B> + Clone,
288        F: FnMut(M) -> (f64, GradientsParams),
289    {
290        // evaluate initial f(x) and df/dx
291        let (mut loss, grads) = closure(module.clone());
292        let mut current_evals = 1;
293        if self.config.max_iter == 0 || current_evals >= self.config.max_eval.unwrap() {
294            return (module, loss);
295        }
296
297        let Some(mut x_flat) = flatten_params_inner::<B, M>(&module) else {
298            return (module, loss);
299        };
300        if x_flat.shape().num_elements() == 0 {
301            return (module, loss);
302        }
303        let mut flat_grad = flatten_grads_inner::<B, M>(&module, &grads);
304
305        let opt_cond =
306            flat_grad.clone().abs().max().into_scalar().to_f64() <= self.config.tolerance_grad;
307        // optimal condition
308        if opt_cond {
309            return (module, loss);
310        }
311
312        // tensors cached in state
313        let mut d = self
314            .state
315            .d
316            .take()
317            .unwrap_or_else(|| flat_grad.clone().neg());
318        let mut t = self.state.t.unwrap_or(lr);
319        let mut prev_flat_grad = self.state.prev_flat_grad.take();
320
321        let mut n_iter = 0;
322
323        // optimize for a max of max_iter iterations
324        while n_iter < self.config.max_iter {
325            // keep track of nb of iterations
326            n_iter += 1;
327            self.state.g_iter += 1;
328
329            // compute gradient descent direction
330            if self.state.g_iter == 1 {
331                d = flat_grad.clone().neg();
332                self.state.history_s.clear();
333                self.state.history_y.clear();
334            } else {
335                // do lbfgs update (update memory)
336                if let Some(pg) = prev_flat_grad.as_ref() {
337                    let y = flat_grad.clone().sub(pg.clone());
338                    let s = d.clone().mul_scalar(t);
339
340                    let ys = y.clone().dot(s.clone()).into_scalar().to_f64();
341
342                    if ys > 1e-10 && self.config.history_size > 0 {
343                        // updating memory
344                        if self.state.history_s.len() >= self.config.history_size {
345                            // shift history by one (limited-memory)
346                            self.state.history_s.remove(0);
347                            self.state.history_y.remove(0);
348                        }
349                        self.state.history_s.push(s);
350                        self.state.history_y.push(y);
351                    }
352                }
353
354                // compute the approximate (L-BFGS) inverse Hessian
355                // multiplied by the gradient
356                let num_old = self.state.history_s.len();
357                let mut q = flat_grad.clone().neg();
358                let mut alphas: Vec<Tensor<B::InnerBackend, 1>> =
359                    vec![Tensor::zeros([1], &flat_grad.device()); num_old];
360
361                if num_old > 0 {
362                    // multiply by initial Hessian
363                    // r/d is the final direction
364                    for i in (0..num_old).rev() {
365                        let s = &self.state.history_s[i];
366                        let y = &self.state.history_y[i];
367                        let rho = y.clone().dot(s.clone()).powf_scalar(-1.0);
368                        let alpha = rho.clone().mul(s.clone().dot(q.clone()));
369                        alphas[i] = alpha.clone();
370                        q = q.sub(y.clone().mul(alpha));
371                    }
372
373                    let last_s = &self.state.history_s[num_old - 1];
374                    let last_y = &self.state.history_y[num_old - 1];
375                    let ys = last_y.clone().dot(last_s.clone());
376                    let yy = last_y.clone().dot(last_y.clone());
377                    let h_diag = ys.div(yy);
378
379                    let mut r = q.mul(h_diag);
380
381                    for ((s, y), alpha) in self
382                        .state
383                        .history_s
384                        .iter()
385                        .zip(self.state.history_y.iter())
386                        .zip(alphas)
387                        .take(num_old)
388                    {
389                        let rho = y.clone().dot(s.clone()).powf_scalar(-1.0);
390
391                        let beta = rho.mul(y.clone().dot(r.clone()));
392
393                        r = r.add(s.clone().mul(alpha.sub(beta)));
394                    }
395                    d = r;
396                } else {
397                    d = q;
398                }
399            }
400
401            prev_flat_grad = Some(flat_grad.clone());
402            let prev_loss_iter = loss;
403
404            // compute step len
405            if self.state.g_iter == 1 {
406                let grad_l1 = flat_grad.clone().abs().sum().into_scalar().to_f64();
407                t = (1.0f64 / grad_l1).min(1.0) * lr;
408            } else {
409                t = lr;
410            }
411
412            // directional derivative
413            let gtd = flat_grad.clone().dot(d.clone()).into_scalar().to_f64();
414
415            if gtd > -self.config.tolerance_change {
416                break;
417            }
418
419            let ls_func_evals;
420
421            if let LineSearchFn::StrongWolfe = self.config.line_search_fn {
422                // perform line search, using user function
423                let mut obj_func =
424                    |current_x: &Tensor<B::InnerBackend, 1>,
425                     step: f64,
426                     dir: &Tensor<B::InnerBackend, 1>| {
427                        let update = dir.clone().mul_scalar(step);
428                        let new_x = current_x.clone().add(update);
429                        let tmp_module = set_params_from_flat_inner::<B, M>(module.clone(), new_x);
430                        let (l, g) = closure(tmp_module);
431                        (l, flatten_grads_inner::<B, M>(&module, &g))
432                    };
433
434                let (ls_f, ls_g, ls_t, evals) = strong_wolfe(
435                    &mut obj_func,
436                    &x_flat,
437                    t,
438                    &d,
439                    loss,
440                    flat_grad.clone(),
441                    gtd,
442                    1e-4,
443                    0.9,
444                    self.config.tolerance_change,
445                    self.config.max_eval.unwrap() - current_evals,
446                );
447
448                loss = ls_f;
449                flat_grad = ls_g;
450                t = ls_t;
451                ls_func_evals = evals;
452
453                x_flat = x_flat.add(d.clone().mul_scalar(t));
454                module = set_params_from_flat_inner::<B, M>(module, x_flat.clone());
455            } else {
456                // no line search, simply move with fixed-step
457                let step_vec = d.clone().mul_scalar(t);
458                x_flat = x_flat.add(step_vec);
459                module = set_params_from_flat_inner::<B, M>(module, x_flat.clone());
460                // re-evaluate function only if not in last iteration
461                // the reason we do this: in a stochastic setting,
462                // no use to re-evaluate that function here
463                let (new_loss, new_grads) = closure(module.clone());
464                loss = new_loss;
465                flat_grad = flatten_grads_inner::<B, M>(&module, &new_grads);
466                ls_func_evals = 1;
467            }
468
469            // update func eval
470            current_evals += ls_func_evals;
471
472            // check conditions
473
474            if current_evals >= self.config.max_eval.unwrap() {
475                break;
476            }
477
478            if flat_grad.clone().abs().max().into_scalar().to_f64() <= self.config.tolerance_grad {
479                break;
480            }
481
482            if d.clone().mul_scalar(t).abs().max().into_scalar().to_f64()
483                <= self.config.tolerance_change
484            {
485                break;
486            }
487
488            if (loss - prev_loss_iter).abs() < self.config.tolerance_change {
489                break;
490            }
491        }
492        self.state.d = Some(d);
493        self.state.t = Some(t);
494        self.state.prev_flat_grad = prev_flat_grad;
495        self.state.prev_loss = Some(loss);
496        (module, loss)
497    }
498    /// Moves the optimizer state to the specified device.
499    pub fn to_device(self, device: &B::Device) -> Self {
500        Self {
501            config: self.config,
502            // History tensors reside in InnerBackend, so we convert the device accordingly
503            state: self.state.to_device(device),
504        }
505    }
506}
507
508#[cfg(test)]
509mod tests;