mini_ode/optimizers/
bfgs.rs1use anyhow::anyhow;
2use std::fmt;
3use tch::Tensor;
4
5use super::Optimizer;
6
7use crate::utils::differentiation;
8use crate::utils::linesearch;
9use crate::utils::validation;
10use crate::utils::warnings::warn;
11
12pub struct BFGS {
23 max_steps: usize,
25 gtol: Option<f64>,
27 ftol: Option<f64>,
29}
30
31impl BFGS {
32 pub fn new(max_steps: usize, gtol: Option<f64>, ftol: Option<f64>) -> Self {
42 Self {
43 max_steps,
44 gtol,
45 ftol,
46 }
47 }
48}
49
50impl Optimizer for BFGS {
51 fn optimize(
52 &self,
53 function: &dyn Fn(&Tensor) -> Tensor,
54 x0: &Tensor,
55 ) -> anyhow::Result<Tensor> {
56 if x0.size().len() != 1 {
58 return Err(anyhow!("`x0` must have rank 1"));
59 }
60
61 let kind = x0.kind();
63 let device = x0.device();
64
65 let mut prev3_step_norm = 0f64;
66 let mut prev2_step_norm = 0f64;
67 let mut prev_step_norm = 0f64;
68
69 let x0_length = x0.size()[0];
70 let identity = match Tensor::f_eye(x0_length, (kind, device)) {
71 Ok(matrix) => matrix,
72 Err(tch::TchError::Torch(_)) => {
76 return Err(anyhow!(
77 "Could not allocate {}x{} matrix. Maybe try less resourcefull algorithm.",
78 x0_length,
79 x0_length
80 ));
81 }
82 e => e.unwrap(),
83 };
84 let mut x = x0.copy();
85 let mut appr_inv_h = identity.copy();
86 let mut curr_grad = match differentiation::differentiate(function, &x) {
87 Ok(grad) => grad,
88 Err(e) => {
89 return Err(anyhow!(
90 "Runtime error: Differentiation failed in BFGS optimizer: {}",
91 e
92 ));
93 }
94 };
95 let mut curr_y = function(&x);
96
97 if curr_y.size() != Vec::<i64>::new() {
99 return Err(anyhow!("Output of function `function` must be scalar"));
100 }
101
102 let mut warned_inv_hess_large = false;
103
104 for _ in 0..self.max_steps {
105 if let Some(gtol) = self.gtol {
107 if curr_grad.f_norm()?.f_double_value(&[])? < gtol {
108 validation::validate_optimizer_output(&x, "BFGS")?;
110 return Ok(x);
111 }
112 } else {
113 if curr_grad.f_norm()?.f_double_value(&[])? == 0. {
117 validation::validate_optimizer_output(&x, "BFGS")?;
119 return Ok(x);
120 }
121 }
122
123 let direction = (-appr_inv_h.f_mm(&curr_grad.f_reshape([-1, 1])?)?).f_reshape([-1])?;
125
126 let linesearch_atol =
128 linesearch::P0.max(prev_step_norm.min(prev2_step_norm).min(prev3_step_norm) / 100.);
129
130 let step =
133 linesearch::choose_step_golden_section(&x, &direction, function, linesearch_atol)?;
134
135 prev3_step_norm = prev2_step_norm;
137 prev2_step_norm = prev_step_norm;
138 prev_step_norm = step.f_norm()?.f_double_value(&[])?;
139
140 x = x + &step;
142
143 let y = function(&x);
145 if let Some(ftol) = self.ftol {
146 if (curr_y.f_double_value(&[])? - y.f_double_value(&[])?) < ftol {
147 validation::validate_optimizer_output(&x, "BFGS")?;
149 return Ok(x);
150 }
151 }
152 curr_y = y;
153
154 let grad = match differentiation::differentiate(function, &x) {
155 Ok(grad) => grad,
156 Err(e) => {
157 return Err(anyhow!(
158 "Runtime error: Differentiation failed in BFGS optimizer: {}",
159 e
160 ));
161 }
162 };
163 let gdiff = &grad - &curr_grad;
164
165 let gamma = {
168 let delta = 0.0001;
169
170 let sty = step.f_dot(&gdiff)?.f_double_value(&[])?;
171 let step_norm_sq = step.f_dot(&step)?.f_double_value(&[])?;
172
173 let theta = if sty >= delta * step_norm_sq {
174 1.
175 } else {
176 let numerator = (1. - delta) * step_norm_sq;
177 let denominator = step_norm_sq - sty;
178
179 if denominator.abs() < 1e-10 {
180 1.
181 } else {
182 (numerator / denominator).min(1.)
183 }
184 };
185
186 let projection_factor = if step_norm_sq < 1e-10 {
187 0.
188 } else {
189 sty / step_norm_sq
190 };
191 let gdiff_prime = &gdiff * theta + &step * ((1. - theta) * projection_factor);
192 let sty_prime = step.f_dot(&gdiff_prime)?.f_double_value(&[])?;
193
194 if sty_prime.abs() < 1e-10 {
195 1. / (delta * step_norm_sq + 1e-10)
196 } else {
197 1. / sty_prime
198 }
199 };
200
201 appr_inv_h = (&identity
203 - gamma * step.f_reshape([-1, 1])?.f_mm(&gdiff.f_reshape([1, -1])?)?)
204 .f_mm(&appr_inv_h)?
205 .f_mm(
206 &(&identity - gamma * gdiff.f_reshape([-1, 1])?.f_mm(&step.f_reshape([1, -1])?)?),
207 )? + gamma * step.f_reshape([-1, 1])?.f_mm(&step.f_reshape([1, -1])?)?;
208
209 let inv_h_norm = appr_inv_h.f_norm()?.f_double_value(&[])?;
211 if !warned_inv_hess_large && inv_h_norm > 1e10 {
212 warn!(
213 "BFGS: inverse Hessian approximation norm reached {:.3e}; problem may be ill-conditioned",
214 inv_h_norm
215 );
216 warned_inv_hess_large = true;
217 }
218
219 curr_grad = grad;
220 }
221
222 validation::validate_optimizer_output(&x, "BFGS")?;
224 Ok(x)
225 }
226}
227
228impl fmt::Display for BFGS {
229 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
230 let mut string = String::from("BFGS(");
231
232 string = string + "max_steps=" + self.max_steps.to_string().as_str();
233 if let Some(gtol) = self.gtol {
234 string = string + ", gtol=" + gtol.to_string().as_str();
235 }
236 if let Some(ftol) = self.ftol {
237 string = string + ", ftol=" + ftol.to_string().as_str();
238 }
239
240 string = string + ")";
241
242 write!(f, "{}", string)
243 }
244}