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
20fn 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 let (min_bound, max_bound) = bounds.unwrap_or(if x1 <= x2 { (x1, x2) } else { (x2, x1) });
35 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}
57struct LineSearchSample {
59 t: f64,
61 f: f64,
63 g: Tensor<1>,
65 gtd: f64,
67}
68
69#[allow(clippy::too_many_arguments)]
70fn strong_wolfe<F>(
71 obj_func: &mut F,
73 x: &Tensor<1>,
74 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 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 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 let mut bracket: Option<[LineSearchSample; 2]> = None;
102 let mut wolfe_bracket: Option<LineSearchSample> = None;
104 while ls_iter < max_ls {
105 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 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 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 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 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 let mut insuf_progress = false;
206
207 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 if diff * d_norm < tolerance_change {
218 break;
219 }
220
221 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 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 (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 (
299 bracket[low_idx].f,
300 bracket[low_idx].g.clone(),
301 bracket[low_idx].t,
302 ls_func_evals,
303 )
304}
305
306#[derive(Clone, Default, Debug, Copy, PartialEq, Eq, Serialize, Deserialize)]
308pub enum LineSearchFn {
309 #[default]
311 None,
312 StrongWolfe,
316}
317
318#[derive(Config, Debug)]
320pub struct LBFGSConfig {
321 #[config(default = 20)]
323 pub max_iter: usize,
324 #[config(default = 100)]
326 pub history_size: usize,
327 #[config(default = 1e-7)]
329 pub tolerance_grad: f64,
330 #[config(default = 1e-9)]
332 pub tolerance_change: f64,
333 #[config(default = "None")]
335 pub max_eval: Option<usize>,
336 #[config(default = "LineSearchFn::None")]
338 pub line_search_fn: LineSearchFn,
339}
340
341impl LBFGSConfig {
342 pub fn init(&self) -> LBFGS {
348 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
364struct 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
379fn 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
404fn 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
418struct 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
443fn 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#[derive(Clone, RecordState)]
455pub struct LBFGSState {
456 pub history_s: Vec<Tensor<1>>,
458 pub history_y: Vec<Tensor<1>>,
460 pub d: Option<Tensor<1>>,
462 pub t: Option<f64>,
464 pub prev_flat_grad: Option<Tensor<1>>,
466 pub prev_loss: Option<f64>,
468 pub g_iter: usize,
470}
471
472impl LBFGSState {
473 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 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#[derive(Clone)]
527pub struct LBFGS {
528 config: LBFGSConfig,
529 state: LBFGSState,
530}
531
532impl LBFGS {
533 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 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 pub fn into_bytes(&self) -> Result<Bytes, RecordError> {
577 self.to_record().into_bytes()
578 }
579
580 pub fn from_bytes(self, bytes: Bytes) -> Result<Self, RecordError> {
582 Ok(self.load_record(OptimizerRecord::from_bytes(bytes)?))
583 }
584
585 #[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 #[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 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 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 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 if opt_cond {
621 return (module, loss);
622 }
623
624 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 while n_iter < self.config.max_iter {
637 n_iter += 1;
639 self.state.g_iter += 1;
640
641 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 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 if self.state.history_s.len() >= self.config.history_size {
657 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 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 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 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 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 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 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 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 current_evals += ls_func_evals;
780
781 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 pub fn to_device(self, device: &Device) -> Self {
809 Self {
810 config: self.config,
811 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 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 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 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 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 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 let f_elements = x4 - x2.mul_scalar(2.0) + curr_x.clone();
912
913 let f_val = f_elements.sum().into_scalar();
914
915 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, 0.9, 1e-9, 10, );
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 #[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 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 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}