1use pounce_common::types::Number;
23
24use crate::block_solve::{BlockSolveError, lu_factor_partial_pivot, lu_solve};
25
26#[derive(Debug, Default, Clone)]
29pub struct ReductionFrame {
30 pub fixed_vars: Vec<usize>,
32 pub fixed_values: Vec<Number>,
34 pub dropped_rows: Vec<usize>,
37 pub var_map: Vec<Option<usize>>,
40 pub row_map: Vec<Option<usize>>,
42}
43
44impl ReductionFrame {
45 pub fn new(
48 n_vars: usize,
49 n_rows: usize,
50 fixed_vars: Vec<usize>,
51 fixed_values: Vec<Number>,
52 dropped_rows: Vec<usize>,
53 ) -> Self {
54 assert_eq!(
55 fixed_vars.len(),
56 fixed_values.len(),
57 "fixed_vars and fixed_values must be the same length"
58 );
59 assert_eq!(
60 fixed_vars.len(),
61 dropped_rows.len(),
62 "fixed_vars and dropped_rows must be the same length (square block)"
63 );
64
65 let mut is_fixed_var = vec![false; n_vars];
68 for &i in &fixed_vars {
69 is_fixed_var[i] = true;
70 }
71 let mut is_dropped_row = vec![false; n_rows];
72 for &i in &dropped_rows {
73 is_dropped_row[i] = true;
74 }
75
76 let mut var_map = vec![None; n_vars];
77 let mut next_reduced = 0;
78 for (i, slot) in var_map.iter_mut().enumerate().take(n_vars) {
79 if is_fixed_var[i] {
80 continue;
81 }
82 *slot = Some(next_reduced);
83 next_reduced += 1;
84 }
85
86 let mut row_map = vec![None; n_rows];
87 let mut next_reduced_row = 0;
88 for (i, slot) in row_map.iter_mut().enumerate().take(n_rows) {
89 if is_dropped_row[i] {
90 continue;
91 }
92 *slot = Some(next_reduced_row);
93 next_reduced_row += 1;
94 }
95
96 Self {
97 fixed_vars,
98 fixed_values,
99 dropped_rows,
100 var_map,
101 row_map,
102 }
103 }
104
105 pub fn n_full_vars(&self) -> usize {
106 self.var_map.len()
107 }
108
109 pub fn n_full_rows(&self) -> usize {
110 self.row_map.len()
111 }
112
113 pub fn n_reduced_vars(&self) -> usize {
114 self.n_full_vars() - self.fixed_vars.len()
115 }
116
117 pub fn n_reduced_rows(&self) -> usize {
118 self.n_full_rows() - self.dropped_rows.len()
119 }
120
121 pub fn project_x(&self, x_full: &[Number]) -> Vec<Number> {
124 assert_eq!(x_full.len(), self.n_full_vars());
125 self.var_map
126 .iter()
127 .zip(x_full.iter())
128 .filter_map(|(slot, &v)| slot.map(|_| v))
129 .collect()
130 }
131
132 pub fn lift_x(&self, x_reduced: &[Number]) -> Vec<Number> {
135 assert_eq!(x_reduced.len(), self.n_reduced_vars());
136 let mut out = vec![0.0; self.n_full_vars()];
137 for (i, slot) in self.var_map.iter().enumerate() {
138 if let Some(r) = slot {
139 out[i] = x_reduced[*r];
140 }
141 }
142 for (k, &i) in self.fixed_vars.iter().enumerate() {
143 out[i] = self.fixed_values[k];
144 }
145 out
146 }
147
148 pub fn project_lambda(&self, lambda_full: &[Number]) -> Vec<Number> {
150 assert_eq!(lambda_full.len(), self.n_full_rows());
151 self.row_map
152 .iter()
153 .zip(lambda_full.iter())
154 .filter_map(|(slot, &v)| slot.map(|_| v))
155 .collect()
156 }
157
158 pub fn lift_lambda(&self, lambda_reduced: &[Number]) -> Vec<Number> {
162 assert_eq!(lambda_reduced.len(), self.n_reduced_rows());
163 let mut out = vec![0.0; self.n_full_rows()];
164 for (i, slot) in self.row_map.iter().enumerate() {
165 if let Some(r) = slot {
166 out[i] = lambda_reduced[*r];
167 }
168 }
169 out
170 }
171
172 pub fn recover_dropped_multipliers(
206 &self,
207 grad_f: &[Number],
208 jac_full_row_major: &[Number],
209 lambda_full: &[Number],
210 ) -> Result<Vec<Number>, BlockSolveError> {
211 let n_vars = self.n_full_vars();
212 let n_rows = self.n_full_rows();
213 assert_eq!(
214 jac_full_row_major.len(),
215 n_rows * n_vars,
216 "jac_full_row_major length mismatch"
217 );
218 self.recover_core(grad_f, lambda_full, |row, col| {
222 jac_full_row_major[row * n_vars + col]
223 })
224 }
225
226 pub fn recover_dropped_multipliers_cols(
242 &self,
243 grad_f: &[Number],
244 jac_cols_row_major: &[Number],
245 n_cols: usize,
246 orig_to_compact: &[usize],
247 lambda_full: &[Number],
248 ) -> Result<Vec<Number>, BlockSolveError> {
249 let n_rows = self.n_full_rows();
250 assert_eq!(
251 jac_cols_row_major.len(),
252 n_rows * n_cols,
253 "jac_cols_row_major length mismatch"
254 );
255 assert_eq!(
256 orig_to_compact.len(),
257 self.n_full_vars(),
258 "orig_to_compact length mismatch"
259 );
260 self.recover_core(grad_f, lambda_full, |row, col| {
261 jac_cols_row_major[row * n_cols + orig_to_compact[col]]
262 })
263 }
264
265 fn recover_core(
270 &self,
271 grad_f: &[Number],
272 lambda_full: &[Number],
273 get: impl Fn(usize, usize) -> Number,
274 ) -> Result<Vec<Number>, BlockSolveError> {
275 let n_rows = self.n_full_rows();
276 let k = self.fixed_vars.len();
277 assert_eq!(grad_f.len(), self.n_full_vars(), "grad_f length mismatch");
278 assert_eq!(lambda_full.len(), n_rows, "lambda_full length mismatch");
279
280 if k == 0 {
281 return Ok(Vec::new());
282 }
283
284 let mut matrix = vec![0.0; k * k];
291 for (i_idx, &i) in self.fixed_vars.iter().enumerate() {
292 for (j_idx, &dr) in self.dropped_rows.iter().enumerate() {
293 matrix[i_idx * k + j_idx] = get(dr, i);
294 }
295 }
296
297 let mut rhs = vec![0.0; k];
298 for (i_idx, &i) in self.fixed_vars.iter().enumerate() {
299 let mut sum = 0.0;
300 for r in 0..n_rows {
301 if self.row_map[r].is_none() {
302 continue;
304 }
305 sum += get(r, i) * lambda_full[r];
306 }
307 rhs[i_idx] = grad_f[i] - sum;
308 }
309
310 let piv = lu_factor_partial_pivot(&mut matrix, k).map_err(|_| BlockSolveError::Singular)?;
311 lu_solve(&matrix, &piv, &mut rhs, k);
312 Ok(rhs)
313 }
314}
315
316#[derive(Debug, Default, Clone)]
320pub struct ReductionStack {
321 frames: Vec<ReductionFrame>,
322}
323
324impl ReductionStack {
325 pub fn is_empty(&self) -> bool {
327 self.frames.is_empty()
328 }
329
330 pub fn len(&self) -> usize {
332 self.frames.len()
333 }
334
335 pub fn push(&mut self, frame: ReductionFrame) {
337 self.frames.push(frame);
338 }
339
340 pub fn top(&self) -> Option<&ReductionFrame> {
342 self.frames.last()
343 }
344
345 pub fn iter_top_down(&self) -> impl Iterator<Item = &ReductionFrame> {
349 self.frames.iter().rev()
350 }
351
352 pub fn iter_bottom_up(&self) -> impl Iterator<Item = &ReductionFrame> {
356 self.frames.iter()
357 }
358}
359
360#[cfg(test)]
361mod tests {
362 use super::*;
363
364 #[test]
365 fn frame_new_builds_maps_correctly() {
366 let frame = ReductionFrame::new(4, 3, vec![1], vec![42.0], vec![0]);
368 assert_eq!(frame.var_map, vec![Some(0), None, Some(1), Some(2)]);
370 assert_eq!(frame.row_map, vec![None, Some(0), Some(1)]);
372 assert_eq!(frame.n_reduced_vars(), 3);
373 assert_eq!(frame.n_reduced_rows(), 2);
374 }
375
376 #[test]
377 fn frame_project_x_drops_fixed() {
378 let frame = ReductionFrame::new(3, 1, vec![1], vec![20.0], vec![0]);
379 let x_full = [10.0, 20.0, 30.0];
380 assert_eq!(frame.project_x(&x_full), vec![10.0, 30.0]);
381 }
382
383 #[test]
384 fn frame_lift_x_splices_fixed_values() {
385 let frame = ReductionFrame::new(3, 1, vec![1], vec![20.0], vec![0]);
386 let x_reduced = [10.0, 30.0];
387 assert_eq!(frame.lift_x(&x_reduced), vec![10.0, 20.0, 30.0]);
388 }
389
390 #[test]
391 fn frame_project_lift_x_roundtrip() {
392 let frame = ReductionFrame::new(4, 2, vec![0, 2], vec![1.0, 9.0], vec![0, 1]);
393 let x_full = [1.0, 5.0, 9.0, 7.0];
394 let reduced = frame.project_x(&x_full);
395 let lifted = frame.lift_x(&reduced);
396 assert_eq!(lifted, x_full);
397 }
398
399 #[test]
400 fn frame_project_lambda_drops_dropped() {
401 let frame = ReductionFrame::new(3, 3, vec![1], vec![20.0], vec![0]);
402 let lambda_full = [1.0, 2.0, 3.0];
403 assert_eq!(frame.project_lambda(&lambda_full), vec![2.0, 3.0]);
404 }
405
406 #[test]
407 fn frame_lift_lambda_zeros_dropped() {
408 let frame = ReductionFrame::new(3, 3, vec![1], vec![20.0], vec![0]);
409 let lambda_reduced = [2.0, 3.0];
410 assert_eq!(frame.lift_lambda(&lambda_reduced), vec![0.0, 2.0, 3.0]);
411 }
412
413 #[test]
414 fn recover_multipliers_singleton_linear() {
415 let frame = ReductionFrame::new(1, 1, vec![0], vec![3.0], vec![0]);
418 let lam = frame
419 .recover_dropped_multipliers(&[4.0], &[1.0], &[0.0])
420 .unwrap();
421 assert_eq!(lam.len(), 1);
422 assert!((lam[0] - 4.0).abs() < 1e-12);
423 }
424
425 #[test]
426 fn recover_multipliers_2x2_linear() {
427 let frame = ReductionFrame::new(2, 2, vec![0, 1], vec![1.0, 2.0], vec![0, 1]);
444 let jac = [1.0, 0.0, 1.0, 1.0]; let grad_f = [2.0, 5.0];
446 let lam = frame
447 .recover_dropped_multipliers(&grad_f, &jac, &[0.0, 0.0])
448 .unwrap();
449 assert!((lam[0] - (-3.0)).abs() < 1e-12, "λ0 was {}", lam[0]);
450 assert!((lam[1] - 5.0).abs() < 1e-12, "λ1 was {}", lam[1]);
451 }
452
453 #[test]
454 fn recover_multipliers_with_kept_rows() {
455 let frame = ReductionFrame::new(2, 2, vec![0], vec![1.0], vec![0]);
463 let jac = [2.0, 3.0, 4.0, 5.0];
464 let grad_f = [10.0, 0.0];
465 let lambda_full = [0.0, 0.5]; let lam = frame
467 .recover_dropped_multipliers(&grad_f, &jac, &lambda_full)
468 .unwrap();
469 assert_eq!(lam.len(), 1);
470 assert!((lam[0] - 4.0).abs() < 1e-12);
471 }
472
473 #[test]
474 fn recover_multipliers_singular_block_jacobian() {
475 let frame = ReductionFrame::new(2, 2, vec![0, 1], vec![0.0, 0.0], vec![0, 1]);
477 let jac = [1.0, 2.0, 2.0, 4.0]; let grad_f = [1.0, 2.0];
479 let err = frame
480 .recover_dropped_multipliers(&grad_f, &jac, &[0.0, 0.0])
481 .unwrap_err();
482 assert_eq!(err, BlockSolveError::Singular);
483 }
484
485 #[test]
486 fn recover_only_reads_fixed_var_columns() {
487 let frame = ReductionFrame::new(3, 3, vec![0, 2], vec![1.0, 2.0], vec![0, 1]);
496 let grad_f = [10.0, 4.0, 7.0];
497 let lambda_full = [0.0, 0.0, 0.5]; let clean = [
499 2.0, 1.0, 0.5, 1.0, -1.0, 3.0, 0.4, 1.0, 0.9, ];
503 let expected = frame
504 .recover_dropped_multipliers(&grad_f, &clean, &lambda_full)
505 .unwrap();
506
507 let mut poisoned = clean;
508 for r in 0..3 {
509 poisoned[r * 3 + 1] = Number::NAN; }
511 let got = frame
512 .recover_dropped_multipliers(&grad_f, &poisoned, &lambda_full)
513 .unwrap();
514
515 assert_eq!(got.len(), expected.len());
516 for (g, e) in got.iter().zip(expected.iter()) {
517 assert!(g.is_finite(), "recovered multiplier went NaN: {g}");
518 assert_eq!(g.to_bits(), e.to_bits(), "got {g}, expected {e}");
519 }
520 }
521
522 #[test]
523 fn recover_cols_matches_dense() {
524 let frame = ReductionFrame::new(3, 3, vec![0, 2], vec![1.0, 2.0], vec![0, 1]);
528 let grad_f = [10.0, 4.0, 7.0];
529 let lambda_full = [0.0, 0.0, 0.5];
530 let dense = [
531 2.0, 1.0, 0.5, 1.0, -1.0, 3.0, 0.4, 1.0, 0.9, ];
535 let dense_lam = frame
536 .recover_dropped_multipliers(&grad_f, &dense, &lambda_full)
537 .unwrap();
538
539 let needed = [0usize, 2];
541 let n_cols = needed.len();
542 let mut orig_to_compact = [usize::MAX; 3];
543 for (cc, &c) in needed.iter().enumerate() {
544 orig_to_compact[c] = cc;
545 }
546 let mut jac_cols = vec![0.0; 3 * n_cols];
547 for r in 0..3 {
548 for (cc, &c) in needed.iter().enumerate() {
549 jac_cols[r * n_cols + cc] = dense[r * 3 + c];
550 }
551 }
552 let cols_lam = frame
553 .recover_dropped_multipliers_cols(
554 &grad_f,
555 &jac_cols,
556 n_cols,
557 &orig_to_compact,
558 &lambda_full,
559 )
560 .unwrap();
561
562 assert_eq!(cols_lam.len(), dense_lam.len());
563 for (c, d) in cols_lam.iter().zip(dense_lam.iter()) {
564 assert_eq!(c.to_bits(), d.to_bits(), "cols {c} != dense {d}");
565 }
566 }
567
568 #[test]
569 fn recover_cols_empty_frame() {
570 let frame = ReductionFrame::new(2, 2, vec![], vec![], vec![]);
572 let lam = frame
573 .recover_dropped_multipliers_cols(&[0.0; 2], &[], 0, &[usize::MAX; 2], &[0.0; 2])
574 .unwrap();
575 assert!(lam.is_empty());
576 }
577
578 #[test]
579 fn recover_multipliers_empty_frame() {
580 let frame = ReductionFrame::new(2, 2, vec![], vec![], vec![]);
581 let lam = frame
582 .recover_dropped_multipliers(&[0.0; 2], &[0.0; 4], &[0.0; 2])
583 .unwrap();
584 assert!(lam.is_empty());
585 }
586
587 #[test]
588 fn kkt_residual_after_recovery_to_1e_minus_12() {
589 let frame = ReductionFrame::new(3, 3, vec![0, 1], vec![2.0 / 3.0, 5.0 / 3.0], vec![0, 1]);
603 let jac = [
605 2.0, 1.0, 0.0, 1.0, -1.0, 0.0, 1.0, 1.0, 1.0, ];
609 let y_star = 8.0 / 3.0;
611 let grad_f = [10.0, 4.0, 2.0 * y_star];
612 let lambda_kept_2 = 2.0 * y_star;
615 let lambda_full = [0.0, 0.0, lambda_kept_2];
616
617 let lam_dropped = frame
618 .recover_dropped_multipliers(&grad_f, &jac, &lambda_full)
619 .unwrap();
620 let mut lambda_recovered = lambda_full;
622 for (k, &r) in frame.dropped_rows.iter().enumerate() {
623 lambda_recovered[r] = lam_dropped[k];
624 }
625 for &i in &frame.fixed_vars {
627 let mut s = grad_f[i];
628 for r in 0..3 {
629 s -= jac[r * 3 + i] * lambda_recovered[r];
630 }
631 assert!(s.abs() < 1e-12, "stationarity at var {i} = {s}");
632 }
633 }
634
635 struct FuzzRng(u64);
641 impl FuzzRng {
642 fn new(seed: u64) -> Self {
643 Self(seed)
644 }
645 fn next_u64(&mut self) -> u64 {
646 self.0 = self
647 .0
648 .wrapping_mul(6364136223846793005)
649 .wrapping_add(1442695040888963407);
650 self.0 >> 32
651 }
652 fn unit(&mut self) -> Number {
653 let raw = (self.next_u64() & 0x3fff_ffff) as Number;
654 raw / (1u64 << 29) as Number - 1.0
655 }
656 }
657
658 #[test]
659 fn frame_fuzz_recover_reproduces_synthetic_lambda() {
660 let mut rng = FuzzRng::new(0xface_b00c_baad_f00d);
661
662 for trial in 0..30 {
663 let n_vars = 2 + (rng.next_u64() % 3) as usize; let n_rows = n_vars;
665 let k = 1 + (rng.next_u64() % n_vars as u64) as usize;
666
667 let mut perm_v: Vec<usize> = (0..n_vars).collect();
668 for i in (1..n_vars).rev() {
669 let j = (rng.next_u64() as usize) % (i + 1);
670 perm_v.swap(i, j);
671 }
672 let mut fixed_vars: Vec<usize> = perm_v[..k].to_vec();
673 fixed_vars.sort_unstable();
674
675 let mut perm_r: Vec<usize> = (0..n_rows).collect();
676 for i in (1..n_rows).rev() {
677 let j = (rng.next_u64() as usize) % (i + 1);
678 perm_r.swap(i, j);
679 }
680 let mut dropped_rows: Vec<usize> = perm_r[..k].to_vec();
681 dropped_rows.sort_unstable();
682
683 let mut jac = vec![0.0; n_rows * n_vars];
684 for r in 0..n_rows {
685 for c in 0..n_vars {
686 jac[r * n_vars + c] = 0.2 * rng.unit();
687 }
688 }
689 for (&r, &c) in dropped_rows.iter().zip(fixed_vars.iter()) {
690 jac[r * n_vars + c] += 2.5;
691 }
692
693 let lambda_star: Vec<Number> = (0..n_rows).map(|_| rng.unit()).collect();
694 let mut grad_f = vec![0.0; n_vars];
695 let fixed_set: std::collections::BTreeSet<usize> = fixed_vars.iter().copied().collect();
696 for i in 0..n_vars {
697 if fixed_set.contains(&i) {
698 let mut s = 0.0;
699 for r in 0..n_rows {
700 s += jac[r * n_vars + i] * lambda_star[r];
701 }
702 grad_f[i] = s;
703 } else {
704 grad_f[i] = rng.unit();
705 }
706 }
707
708 let dropped_set: std::collections::BTreeSet<usize> =
709 dropped_rows.iter().copied().collect();
710 let mut lambda_given = vec![0.0; n_rows];
711 for r in 0..n_rows {
712 if !dropped_set.contains(&r) {
713 lambda_given[r] = lambda_star[r];
714 }
715 }
716
717 let frame = ReductionFrame::new(
718 n_vars,
719 n_rows,
720 fixed_vars.clone(),
721 vec![0.0; k],
722 dropped_rows.clone(),
723 );
724
725 let lam_dropped = frame
726 .recover_dropped_multipliers(&grad_f, &jac, &lambda_given)
727 .unwrap_or_else(|e| panic!("trial {trial}: {e:?}"));
728
729 for (idx, &r) in dropped_rows.iter().enumerate() {
730 let expected = lambda_star[r];
731 let got = lam_dropped[idx];
732 assert!(
733 (expected - got).abs() < 1e-10,
734 "trial {trial}: λ[{r}] expected {expected:.6}, got {got:.6}"
735 );
736 }
737 }
738 }
739
740 #[test]
741 fn reduction_stack_push_top_iter() {
742 let mut stack = ReductionStack::default();
743 assert!(stack.is_empty());
744 let f1 = ReductionFrame::new(2, 2, vec![0], vec![1.0], vec![0]);
745 let f2 = ReductionFrame::new(2, 2, vec![1], vec![2.0], vec![1]);
746 stack.push(f1.clone());
747 stack.push(f2.clone());
748 assert_eq!(stack.len(), 2);
749 let top = stack.top().expect("non-empty");
750 assert_eq!(top.fixed_vars, f2.fixed_vars);
751 let order: Vec<_> = stack.iter_top_down().map(|f| f.fixed_vars[0]).collect();
753 assert_eq!(order, vec![1, 0]);
754 let order_up: Vec<_> = stack.iter_bottom_up().map(|f| f.fixed_vars[0]).collect();
755 assert_eq!(order_up, vec![0, 1]);
756 }
757
758 #[test]
764 fn frame_project_lift_lambda_roundtrip() {
765 let frame = ReductionFrame::new(4, 3, vec![0, 2], vec![1.0, 9.0], vec![0, 1]);
766 let lambda_full = [4.0, 5.0, 6.0];
768 let reduced = frame.project_lambda(&lambda_full);
769 assert_eq!(reduced, vec![6.0]);
771 let lifted = frame.lift_lambda(&reduced);
773 assert_eq!(lifted, vec![0.0, 0.0, 6.0]);
774 let reduced_again = frame.project_lambda(&lifted);
777 assert_eq!(reduced_again, reduced);
778 }
779
780 #[test]
785 fn reduction_stack_multi_frame_roundtrip() {
786 let f1 = ReductionFrame::new(4, 4, vec![0], vec![10.0], vec![0]);
790 let f2 = ReductionFrame::new(4, 4, vec![2], vec![30.0], vec![2]);
791 let mut stack = ReductionStack::default();
792 stack.push(f1.clone());
793 stack.push(f2.clone());
794
795 let x_full_expected = vec![10.0, 7.0, 30.0, 5.0];
798 let lambda_full_expected = vec![0.0, 8.0, 0.0, 6.0];
799
800 for frame in stack.iter_top_down() {
810 let reduced_x = frame.project_x(&x_full_expected);
811 let lifted_x = frame.lift_x(&reduced_x);
812 assert_eq!(lifted_x, x_full_expected);
813 let reduced_l = frame.project_lambda(&lambda_full_expected);
814 let lifted_l = frame.lift_lambda(&reduced_l);
815 for r in 0..4 {
818 if frame.row_map[r].is_some() {
819 assert_eq!(lifted_l[r], lambda_full_expected[r]);
820 } else {
821 assert_eq!(lifted_l[r], 0.0);
822 }
823 }
824 }
825 }
826}