pounce_algorithm/line_search/
penalty_acceptor.rs1use crate::ipopt_cq::IpoptCqHandle;
35use crate::ipopt_data::IpoptDataHandle;
36use crate::iterates_vector::IteratesVector;
37use crate::line_search::filter_acceptor::AcceptDecision;
38use crate::line_search::ls_acceptor::BacktrackingLsAcceptor;
39use pounce_common::types::Number;
40use pounce_common::utils::compare_le;
41use pounce_linalg::Vector;
42use std::rc::Rc;
43
44pub struct PenaltyLsAcceptor {
45 pub rho: Number,
48 pub nu_inc: Number,
50 pub nu_init: Number,
52 pub nu_max: Number,
54 pub eta_penalty: Number,
56 nu: Number,
57 last_nu: Number,
58 cache: Option<RefCache>,
61}
62
63struct RefCache {
65 theta_ref: Number,
66 barr_ref: Number,
67 grad_barr_t_delta: Number,
68 dwd: Number,
69 c_ref: Rc<dyn Vector>,
71 d_minus_s_ref: Rc<dyn Vector>,
73 jac_c_delta: Rc<dyn Vector>,
75 jac_d_delta_minus_ds: Rc<dyn Vector>,
77}
78
79impl Default for PenaltyLsAcceptor {
80 fn default() -> Self {
81 Self {
82 rho: 0.1,
83 nu_inc: 1e-4,
84 nu_init: 1e-6,
85 nu_max: 1e40,
86 eta_penalty: 1e-8,
87 nu: 1e-6,
88 last_nu: 1e-6,
89 cache: None,
90 }
91 }
92}
93
94impl PenaltyLsAcceptor {
95 pub fn new() -> Self {
96 Self::default()
97 }
98
99 pub fn nu(&self) -> Number {
100 self.nu
101 }
102
103 pub fn last_nu(&self) -> Number {
104 self.last_nu
105 }
106
107 pub fn reset(&mut self) {
110 self.nu = self.nu_init;
111 self.last_nu = self.nu_init;
112 self.cache = None;
113 }
114
115 pub fn update_nu(
124 &mut self,
125 grad_barr_t_delta: Number,
126 delta_w_delta: Number,
127 reference_theta: Number,
128 ) {
129 self.last_nu = self.nu;
130 if reference_theta > 0.0 {
131 let nu_plus =
132 (grad_barr_t_delta + 0.5 * delta_w_delta) / ((1.0 - self.rho) * reference_theta);
133 if self.nu < nu_plus {
134 self.nu = nu_plus + self.nu_inc;
135 }
136 }
137 }
138
139 fn calc_pred(&self, alpha: Number) -> Number {
143 let cache = self
144 .cache
145 .as_ref()
146 .expect("calc_pred called before init_this_line_search");
147 let mut tmp_c = cache.c_ref.make_new();
149 tmp_c.set(0.0);
150 tmp_c.add_two_vectors(1.0, &*cache.c_ref, alpha, &*cache.jac_c_delta, 0.0);
151 let mut tmp_d = cache.d_minus_s_ref.make_new();
152 tmp_d.set(0.0);
153 tmp_d.add_two_vectors(
154 1.0,
155 &*cache.d_minus_s_ref,
156 alpha,
157 &*cache.jac_d_delta_minus_ds,
158 0.0,
159 );
160 let theta_2 = tmp_c.asum() + tmp_d.asum();
161
162 let pred = -alpha * cache.grad_barr_t_delta - 0.5 * alpha * alpha * cache.dwd
163 + self.nu * (cache.theta_ref - theta_2);
164 if pred < 0.0 { 0.0 } else { pred }
165 }
166}
167
168impl BacktrackingLsAcceptor for PenaltyLsAcceptor {
169 fn reset(&mut self) {
170 PenaltyLsAcceptor::reset(self);
171 }
172
173 fn init_this_line_search(
177 &mut self,
178 _data: &IpoptDataHandle,
179 cq: &IpoptCqHandle,
180 delta: &IteratesVector,
181 ) {
182 let cqr = cq.borrow();
183 let theta_ref = cqr.curr_constraint_violation();
184 let barr_ref = cqr.curr_barrier_obj();
185 let grad_barr_t_delta = cqr.curr_grad_barr_t_delta(&*delta.x, &*delta.s);
186 let dwd = cqr.curr_dwd(&*delta.x, &*delta.s);
187
188 let c_ref = cqr.curr_c();
190 let d_minus_s_ref = cqr.curr_d_minus_s();
191 let jac_c_delta = cqr.curr_jac_c_times_vec(&*delta.x);
192 let jac_d_delta = cqr.curr_jac_d_times_vec(&*delta.x);
194 let mut tmp = jac_d_delta.make_new();
195 tmp.set(0.0);
196 tmp.add_two_vectors(1.0, &*jac_d_delta, -1.0, &*delta.s, 0.0);
197 let jac_d_delta_minus_ds: Rc<dyn Vector> = Rc::from(tmp);
198 drop(cqr);
199
200 self.cache = Some(RefCache {
201 theta_ref,
202 barr_ref,
203 grad_barr_t_delta,
204 dwd,
205 c_ref,
206 d_minus_s_ref,
207 jac_c_delta,
208 jac_d_delta_minus_ds,
209 });
210
211 self.update_nu(grad_barr_t_delta, dwd, theta_ref);
213 }
214
215 fn check_trial_point(
229 &mut self,
230 alpha_primal: Number,
231 _theta: Number,
232 _phi: Number,
233 _d_phi: Number,
234 theta_trial: Number,
235 phi_trial: Number,
236 ) -> AcceptDecision {
237 let cache = match &self.cache {
241 Some(c) => c,
242 None => return AcceptDecision::Accept,
243 };
244
245 let pred = self.calc_pred(alpha_primal);
246 let ref_merit = cache.barr_ref + self.nu * cache.theta_ref;
247 let ared = ref_merit - (phi_trial + self.nu * theta_trial);
248
249 if compare_le(self.eta_penalty * pred, ared, ref_merit.abs()) {
250 AcceptDecision::Accept
251 } else {
252 AcceptDecision::Reject
253 }
254 }
255}
256
257#[cfg(test)]
258mod tests {
259 use super::*;
260
261 #[test]
262 fn no_bump_when_theta_zero() {
263 let mut a = PenaltyLsAcceptor::new();
264 let nu0 = a.nu();
265 a.update_nu(10.0, 5.0, 0.0);
266 assert_eq!(a.nu(), nu0);
267 assert_eq!(a.last_nu(), nu0);
268 }
269
270 #[test]
271 fn bump_when_nu_plus_exceeds_current() {
272 let mut a = PenaltyLsAcceptor {
273 rho: 0.1,
274 nu_inc: 1e-4,
275 nu: 0.0,
276 last_nu: 0.0,
277 ..Default::default()
278 };
279 a.update_nu(1.0, 0.0, 1.0);
282 assert!(a.last_nu() == 0.0);
283 let expected = 1.0 / 0.9 + 1e-4;
284 assert!((a.nu() - expected).abs() < 1e-12);
285 }
286
287 #[test]
288 fn no_bump_when_already_above_nu_plus() {
289 let mut a = PenaltyLsAcceptor {
290 rho: 0.1,
291 nu_inc: 1e-4,
292 nu: 1e6,
293 last_nu: 1e6,
294 ..Default::default()
295 };
296 a.update_nu(1.0, 0.0, 1.0);
297 assert_eq!(a.nu(), 1e6);
298 }
299
300 #[test]
301 fn reset_restores_init() {
302 let mut a = PenaltyLsAcceptor::new();
303 a.update_nu(10.0, 0.0, 1.0); let bumped = a.nu();
305 assert!(bumped > a.nu_init);
306 PenaltyLsAcceptor::reset(&mut a);
307 assert_eq!(a.nu(), a.nu_init);
308 }
309
310 #[test]
311 fn check_trial_point_without_cache_accepts() {
312 let mut a = PenaltyLsAcceptor::new();
314 assert_eq!(
315 a.check_trial_point(1.0, 1.0, 10.0, -1.0, 0.5, 8.0),
316 AcceptDecision::Accept
317 );
318 }
319
320 fn cache_for_test(
323 theta_ref: Number,
324 barr_ref: Number,
325 grad_barr_t_delta: Number,
326 dwd: Number,
327 c_ref: Vec<Number>,
328 d_minus_s_ref: Vec<Number>,
329 jac_c_delta: Vec<Number>,
330 jac_d_delta_minus_ds: Vec<Number>,
331 ) -> RefCache {
332 use pounce_linalg::Vector;
333 use pounce_linalg::dense_vector::DenseVectorSpace;
334 let mkr = |v: Vec<Number>| -> Rc<dyn Vector> {
335 let mut x = DenseVectorSpace::new(v.len() as i32).make_new_dense();
336 x.values_mut().copy_from_slice(&v);
337 Rc::new(x)
338 };
339 RefCache {
340 theta_ref,
341 barr_ref,
342 grad_barr_t_delta,
343 dwd,
344 c_ref: mkr(c_ref),
345 d_minus_s_ref: mkr(d_minus_s_ref),
346 jac_c_delta: mkr(jac_c_delta),
347 jac_d_delta_minus_ds: mkr(jac_d_delta_minus_ds),
348 }
349 }
350
351 #[test]
352 fn calc_pred_matches_closed_form() {
353 let mut a = PenaltyLsAcceptor::new();
359 a.nu = 0.5;
360 a.cache = Some(cache_for_test(
361 3.0,
362 0.0,
363 2.0,
364 4.0,
365 vec![1.0, 2.0],
366 vec![4.0],
367 vec![-1.0, -1.0],
368 vec![-2.0],
369 ));
370 assert!((a.calc_pred(0.5) - 0.0).abs() < 1e-12);
371 }
372
373 #[test]
374 fn calc_pred_positive_when_directions_align() {
375 let mut a = PenaltyLsAcceptor::new();
379 a.nu = 1.0;
380 a.cache = Some(cache_for_test(
381 3.0,
382 0.0,
383 -2.0,
384 0.0,
385 vec![1.0, 2.0],
386 vec![0.0],
387 vec![-1.0, -2.0],
388 vec![0.0],
389 ));
390 assert!((a.calc_pred(1.0) - 5.0).abs() < 1e-12);
391 }
392
393 #[test]
394 fn check_trial_point_accepts_when_ared_meets_pred() {
395 let mut a = PenaltyLsAcceptor::new();
399 a.nu = 1.0;
400 a.eta_penalty = 0.5;
401 a.cache = Some(cache_for_test(
402 3.0,
403 0.0,
404 -2.0,
405 0.0,
406 vec![1.0, 2.0],
407 vec![0.0],
408 vec![-1.0, -2.0],
409 vec![0.0],
410 ));
411 assert_eq!(
412 a.check_trial_point(1.0, 3.0, 0.0, -2.0, 0.0, -3.0),
413 AcceptDecision::Accept
414 );
415 }
416
417 #[test]
418 fn check_trial_point_rejects_insufficient_decrease() {
419 let mut a = PenaltyLsAcceptor::new();
422 a.nu = 1.0;
423 a.eta_penalty = 0.5;
424 a.cache = Some(cache_for_test(
425 3.0,
426 0.0,
427 -2.0,
428 0.0,
429 vec![1.0, 2.0],
430 vec![0.0],
431 vec![-1.0, -2.0],
432 vec![0.0],
433 ));
434 assert_eq!(
435 a.check_trial_point(1.0, 3.0, 0.0, -2.0, 2.999, 0.0),
436 AcceptDecision::Reject
437 );
438 }
439}