Skip to main content

burn_optim/optim/
lbfgs.rs

1#![allow(clippy::excessive_precision)]
2
3use burn_core as burn;
4
5use super::GradientsParams;
6use crate::{LearningRate, OptimizerRecord};
7use crate::{RecordState, StateSink, StateSource};
8use burn::config::Config;
9use burn::module::{AutodiffModule, Module, ModuleMapper, ModuleVisitor, Param};
10use burn::store::RecordError;
11use burn::tensor::{Bytes, Device, Tensor, TensorData};
12use serde::{Deserialize, Serialize};
13
14use alloc::vec;
15use alloc::vec::Vec;
16#[cfg(not(feature = "std"))]
17#[allow(unused_imports)]
18use num_traits::Float as _;
19
20/// Cubic Interpolate
21///
22/// Uses two points (x1, f1), (x2, f2) and their first derivatives g1,g2 to construct
23/// a cubic interpolant and return its minimum within the given bounds.
24fn cubic_interpolate(
25    x1: f64,
26    f1: f64,
27    g1: f64,
28    x2: f64,
29    f2: f64,
30    g2: f64,
31    bounds: Option<(f64, f64)>,
32) -> f64 {
33    // Compute bounds of interpolation area
34    let (min_bound, max_bound) = bounds.unwrap_or(if x1 <= x2 { (x1, x2) } else { (x2, x1) });
35    // Code for most common case: cubic interpolation of 2 points
36    // with function and derivative values for both
37    // Solution in this case (where x2 is the farthest point)
38    // d1 = g1 + g2 - 3*(f1 - f2) / (x1-x2);
39    // d2 = sqrt(d1^2 - g1 * g2);
40    // min_pos = x2 - (x2 - x1)*((g2 + d2 - d1)/(g2 - g1 + 2*d2));
41    // t_new = min(max(min_pos,min_bound), max_bound);
42    let d1 = g1 + g2 - 3.0 * (f1 - f2) / (x1 - x2);
43    let d2_square = d1 * d1 - g1 * g2;
44
45    if d2_square >= 0.0 {
46        let d2 = d2_square.sqrt();
47        let min_pos = if x1 <= x2 {
48            x2 - (x2 - x1) * ((g2 + d2 - d1) / (g2 - g1 + 2.0 * d2))
49        } else {
50            x1 - (x1 - x2) * ((g1 + d2 - d1) / (g1 - g2 + 2.0 * d2))
51        };
52        min_pos.max(min_bound).min(max_bound)
53    } else {
54        (min_bound + max_bound) / 2.0
55    }
56}
57/// Auxiliary Struct For Strong_Wolfe
58struct LineSearchSample {
59    // step size
60    t: f64,
61    // loss
62    f: f64,
63    // gradient
64    g: Tensor<1>,
65    // directional derivative
66    gtd: f64,
67}
68
69#[allow(clippy::too_many_arguments)]
70fn strong_wolfe<F>(
71    // obj_func(x,step size,direction) -> (loss,grad)
72    obj_func: &mut F,
73    x: &Tensor<1>,
74    // initial step size
75    mut t: f64,
76    d: &Tensor<1>,
77    f: f64,
78    g: Tensor<1>,
79    gtd: f64,
80    c1: f64,
81    c2: f64,
82    tolerance_change: f64,
83    max_ls: usize,
84) -> (f64, Tensor<1>, f64, usize)
85where
86    F: FnMut(&Tensor<1>, f64, &Tensor<1>) -> (f64, Tensor<1>),
87{
88    let d_norm: f64 = d.clone().abs().max().into_scalar();
89
90    // evaluate objective and gradient using initial step
91    let (mut f_new, mut g_new) = obj_func(x, t, d);
92    let mut ls_func_evals = 1;
93    let mut gtd_new = g_new.clone().dot(d.clone()).into_scalar();
94
95    // bracket an interval [t_prev,t] containing a point satisfying the Wolfe criteria
96    let (mut t_prev, mut f_prev, mut g_prev, mut gtd_prev) = (0.0, f, g.clone(), gtd);
97    let mut done = false;
98    let mut ls_iter = 0;
99
100    // the interval [low,high] using for Zoom phase
101    let mut bracket: Option<[LineSearchSample; 2]> = None;
102    // point which satisfy the wolfe condition
103    let mut wolfe_bracket: Option<LineSearchSample> = None;
104    while ls_iter < max_ls {
105        // Checking Conditions.
106
107        // Checking the Armijo Condition and function value increasing condition.
108        // Armijo: f(x+t*d) <= f(x) + c_1 t gtd
109        if f_new > (f + c1 * t * gtd) || (ls_iter > 1 && f_new >= f_prev) {
110            bracket = Some([
111                LineSearchSample {
112                    t: t_prev,
113                    f: f_prev,
114                    g: g_prev,
115                    gtd: gtd_prev,
116                },
117                LineSearchSample {
118                    t,
119                    f: f_new,
120                    g: g_new.clone(),
121                    gtd: gtd_new,
122                },
123            ]);
124            break;
125        }
126
127        // Checking Strong Wolfe Condition
128        // |gtd_new| <= -c_2 gtd
129        if gtd_new.abs() <= -c2 * gtd {
130            wolfe_bracket = Some(LineSearchSample {
131                t,
132                f: f_new,
133                g: g_new.clone(),
134                gtd: gtd_new,
135            });
136            done = true;
137            break;
138        }
139
140        // gtd_new >=0 , there must be a local minimum in the interval.
141        if gtd_new >= 0.0 {
142            bracket = Some([
143                LineSearchSample {
144                    t: t_prev,
145                    f: f_prev,
146                    g: g_prev,
147                    gtd: gtd_prev,
148                },
149                LineSearchSample {
150                    t,
151                    f: f_new,
152                    g: g_new.clone(),
153                    gtd: gtd_new,
154                },
155            ]);
156            break;
157        }
158
159        // interpolate
160        let min_step = t + 0.01 * (t - t_prev);
161        let max_step = t * 10.0;
162        let t_next = cubic_interpolate(
163            t_prev,
164            f_prev,
165            gtd_prev,
166            t,
167            f_new,
168            gtd_new,
169            Some((min_step, max_step)),
170        );
171        t_prev = t;
172        f_prev = f_new;
173        g_prev = g_new;
174        gtd_prev = gtd_new;
175
176        // next step
177        t = t_next;
178        (f_new, g_new) = obj_func(x, t, d);
179        ls_func_evals += 1;
180        gtd_new = g_new.clone().dot(d.clone()).into_scalar();
181        ls_iter += 1;
182    }
183    if let Some(sample) = wolfe_bracket {
184        return (sample.f, sample.g, sample.t, ls_func_evals);
185    }
186
187    let mut bracket = bracket.unwrap_or_else(|| {
188        [
189            LineSearchSample {
190                t: 0.0,
191                f,
192                g: g.clone(),
193                gtd,
194            },
195            LineSearchSample {
196                t,
197                f: f_new,
198                g: g_new.clone(),
199                gtd: gtd_new,
200            },
201        ]
202    });
203
204    // zoom phase
205    let mut insuf_progress = false;
206
207    // find high and low points in bracket
208    let (mut low_idx, mut high_idx) = if bracket[0].f <= bracket[1].f {
209        (0, 1)
210    } else {
211        (1, 0)
212    };
213
214    while !done && ls_iter < max_ls {
215        let diff = (bracket[1].t - bracket[0].t).abs();
216        // line-search bracket is so small
217        if diff * d_norm < tolerance_change {
218            break;
219        }
220
221        // compute new trial value
222        t = cubic_interpolate(
223            bracket[0].t,
224            bracket[0].f,
225            bracket[0].gtd,
226            bracket[1].t,
227            bracket[1].f,
228            bracket[1].gtd,
229            None,
230        );
231
232        let b_min = bracket[0].t.min(bracket[1].t);
233        let b_max = bracket[0].t.max(bracket[1].t);
234        let eps = 0.1 * (b_max - b_min);
235
236        if (b_max - t).min(t - b_min) < eps {
237            // interpolation close to boundary
238            if insuf_progress || t >= b_max || t <= b_min {
239                t = if (t - b_max).abs() < (t - b_min).abs() {
240                    b_max - eps
241                } else {
242                    b_min + eps
243                };
244                insuf_progress = false;
245            } else {
246                insuf_progress = true;
247            }
248        } else {
249            insuf_progress = false;
250        }
251
252        // Evaluate new point
253        (f_new, g_new) = obj_func(x, t, d);
254
255        ls_func_evals += 1;
256        gtd_new = g_new.clone().dot(d.clone()).into_scalar();
257        ls_iter += 1;
258
259        let armijo_holds = f_new <= (f + c1 * t * gtd) && f_new < bracket[low_idx].f;
260
261        if !armijo_holds {
262            bracket[high_idx] = LineSearchSample {
263                t,
264                f: f_new,
265                g: g_new,
266                gtd: gtd_new,
267            };
268        } else {
269            if gtd_new.abs() <= -c2 * gtd {
270                return (f_new, g_new, t, ls_func_evals);
271            }
272
273            if gtd_new * (bracket[high_idx].t - bracket[low_idx].t) >= 0.0 {
274                bracket[high_idx] = LineSearchSample {
275                    t: bracket[low_idx].t,
276                    f: bracket[low_idx].f,
277                    g: bracket[low_idx].g.clone(),
278                    gtd: bracket[low_idx].gtd,
279                };
280            }
281            bracket[low_idx] = LineSearchSample {
282                t,
283                f: f_new,
284                g: g_new,
285                gtd: gtd_new,
286            };
287        }
288
289        if bracket[0].f <= bracket[1].f {
290            low_idx = 0;
291            high_idx = 1;
292        } else {
293            low_idx = 1;
294            high_idx = 0;
295        }
296    }
297    // return stuff
298    (
299        bracket[low_idx].f,
300        bracket[low_idx].g.clone(),
301        bracket[low_idx].t,
302        ls_func_evals,
303    )
304}
305
306/// Strategy for the line search optimization phase
307#[derive(Clone, Default, Debug, Copy, PartialEq, Eq, Serialize, Deserialize)]
308pub enum LineSearchFn {
309    /// No line search performed
310    #[default]
311    None,
312    /// strong wolfe conditions
313    ///
314    /// See: <https://en.wikipedia.org/wiki/Wolfe_conditions>
315    StrongWolfe,
316}
317
318/// LBFGS Configuration.
319#[derive(Config, Debug)]
320pub struct LBFGSConfig {
321    /// Maximal number of iterations per optimization step (default: 20)
322    #[config(default = 20)]
323    pub max_iter: usize,
324    /// Update history size (default: 100).
325    #[config(default = 100)]
326    pub history_size: usize,
327    /// Termination tolerance on first order optimality (default: 1e-7).
328    #[config(default = 1e-7)]
329    pub tolerance_grad: f64,
330    /// Termination tolerance on function value/parameter changes (default: 1e-9).
331    #[config(default = 1e-9)]
332    pub tolerance_change: f64,
333    /// Maximal number of function evaluations per optimization step (default: max_iter * 1.25).
334    #[config(default = "None")]
335    pub max_eval: Option<usize>,
336    /// Either ‘strong_wolfe’ or None (default: None).
337    #[config(default = "LineSearchFn::None")]
338    pub line_search_fn: LineSearchFn,
339}
340
341impl LBFGSConfig {
342    /// Initialize LBFGS optimizer.
343    ///
344    /// # Returns
345    ///
346    /// Returns an optimizer that can be used to optimize a module
347    pub fn init(&self) -> LBFGS {
348        // by default max_eval = max_iter * 5/4
349        let max_eval = self.max_eval.unwrap_or(self.max_iter * 5 / 4);
350        LBFGS {
351            config: LBFGSConfig {
352                max_iter: self.max_iter,
353                history_size: self.history_size,
354                tolerance_grad: self.tolerance_grad,
355                tolerance_change: self.tolerance_change,
356                max_eval: Some(max_eval),
357                line_search_fn: self.line_search_fn,
358            },
359            state: Default::default(),
360        }
361    }
362}
363
364/// Collects gradients in module visit order.
365struct FlattenGradsVisitorInner<'a> {
366    grads: &'a GradientsParams,
367    tensors: &'a mut Vec<Tensor<1>>,
368}
369
370impl ModuleVisitor for FlattenGradsVisitorInner<'_> {
371    fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<D>>) {
372        if let Some(g) = self.grads.get::<D>(param.id) {
373            let numel = g.shape().num_elements();
374            self.tensors.push(g.reshape([numel]));
375        }
376    }
377}
378
379/// Flatten params to inner backend 1D tensor.
380fn flatten_params_inner<M: Module>(module: &M) -> Tensor<1> {
381    let mut tensors = Vec::new();
382    let mut visitor = FlattenParamsVisitorInner {
383        tensors: &mut tensors,
384    };
385    module.visit(&mut visitor);
386    if tensors.is_empty() {
387        return Tensor::empty([0], &module.devices()[0].clone().inner());
388    }
389    Tensor::cat(tensors, 0)
390}
391
392struct FlattenParamsVisitorInner<'a> {
393    tensors: &'a mut Vec<Tensor<1>>,
394}
395
396impl ModuleVisitor for FlattenParamsVisitorInner<'_> {
397    fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<D>>) {
398        let t = param.val().inner();
399        let numel = t.shape().num_elements();
400        self.tensors.push(t.reshape([numel]));
401    }
402}
403
404/// Flatten gradients for a module.
405fn flatten_grads_inner<M: Module>(module: &M, grads: &GradientsParams) -> Tensor<1> {
406    let mut tensors = Vec::new();
407    let mut visitor = FlattenGradsVisitorInner {
408        grads,
409        tensors: &mut tensors,
410    };
411    module.visit(&mut visitor);
412    if tensors.is_empty() {
413        return Tensor::empty([0], &module.devices()[0].clone().inner());
414    }
415    Tensor::cat(tensors, 0)
416}
417
418/// Mapper that assigns each float param from a flat inner-backend 1D tensor.
419struct ParamsFromFlatMapperInner<'a> {
420    flat: &'a Tensor<1>,
421    offset: &'a mut usize,
422}
423
424impl ParamsFromFlatMapperInner<'_> {
425    fn take_slice(&mut self, numel: usize) -> Tensor<1> {
426        let start = *self.offset;
427        *self.offset += numel;
428        self.flat.clone().slice(start..*self.offset)
429    }
430}
431
432impl ModuleMapper for ParamsFromFlatMapperInner<'_> {
433    fn map_float<const D: usize>(&mut self, param: Param<Tensor<D>>) -> Param<Tensor<D>> {
434        let (id, tensor, mapper) = param.consume();
435        let numel = tensor.shape().num_elements();
436        let slice_1d = self.take_slice(numel);
437        let new_inner = slice_1d.reshape(tensor.shape());
438        let new_tensor = Tensor::from_inner(new_inner).require_grad();
439        Param::from_mapped_value(id, new_tensor, mapper)
440    }
441}
442
443/// Overwrite module parameters from a flat inner-backend 1D tensor
444fn set_params_from_flat_inner<M: Module>(module: M, flat: Tensor<1>) -> M {
445    let mut offset = 0;
446    let mut mapper = ParamsFromFlatMapperInner {
447        flat: &flat,
448        offset: &mut offset,
449    };
450    module.map(&mut mapper)
451}
452
453/// L-BFGS optimizer state
454#[derive(Clone, RecordState)]
455pub struct LBFGSState {
456    /// Historical displacement vectors
457    pub history_s: Vec<Tensor<1>>,
458    /// Historical gradient difference vectors
459    pub history_y: Vec<Tensor<1>>,
460    /// Search direction
461    pub d: Option<Tensor<1>>,
462    /// Step size from the previous iteration
463    pub t: Option<f64>,
464    /// Flattened gradient from the previous iteration
465    pub prev_flat_grad: Option<Tensor<1>>,
466    /// Loss value from the previous iteration
467    pub prev_loss: Option<f64>,
468    /// Global iteration count
469    pub g_iter: usize,
470}
471
472impl LBFGSState {
473    /// The device of the state's tensors, if any have been populated.
474    fn current_device(&self) -> Option<Device> {
475        self.prev_flat_grad
476            .as_ref()
477            .or(self.d.as_ref())
478            .or(self.history_s.first())
479            .map(|t| t.device())
480    }
481
482    /// Moves all historical tensors to the target device.
483    pub fn to_device(self, device: &Device) -> Self {
484        Self {
485            history_s: self
486                .history_s
487                .into_iter()
488                .map(|t| t.to_device(device))
489                .collect(),
490            history_y: self
491                .history_y
492                .into_iter()
493                .map(|t| t.to_device(device))
494                .collect(),
495            d: self.d.map(|t| t.to_device(device)),
496            t: self.t,
497            prev_flat_grad: self.prev_flat_grad.map(|t| t.to_device(device)),
498            prev_loss: self.prev_loss,
499            g_iter: self.g_iter,
500        }
501    }
502}
503impl Default for LBFGSState {
504    fn default() -> Self {
505        Self {
506            history_s: Vec::new(),
507            history_y: Vec::new(),
508            d: None,
509            t: Some(1.0),
510            prev_flat_grad: None,
511            prev_loss: None,
512            g_iter: 0,
513        }
514    }
515}
516
517/// L-BFGS optimizer.
518///
519/// 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)
520///
521/// See also:
522/// - [L-BFGS](https://en.wikipedia.org/wiki/Limited-memory_BFGS)
523///
524/// # Note
525/// This optimizer is memory intensive
526#[derive(Clone)]
527pub struct LBFGS {
528    config: LBFGSConfig,
529    state: LBFGSState,
530}
531
532impl LBFGS {
533    /// Decompose the optimizer state into a serializable [`OptimizerRecord`] (burnpack format).
534    ///
535    /// L-BFGS keeps a single global state rather than per-parameter state, so its tensors are
536    /// named directly (e.g. `history_s.0`) and carry no parameter id.
537    pub fn to_record(&self) -> OptimizerRecord {
538        let mut sink = StateSink::default();
539        RecordState::state_flatten(&self.state, "", &mut sink);
540
541        let tensors = sink
542            .tensors
543            .into_iter()
544            .map(|(name, data)| {
545                burn_pack::Tensor::new(name, data.dtype, data.shape, None, data.bytes)
546            })
547            .collect();
548        let scalars = sink.scalars.into_iter().collect();
549
550        OptimizerRecord {
551            tensors,
552            scalars,
553            paths: Default::default(),
554        }
555    }
556
557    /// Load the optimizer state from an [`OptimizerRecord`].
558    ///
559    /// State tensors are materialized on the default device; the state is migrated to the gradient
560    /// device on the next [`step`](LBFGS::step), so no device argument is needed.
561    pub fn load_record(mut self, record: OptimizerRecord) -> Self {
562        let device = Device::default();
563        let mut source = StateSource::new(record.scalars);
564        for tensor in record.tensors {
565            let (name, dtype, shape, _, bytes) =
566                tensor.into_parts().expect("record tensors are resident");
567            source.insert_tensor(name, TensorData::from_bytes(bytes, shape, dtype));
568        }
569        if let Some(state) = LBFGSState::state_unflatten("", &mut source, &device) {
570            self.state = state;
571        }
572        self
573    }
574
575    /// Serialize the optimizer state to an in-memory burnpack byte buffer.
576    pub fn into_bytes(&self) -> Result<Bytes, RecordError> {
577        self.to_record().into_bytes()
578    }
579
580    /// Load the optimizer state from an in-memory burnpack byte buffer.
581    pub fn from_bytes(self, bytes: Bytes) -> Result<Self, RecordError> {
582        Ok(self.load_record(OptimizerRecord::from_bytes(bytes)?))
583    }
584
585    /// Save the optimizer state to a burnpack file on disk.
586    #[cfg(feature = "std")]
587    pub fn save<P: AsRef<std::path::Path>>(&self, path: P) -> Result<(), RecordError> {
588        self.to_record().save(path)
589    }
590
591    /// Load the optimizer state from a burnpack file on disk.
592    #[cfg(feature = "std")]
593    pub fn load<P: AsRef<std::path::Path>>(self, path: P) -> Result<Self, RecordError> {
594        Ok(self.load_record(OptimizerRecord::load(path)?))
595    }
596
597    /// A single optimization step for any tensor that represents the parameters of a model.
598    pub fn step<M, F>(&mut self, lr: LearningRate, mut module: M, mut closure: F) -> (M, f64)
599    where
600        M: AutodiffModule + Clone,
601        F: FnMut(M) -> (f64, GradientsParams),
602    {
603        // evaluate initial f(x) and df/dx
604        let (mut loss, grads) = closure(module.clone());
605        let mut current_evals = 1;
606
607        let mut flat_grad = flatten_grads_inner::<M>(&module, &grads);
608        let mut x_flat = flatten_params_inner::<M>(&module);
609
610        // Migrate the state to the gradient's device when they differ (e.g. just after loading a
611        // record on the default device). This is a no-op once the state is built from gradients.
612        let device = flat_grad.device();
613        if self.state.current_device().is_some_and(|d| d != device) {
614            self.state = core::mem::take(&mut self.state).to_device(&device);
615        }
616
617        let opt_cond =
618            flat_grad.clone().abs().max().into_scalar::<f64>() <= self.config.tolerance_grad;
619        // optimal condition
620        if opt_cond {
621            return (module, loss);
622        }
623
624        // tensors cached in state
625        let mut d = self
626            .state
627            .d
628            .take()
629            .unwrap_or_else(|| flat_grad.clone().neg());
630        let mut t = self.state.t.unwrap_or(lr);
631        let mut prev_flat_grad = self.state.prev_flat_grad.take();
632
633        let mut n_iter = 0;
634
635        // optimize for a max of max_iter iterations
636        while n_iter < self.config.max_iter {
637            // keep track of nb of iterations
638            n_iter += 1;
639            self.state.g_iter += 1;
640
641            // compute gradient descent direction
642            if self.state.g_iter == 1 {
643                d = flat_grad.clone().neg();
644                self.state.history_s.clear();
645                self.state.history_y.clear();
646            } else {
647                // do lbfgs update (update memory)
648                if let Some(pg) = prev_flat_grad.as_ref() {
649                    let y = flat_grad.clone().sub(pg.clone());
650                    let s = d.clone().mul_scalar(t);
651
652                    let ys: f64 = y.clone().dot(s.clone()).into_scalar();
653
654                    if ys > 1e-10 {
655                        // updating memory
656                        if self.state.history_s.len() >= self.config.history_size {
657                            // shift history by one (limited-memory)
658                            self.state.history_s.remove(0);
659                            self.state.history_y.remove(0);
660                        }
661                        self.state.history_s.push(s);
662                        self.state.history_y.push(y);
663                    }
664                }
665
666                // compute the approximate (L-BFGS) inverse Hessian
667                // multiplied by the gradient
668                let num_old = self.state.history_s.len();
669                let mut q = flat_grad.clone().neg();
670                let mut alphas: Vec<Tensor<1>> =
671                    vec![Tensor::zeros([1], &flat_grad.device().inner()); num_old];
672
673                if num_old > 0 {
674                    // multiply by initial Hessian
675                    // r/d is the final direction
676                    for i in (0..num_old).rev() {
677                        let s = &self.state.history_s[i];
678                        let y = &self.state.history_y[i];
679                        let rho = y.clone().dot(s.clone()).powf_scalar(-1.0);
680                        let alpha = rho.clone().mul(s.clone().dot(q.clone()));
681                        alphas[i] = alpha.clone();
682                        q = q.sub(y.clone().mul(alpha));
683                    }
684
685                    let last_s = &self.state.history_s[num_old - 1];
686                    let last_y = &self.state.history_y[num_old - 1];
687                    let ys = last_y.clone().dot(last_s.clone());
688                    let yy = last_y.clone().dot(last_y.clone());
689                    let h_diag = ys.div(yy);
690
691                    let mut r = q.mul(h_diag);
692
693                    for ((s, y), alpha) in self
694                        .state
695                        .history_s
696                        .iter()
697                        .zip(self.state.history_y.iter())
698                        .zip(alphas)
699                        .take(num_old)
700                    {
701                        let rho = y.clone().dot(s.clone()).powf_scalar(-1.0);
702
703                        let beta = rho.mul(y.clone().dot(r.clone()));
704
705                        r = r.add(s.clone().mul(alpha.sub(beta)));
706                    }
707                    d = r;
708                } else {
709                    d = q;
710                }
711            }
712
713            prev_flat_grad = Some(flat_grad.clone());
714            let prev_loss_iter = loss;
715
716            // compute step len
717            if self.state.g_iter == 1 {
718                let grad_l1: f64 = flat_grad.clone().abs().sum().into_scalar();
719                t = (1.0f64 / grad_l1).min(1.0) * lr;
720            } else {
721                t = lr;
722            }
723
724            // directional derivative
725            let gtd = flat_grad.clone().dot(d.clone()).into_scalar();
726
727            if gtd > -self.config.tolerance_change {
728                break;
729            }
730
731            let ls_func_evals;
732
733            if let LineSearchFn::StrongWolfe = self.config.line_search_fn {
734                // perform line search, using user function
735                let mut obj_func = |current_x: &Tensor<1>, step: f64, dir: &Tensor<1>| {
736                    let update = dir.clone().mul_scalar(step);
737                    let new_x = current_x.clone().add(update);
738                    let tmp_module = set_params_from_flat_inner::<M>(module.clone(), new_x);
739                    let (l, g) = closure(tmp_module);
740                    (l, flatten_grads_inner::<M>(&module, &g))
741                };
742
743                let (ls_f, ls_g, ls_t, evals) = strong_wolfe(
744                    &mut obj_func,
745                    &x_flat,
746                    t,
747                    &d,
748                    loss,
749                    flat_grad.clone(),
750                    gtd,
751                    1e-4,
752                    0.9,
753                    self.config.tolerance_change,
754                    self.config.max_eval.unwrap() - current_evals,
755                );
756
757                loss = ls_f;
758                flat_grad = ls_g;
759                t = ls_t;
760                ls_func_evals = evals;
761
762                x_flat = x_flat.add(d.clone().mul_scalar(t));
763                module = set_params_from_flat_inner::<M>(module, x_flat.clone());
764            } else {
765                // no line search, simply move with fixed-step
766                let step_vec = d.clone().mul_scalar(t);
767                x_flat = x_flat.add(step_vec);
768                module = set_params_from_flat_inner::<M>(module, x_flat.clone());
769                // re-evaluate function only if not in last iteration
770                // the reason we do this: in a stochastic setting,
771                // no use to re-evaluate that function here
772                let (new_loss, new_grads) = closure(module.clone());
773                loss = new_loss;
774                flat_grad = flatten_grads_inner::<M>(&module, &new_grads);
775                ls_func_evals = 1;
776            }
777
778            // update func eval
779            current_evals += ls_func_evals;
780
781            // check conditions
782
783            if current_evals >= self.config.max_eval.unwrap() {
784                break;
785            }
786
787            if flat_grad.clone().abs().max().into_scalar::<f64>() <= self.config.tolerance_grad {
788                break;
789            }
790
791            if d.clone().mul_scalar(t).abs().max().into_scalar::<f64>()
792                <= self.config.tolerance_change
793            {
794                break;
795            }
796
797            if (loss - prev_loss_iter).abs() < self.config.tolerance_change {
798                break;
799            }
800        }
801        self.state.d = Some(d);
802        self.state.t = Some(t);
803        self.state.prev_flat_grad = prev_flat_grad;
804        self.state.prev_loss = Some(loss);
805        (module, loss)
806    }
807    /// Moves the optimizer state to the specified device.
808    pub fn to_device(self, device: &Device) -> Self {
809        Self {
810            config: self.config,
811            // History tensors reside in InnerBackend, so we convert the device accordingly
812            state: self.state.to_device(device),
813        }
814    }
815}
816
817#[cfg(test)]
818mod tests {
819
820    use super::*;
821    use crate::GradientsParams;
822    use burn::module::Param;
823    use burn::tensor::{Tensor, TensorData};
824    use burn_nn::Linear;
825
826    fn given_linear_layer(weight: TensorData, bias: TensorData, device: &Device) -> Linear {
827        Linear {
828            weight: Param::from_data(weight, device),
829            bias: Some(Param::from_data(bias, device)),
830        }
831    }
832    #[test]
833    fn test_cubic_interpolate() {
834        let tolerance = 1e-8;
835
836        // basic
837        let (x1, f1, g1, x2, f2, g2) = (-1.0, 1.0, -2.0, 1.0, 1.0, 2.0);
838        let result = cubic_interpolate(x1, f1, g1, x2, f2, g2, None);
839        assert!(
840            (result - 0.00000).abs() < tolerance,
841            "Basic: Result {} should be close to 0.0",
842            result
843        );
844
845        // bound
846        let (x1, f1, g1, x2, f2, g2) = (0.0, 0.25, -1.0, 1.0, 0.25, 1.0);
847        let bounds = Some((0.6, 1.0));
848        let result = cubic_interpolate(x1, f1, g1, x2, f2, g2, bounds);
849        assert!(
850            (result - 0.6000000000).abs() < tolerance,
851            "Bound: Result {} should be clamped to 0.6",
852            result
853        );
854
855        // d2_square < 0,should return mid value
856        let (x1, f1, g1, x2, f2, g2) = (0.0, 0.0, 10.0, 1.0, 5.0, 10.0);
857        let result = cubic_interpolate(x1, f1, g1, x2, f2, g2, Some((0.0, 1.0)));
858        assert!(
859            (result - 0.5000000).abs() < tolerance,
860            "Fallback: Result {} should be midpoint 0.5",
861            result
862        );
863
864        // asymmetric
865        let (x1, f1, g1, x2, f2, g2) = (0.0, 1.0, -5.0, 1.0, 0.5, 1.0);
866        let result = cubic_interpolate(x1, f1, g1, x2, f2, g2, None);
867        assert!(
868            (result - 0.4606553370833684).abs() < tolerance,
869            "Asymmetric: Result {} should be 0.4606553370833684",
870            result
871        );
872
873        // not good value
874        let (x1, f1, g1, x2, f2, g2) = (
875            1.231232145,
876            -0.12567458754,
877            9.1231243007,
878            8.239105015,
879            -100.9012398021,
880            123201321.0293982,
881        );
882        let result_1 = cubic_interpolate(x1, f1, g1, x2, f2, g2, None);
883        let result_2 = cubic_interpolate(x1, f1, g1, x2, f2, g2, Some((-4.4, 4.4)));
884        assert!(
885            (result_1 - 5.9031480234724434).abs() < tolerance,
886            "not good value 1: Result {} should be 5.9031480234724434",
887            result
888        );
889        assert!(
890            (result_2 - 4.4000000000000004).abs() < tolerance,
891            "not good value 2: Result {} should be 4.4000000000000004",
892            result
893        );
894    }
895    #[test]
896    fn test_strong_wolfe_direct_comparison() {
897        let device = Device::default().autodiff();
898        let tol = 1e-6;
899
900        {
901            let x = Tensor::<1>::from_floats([2.1321912957_f64], &device);
902            let d = Tensor::<1>::from_floats([0.91312321_f64], &device);
903            let t_initial = 1.213132_f64;
904            fn func(x_base: &Tensor<1>, t_val: f64, d_vec: &Tensor<1>) -> (f64, Tensor<1>) {
905                let curr_x = x_base.clone().add(d_vec.clone().mul_scalar(t_val));
906                let x2 = curr_x.clone().mul(curr_x.clone());
907                let x3 = x2.clone().mul(curr_x.clone());
908                let x4 = x2.clone().mul(x2.clone());
909
910                // f(x) = x^4 - 2*x^2 + x
911                let f_elements = x4 - x2.mul_scalar(2.0) + curr_x.clone();
912
913                let f_val = f_elements.sum().into_scalar();
914
915                // g(x) = 4*x^3 - 4*x + 1
916                let g = x3.mul_scalar(4.0) - curr_x.clone().mul_scalar(4.0)
917                    + Tensor::ones_like(&curr_x);
918
919                (f_val, g)
920            }
921            let (f_init, g_init) = func(&x, 0.0, &d);
922            let gtd_init = g_init.clone().dot(d.clone()).into_scalar::<f64>();
923            println!("Initial State: f={},gtd = {}", f_init, gtd_init);
924            assert!((f_init - 13.7080059052).abs() < tol);
925            assert!((gtd_init - 28.5305728912).abs() < tol);
926            let mut obj_func = |xb: &Tensor<1>, tv: f64, dv: &Tensor<1>| func(xb, tv, dv);
927
928            let (f_final, _g_final, t_final, evals) = strong_wolfe(
929                &mut obj_func,
930                &x,
931                t_initial,
932                &d,
933                f_init,
934                g_init,
935                gtd_init,
936                1e-4, // c1
937                0.9,  // c2
938                1e-9, // tolerance_change
939                10,   // max_ls
940            );
941            let g_f = _g_final.into_scalar::<f64>();
942            println!(
943                "f_final:{:?},_g_final:{:?},t_final:{:?},evals:{:?}",
944                f_final, g_f, t_final, evals
945            );
946            assert!((f_final - 13.708005905151367).abs() < tol);
947            assert!((g_f - 31.2450428009).abs() < tol);
948            assert!((t_final - 0.0).abs() < tol);
949            assert!((evals == 11));
950        }
951    }
952    #[test]
953    fn test_lbfgs_strong_wolfe_comparison() {
954        let device = Device::default().autodiff();
955        let tol = 1e-5;
956        let x_data = Tensor::<2>::from_data([[1.0], [2.0], [3.0]], &device);
957        let y_true = Tensor::<2>::from_data([[3.0], [5.0], [7.0]], &device);
958        let weight = TensorData::from([[0.5f64]]);
959        let bias = TensorData::from([0.1f64]);
960        let module = given_linear_layer(weight, bias, &device);
961
962        let mut optimizer = LBFGSConfig::new()
963            .with_line_search_fn(LineSearchFn::StrongWolfe)
964            .init();
965        let mut closure = |mod_in: Linear| {
966            let output = mod_in.forward(x_data.clone());
967            let loss = burn_nn::loss::MseLoss::new().forward(
968                output,
969                y_true.clone(),
970                burn_nn::loss::Reduction::Sum,
971            );
972
973            let grads = loss.backward();
974            let grads_params = GradientsParams::from_grads(grads, &mod_in);
975
976            (loss.into_scalar::<f64>(), grads_params)
977        };
978        let initial_loss = closure(module.clone()).0;
979        assert!((initial_loss - 50.1300048828).abs() < tol);
980        let (updated_module, final_loss) = optimizer.step(0.001, module, &mut closure);
981        assert!((final_loss - 0.0234732367).abs() < tol);
982        let optimized_data: f64 = updated_module.weight.val().into_scalar();
983        let optimized_bias: f64 = updated_module.bias.as_ref().unwrap().val().into_scalar();
984        assert!((optimized_data - 2.0570652485).abs() < tol);
985        assert!((optimized_bias - 0.8106800914).abs() < tol);
986    }
987
988    // A burnpack round-trip of the L-BFGS state (which holds `Vec<Tensor>` history buffers, optional
989    // tensors and optional scalars) must restore enough that a further step agrees with the original.
990    #[test]
991    fn test_lbfgs_burnpack_round_trip() {
992        let device = Device::default().autodiff();
993        let tol = 1e-6;
994        let x_data = Tensor::<2>::from_data([[1.0], [2.0], [3.0]], &device);
995        let y_true = Tensor::<2>::from_data([[3.0], [5.0], [7.0]], &device);
996        let module = given_linear_layer(
997            TensorData::from([[0.5f64]]),
998            TensorData::from([0.1f64]),
999            &device,
1000        );
1001
1002        let make_closure = || {
1003            let x = x_data.clone();
1004            let y = y_true.clone();
1005            move |mod_in: Linear| {
1006                let output = mod_in.forward(x.clone());
1007                let loss = burn_nn::loss::MseLoss::new().forward(
1008                    output,
1009                    y.clone(),
1010                    burn_nn::loss::Reduction::Sum,
1011                );
1012                let grads = loss.backward();
1013                let grads_params = GradientsParams::from_grads(grads, &mod_in);
1014                (loss.into_scalar::<f64>(), grads_params)
1015            }
1016        };
1017
1018        let mut optimizer = LBFGSConfig::new()
1019            .with_line_search_fn(LineSearchFn::StrongWolfe)
1020            .init();
1021        let (module, _) = optimizer.step(0.001, module, &mut make_closure());
1022
1023        // Round-trip the optimizer state. State tensors live on the inner (non-autodiff) backend.
1024        let bytes = optimizer.into_bytes().unwrap();
1025        let mut reloaded = LBFGSConfig::new()
1026            .with_line_search_fn(LineSearchFn::StrongWolfe)
1027            .init()
1028            .from_bytes(bytes)
1029            .unwrap();
1030
1031        // A further identical step on each optimizer must agree — exercising the restored history.
1032        let (_, loss_original) = optimizer.step(0.001, module.clone(), &mut make_closure());
1033        let (_, loss_reloaded) = reloaded.step(0.001, module, &mut make_closure());
1034        assert!(
1035            (loss_original - loss_reloaded).abs() < tol,
1036            "losses differ after burnpack round-trip: {loss_original} vs {loss_reloaded}"
1037        );
1038    }
1039
1040    #[test]
1041    fn test_lbfgs_no_strong_wolfe_comparison() {
1042        let device = Device::default().autodiff();
1043        let tol = 1e-5;
1044        let x_data = Tensor::<2>::from_data([[1.0], [2.0], [3.0]], &device);
1045        let y_true = Tensor::<2>::from_data([[3.0], [5.0], [7.0]], &device);
1046        let weight = TensorData::from([[0.5f64]]);
1047        let bias = TensorData::from([0.1f64]);
1048        let module = given_linear_layer(weight, bias, &device);
1049
1050        let mut optimizer = LBFGSConfig::new()
1051            .with_line_search_fn(LineSearchFn::None)
1052            .init();
1053        let mut closure = |mod_in: Linear| {
1054            let output = mod_in.forward(x_data.clone());
1055            let loss = burn_nn::loss::MseLoss::new().forward(
1056                output,
1057                y_true.clone(),
1058                burn_nn::loss::Reduction::Sum,
1059            );
1060
1061            let grads = loss.backward();
1062            let grads_params = GradientsParams::from_grads(grads, &mod_in);
1063
1064            (loss.into_scalar::<f64>(), grads_params)
1065        };
1066        let initial_loss = closure(module.clone()).0;
1067        assert!((initial_loss - 50.1300048828).abs() < tol);
1068        let (updated_module, final_loss) = optimizer.step(0.001, module, &mut closure);
1069        assert!((final_loss - 48.2181930542).abs() < tol);
1070        let optimized_data: f64 = updated_module.weight.val().into_scalar();
1071        let optimized_bias: f64 = updated_module.bias.as_ref().unwrap().val().into_scalar();
1072
1073        assert!((optimized_data - 0.5302446192).abs() < tol);
1074        assert!((optimized_bias - 0.1142520783).abs() < tol);
1075    }
1076}