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#[derive(Clone, Default, Debug, Copy, PartialEq, Eq, Serialize, Deserialize)]
28pub enum LineSearchFn {
29 #[default]
31 None,
32 StrongWolfe,
36}
37
38#[derive(Config, Debug)]
40pub struct LBFGSConfig {
41 #[config(default = 20)]
43 pub max_iter: usize,
44 #[config(default = 100)]
46 pub history_size: usize,
47 #[config(default = 1e-7)]
49 pub tolerance_grad: f64,
50 #[config(default = 1e-9)]
52 pub tolerance_change: f64,
53 #[config(default = "None")]
55 pub max_eval: Option<usize>,
56 #[config(default = "LineSearchFn::None")]
58 pub line_search_fn: LineSearchFn,
59}
60
61impl LBFGSConfig {
62 pub fn init<B: AutodiffBackend>(&self) -> LBFGS<B> {
68 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
84struct 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
104fn 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
137fn 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
155struct 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
188fn 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#[derive(Clone, Record)]
204pub struct LBFGSState<B: Backend> {
205 pub history_s: Vec<Tensor<B, 1>>,
207 pub history_y: Vec<Tensor<B, 1>>,
209 pub d: Option<Tensor<B, 1>>,
211 pub t: Option<f64>,
213 pub prev_flat_grad: Option<Tensor<B, 1>>,
215 pub prev_loss: Option<f64>,
217 pub g_iter: usize,
219}
220
221impl<B: Backend> LBFGSState<B> {
222 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#[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 pub fn to_record(&self) -> LBFGSState<B::InnerBackend> {
275 self.state.clone()
276 }
277
278 pub fn load_record(mut self, record: LBFGSState<B::InnerBackend>) -> Self {
280 self.state = record;
281 self
282 }
283
284 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 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 if opt_cond {
309 return (module, loss);
310 }
311
312 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 while n_iter < self.config.max_iter {
325 n_iter += 1;
327 self.state.g_iter += 1;
328
329 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 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 if self.state.history_s.len() >= self.config.history_size {
345 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 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 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 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 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 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 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 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 current_evals += ls_func_evals;
471
472 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 pub fn to_device(self, device: &B::Device) -> Self {
500 Self {
501 config: self.config,
502 state: self.state.to_device(device),
504 }
505 }
506}
507
508#[cfg(test)]
509mod tests;