use std::rc::Rc;
use crate::kkt::aug_system_solver::{AugSysCoeffs, AugSysRhs, AugSysSol, AugSystemSolver};
use crate::kkt::std_aug_system_solver::StdAugSystemSolver;
use pounce_common::diagnostics::DiagnosticsState;
use pounce_common::timing::TimingStatistics;
use pounce_common::types::{Index, Number};
use pounce_feral::{FeralConfig, FeralSchurSolver};
use pounce_linsol::{ESymSolverStatus, FactorPattern};
const DEFAULT_MAX_SCHUR_FRAC: f64 = 0.5;
pub struct SchurAugSystemSolver {
inner: StdAugSystemSolver,
schur: FeralSchurSolver,
schur_indices: Vec<usize>,
max_schur_frac: f64,
decided_for_dim: Option<Index>,
use_schur: bool,
have_factor: bool,
negevals: Index,
last_status: ESymSolverStatus,
timing: Option<Rc<TimingStatistics>>,
}
impl SchurAugSystemSolver {
pub fn new(inner: StdAugSystemSolver, schur_indices: Vec<usize>, cfg: FeralConfig) -> Self {
Self {
inner,
schur: FeralSchurSolver::new(cfg),
schur_indices,
max_schur_frac: DEFAULT_MAX_SCHUR_FRAC,
decided_for_dim: None,
use_schur: false,
have_factor: false,
negevals: 0,
last_status: ESymSolverStatus::Success,
timing: None,
}
}
fn decide(&mut self, dim: Index) {
if self.decided_for_dim == Some(dim) {
return;
}
self.decided_for_dim = Some(dim);
self.use_schur = false;
let n_s = self.schur_indices.len();
let d = dim as usize;
if n_s == 0 || n_s >= d {
return;
}
if (n_s as f64) / (d as f64) > self.max_schur_frac {
tracing::warn!(
target: "pounce::kkt",
n_schur = n_s, dim = d, max_frac = self.max_schur_frac,
"Schur block too large relative to the KKT; using the standard solver"
);
return;
}
let (irn, jcn) = {
let (a, b, _v) = self.inner.assembled_triplet();
(a.to_vec(), b.to_vec())
};
let st = self
.schur
.initialize_structure(dim, &irn, &jcn, &self.schur_indices);
if st == ESymSolverStatus::Success {
self.use_schur = true;
} else {
tracing::warn!(
target: "pounce::kkt",
"Schur partition rejected by the backend; using the standard solver"
);
}
}
fn schur_solve_one(
&mut self,
rhs: &AugSysRhs<'_>,
sol: &mut AugSysSol<'_>,
check_neg_evals: bool,
num_neg_evals: Index,
) -> ESymSolverStatus {
let dim = self.inner.assembled_dim() as usize;
let vals = self.inner.assembled_triplet().2.to_vec();
self.schur.values_array_mut().copy_from_slice(&vals);
let status = {
let _g = self
.timing
.as_deref()
.map(|t| t.linear_system_factorization.guard());
self.schur.factor(check_neg_evals, num_neg_evals)
};
self.last_status = status;
match status {
ESymSolverStatus::Success => {
self.negevals = self.schur.number_of_neg_evals();
let mut packed = vec![0.0; dim];
self.inner.pack_rhs(rhs, &mut packed);
let bstat = {
let _g = self
.timing
.as_deref()
.map(|t| t.linear_system_back_solve.guard());
self.schur.backsolve(1, &mut packed)
};
if bstat != ESymSolverStatus::Success {
self.have_factor = false;
self.last_status = bstat;
return bstat;
}
self.inner.unpack_sol(&packed, sol);
self.have_factor = true;
ESymSolverStatus::Success
}
ESymSolverStatus::WrongInertia => {
self.negevals = self.schur.number_of_neg_evals();
self.have_factor = false;
status
}
other => {
self.have_factor = false;
other
}
}
}
}
impl AugSystemSolver for SchurAugSystemSolver {
fn provides_inertia(&self) -> bool {
self.inner.provides_inertia()
}
fn number_of_neg_evals(&self) -> Index {
if self.use_schur {
self.negevals
} else {
self.inner.number_of_neg_evals()
}
}
fn system_dim(&self) -> Index {
self.inner.system_dim()
}
fn kkt_triplets(&self) -> Option<(Index, Vec<Index>, Vec<Index>, Vec<Number>)> {
self.inner.kkt_triplets()
}
fn l_factor(&self, want_values: bool) -> Option<FactorPattern> {
if self.use_schur {
None
} else {
self.inner.l_factor(want_values)
}
}
fn increase_quality(&mut self) -> bool {
self.have_factor = false;
if self.use_schur {
self.schur.increase_quality()
} else {
self.inner.increase_quality()
}
}
fn last_solve_status(&self) -> ESymSolverStatus {
if self.use_schur {
self.last_status
} else {
self.inner.last_solve_status()
}
}
fn set_timing_stats(&mut self, timing: Rc<TimingStatistics>) {
self.timing = Some(Rc::clone(&timing));
self.inner.set_timing_stats(timing);
}
fn set_diagnostics(&mut self, diag: Rc<DiagnosticsState>) {
self.inner.set_diagnostics(diag);
}
fn solve(
&mut self,
coeffs: &AugSysCoeffs<'_>,
rhs: &AugSysRhs<'_>,
sol: &mut AugSysSol<'_>,
check_neg_evals: bool,
num_neg_evals: Index,
) -> ESymSolverStatus {
let s = self.inner.assemble(coeffs);
if s != ESymSolverStatus::Success {
self.last_status = s;
return s;
}
let dim = self.inner.assembled_dim();
self.decide(dim);
if self.use_schur {
let st = self.schur_solve_one(rhs, sol, check_neg_evals, num_neg_evals);
match st {
ESymSolverStatus::Success | ESymSolverStatus::WrongInertia => return st,
_ => {
tracing::warn!(
target: "pounce::kkt",
status = ?st,
"Schur backend could not factor this KKT; falling back to the standard solver"
);
self.use_schur = false;
return self
.inner
.solve(coeffs, rhs, sol, check_neg_evals, num_neg_evals);
}
}
}
self.inner
.solve(coeffs, rhs, sol, check_neg_evals, num_neg_evals)
}
fn resolve(
&mut self,
coeffs: &AugSysCoeffs<'_>,
rhs: &AugSysRhs<'_>,
sol: &mut AugSysSol<'_>,
) -> ESymSolverStatus {
if self.use_schur {
if self.have_factor {
let dim = self.inner.assembled_dim() as usize;
let mut packed = vec![0.0; dim];
self.inner.pack_rhs(rhs, &mut packed);
let bstat = {
let _g = self
.timing
.as_deref()
.map(|t| t.linear_system_back_solve.guard());
self.schur.backsolve(1, &mut packed)
};
if bstat == ESymSolverStatus::Success {
self.inner.unpack_sol(&packed, sol);
}
return bstat;
}
return self.solve(coeffs, rhs, sol, false, 0);
}
self.inner.resolve(coeffs, rhs, sol)
}
}