1use crate::kkt::aug_system_solver::{AugSysCoeffs, AugSysRhs, AugSysSol, AugSystemSolver};
30use pounce_common::diagnostics::{DiagCategory, DiagnosticsState};
31use pounce_common::timing::TimingStatistics;
32use pounce_common::types::{Index, Number};
33use pounce_linalg::Vector;
34use pounce_linalg::compound_vector::CompoundVector;
35use pounce_linalg::dense_vector::DenseVector;
36use pounce_linalg::diag_matrix::DiagMatrix;
37use pounce_linalg::triplet::{GenTMatrix, SymTMatrix};
38use pounce_linsol::{ESymSolverStatus, FactorPattern, SymLinearSolver, TSymLinearSolver};
39use std::ops::Range;
40use std::rc::Rc;
41
42pub struct StdAugSystemSolver {
44 linsol: TSymLinearSolver,
45
46 initialized: bool,
48 struct_sig: Option<(usize, usize, usize, Index, Index, Index)>,
57 n_x: Index,
58 n_s: Index,
59 n_c: Index,
60 n_d: Index,
61 dim: Index,
63
64 irn: Vec<Index>,
66 jcn: Vec<Index>,
68 vals: Vec<Number>,
70
71 w_range: Range<usize>,
73 dx_range: Range<usize>,
74 ds_range: Range<usize>,
75 jc_range: Range<usize>,
76 dc_range: Range<usize>,
77 jd_range: Range<usize>,
78 minus_i_range: Range<usize>,
79 dd_range: Range<usize>,
80
81 last_neg_evals: Index,
82 last_status: Option<ESymSolverStatus>,
83
84 have_factor: bool,
88
89 timing: Option<Rc<TimingStatistics>>,
93
94 diagnostics: Option<Rc<DiagnosticsState>>,
100
101 legacy_dump_path: std::cell::OnceCell<Option<String>>,
109}
110
111impl std::fmt::Debug for StdAugSystemSolver {
112 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
113 f.debug_struct("StdAugSystemSolver")
114 .field("dim", &self.dim)
115 .field("nnz", &self.vals.len())
116 .field("initialized", &self.initialized)
117 .field("last_neg_evals", &self.last_neg_evals)
118 .field("last_status", &self.last_status)
119 .finish_non_exhaustive()
120 }
121}
122
123impl StdAugSystemSolver {
124 pub fn new(linsol: TSymLinearSolver) -> Self {
126 Self {
127 linsol,
128 initialized: false,
129 struct_sig: None,
130 legacy_dump_path: std::cell::OnceCell::new(),
131 n_x: 0,
132 n_s: 0,
133 n_c: 0,
134 n_d: 0,
135 dim: 0,
136 irn: Vec::new(),
137 jcn: Vec::new(),
138 vals: Vec::new(),
139 w_range: 0..0,
140 dx_range: 0..0,
141 ds_range: 0..0,
142 jc_range: 0..0,
143 dc_range: 0..0,
144 jd_range: 0..0,
145 minus_i_range: 0..0,
146 dd_range: 0..0,
147 last_neg_evals: 0,
148 last_status: None,
149 have_factor: false,
150 timing: None,
151 diagnostics: None,
152 }
153 }
154
155 fn build_structure(&mut self, coeffs: &AugSysCoeffs<'_>) -> ESymSolverStatus {
156 let n_x = coeffs.j_c.n_cols();
157 let n_c = coeffs.j_c.n_rows();
158 let n_d = coeffs.j_d.n_rows();
159 debug_assert_eq!(coeffs.j_d.n_cols(), n_x);
160 let n_s = n_d;
161
162 let w_nnz = match coeffs.w {
163 None => 0_usize,
164 Some(w) => w_nonzeros(w),
165 };
166 let jc_nnz = gen_t_downcast(coeffs.j_c).nonzeros() as usize;
167 let jd_nnz = gen_t_downcast(coeffs.j_d).nonzeros() as usize;
168
169 let total = w_nnz
170 + (n_x as usize) + (n_s as usize) + jc_nnz
173 + (n_c as usize) + jd_nnz
175 + (n_s as usize) + (n_d as usize); self.irn = Vec::with_capacity(total);
179 self.jcn = Vec::with_capacity(total);
180 self.vals = vec![0.0; total];
181
182 let w_start = self.irn.len();
184 if let Some(w) = coeffs.w {
185 if let Some(t) = w.as_any().downcast_ref::<SymTMatrix>() {
186 self.irn.extend_from_slice(t.irows());
187 self.jcn.extend_from_slice(t.jcols());
188 } else if let Some(dm) = w.as_any().downcast_ref::<DiagMatrix>() {
189 let n = w_diag_dim(dm);
192 for i in 0..n {
193 self.irn.push(i + 1);
194 self.jcn.push(i + 1);
195 }
196 } else {
197 unreachable!("StdAugSystemSolver: W must be a SymTMatrix or DiagMatrix in v1.0")
198 }
199 }
200 self.w_range = w_start..self.irn.len();
201
202 let dx_start = self.irn.len();
204 for i in 0..n_x {
205 self.irn.push(i + 1);
206 self.jcn.push(i + 1);
207 }
208 self.dx_range = dx_start..self.irn.len();
209
210 let ds_start = self.irn.len();
212 for i in 0..n_s {
213 let r = n_x + i + 1;
214 self.irn.push(r);
215 self.jcn.push(r);
216 }
217 self.ds_range = ds_start..self.irn.len();
218
219 let jc_start = self.irn.len();
221 let j_c = gen_t_downcast(coeffs.j_c);
222 let row_off_c = n_x + n_s;
223 for (&i, &j) in j_c.irows().iter().zip(j_c.jcols().iter()) {
224 self.irn.push(row_off_c + i);
227 self.jcn.push(j);
228 }
229 self.jc_range = jc_start..self.irn.len();
230
231 let dc_start = self.irn.len();
233 for i in 0..n_c {
234 let r = n_x + n_s + i + 1;
235 self.irn.push(r);
236 self.jcn.push(r);
237 }
238 self.dc_range = dc_start..self.irn.len();
239
240 let jd_start = self.irn.len();
242 let j_d = gen_t_downcast(coeffs.j_d);
243 let row_off_d = n_x + n_s + n_c;
244 for (&i, &j) in j_d.irows().iter().zip(j_d.jcols().iter()) {
245 self.irn.push(row_off_d + i);
246 self.jcn.push(j);
247 }
248 self.jd_range = jd_start..self.irn.len();
249
250 let mi_start = self.irn.len();
252 for i in 0..n_s {
253 self.irn.push(n_x + n_s + n_c + i + 1);
254 self.jcn.push(n_x + i + 1);
255 }
256 self.minus_i_range = mi_start..self.irn.len();
257
258 let dd_start = self.irn.len();
260 for i in 0..n_d {
261 let r = n_x + n_s + n_c + i + 1;
262 self.irn.push(r);
263 self.jcn.push(r);
264 }
265 self.dd_range = dd_start..self.irn.len();
266
267 debug_assert_eq!(self.irn.len(), total);
268 debug_assert_eq!(self.jcn.len(), total);
269
270 self.n_x = n_x;
271 self.n_s = n_s;
272 self.n_c = n_c;
273 self.n_d = n_d;
274 self.dim = n_x + n_s + n_c + n_d;
275
276 let status = self
277 .linsol
278 .initialize_structure(self.dim, &self.irn, &self.jcn);
279 if status == ESymSolverStatus::Success {
280 self.initialized = true;
281 }
282 status
283 }
284
285 fn refill_values(&mut self, coeffs: &AugSysCoeffs<'_>) {
286 if !self.w_range.is_empty() {
288 let Some(w_dyn) = coeffs.w else {
289 unreachable!("structure pinned with W; W cannot be None now")
290 };
291 let dst = &mut self.vals[self.w_range.clone()];
292 if let Some(t) = w_dyn.as_any().downcast_ref::<SymTMatrix>() {
293 for (d, &v) in dst.iter_mut().zip(t.values().iter()) {
294 *d = coeffs.w_factor * v;
295 }
296 } else if let Some(dm) = w_dyn.as_any().downcast_ref::<DiagMatrix>() {
297 let diag = w_diag_values(dm);
298 for (d, &v) in dst.iter_mut().zip(diag.iter()) {
299 *d = coeffs.w_factor * v;
300 }
301 } else {
302 unreachable!("StdAugSystemSolver: W must be a SymTMatrix or DiagMatrix in v1.0")
303 }
304 }
305 fill_diag(
307 &mut self.vals[self.dx_range.clone()],
308 coeffs.d_x,
309 coeffs.delta_x,
310 1.0,
311 );
312 fill_diag(
314 &mut self.vals[self.ds_range.clone()],
315 coeffs.d_s,
316 coeffs.delta_s,
317 1.0,
318 );
319 {
321 let j_c = gen_t_downcast(coeffs.j_c);
322 self.vals[self.jc_range.clone()].copy_from_slice(j_c.values());
323 }
324 fill_diag(
326 &mut self.vals[self.dc_range.clone()],
327 coeffs.d_c,
328 coeffs.delta_c,
329 -1.0,
330 );
331 {
333 let j_d = gen_t_downcast(coeffs.j_d);
334 self.vals[self.jd_range.clone()].copy_from_slice(j_d.values());
335 }
336 for v in &mut self.vals[self.minus_i_range.clone()] {
338 *v = -1.0;
339 }
340 fill_diag(
342 &mut self.vals[self.dd_range.clone()],
343 coeffs.d_d,
344 coeffs.delta_d,
345 -1.0,
346 );
347 }
348
349 pub(crate) fn assemble(&mut self, coeffs: &AugSysCoeffs<'_>) -> ESymSolverStatus {
358 let sig = {
359 let w_nnz = coeffs.w.map(w_nonzeros).unwrap_or(0);
360 let jc_nnz = gen_t_downcast(coeffs.j_c).nonzeros() as usize;
361 let jd_nnz = gen_t_downcast(coeffs.j_d).nonzeros() as usize;
362 (
363 w_nnz,
364 jc_nnz,
365 jd_nnz,
366 coeffs.j_c.n_cols(),
367 coeffs.j_c.n_rows(),
368 coeffs.j_d.n_rows(),
369 )
370 };
371 if !self.initialized || self.struct_sig != Some(sig) {
372 let timing = self.timing.clone();
384 let _sym_guard = timing
385 .as_deref()
386 .map(|t| t.linear_system_symbolic_factorization.guard());
387 let s = self.build_structure(coeffs);
388 if s != ESymSolverStatus::Success {
389 self.last_status = Some(s);
390 return s;
391 }
392 self.struct_sig = Some(sig);
393 }
394 self.refill_values(coeffs);
395 ESymSolverStatus::Success
396 }
397
398 pub(crate) fn assembled_dim(&self) -> Index {
400 self.dim
401 }
402 pub(crate) fn assembled_triplet(&self) -> (&[Index], &[Index], &[Number]) {
405 (&self.irn, &self.jcn, &self.vals)
406 }
407
408 fn legacy_dump_path(&self) -> Option<&String> {
411 self.legacy_dump_path
412 .get_or_init(|| std::env::var("POUNCE_DUMP_KKT").ok())
413 .as_ref()
414 }
415
416 fn kkt_dump_active(&self) -> bool {
420 self.diagnostics
421 .as_deref()
422 .is_some_and(|d| d.want(DiagCategory::Kkt))
423 || self.legacy_dump_path().is_some()
424 }
425
426 pub(crate) fn pack_rhs(&self, rhs: &AugSysRhs<'_>, packed: &mut [Number]) {
427 let n_x = self.n_x as usize;
428 let n_s = self.n_s as usize;
429 let n_c = self.n_c as usize;
430 let n_d = self.n_d as usize;
431 copy_vec(rhs.rhs_x, &mut packed[..n_x]);
432 copy_vec(rhs.rhs_s, &mut packed[n_x..n_x + n_s]);
433 copy_vec(rhs.rhs_c, &mut packed[n_x + n_s..n_x + n_s + n_c]);
434 copy_vec(
435 rhs.rhs_d,
436 &mut packed[n_x + n_s + n_c..n_x + n_s + n_c + n_d],
437 );
438 }
439
440 pub(crate) fn unpack_sol(&self, packed: &[Number], sol: &mut AugSysSol<'_>) {
441 let n_x = self.n_x as usize;
442 let n_s = self.n_s as usize;
443 let n_c = self.n_c as usize;
444 let n_d = self.n_d as usize;
445 write_vec(sol.sol_x, &packed[..n_x]);
446 write_vec(sol.sol_s, &packed[n_x..n_x + n_s]);
447 write_vec(sol.sol_c, &packed[n_x + n_s..n_x + n_s + n_c]);
448 write_vec(sol.sol_d, &packed[n_x + n_s + n_c..n_x + n_s + n_c + n_d]);
449 }
450}
451
452impl AugSystemSolver for StdAugSystemSolver {
453 fn provides_inertia(&self) -> bool {
454 self.linsol.provides_inertia()
455 }
456
457 fn number_of_neg_evals(&self) -> Index {
458 self.last_neg_evals
459 }
460
461 fn system_dim(&self) -> Index {
462 self.dim
463 }
464
465 fn kkt_triplets(&self) -> Option<(Index, Vec<Index>, Vec<Index>, Vec<Number>)> {
466 if self.irn.is_empty() {
467 return None;
468 }
469 Some((
470 self.dim,
471 self.irn.clone(),
472 self.jcn.clone(),
473 self.vals.clone(),
474 ))
475 }
476
477 fn l_factor(&self, want_values: bool) -> Option<FactorPattern> {
478 self.linsol.factor_pattern(want_values)
479 }
480
481 fn increase_quality(&mut self) -> bool {
482 self.have_factor = false;
486 self.linsol.increase_quality()
487 }
488
489 fn last_solve_status(&self) -> ESymSolverStatus {
490 self.last_status.unwrap_or(ESymSolverStatus::FatalError)
491 }
492
493 fn solve(
494 &mut self,
495 coeffs: &AugSysCoeffs<'_>,
496 rhs: &AugSysRhs<'_>,
497 sol: &mut AugSysSol<'_>,
498 check_neg_evals: bool,
499 num_neg_evals: Index,
500 ) -> ESymSolverStatus {
501 let s = self.assemble(coeffs);
509 if s != ESymSolverStatus::Success {
510 return s;
511 }
512
513 let mut packed = vec![0.0; self.dim as usize];
514 self.pack_rhs(rhs, &mut packed);
515
516 let dump_rhs = packed.clone();
517
518 let _factor_guard = self
522 .timing
523 .as_deref()
524 .map(|t| t.linear_system_factorization.guard());
525 let status = self.linsol.multi_solve(
526 &self.vals,
527 true,
528 1,
529 &mut packed,
530 check_neg_evals,
531 num_neg_evals,
532 );
533 drop(_factor_guard);
534 self.last_status = Some(status);
535 if self.linsol.provides_inertia()
547 && matches!(
548 status,
549 ESymSolverStatus::Success
550 | ESymSolverStatus::WrongInertia
551 | ESymSolverStatus::Singular
552 )
553 {
554 self.last_neg_evals = self.linsol.number_of_neg_evals();
555 }
556 if status == ESymSolverStatus::Success {
557 self.unpack_sol(&packed, sol);
558 self.have_factor = true;
559 }
560
561 if let Some(diag) = self.diagnostics.clone() {
566 if diag.want(DiagCategory::Kkt) {
567 let solve_idx = diag.next_solve_index();
568 let filename = format!("kkt_solve_{solve_idx:03}.jsonl");
569 let variant = diag.config.kkt_variant;
576 let factor_pattern =
577 if status == ESymSolverStatus::Success && variant.wants_l_pattern() {
578 self.linsol.factor_pattern(variant.wants_l_values())
579 } else {
580 None
581 };
582 if let Some(mut w) = diag.open_writer(&filename) {
583 let _ = write_kkt_record(
584 &mut w,
585 self.dim,
586 &self.irn,
587 &self.jcn,
588 &self.vals,
589 &dump_rhs,
590 &packed,
591 check_neg_evals,
592 num_neg_evals,
593 status,
594 self.last_neg_evals,
595 factor_pattern.as_ref(),
596 );
597 }
598 }
599 }
600 if let Some(path) = self.legacy_dump_path().cloned() {
601 use std::sync::atomic::{AtomicBool, Ordering};
602 static WARNED: AtomicBool = AtomicBool::new(false);
603 if !WARNED.swap(true, Ordering::SeqCst) {
604 tracing::warn!(target: "pounce::linsol",
605 "warning: POUNCE_DUMP_KKT is deprecated; prefer `--dump kkt:<iter-spec>` (see pounce --help)"
606 );
607 }
608 dump_kkt(
609 &path,
610 self.dim,
611 &self.irn,
612 &self.jcn,
613 &self.vals,
614 &dump_rhs,
615 &packed,
616 check_neg_evals,
617 num_neg_evals,
618 status,
619 self.last_neg_evals,
620 );
621 }
622
623 status
624 }
625
626 fn resolve(
627 &mut self,
628 coeffs: &AugSysCoeffs<'_>,
629 rhs: &AugSysRhs<'_>,
630 sol: &mut AugSysSol<'_>,
631 ) -> ESymSolverStatus {
632 if !self.have_factor {
639 return self.solve(coeffs, rhs, sol, false, 0);
640 }
641
642 let mut packed = vec![0.0; self.dim as usize];
643 self.pack_rhs(rhs, &mut packed);
644
645 let _back_guard = self
648 .timing
649 .as_deref()
650 .map(|t| t.linear_system_back_solve.guard());
651 let status = self
652 .linsol
653 .multi_solve(&self.vals, false, 1, &mut packed, false, 0);
654 drop(_back_guard);
655 self.last_status = Some(status);
656 if status == ESymSolverStatus::Success {
657 self.unpack_sol(&packed, sol);
658 }
659 status
660 }
661
662 fn set_diagnostics(&mut self, diag: Rc<DiagnosticsState>) {
663 self.diagnostics = Some(diag);
664 }
665
666 fn set_slack_scaling(&mut self, nx: Index, s_scale: &[Number]) {
667 self.linsol.set_slack_scaling(nx, s_scale);
668 }
669
670 fn set_timing_stats(&mut self, timing: Rc<TimingStatistics>) {
671 self.timing = Some(timing);
672 }
673
674 fn multi_solve_matches_single_solve(&self, nrhs: usize) -> bool {
675 self.linsol.multi_solve_matches_single_solve(nrhs)
676 }
677
678 fn try_solve_many_flat(
679 &mut self,
680 coeffs: &AugSysCoeffs<'_>,
681 packed_rhs: &mut [Number],
682 nrhs: usize,
683 check_neg_evals: bool,
684 num_neg_evals: Index,
685 ) -> Option<ESymSolverStatus> {
686 if self.kkt_dump_active() {
690 return None;
691 }
692 let s = self.assemble(coeffs);
693 if s != ESymSolverStatus::Success {
694 self.last_status = Some(s);
695 return Some(s);
696 }
697 if packed_rhs.len() != (self.dim as usize) * nrhs {
703 return None;
704 }
705
706 let _factor_guard = self
713 .timing
714 .as_deref()
715 .map(|t| t.linear_system_factorization.guard());
716 let status = self.linsol.multi_solve(
717 &self.vals,
718 true,
719 nrhs as Index,
720 packed_rhs,
721 check_neg_evals,
722 num_neg_evals,
723 );
724 drop(_factor_guard);
725 self.last_status = Some(status);
726 if self.linsol.provides_inertia()
729 && matches!(
730 status,
731 ESymSolverStatus::Success
732 | ESymSolverStatus::WrongInertia
733 | ESymSolverStatus::Singular
734 )
735 {
736 self.last_neg_evals = self.linsol.number_of_neg_evals();
737 }
738 if status == ESymSolverStatus::Success {
739 self.have_factor = true;
740 }
741 Some(status)
742 }
743
744 fn try_resolve_many_flat(
745 &mut self,
746 _coeffs: &AugSysCoeffs<'_>,
747 packed_rhs: &mut [Number],
748 nrhs: usize,
749 ) -> Option<ESymSolverStatus> {
750 if !self.have_factor {
755 return None;
756 }
757 if packed_rhs.len() != (self.dim as usize) * nrhs {
758 return Some(ESymSolverStatus::FatalError);
759 }
760 let _back_guard = self
761 .timing
762 .as_deref()
763 .map(|t| t.linear_system_back_solve.guard());
764 let status =
765 self.linsol
766 .multi_solve(&self.vals, false, nrhs as Index, packed_rhs, false, 0);
767 drop(_back_guard);
768 self.last_status = Some(status);
769 Some(status)
770 }
771}
772
773#[allow(clippy::too_many_arguments)]
776fn write_kkt_record(
781 w: &mut dyn std::io::Write,
782 dim: Index,
783 irn: &[Index],
784 jcn: &[Index],
785 vals: &[Number],
786 rhs: &[Number],
787 sol: &[Number],
788 check_neg_evals: bool,
789 num_neg_evals: Index,
790 status: ESymSolverStatus,
791 last_neg_evals: Index,
792 factor_pattern: Option<&FactorPattern>,
793) -> std::io::Result<()> {
794 use std::fmt::Write as _;
795
796 let mut line = String::with_capacity(64 * vals.len());
797 line.push('{');
798 let _ = write!(line, "\"n\":{dim},");
799 let _ = write!(line, "\"check_neg_evals\":{check_neg_evals},");
800 let _ = write!(line, "\"num_neg_evals_expected\":{num_neg_evals},");
801 let _ = write!(line, "\"num_neg_evals_actual\":{last_neg_evals},");
802 let _ = write!(line, "\"status\":\"{status:?}\",");
803
804 line.push_str("\"irn\":[");
805 for (i, v) in irn.iter().enumerate() {
806 if i > 0 {
807 line.push(',');
808 }
809 let _ = write!(line, "{v}");
810 }
811 line.push_str("],\"jcn\":[");
812 for (i, v) in jcn.iter().enumerate() {
813 if i > 0 {
814 line.push(',');
815 }
816 let _ = write!(line, "{v}");
817 }
818 line.push_str("],\"vals\":[");
819 for (i, v) in vals.iter().enumerate() {
820 if i > 0 {
821 line.push(',');
822 }
823 let _ = write!(line, "{v:.17e}");
824 }
825 line.push_str("],\"rhs\":[");
826 for (i, v) in rhs.iter().enumerate() {
827 if i > 0 {
828 line.push(',');
829 }
830 let _ = write!(line, "{v:.17e}");
831 }
832 line.push_str("],\"sol\":[");
833 for (i, v) in sol.iter().enumerate() {
834 if i > 0 {
835 line.push(',');
836 }
837 let _ = write!(line, "{v:.17e}");
838 }
839 line.push(']');
840
841 if let Some(fp) = factor_pattern {
846 line.push_str(",\"L_irn\":[");
847 for (i, v) in fp.l_irn.iter().enumerate() {
848 if i > 0 {
849 line.push(',');
850 }
851 let _ = write!(line, "{v}");
852 }
853 line.push_str("],\"L_jcn\":[");
854 for (i, v) in fp.l_jcn.iter().enumerate() {
855 if i > 0 {
856 line.push(',');
857 }
858 let _ = write!(line, "{v}");
859 }
860 line.push_str("],\"perm\":[");
861 for (i, v) in fp.perm.iter().enumerate() {
862 if i > 0 {
863 line.push(',');
864 }
865 let _ = write!(line, "{v}");
866 }
867 line.push(']');
868 if let Some(vals) = fp.l_vals.as_ref() {
869 line.push_str(",\"L_vals\":[");
870 for (i, v) in vals.iter().enumerate() {
871 if i > 0 {
872 line.push(',');
873 }
874 let _ = write!(line, "{v:.17e}");
875 }
876 line.push(']');
877 }
878 }
879
880 line.push_str("}\n");
881
882 w.write_all(line.as_bytes())
883}
884
885fn dump_kkt(
886 path: &str,
887 dim: Index,
888 irn: &[Index],
889 jcn: &[Index],
890 vals: &[Number],
891 rhs: &[Number],
892 sol: &[Number],
893 check_neg_evals: bool,
894 num_neg_evals: Index,
895 status: ESymSolverStatus,
896 last_neg_evals: Index,
897) {
898 if let Ok(mut f) = std::fs::OpenOptions::new()
899 .create(true)
900 .append(true)
901 .open(path)
902 {
903 let _ = write_kkt_record(
904 &mut f,
905 dim,
906 irn,
907 jcn,
908 vals,
909 rhs,
910 sol,
911 check_neg_evals,
912 num_neg_evals,
913 status,
914 last_neg_evals,
915 None, );
917 }
918}
919
920fn w_nonzeros(w: &dyn pounce_linalg::SymMatrix) -> usize {
925 if let Some(t) = w.as_any().downcast_ref::<SymTMatrix>() {
926 t.nonzeros() as usize
927 } else if let Some(dm) = w.as_any().downcast_ref::<DiagMatrix>() {
928 w_diag_dim(dm) as usize
929 } else {
930 unreachable!("StdAugSystemSolver: W must be a SymTMatrix or DiagMatrix in v1.0")
931 }
932}
933
934fn w_diag_dim(dm: &DiagMatrix) -> Index {
935 dm.get_diag()
936 .expect("DiagMatrix W has no diagonal set")
937 .dim()
938}
939
940fn w_diag_values(dm: &DiagMatrix) -> Vec<Number> {
941 let diag = dm.get_diag().expect("DiagMatrix W has no diagonal set");
942 diag.as_any()
943 .downcast_ref::<DenseVector>()
944 .expect("StdAugSystemSolver: DiagMatrix W diagonal must be DenseVector in v1.0")
945 .expanded_values()
946}
947
948fn gen_t_downcast(m: &dyn pounce_linalg::Matrix) -> &GenTMatrix {
949 let Some(t) = m.as_any().downcast_ref::<GenTMatrix>() else {
950 unreachable!("StdAugSystemSolver: J_c / J_d must be GenTMatrix in v1.0")
951 };
952 t
953}
954
955fn flat_read(v: &dyn Vector) -> Vec<Number> {
961 if let Some(dv) = v.as_any().downcast_ref::<DenseVector>() {
962 return dv.expanded_values();
963 }
964 if let Some(cv) = v.as_any().downcast_ref::<CompoundVector>() {
965 let mut out = Vec::with_capacity(cv.dim() as usize);
966 for k in 0..cv.n_comps() {
967 let blk = cv.comp(k);
968 let dblk = blk
969 .as_any()
970 .downcast_ref::<DenseVector>()
971 .expect("StdAugSystemSolver: CompoundVector blocks must be DenseVectors");
972 out.extend_from_slice(&dblk.expanded_values());
973 }
974 return out;
975 }
976 unreachable!(
977 "StdAugSystemSolver: D_*/rhs/sol must be DenseVector or CompoundVector of DenseVectors in v1.0"
978 )
979}
980
981fn flat_write(dst: &mut dyn Vector, src: &[Number]) {
983 if let Some(dv) = dst.as_any_mut().downcast_mut::<DenseVector>() {
984 dv.set_values(src);
985 return;
986 }
987 if let Some(cv) = dst.as_any_mut().downcast_mut::<CompoundVector>() {
988 let mut off = 0usize;
989 for k in 0..cv.n_comps() {
990 let blk = cv.comp_mut(k);
991 let dim = blk.dim() as usize;
992 let dblk = blk
993 .as_any_mut()
994 .downcast_mut::<DenseVector>()
995 .expect("StdAugSystemSolver: CompoundVector blocks must be DenseVectors");
996 dblk.set_values(&src[off..off + dim]);
997 off += dim;
998 }
999 return;
1000 }
1001 unreachable!(
1002 "StdAugSystemSolver: sol must be DenseVector or CompoundVector of DenseVectors in v1.0"
1003 )
1004}
1005
1006fn fill_diag(dst: &mut [Number], d: Option<&dyn Vector>, delta: Number, sign: Number) {
1009 match d {
1010 None => {
1011 for v in dst.iter_mut() {
1012 *v = sign * delta;
1013 }
1014 }
1015 Some(d) => {
1016 let xs = flat_read(d);
1017 debug_assert_eq!(xs.len(), dst.len());
1018 for (out, &x) in dst.iter_mut().zip(xs.iter()) {
1019 *out = sign * (x + delta);
1020 }
1021 }
1022 }
1023}
1024
1025fn copy_vec(src: &dyn Vector, dst: &mut [Number]) {
1026 let xs = flat_read(src);
1027 debug_assert_eq!(xs.len(), dst.len());
1028 dst.copy_from_slice(&xs);
1029}
1030
1031fn write_vec(dst: &mut dyn Vector, src: &[Number]) {
1032 flat_write(dst, src);
1033}
1034
1035#[cfg(test)]
1036mod tests {
1037 use super::*;
1038 use pounce_common::types::{Index, Number};
1039 use pounce_linalg::dense_vector::DenseVectorSpace;
1040 use pounce_linalg::triplet::{GenTMatrixSpace, SymTMatrixSpace};
1041 use pounce_linsol::EMatrixFormat;
1042 use pounce_linsol::sparse_sym_iface::SparseSymLinearSolverInterface;
1043
1044 struct DenseMock {
1047 dim: Index,
1048 nz: Index,
1049 a: Vec<Number>,
1050 last_factor: Vec<Number>, neg_evals: Index,
1052 }
1053
1054 impl DenseMock {
1055 fn new() -> Self {
1056 Self {
1057 dim: 0,
1058 nz: 0,
1059 a: Vec::new(),
1060 last_factor: Vec::new(),
1061 neg_evals: 0,
1062 }
1063 }
1064 }
1065
1066 impl SparseSymLinearSolverInterface for DenseMock {
1067 fn initialize_structure(
1068 &mut self,
1069 dim: Index,
1070 nz: Index,
1071 _ia: &[Index],
1072 _ja: &[Index],
1073 ) -> ESymSolverStatus {
1074 self.dim = dim;
1075 self.nz = nz;
1076 self.a = vec![0.0; nz as usize];
1077 ESymSolverStatus::Success
1078 }
1079 fn values_array_mut(&mut self) -> &mut [Number] {
1080 &mut self.a
1081 }
1082 fn multi_solve(
1083 &mut self,
1084 new_matrix: bool,
1085 ia: &[Index],
1086 ja: &[Index],
1087 nrhs: Index,
1088 rhs_vals: &mut [Number],
1089 _check: bool,
1090 _nev: Index,
1091 ) -> ESymSolverStatus {
1092 let n = self.dim as usize;
1093 if new_matrix {
1094 let mut dense = vec![0.0; n * n];
1097 for k in 0..self.nz as usize {
1098 let i = (ia[k] - 1) as usize;
1099 let j = (ja[k] - 1) as usize;
1100 dense[i * n + j] += self.a[k];
1101 if i != j {
1102 dense[j * n + i] += self.a[k];
1103 }
1104 }
1105 self.last_factor = dense;
1106 }
1107 for col in 0..nrhs as usize {
1109 let mut a = self.last_factor.clone();
1110 let b = &mut rhs_vals[col * n..col * n + n];
1111 let mut neg = 0_i32;
1112 for k in 0..n {
1113 let mut piv = k;
1115 let mut piv_abs = a[k * n + k].abs();
1116 for r in (k + 1)..n {
1117 let av = a[r * n + k].abs();
1118 if av > piv_abs {
1119 piv_abs = av;
1120 piv = r;
1121 }
1122 }
1123 if piv != k {
1124 for c in 0..n {
1125 a.swap(k * n + c, piv * n + c);
1126 }
1127 b.swap(k, piv);
1128 }
1129 let p = a[k * n + k];
1130 if p.abs() < 1e-14 {
1131 return ESymSolverStatus::Singular;
1132 }
1133 if p < 0.0 {
1134 neg += 1;
1135 }
1136 for r in (k + 1)..n {
1137 let f = a[r * n + k] / p;
1138 for c in k..n {
1139 a[r * n + c] -= f * a[k * n + c];
1140 }
1141 b[r] -= f * b[k];
1142 }
1143 }
1144 for k in (0..n).rev() {
1146 let mut s = b[k];
1147 for c in (k + 1)..n {
1148 s -= a[k * n + c] * b[c];
1149 }
1150 b[k] = s / a[k * n + k];
1151 }
1152 self.neg_evals = neg;
1153 }
1154 ESymSolverStatus::Success
1155 }
1156 fn number_of_neg_evals(&self) -> Index {
1157 self.neg_evals
1158 }
1159 fn increase_quality(&mut self) -> bool {
1160 false
1161 }
1162 fn provides_inertia(&self) -> bool {
1163 true
1164 }
1165 fn matrix_format(&self) -> EMatrixFormat {
1166 EMatrixFormat::TripletFormat
1167 }
1168 }
1169
1170 #[test]
1182 fn solves_5x5_kkt_through_dense_mock() {
1183 let w_space = SymTMatrixSpace::new(2, vec![1, 2], vec![1, 2]);
1185 let mut w = SymTMatrix::new(w_space);
1186 w.set_values(&[2.0, 3.0]);
1187
1188 let jc_space = GenTMatrixSpace::new(1, 2, vec![1, 1], vec![1, 2]);
1190 let mut j_c = GenTMatrix::new(jc_space);
1191 j_c.set_values(&[1.0, 1.0]);
1192
1193 let jd_space = GenTMatrixSpace::new(1, 2, vec![1], vec![1]);
1195 let mut j_d = GenTMatrix::new(jd_space);
1196 j_d.set_values(&[1.0]);
1197
1198 let s_space = DenseVectorSpace::new(1);
1200 let mut d_s = s_space.make_new_dense();
1201 d_s.set_values(&[1.0]);
1202
1203 let xs = DenseVectorSpace::new(2);
1211 let mut rx = xs.make_new_dense();
1212 rx.set_values(&[4.0, 4.0]);
1213 let mut rs = s_space.make_new_dense();
1214 rs.set_values(&[0.0]);
1215 let cs = DenseVectorSpace::new(1);
1216 let mut rc = cs.make_new_dense();
1217 rc.set_values(&[2.0]);
1218 let ds_space = DenseVectorSpace::new(1);
1219 let mut rd = ds_space.make_new_dense();
1220 rd.set_values(&[0.0]);
1221
1222 let mut sx = xs.make_new_dense();
1223 let mut ss = s_space.make_new_dense();
1224 let mut sc = cs.make_new_dense();
1225 let mut sd = ds_space.make_new_dense();
1226
1227 let linsol = TSymLinearSolver::new(Box::new(DenseMock::new()), None, false);
1228 let mut solver = StdAugSystemSolver::new(linsol);
1229
1230 let coeffs = AugSysCoeffs {
1231 w: Some(&w),
1232 w_factor: 1.0,
1233 d_x: None,
1234 delta_x: 0.0,
1235 d_s: Some(&d_s),
1236 delta_s: 0.0,
1237 j_c: &j_c,
1238 d_c: None,
1239 delta_c: 0.0,
1240 j_d: &j_d,
1241 d_d: None,
1242 delta_d: 0.0,
1243 };
1244 let rhs = AugSysRhs {
1245 rhs_x: &rx,
1246 rhs_s: &rs,
1247 rhs_c: &rc,
1248 rhs_d: &rd,
1249 };
1250 let mut sol = AugSysSol {
1251 sol_x: &mut sx,
1252 sol_s: &mut ss,
1253 sol_c: &mut sc,
1254 sol_d: &mut sd,
1255 };
1256 let status = solver.solve(&coeffs, &rhs, &mut sol, false, 0);
1257 assert_eq!(status, ESymSolverStatus::Success);
1258
1259 for v in sx.values() {
1260 assert!((v - 1.0).abs() < 1e-10, "sol_x = {v}");
1261 }
1262 for v in ss.values() {
1263 assert!((v - 1.0).abs() < 1e-10, "sol_s = {v}");
1264 }
1265 for v in sc.values() {
1266 assert!((v - 1.0).abs() < 1e-10, "sol_c = {v}");
1267 }
1268 for v in sd.values() {
1269 assert!((v - 1.0).abs() < 1e-10, "sol_d = {v}");
1270 }
1271 }
1272
1273 #[test]
1283 fn lowrank_smw_matches_dense_w_on_constrained_system() {
1284 use crate::kkt::low_rank_aug_system_solver::LowRankAugSystemSolver;
1285 use pounce_linalg::diag_matrix::DiagMatrix;
1286 use pounce_linalg::low_rank_update_sym_matrix::LowRankUpdateSymMatrixSpace;
1287 use pounce_linalg::multi_vector_matrix::MultiVectorMatrixSpace;
1288
1289 let n = 4usize;
1290 let sigma = 2.0;
1291 let vcols = [
1296 vec![0.6, 0.1, -0.2, 0.3],
1297 vec![0.2, 0.5, 0.1, -0.1],
1298 vec![-0.1, 0.2, 0.4, 0.2],
1299 vec![0.3, -0.2, 0.1, 0.4],
1300 vec![0.15, 0.25, -0.3, 0.1],
1301 vec![-0.2, 0.1, 0.2, 0.35],
1302 ];
1303 let ucols = [
1304 vec![0.3, -0.1, 0.2, 0.1],
1305 vec![0.1, 0.3, -0.2, 0.2],
1306 vec![0.2, 0.1, 0.1, -0.3],
1307 vec![-0.1, 0.2, 0.15, 0.1],
1308 vec![0.25, -0.15, 0.1, 0.2],
1309 vec![0.1, 0.2, -0.25, 0.15],
1310 ];
1311 let mut wfull = vec![0.0_f64; n * n];
1313 for i in 0..n {
1314 wfull[i * n + i] = sigma;
1315 }
1316 for c in vcols.iter() {
1317 for i in 0..n {
1318 for j in 0..n {
1319 wfull[i * n + j] += c[i] * c[j];
1320 }
1321 }
1322 }
1323 for c in ucols.iter() {
1324 for i in 0..n {
1325 for j in 0..n {
1326 wfull[i * n + j] -= c[i] * c[j];
1327 }
1328 }
1329 }
1330
1331 let make_jc = || {
1333 let sp = GenTMatrixSpace::new(1, 4, vec![1, 1, 1, 1], vec![1, 2, 3, 4]);
1334 let mut m = GenTMatrix::new(sp);
1335 m.set_values(&[1.0, 1.0, 1.0, 1.0]);
1336 m
1337 };
1338 let make_jd = || {
1342 let sp = GenTMatrixSpace::new(1, 4, vec![1, 1], vec![1, 3]);
1343 let mut m = GenTMatrix::new(sp);
1344 m.set_values(&[1.0, 1.0]);
1345 m
1346 };
1347
1348 let xs = DenseVectorSpace::new(4);
1349 let cs = DenseVectorSpace::new(1);
1350 let mk = |sp: &Rc<DenseVectorSpace>, vals: &[Number]| {
1351 let mut d = sp.make_new_dense();
1352 d.set_values(vals);
1353 d
1354 };
1355
1356 let solve_with = |w: &dyn pounce_linalg::SymMatrix,
1357 aug: &mut dyn AugSystemSolver|
1358 -> (Vec<Number>, Vec<Number>) {
1359 let j_c = make_jc();
1360 let j_d = make_jd();
1361 let rx = mk(&xs, &[1.0, 2.0, -1.0, 0.5]);
1362 let rs = mk(&cs, &[0.4]);
1363 let rc = mk(&cs, &[3.0]);
1364 let rd = mk(&cs, &[0.7]);
1365 let mut sx = mk(&xs, &[0.0, 0.0, 0.0, 0.0]);
1366 let mut ss = mk(&cs, &[0.0]);
1367 let mut sc = mk(&cs, &[0.0]);
1368 let mut sd = mk(&cs, &[0.0]);
1369 let d_s = mk(&cs, &[1.5]);
1370 let coeffs = AugSysCoeffs {
1371 w: Some(w),
1372 w_factor: 1.0,
1373 d_x: None,
1374 delta_x: 0.0,
1375 d_s: Some(&d_s),
1376 delta_s: 0.0,
1377 j_c: &j_c,
1378 d_c: None,
1379 delta_c: 0.0,
1380 j_d: &j_d,
1381 d_d: None,
1382 delta_d: 0.0,
1383 };
1384 let rhs = AugSysRhs {
1385 rhs_x: &rx,
1386 rhs_s: &rs,
1387 rhs_c: &rc,
1388 rhs_d: &rd,
1389 };
1390 let mut sol = AugSysSol {
1391 sol_x: &mut sx,
1392 sol_s: &mut ss,
1393 sol_c: &mut sc,
1394 sol_d: &mut sd,
1395 };
1396 let status = aug.solve(&coeffs, &rhs, &mut sol, false, 1);
1397 assert_eq!(status, ESymSolverStatus::Success);
1398 (sx.expanded_values(), sc.expanded_values())
1399 };
1400
1401 let mut wi = Vec::new();
1403 let mut wj = Vec::new();
1404 let mut wv = Vec::new();
1405 for i in 0..n {
1406 for j in 0..=i {
1407 wi.push(i as Index + 1);
1408 wj.push(j as Index + 1);
1409 wv.push(wfull[i * n + j]);
1410 }
1411 }
1412 let w_space = SymTMatrixSpace::new(4, wi, wj);
1413 let mut w_dense = SymTMatrix::new(w_space);
1414 w_dense.set_values(&wv);
1415 let mut std_solver = StdAugSystemSolver::new(TSymLinearSolver::new(
1416 Box::new(pounce_feral::FeralSolverInterface::new()),
1417 None,
1418 false,
1419 ));
1420 let (ref_x, ref_c) = solve_with(&w_dense, &mut std_solver);
1421
1422 let lr_space = LowRankUpdateSymMatrixSpace::new(4, None, false);
1424 let mut lr = lr_space.make_new_low_rank();
1425 let mut diag = xs.make_new_dense();
1426 diag.set_values(&[sigma; 4]);
1427 lr.set_diag(Rc::new(diag) as Rc<dyn Vector>);
1428 let build_mvm = |cols: &[Vec<Number>]| {
1429 let sp = MultiVectorMatrixSpace::new(cols.len() as Index, Rc::clone(&xs));
1430 let mut mvm = sp.make_new_multi_vector();
1431 for (k, c) in cols.iter().enumerate() {
1432 let mut cv = xs.make_new_dense();
1433 cv.set_values(c);
1434 mvm.set_vector(k as Index, Rc::new(cv) as Rc<dyn Vector>);
1435 }
1436 mvm
1437 };
1438 lr.set_v(Rc::new(build_mvm(&vcols)));
1439 lr.set_u(Rc::new(build_mvm(&ucols)));
1440 let _ = DiagMatrix::new(4); let mut lr_solver =
1443 LowRankAugSystemSolver::new(Box::new(StdAugSystemSolver::new(TSymLinearSolver::new(
1444 Box::new(pounce_feral::FeralSolverInterface::new()),
1445 None,
1446 false,
1447 ))));
1448 let (lr_x, lr_c) = solve_with(&lr, &mut lr_solver);
1449
1450 for (a, b) in ref_x.iter().zip(lr_x.iter()) {
1451 assert!((a - b).abs() < 1e-9, "sol_x mismatch: dense={a} smw={b}");
1452 }
1453 for (a, b) in ref_c.iter().zip(lr_c.iter()) {
1454 assert!((a - b).abs() < 1e-9, "sol_c mismatch: dense={a} smw={b}");
1455 }
1456 }
1457}