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
102impl std::fmt::Debug for StdAugSystemSolver {
103 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
104 f.debug_struct("StdAugSystemSolver")
105 .field("dim", &self.dim)
106 .field("nnz", &self.vals.len())
107 .field("initialized", &self.initialized)
108 .field("last_neg_evals", &self.last_neg_evals)
109 .field("last_status", &self.last_status)
110 .finish_non_exhaustive()
111 }
112}
113
114impl StdAugSystemSolver {
115 pub fn new(linsol: TSymLinearSolver) -> Self {
117 Self {
118 linsol,
119 initialized: false,
120 struct_sig: None,
121 n_x: 0,
122 n_s: 0,
123 n_c: 0,
124 n_d: 0,
125 dim: 0,
126 irn: Vec::new(),
127 jcn: Vec::new(),
128 vals: Vec::new(),
129 w_range: 0..0,
130 dx_range: 0..0,
131 ds_range: 0..0,
132 jc_range: 0..0,
133 dc_range: 0..0,
134 jd_range: 0..0,
135 minus_i_range: 0..0,
136 dd_range: 0..0,
137 last_neg_evals: 0,
138 last_status: None,
139 have_factor: false,
140 timing: None,
141 diagnostics: None,
142 }
143 }
144
145 fn build_structure(&mut self, coeffs: &AugSysCoeffs<'_>) -> ESymSolverStatus {
146 let n_x = coeffs.j_c.n_cols();
147 let n_c = coeffs.j_c.n_rows();
148 let n_d = coeffs.j_d.n_rows();
149 debug_assert_eq!(coeffs.j_d.n_cols(), n_x);
150 let n_s = n_d;
151
152 let w_nnz = match coeffs.w {
153 None => 0_usize,
154 Some(w) => w_nonzeros(w),
155 };
156 let jc_nnz = gen_t_downcast(coeffs.j_c).nonzeros() as usize;
157 let jd_nnz = gen_t_downcast(coeffs.j_d).nonzeros() as usize;
158
159 let total = w_nnz
160 + (n_x as usize) + (n_s as usize) + jc_nnz
163 + (n_c as usize) + jd_nnz
165 + (n_s as usize) + (n_d as usize); self.irn = Vec::with_capacity(total);
169 self.jcn = Vec::with_capacity(total);
170 self.vals = vec![0.0; total];
171
172 let w_start = self.irn.len();
174 if let Some(w) = coeffs.w {
175 if let Some(t) = w.as_any().downcast_ref::<SymTMatrix>() {
176 self.irn.extend_from_slice(t.irows());
177 self.jcn.extend_from_slice(t.jcols());
178 } else if let Some(dm) = w.as_any().downcast_ref::<DiagMatrix>() {
179 let n = w_diag_dim(dm);
182 for i in 0..n {
183 self.irn.push(i + 1);
184 self.jcn.push(i + 1);
185 }
186 } else {
187 unreachable!("StdAugSystemSolver: W must be a SymTMatrix or DiagMatrix in v1.0")
188 }
189 }
190 self.w_range = w_start..self.irn.len();
191
192 let dx_start = self.irn.len();
194 for i in 0..n_x {
195 self.irn.push(i + 1);
196 self.jcn.push(i + 1);
197 }
198 self.dx_range = dx_start..self.irn.len();
199
200 let ds_start = self.irn.len();
202 for i in 0..n_s {
203 let r = n_x + i + 1;
204 self.irn.push(r);
205 self.jcn.push(r);
206 }
207 self.ds_range = ds_start..self.irn.len();
208
209 let jc_start = self.irn.len();
211 let j_c = gen_t_downcast(coeffs.j_c);
212 let row_off_c = n_x + n_s;
213 for (&i, &j) in j_c.irows().iter().zip(j_c.jcols().iter()) {
214 self.irn.push(row_off_c + i);
217 self.jcn.push(j);
218 }
219 self.jc_range = jc_start..self.irn.len();
220
221 let dc_start = self.irn.len();
223 for i in 0..n_c {
224 let r = n_x + n_s + i + 1;
225 self.irn.push(r);
226 self.jcn.push(r);
227 }
228 self.dc_range = dc_start..self.irn.len();
229
230 let jd_start = self.irn.len();
232 let j_d = gen_t_downcast(coeffs.j_d);
233 let row_off_d = n_x + n_s + n_c;
234 for (&i, &j) in j_d.irows().iter().zip(j_d.jcols().iter()) {
235 self.irn.push(row_off_d + i);
236 self.jcn.push(j);
237 }
238 self.jd_range = jd_start..self.irn.len();
239
240 let mi_start = self.irn.len();
242 for i in 0..n_s {
243 self.irn.push(n_x + n_s + n_c + i + 1);
244 self.jcn.push(n_x + i + 1);
245 }
246 self.minus_i_range = mi_start..self.irn.len();
247
248 let dd_start = self.irn.len();
250 for i in 0..n_d {
251 let r = n_x + n_s + n_c + i + 1;
252 self.irn.push(r);
253 self.jcn.push(r);
254 }
255 self.dd_range = dd_start..self.irn.len();
256
257 debug_assert_eq!(self.irn.len(), total);
258 debug_assert_eq!(self.jcn.len(), total);
259
260 self.n_x = n_x;
261 self.n_s = n_s;
262 self.n_c = n_c;
263 self.n_d = n_d;
264 self.dim = n_x + n_s + n_c + n_d;
265
266 let status = self
267 .linsol
268 .initialize_structure(self.dim, &self.irn, &self.jcn);
269 if status == ESymSolverStatus::Success {
270 self.initialized = true;
271 }
272 status
273 }
274
275 fn refill_values(&mut self, coeffs: &AugSysCoeffs<'_>) {
276 if !self.w_range.is_empty() {
278 let Some(w_dyn) = coeffs.w else {
279 unreachable!("structure pinned with W; W cannot be None now")
280 };
281 let dst = &mut self.vals[self.w_range.clone()];
282 if let Some(t) = w_dyn.as_any().downcast_ref::<SymTMatrix>() {
283 for (d, &v) in dst.iter_mut().zip(t.values().iter()) {
284 *d = coeffs.w_factor * v;
285 }
286 } else if let Some(dm) = w_dyn.as_any().downcast_ref::<DiagMatrix>() {
287 let diag = w_diag_values(dm);
288 for (d, &v) in dst.iter_mut().zip(diag.iter()) {
289 *d = coeffs.w_factor * v;
290 }
291 } else {
292 unreachable!("StdAugSystemSolver: W must be a SymTMatrix or DiagMatrix in v1.0")
293 }
294 }
295 fill_diag(
297 &mut self.vals[self.dx_range.clone()],
298 coeffs.d_x,
299 coeffs.delta_x,
300 1.0,
301 );
302 fill_diag(
304 &mut self.vals[self.ds_range.clone()],
305 coeffs.d_s,
306 coeffs.delta_s,
307 1.0,
308 );
309 {
311 let j_c = gen_t_downcast(coeffs.j_c);
312 self.vals[self.jc_range.clone()].copy_from_slice(j_c.values());
313 }
314 fill_diag(
316 &mut self.vals[self.dc_range.clone()],
317 coeffs.d_c,
318 coeffs.delta_c,
319 -1.0,
320 );
321 {
323 let j_d = gen_t_downcast(coeffs.j_d);
324 self.vals[self.jd_range.clone()].copy_from_slice(j_d.values());
325 }
326 for v in &mut self.vals[self.minus_i_range.clone()] {
328 *v = -1.0;
329 }
330 fill_diag(
332 &mut self.vals[self.dd_range.clone()],
333 coeffs.d_d,
334 coeffs.delta_d,
335 -1.0,
336 );
337 }
338
339 pub(crate) fn assemble(&mut self, coeffs: &AugSysCoeffs<'_>) -> ESymSolverStatus {
348 let sig = {
349 let w_nnz = coeffs.w.map(w_nonzeros).unwrap_or(0);
350 let jc_nnz = gen_t_downcast(coeffs.j_c).nonzeros() as usize;
351 let jd_nnz = gen_t_downcast(coeffs.j_d).nonzeros() as usize;
352 (
353 w_nnz,
354 jc_nnz,
355 jd_nnz,
356 coeffs.j_c.n_cols(),
357 coeffs.j_c.n_rows(),
358 coeffs.j_d.n_rows(),
359 )
360 };
361 if !self.initialized || self.struct_sig != Some(sig) {
362 let s = self.build_structure(coeffs);
363 if s != ESymSolverStatus::Success {
364 self.last_status = Some(s);
365 return s;
366 }
367 self.struct_sig = Some(sig);
368 }
369 self.refill_values(coeffs);
370 ESymSolverStatus::Success
371 }
372
373 pub(crate) fn assembled_dim(&self) -> Index {
375 self.dim
376 }
377 pub(crate) fn assembled_triplet(&self) -> (&[Index], &[Index], &[Number]) {
380 (&self.irn, &self.jcn, &self.vals)
381 }
382
383 pub(crate) fn pack_rhs(&self, rhs: &AugSysRhs<'_>, packed: &mut [Number]) {
384 let n_x = self.n_x as usize;
385 let n_s = self.n_s as usize;
386 let n_c = self.n_c as usize;
387 let n_d = self.n_d as usize;
388 copy_vec(rhs.rhs_x, &mut packed[..n_x]);
389 copy_vec(rhs.rhs_s, &mut packed[n_x..n_x + n_s]);
390 copy_vec(rhs.rhs_c, &mut packed[n_x + n_s..n_x + n_s + n_c]);
391 copy_vec(
392 rhs.rhs_d,
393 &mut packed[n_x + n_s + n_c..n_x + n_s + n_c + n_d],
394 );
395 }
396
397 pub(crate) fn unpack_sol(&self, packed: &[Number], sol: &mut AugSysSol<'_>) {
398 let n_x = self.n_x as usize;
399 let n_s = self.n_s as usize;
400 let n_c = self.n_c as usize;
401 let n_d = self.n_d as usize;
402 write_vec(sol.sol_x, &packed[..n_x]);
403 write_vec(sol.sol_s, &packed[n_x..n_x + n_s]);
404 write_vec(sol.sol_c, &packed[n_x + n_s..n_x + n_s + n_c]);
405 write_vec(sol.sol_d, &packed[n_x + n_s + n_c..n_x + n_s + n_c + n_d]);
406 }
407}
408
409impl AugSystemSolver for StdAugSystemSolver {
410 fn provides_inertia(&self) -> bool {
411 self.linsol.provides_inertia()
412 }
413
414 fn number_of_neg_evals(&self) -> Index {
415 self.last_neg_evals
416 }
417
418 fn system_dim(&self) -> Index {
419 self.dim
420 }
421
422 fn kkt_triplets(&self) -> Option<(Index, Vec<Index>, Vec<Index>, Vec<Number>)> {
423 if self.irn.is_empty() {
424 return None;
425 }
426 Some((
427 self.dim,
428 self.irn.clone(),
429 self.jcn.clone(),
430 self.vals.clone(),
431 ))
432 }
433
434 fn l_factor(&self, want_values: bool) -> Option<FactorPattern> {
435 self.linsol.factor_pattern(want_values)
436 }
437
438 fn increase_quality(&mut self) -> bool {
439 self.have_factor = false;
443 self.linsol.increase_quality()
444 }
445
446 fn last_solve_status(&self) -> ESymSolverStatus {
447 self.last_status.unwrap_or(ESymSolverStatus::FatalError)
448 }
449
450 fn solve(
451 &mut self,
452 coeffs: &AugSysCoeffs<'_>,
453 rhs: &AugSysRhs<'_>,
454 sol: &mut AugSysSol<'_>,
455 check_neg_evals: bool,
456 num_neg_evals: Index,
457 ) -> ESymSolverStatus {
458 let s = self.assemble(coeffs);
466 if s != ESymSolverStatus::Success {
467 return s;
468 }
469
470 let mut packed = vec![0.0; self.dim as usize];
471 self.pack_rhs(rhs, &mut packed);
472
473 let dump_rhs = packed.clone();
474
475 let _factor_guard = self
479 .timing
480 .as_deref()
481 .map(|t| t.linear_system_factorization.guard());
482 let status = self.linsol.multi_solve(
483 &self.vals,
484 true,
485 1,
486 &mut packed,
487 check_neg_evals,
488 num_neg_evals,
489 );
490 drop(_factor_guard);
491 self.last_status = Some(status);
492 if self.linsol.provides_inertia()
504 && matches!(
505 status,
506 ESymSolverStatus::Success
507 | ESymSolverStatus::WrongInertia
508 | ESymSolverStatus::Singular
509 )
510 {
511 self.last_neg_evals = self.linsol.number_of_neg_evals();
512 }
513 if status == ESymSolverStatus::Success {
514 self.unpack_sol(&packed, sol);
515 self.have_factor = true;
516 }
517
518 if let Some(diag) = self.diagnostics.clone() {
523 if diag.want(DiagCategory::Kkt) {
524 let solve_idx = diag.next_solve_index();
525 let filename = format!("kkt_solve_{solve_idx:03}.jsonl");
526 let variant = diag.config.kkt_variant;
533 let factor_pattern =
534 if status == ESymSolverStatus::Success && variant.wants_l_pattern() {
535 self.linsol.factor_pattern(variant.wants_l_values())
536 } else {
537 None
538 };
539 if let Some(mut w) = diag.open_writer(&filename) {
540 let _ = write_kkt_record(
541 &mut w,
542 self.dim,
543 &self.irn,
544 &self.jcn,
545 &self.vals,
546 &dump_rhs,
547 &packed,
548 check_neg_evals,
549 num_neg_evals,
550 status,
551 self.last_neg_evals,
552 factor_pattern.as_ref(),
553 );
554 }
555 }
556 }
557 if let Ok(path) = std::env::var("POUNCE_DUMP_KKT") {
558 use std::sync::atomic::{AtomicBool, Ordering};
559 static WARNED: AtomicBool = AtomicBool::new(false);
560 if !WARNED.swap(true, Ordering::SeqCst) {
561 tracing::warn!(target: "pounce::linsol",
562 "warning: POUNCE_DUMP_KKT is deprecated; prefer `--dump kkt:<iter-spec>` (see pounce --help)"
563 );
564 }
565 dump_kkt(
566 &path,
567 self.dim,
568 &self.irn,
569 &self.jcn,
570 &self.vals,
571 &dump_rhs,
572 &packed,
573 check_neg_evals,
574 num_neg_evals,
575 status,
576 self.last_neg_evals,
577 );
578 }
579
580 status
581 }
582
583 fn resolve(
584 &mut self,
585 coeffs: &AugSysCoeffs<'_>,
586 rhs: &AugSysRhs<'_>,
587 sol: &mut AugSysSol<'_>,
588 ) -> ESymSolverStatus {
589 if !self.have_factor {
596 return self.solve(coeffs, rhs, sol, false, 0);
597 }
598
599 let mut packed = vec![0.0; self.dim as usize];
600 self.pack_rhs(rhs, &mut packed);
601
602 let _back_guard = self
605 .timing
606 .as_deref()
607 .map(|t| t.linear_system_back_solve.guard());
608 let status = self
609 .linsol
610 .multi_solve(&self.vals, false, 1, &mut packed, false, 0);
611 drop(_back_guard);
612 self.last_status = Some(status);
613 if status == ESymSolverStatus::Success {
614 self.unpack_sol(&packed, sol);
615 }
616 status
617 }
618
619 fn set_diagnostics(&mut self, diag: Rc<DiagnosticsState>) {
620 self.diagnostics = Some(diag);
621 }
622
623 fn set_timing_stats(&mut self, timing: Rc<TimingStatistics>) {
624 self.timing = Some(timing);
625 }
626
627 fn try_resolve_many_flat(
628 &mut self,
629 _coeffs: &AugSysCoeffs<'_>,
630 packed_rhs: &mut [Number],
631 nrhs: usize,
632 ) -> Option<ESymSolverStatus> {
633 if !self.have_factor {
638 return None;
639 }
640 if packed_rhs.len() != (self.dim as usize) * nrhs {
641 return Some(ESymSolverStatus::FatalError);
642 }
643 let _back_guard = self
644 .timing
645 .as_deref()
646 .map(|t| t.linear_system_back_solve.guard());
647 let status =
648 self.linsol
649 .multi_solve(&self.vals, false, nrhs as Index, packed_rhs, false, 0);
650 drop(_back_guard);
651 self.last_status = Some(status);
652 Some(status)
653 }
654}
655
656#[allow(clippy::too_many_arguments)]
659fn write_kkt_record(
664 w: &mut dyn std::io::Write,
665 dim: Index,
666 irn: &[Index],
667 jcn: &[Index],
668 vals: &[Number],
669 rhs: &[Number],
670 sol: &[Number],
671 check_neg_evals: bool,
672 num_neg_evals: Index,
673 status: ESymSolverStatus,
674 last_neg_evals: Index,
675 factor_pattern: Option<&FactorPattern>,
676) -> std::io::Result<()> {
677 use std::fmt::Write as _;
678
679 let mut line = String::with_capacity(64 * vals.len());
680 line.push('{');
681 let _ = write!(line, "\"n\":{dim},");
682 let _ = write!(line, "\"check_neg_evals\":{check_neg_evals},");
683 let _ = write!(line, "\"num_neg_evals_expected\":{num_neg_evals},");
684 let _ = write!(line, "\"num_neg_evals_actual\":{last_neg_evals},");
685 let _ = write!(line, "\"status\":\"{status:?}\",");
686
687 line.push_str("\"irn\":[");
688 for (i, v) in irn.iter().enumerate() {
689 if i > 0 {
690 line.push(',');
691 }
692 let _ = write!(line, "{v}");
693 }
694 line.push_str("],\"jcn\":[");
695 for (i, v) in jcn.iter().enumerate() {
696 if i > 0 {
697 line.push(',');
698 }
699 let _ = write!(line, "{v}");
700 }
701 line.push_str("],\"vals\":[");
702 for (i, v) in vals.iter().enumerate() {
703 if i > 0 {
704 line.push(',');
705 }
706 let _ = write!(line, "{v:.17e}");
707 }
708 line.push_str("],\"rhs\":[");
709 for (i, v) in rhs.iter().enumerate() {
710 if i > 0 {
711 line.push(',');
712 }
713 let _ = write!(line, "{v:.17e}");
714 }
715 line.push_str("],\"sol\":[");
716 for (i, v) in sol.iter().enumerate() {
717 if i > 0 {
718 line.push(',');
719 }
720 let _ = write!(line, "{v:.17e}");
721 }
722 line.push(']');
723
724 if let Some(fp) = factor_pattern {
729 line.push_str(",\"L_irn\":[");
730 for (i, v) in fp.l_irn.iter().enumerate() {
731 if i > 0 {
732 line.push(',');
733 }
734 let _ = write!(line, "{v}");
735 }
736 line.push_str("],\"L_jcn\":[");
737 for (i, v) in fp.l_jcn.iter().enumerate() {
738 if i > 0 {
739 line.push(',');
740 }
741 let _ = write!(line, "{v}");
742 }
743 line.push_str("],\"perm\":[");
744 for (i, v) in fp.perm.iter().enumerate() {
745 if i > 0 {
746 line.push(',');
747 }
748 let _ = write!(line, "{v}");
749 }
750 line.push(']');
751 if let Some(vals) = fp.l_vals.as_ref() {
752 line.push_str(",\"L_vals\":[");
753 for (i, v) in vals.iter().enumerate() {
754 if i > 0 {
755 line.push(',');
756 }
757 let _ = write!(line, "{v:.17e}");
758 }
759 line.push(']');
760 }
761 }
762
763 line.push_str("}\n");
764
765 w.write_all(line.as_bytes())
766}
767
768fn dump_kkt(
769 path: &str,
770 dim: Index,
771 irn: &[Index],
772 jcn: &[Index],
773 vals: &[Number],
774 rhs: &[Number],
775 sol: &[Number],
776 check_neg_evals: bool,
777 num_neg_evals: Index,
778 status: ESymSolverStatus,
779 last_neg_evals: Index,
780) {
781 if let Ok(mut f) = std::fs::OpenOptions::new()
782 .create(true)
783 .append(true)
784 .open(path)
785 {
786 let _ = write_kkt_record(
787 &mut f,
788 dim,
789 irn,
790 jcn,
791 vals,
792 rhs,
793 sol,
794 check_neg_evals,
795 num_neg_evals,
796 status,
797 last_neg_evals,
798 None, );
800 }
801}
802
803fn w_nonzeros(w: &dyn pounce_linalg::SymMatrix) -> usize {
808 if let Some(t) = w.as_any().downcast_ref::<SymTMatrix>() {
809 t.nonzeros() as usize
810 } else if let Some(dm) = w.as_any().downcast_ref::<DiagMatrix>() {
811 w_diag_dim(dm) as usize
812 } else {
813 unreachable!("StdAugSystemSolver: W must be a SymTMatrix or DiagMatrix in v1.0")
814 }
815}
816
817fn w_diag_dim(dm: &DiagMatrix) -> Index {
818 dm.get_diag()
819 .expect("DiagMatrix W has no diagonal set")
820 .dim()
821}
822
823fn w_diag_values(dm: &DiagMatrix) -> Vec<Number> {
824 let diag = dm.get_diag().expect("DiagMatrix W has no diagonal set");
825 diag.as_any()
826 .downcast_ref::<DenseVector>()
827 .expect("StdAugSystemSolver: DiagMatrix W diagonal must be DenseVector in v1.0")
828 .expanded_values()
829}
830
831fn gen_t_downcast(m: &dyn pounce_linalg::Matrix) -> &GenTMatrix {
832 let Some(t) = m.as_any().downcast_ref::<GenTMatrix>() else {
833 unreachable!("StdAugSystemSolver: J_c / J_d must be GenTMatrix in v1.0")
834 };
835 t
836}
837
838fn flat_read(v: &dyn Vector) -> Vec<Number> {
844 if let Some(dv) = v.as_any().downcast_ref::<DenseVector>() {
845 return dv.expanded_values();
846 }
847 if let Some(cv) = v.as_any().downcast_ref::<CompoundVector>() {
848 let mut out = Vec::with_capacity(cv.dim() as usize);
849 for k in 0..cv.n_comps() {
850 let blk = cv.comp(k);
851 let dblk = blk
852 .as_any()
853 .downcast_ref::<DenseVector>()
854 .expect("StdAugSystemSolver: CompoundVector blocks must be DenseVectors");
855 out.extend_from_slice(&dblk.expanded_values());
856 }
857 return out;
858 }
859 unreachable!(
860 "StdAugSystemSolver: D_*/rhs/sol must be DenseVector or CompoundVector of DenseVectors in v1.0"
861 )
862}
863
864fn flat_write(dst: &mut dyn Vector, src: &[Number]) {
866 if let Some(dv) = dst.as_any_mut().downcast_mut::<DenseVector>() {
867 dv.set_values(src);
868 return;
869 }
870 if let Some(cv) = dst.as_any_mut().downcast_mut::<CompoundVector>() {
871 let mut off = 0usize;
872 for k in 0..cv.n_comps() {
873 let blk = cv.comp_mut(k);
874 let dim = blk.dim() as usize;
875 let dblk = blk
876 .as_any_mut()
877 .downcast_mut::<DenseVector>()
878 .expect("StdAugSystemSolver: CompoundVector blocks must be DenseVectors");
879 dblk.set_values(&src[off..off + dim]);
880 off += dim;
881 }
882 return;
883 }
884 unreachable!(
885 "StdAugSystemSolver: sol must be DenseVector or CompoundVector of DenseVectors in v1.0"
886 )
887}
888
889fn fill_diag(dst: &mut [Number], d: Option<&dyn Vector>, delta: Number, sign: Number) {
892 match d {
893 None => {
894 for v in dst.iter_mut() {
895 *v = sign * delta;
896 }
897 }
898 Some(d) => {
899 let xs = flat_read(d);
900 debug_assert_eq!(xs.len(), dst.len());
901 for (out, &x) in dst.iter_mut().zip(xs.iter()) {
902 *out = sign * (x + delta);
903 }
904 }
905 }
906}
907
908fn copy_vec(src: &dyn Vector, dst: &mut [Number]) {
909 let xs = flat_read(src);
910 debug_assert_eq!(xs.len(), dst.len());
911 dst.copy_from_slice(&xs);
912}
913
914fn write_vec(dst: &mut dyn Vector, src: &[Number]) {
915 flat_write(dst, src);
916}
917
918#[cfg(test)]
919mod tests {
920 use super::*;
921 use pounce_common::types::{Index, Number};
922 use pounce_linalg::dense_vector::DenseVectorSpace;
923 use pounce_linalg::triplet::{GenTMatrixSpace, SymTMatrixSpace};
924 use pounce_linsol::EMatrixFormat;
925 use pounce_linsol::sparse_sym_iface::SparseSymLinearSolverInterface;
926
927 struct DenseMock {
930 dim: Index,
931 nz: Index,
932 a: Vec<Number>,
933 last_factor: Vec<Number>, neg_evals: Index,
935 }
936
937 impl DenseMock {
938 fn new() -> Self {
939 Self {
940 dim: 0,
941 nz: 0,
942 a: Vec::new(),
943 last_factor: Vec::new(),
944 neg_evals: 0,
945 }
946 }
947 }
948
949 impl SparseSymLinearSolverInterface for DenseMock {
950 fn initialize_structure(
951 &mut self,
952 dim: Index,
953 nz: Index,
954 _ia: &[Index],
955 _ja: &[Index],
956 ) -> ESymSolverStatus {
957 self.dim = dim;
958 self.nz = nz;
959 self.a = vec![0.0; nz as usize];
960 ESymSolverStatus::Success
961 }
962 fn values_array_mut(&mut self) -> &mut [Number] {
963 &mut self.a
964 }
965 fn multi_solve(
966 &mut self,
967 new_matrix: bool,
968 ia: &[Index],
969 ja: &[Index],
970 nrhs: Index,
971 rhs_vals: &mut [Number],
972 _check: bool,
973 _nev: Index,
974 ) -> ESymSolverStatus {
975 let n = self.dim as usize;
976 if new_matrix {
977 let mut dense = vec![0.0; n * n];
980 for k in 0..self.nz as usize {
981 let i = (ia[k] - 1) as usize;
982 let j = (ja[k] - 1) as usize;
983 dense[i * n + j] += self.a[k];
984 if i != j {
985 dense[j * n + i] += self.a[k];
986 }
987 }
988 self.last_factor = dense;
989 }
990 for col in 0..nrhs as usize {
992 let mut a = self.last_factor.clone();
993 let b = &mut rhs_vals[col * n..col * n + n];
994 let mut neg = 0_i32;
995 for k in 0..n {
996 let mut piv = k;
998 let mut piv_abs = a[k * n + k].abs();
999 for r in (k + 1)..n {
1000 let av = a[r * n + k].abs();
1001 if av > piv_abs {
1002 piv_abs = av;
1003 piv = r;
1004 }
1005 }
1006 if piv != k {
1007 for c in 0..n {
1008 a.swap(k * n + c, piv * n + c);
1009 }
1010 b.swap(k, piv);
1011 }
1012 let p = a[k * n + k];
1013 if p.abs() < 1e-14 {
1014 return ESymSolverStatus::Singular;
1015 }
1016 if p < 0.0 {
1017 neg += 1;
1018 }
1019 for r in (k + 1)..n {
1020 let f = a[r * n + k] / p;
1021 for c in k..n {
1022 a[r * n + c] -= f * a[k * n + c];
1023 }
1024 b[r] -= f * b[k];
1025 }
1026 }
1027 for k in (0..n).rev() {
1029 let mut s = b[k];
1030 for c in (k + 1)..n {
1031 s -= a[k * n + c] * b[c];
1032 }
1033 b[k] = s / a[k * n + k];
1034 }
1035 self.neg_evals = neg;
1036 }
1037 ESymSolverStatus::Success
1038 }
1039 fn number_of_neg_evals(&self) -> Index {
1040 self.neg_evals
1041 }
1042 fn increase_quality(&mut self) -> bool {
1043 false
1044 }
1045 fn provides_inertia(&self) -> bool {
1046 true
1047 }
1048 fn matrix_format(&self) -> EMatrixFormat {
1049 EMatrixFormat::TripletFormat
1050 }
1051 }
1052
1053 #[test]
1065 fn solves_5x5_kkt_through_dense_mock() {
1066 let w_space = SymTMatrixSpace::new(2, vec![1, 2], vec![1, 2]);
1068 let mut w = SymTMatrix::new(w_space);
1069 w.set_values(&[2.0, 3.0]);
1070
1071 let jc_space = GenTMatrixSpace::new(1, 2, vec![1, 1], vec![1, 2]);
1073 let mut j_c = GenTMatrix::new(jc_space);
1074 j_c.set_values(&[1.0, 1.0]);
1075
1076 let jd_space = GenTMatrixSpace::new(1, 2, vec![1], vec![1]);
1078 let mut j_d = GenTMatrix::new(jd_space);
1079 j_d.set_values(&[1.0]);
1080
1081 let s_space = DenseVectorSpace::new(1);
1083 let mut d_s = s_space.make_new_dense();
1084 d_s.set_values(&[1.0]);
1085
1086 let xs = DenseVectorSpace::new(2);
1094 let mut rx = xs.make_new_dense();
1095 rx.set_values(&[4.0, 4.0]);
1096 let mut rs = s_space.make_new_dense();
1097 rs.set_values(&[0.0]);
1098 let cs = DenseVectorSpace::new(1);
1099 let mut rc = cs.make_new_dense();
1100 rc.set_values(&[2.0]);
1101 let ds_space = DenseVectorSpace::new(1);
1102 let mut rd = ds_space.make_new_dense();
1103 rd.set_values(&[0.0]);
1104
1105 let mut sx = xs.make_new_dense();
1106 let mut ss = s_space.make_new_dense();
1107 let mut sc = cs.make_new_dense();
1108 let mut sd = ds_space.make_new_dense();
1109
1110 let linsol = TSymLinearSolver::new(Box::new(DenseMock::new()), None, false);
1111 let mut solver = StdAugSystemSolver::new(linsol);
1112
1113 let coeffs = AugSysCoeffs {
1114 w: Some(&w),
1115 w_factor: 1.0,
1116 d_x: None,
1117 delta_x: 0.0,
1118 d_s: Some(&d_s),
1119 delta_s: 0.0,
1120 j_c: &j_c,
1121 d_c: None,
1122 delta_c: 0.0,
1123 j_d: &j_d,
1124 d_d: None,
1125 delta_d: 0.0,
1126 };
1127 let rhs = AugSysRhs {
1128 rhs_x: &rx,
1129 rhs_s: &rs,
1130 rhs_c: &rc,
1131 rhs_d: &rd,
1132 };
1133 let mut sol = AugSysSol {
1134 sol_x: &mut sx,
1135 sol_s: &mut ss,
1136 sol_c: &mut sc,
1137 sol_d: &mut sd,
1138 };
1139 let status = solver.solve(&coeffs, &rhs, &mut sol, false, 0);
1140 assert_eq!(status, ESymSolverStatus::Success);
1141
1142 for v in sx.values() {
1143 assert!((v - 1.0).abs() < 1e-10, "sol_x = {v}");
1144 }
1145 for v in ss.values() {
1146 assert!((v - 1.0).abs() < 1e-10, "sol_s = {v}");
1147 }
1148 for v in sc.values() {
1149 assert!((v - 1.0).abs() < 1e-10, "sol_c = {v}");
1150 }
1151 for v in sd.values() {
1152 assert!((v - 1.0).abs() < 1e-10, "sol_d = {v}");
1153 }
1154 }
1155
1156 #[test]
1166 fn lowrank_smw_matches_dense_w_on_constrained_system() {
1167 use crate::kkt::low_rank_aug_system_solver::LowRankAugSystemSolver;
1168 use pounce_linalg::diag_matrix::DiagMatrix;
1169 use pounce_linalg::low_rank_update_sym_matrix::LowRankUpdateSymMatrixSpace;
1170 use pounce_linalg::multi_vector_matrix::MultiVectorMatrixSpace;
1171
1172 let n = 4usize;
1173 let sigma = 2.0;
1174 let vcols = [
1179 vec![0.6, 0.1, -0.2, 0.3],
1180 vec![0.2, 0.5, 0.1, -0.1],
1181 vec![-0.1, 0.2, 0.4, 0.2],
1182 vec![0.3, -0.2, 0.1, 0.4],
1183 vec![0.15, 0.25, -0.3, 0.1],
1184 vec![-0.2, 0.1, 0.2, 0.35],
1185 ];
1186 let ucols = [
1187 vec![0.3, -0.1, 0.2, 0.1],
1188 vec![0.1, 0.3, -0.2, 0.2],
1189 vec![0.2, 0.1, 0.1, -0.3],
1190 vec![-0.1, 0.2, 0.15, 0.1],
1191 vec![0.25, -0.15, 0.1, 0.2],
1192 vec![0.1, 0.2, -0.25, 0.15],
1193 ];
1194 let mut wfull = vec![0.0_f64; n * n];
1196 for i in 0..n {
1197 wfull[i * n + i] = sigma;
1198 }
1199 for c in vcols.iter() {
1200 for i in 0..n {
1201 for j in 0..n {
1202 wfull[i * n + j] += c[i] * c[j];
1203 }
1204 }
1205 }
1206 for c in ucols.iter() {
1207 for i in 0..n {
1208 for j in 0..n {
1209 wfull[i * n + j] -= c[i] * c[j];
1210 }
1211 }
1212 }
1213
1214 let make_jc = || {
1216 let sp = GenTMatrixSpace::new(1, 4, vec![1, 1, 1, 1], vec![1, 2, 3, 4]);
1217 let mut m = GenTMatrix::new(sp);
1218 m.set_values(&[1.0, 1.0, 1.0, 1.0]);
1219 m
1220 };
1221 let make_jd = || {
1225 let sp = GenTMatrixSpace::new(1, 4, vec![1, 1], vec![1, 3]);
1226 let mut m = GenTMatrix::new(sp);
1227 m.set_values(&[1.0, 1.0]);
1228 m
1229 };
1230
1231 let xs = DenseVectorSpace::new(4);
1232 let cs = DenseVectorSpace::new(1);
1233 let mk = |sp: &Rc<DenseVectorSpace>, vals: &[Number]| {
1234 let mut d = sp.make_new_dense();
1235 d.set_values(vals);
1236 d
1237 };
1238
1239 let solve_with = |w: &dyn pounce_linalg::SymMatrix,
1240 aug: &mut dyn AugSystemSolver|
1241 -> (Vec<Number>, Vec<Number>) {
1242 let j_c = make_jc();
1243 let j_d = make_jd();
1244 let rx = mk(&xs, &[1.0, 2.0, -1.0, 0.5]);
1245 let rs = mk(&cs, &[0.4]);
1246 let rc = mk(&cs, &[3.0]);
1247 let rd = mk(&cs, &[0.7]);
1248 let mut sx = mk(&xs, &[0.0, 0.0, 0.0, 0.0]);
1249 let mut ss = mk(&cs, &[0.0]);
1250 let mut sc = mk(&cs, &[0.0]);
1251 let mut sd = mk(&cs, &[0.0]);
1252 let d_s = mk(&cs, &[1.5]);
1253 let coeffs = AugSysCoeffs {
1254 w: Some(w),
1255 w_factor: 1.0,
1256 d_x: None,
1257 delta_x: 0.0,
1258 d_s: Some(&d_s),
1259 delta_s: 0.0,
1260 j_c: &j_c,
1261 d_c: None,
1262 delta_c: 0.0,
1263 j_d: &j_d,
1264 d_d: None,
1265 delta_d: 0.0,
1266 };
1267 let rhs = AugSysRhs {
1268 rhs_x: &rx,
1269 rhs_s: &rs,
1270 rhs_c: &rc,
1271 rhs_d: &rd,
1272 };
1273 let mut sol = AugSysSol {
1274 sol_x: &mut sx,
1275 sol_s: &mut ss,
1276 sol_c: &mut sc,
1277 sol_d: &mut sd,
1278 };
1279 let status = aug.solve(&coeffs, &rhs, &mut sol, false, 1);
1280 assert_eq!(status, ESymSolverStatus::Success);
1281 (sx.expanded_values(), sc.expanded_values())
1282 };
1283
1284 let mut wi = Vec::new();
1286 let mut wj = Vec::new();
1287 let mut wv = Vec::new();
1288 for i in 0..n {
1289 for j in 0..=i {
1290 wi.push(i as Index + 1);
1291 wj.push(j as Index + 1);
1292 wv.push(wfull[i * n + j]);
1293 }
1294 }
1295 let w_space = SymTMatrixSpace::new(4, wi, wj);
1296 let mut w_dense = SymTMatrix::new(w_space);
1297 w_dense.set_values(&wv);
1298 let mut std_solver = StdAugSystemSolver::new(TSymLinearSolver::new(
1299 Box::new(pounce_feral::FeralSolverInterface::new()),
1300 None,
1301 false,
1302 ));
1303 let (ref_x, ref_c) = solve_with(&w_dense, &mut std_solver);
1304
1305 let lr_space = LowRankUpdateSymMatrixSpace::new(4, None, false);
1307 let mut lr = lr_space.make_new_low_rank();
1308 let mut diag = xs.make_new_dense();
1309 diag.set_values(&[sigma; 4]);
1310 lr.set_diag(Rc::new(diag) as Rc<dyn Vector>);
1311 let build_mvm = |cols: &[Vec<Number>]| {
1312 let sp = MultiVectorMatrixSpace::new(cols.len() as Index, Rc::clone(&xs));
1313 let mut mvm = sp.make_new_multi_vector();
1314 for (k, c) in cols.iter().enumerate() {
1315 let mut cv = xs.make_new_dense();
1316 cv.set_values(c);
1317 mvm.set_vector(k as Index, Rc::new(cv) as Rc<dyn Vector>);
1318 }
1319 mvm
1320 };
1321 lr.set_v(Rc::new(build_mvm(&vcols)));
1322 lr.set_u(Rc::new(build_mvm(&ucols)));
1323 let _ = DiagMatrix::new(4); let mut lr_solver =
1326 LowRankAugSystemSolver::new(Box::new(StdAugSystemSolver::new(TSymLinearSolver::new(
1327 Box::new(pounce_feral::FeralSolverInterface::new()),
1328 None,
1329 false,
1330 ))));
1331 let (lr_x, lr_c) = solve_with(&lr, &mut lr_solver);
1332
1333 for (a, b) in ref_x.iter().zip(lr_x.iter()) {
1334 assert!((a - b).abs() < 1e-9, "sol_x mismatch: dense={a} smw={b}");
1335 }
1336 for (a, b) in ref_c.iter().zip(lr_c.iter()) {
1337 assert!((a - b).abs() < 1e-9, "sol_c mismatch: dense={a} smw={b}");
1338 }
1339 }
1340}