use super::*;
use faer::Unbind;
pub(crate) const DENSE_OUTER_MAX_P: usize = 1024;
pub(crate) const DENSE_OUTER_PARALLEL_FLOP_THRESHOLD: u64 = 100_000;
pub(crate) enum XtWxBackend {
Dense(DenseOuterState),
Sparse(SparseSpGemmState),
}
pub(crate) struct DenseOuterState {
pub(crate) xtwx_dense: Array2<f64>,
pub(crate) thread_buffers: Vec<Array2<f64>>,
}
pub(crate) struct SparseSpGemmState {
pub(crate) wxvalues: Vec<f64>,
pub(crate) wx_tvalues: Vec<f64>,
pub(crate) info: SparseMatMulInfo,
pub(crate) scratch: MemBuffer,
pub(crate) par: Par,
}
pub(crate) struct SparseXtWxCache {
pub(crate) xtwx_symbolic: SymbolicSparseColMat<usize>,
pub(crate) xtwxvalues: Vec<f64>,
pub(crate) nrows: usize,
pub(crate) ncols: usize,
pub(crate) nnz: usize,
pub(crate) x_col_ptr: Vec<usize>,
pub(crate) xrow_idx: Vec<usize>,
pub(crate) x_t_csc: SparseColMat<usize, f64>,
pub(crate) backend: XtWxBackend,
}
impl SparseXtWxCache {
pub(crate) fn new(x: &SparseColMat<usize, f64>) -> Result<Self, EstimationError> {
let x_t_csc =
x.as_ref().transpose().to_col_major().map_err(|_| {
EstimationError::InvalidInput("failed to transpose to CSC".to_string())
})?;
let (xtwx_symbolic, info) = sparse_sparse_matmul_symbolic(x_t_csc.symbolic(), x.symbolic())
.map_err(|_| {
EstimationError::InvalidInput("failed to build symbolic XtWX cache".to_string())
})?;
let xtwxvalues = vec![0.0; xtwx_symbolic.row_idx().len()];
let backend = if x.ncols() <= DENSE_OUTER_MAX_P {
XtWxBackend::Dense(DenseOuterState {
xtwx_dense: Array2::<f64>::zeros((x.ncols(), x.ncols())),
thread_buffers: Vec::new(),
})
} else {
let par = get_global_parallelism();
let scratch = MemBuffer::new(sparse_sparse_matmul_numeric_scratch::<usize, f64>(
xtwx_symbolic.as_ref(),
par,
));
XtWxBackend::Sparse(SparseSpGemmState {
wxvalues: vec![0.0; x.val().len()],
wx_tvalues: vec![0.0; x_t_csc.val().len()],
info,
scratch,
par,
})
};
Ok(Self {
xtwx_symbolic,
xtwxvalues,
nrows: x.nrows(),
ncols: x.ncols(),
nnz: x.val().len(),
x_col_ptr: x.symbolic().col_ptr().to_vec(),
xrow_idx: x.symbolic().row_idx().to_vec(),
x_t_csc,
backend,
})
}
pub(crate) fn matches(&self, x: &SparseColMat<usize, f64>) -> bool {
if self.nrows != x.nrows() || self.ncols != x.ncols() || self.nnz != x.val().len() {
return false;
}
let sym = x.symbolic();
self.x_col_ptr.as_slice() == sym.col_ptr() && self.xrow_idx.as_slice() == sym.row_idx()
}
pub(crate) fn compute_numeric(
&mut self,
x: &SparseColMat<usize, f64>,
weights: &Array1<f64>,
) -> Result<(), EstimationError> {
if weights.len() != self.nrows {
crate::bail_invalid_estim!(
"weights length {} does not match design rows {}",
weights.len(),
self.nrows
);
}
match &mut self.backend {
XtWxBackend::Dense(state) => {
state.compute(self.x_t_csc.as_ref(), weights, self.nrows, self.ncols);
let col_ptr = self.xtwx_symbolic.col_ptr();
let row_idx = self.xtwx_symbolic.row_idx();
let dense = &state.xtwx_dense;
for col in 0..self.ncols {
let start = col_ptr[col];
let end = col_ptr[col + 1];
for idx in start..end {
let row = row_idx[idx];
if row <= col {
self.xtwxvalues[idx] = dense[[row, col]];
}
}
}
}
XtWxBackend::Sparse(state) => state.compute(
x,
self.x_t_csc.as_ref(),
weights,
self.ncols,
self.xtwx_symbolic.as_ref(),
&mut self.xtwxvalues,
),
}
Ok(())
}
}
impl DenseOuterState {
pub(crate) fn compute(
&mut self,
x_t: SparseColMatRef<'_, usize, f64>,
weights: &Array1<f64>,
n: usize,
p: usize,
) {
assert_eq!(self.xtwx_dense.dim(), (p, p));
self.xtwx_dense.fill(0.0);
if n == 0 || p == 0 {
return;
}
let xtwx_start = std::time::Instant::now();
let nnz_total = x_t.symbolic().row_idx().len() as u64;
let work = nnz_total
.saturating_mul(nnz_total)
.checked_div(n as u64)
.unwrap_or(u64::MAX);
let n_threads = rayon::current_num_threads();
let parallelize = n_threads > 1 && work >= DENSE_OUTER_PARALLEL_FLOP_THRESHOLD;
if !parallelize {
accumulate_outer_upper(&mut self.xtwx_dense, x_t, weights, 0..n);
log::info!(
"[STAGE] PIRLS dense XᵀWX assembly (serial) n={} p={} flops~{} elapsed={:.3}s",
n,
p,
(n as u64).saturating_mul((p as u64).saturating_mul(p as u64)),
xtwx_start.elapsed().as_secs_f64(),
);
return;
}
if self.thread_buffers.len() != n_threads {
self.thread_buffers
.resize_with(n_threads, || Array2::<f64>::zeros((p, p)));
}
let chunk = n.div_ceil(n_threads);
self.thread_buffers
.par_iter_mut()
.enumerate()
.for_each(|(t, buf)| {
buf.fill(0.0);
let start = t * chunk;
let end = (start + chunk).min(n);
if start < end {
accumulate_outer_upper(buf, x_t, weights, start..end);
}
});
for buf in &self.thread_buffers {
self.xtwx_dense += buf;
}
log::info!(
"[STAGE] PIRLS dense XᵀWX assembly (parallel, threads={}) n={} p={} flops~{} elapsed={:.3}s",
rayon::current_num_threads(),
n,
p,
(n as u64).saturating_mul((p as u64).saturating_mul(p as u64)),
xtwx_start.elapsed().as_secs_f64(),
);
}
}
impl SparseSpGemmState {
pub(crate) fn compute(
&mut self,
x: &SparseColMat<usize, f64>,
x_t: SparseColMatRef<'_, usize, f64>,
weights: &Array1<f64>,
p: usize,
xtwx_symbolic: SymbolicSparseColMatRef<'_, usize>,
xtwxvalues: &mut [f64],
) {
let n = x_t.ncols();
assert_eq!(weights.len(), n);
assert!(weights.iter().all(|w| w.is_finite()));
let x_ref = x.as_ref();
for col in 0..p {
let xvals = x_ref.val_of_col(col);
let range = x_ref.col_range(col);
let dst = &mut self.wxvalues[range];
dst.copy_from_slice(xvals);
}
for col in 0..n {
let w = weights[col];
let xvals = x_t.val_of_col(col);
let range = x_t.col_range(col);
let dst = &mut self.wx_tvalues[range];
for (d, &s) in dst.iter_mut().zip(xvals.iter()) {
*d = s * w;
}
}
let wx_ref = SparseColMatRef::new(x.symbolic(), &self.wxvalues[..]);
let wx_t_ref = SparseColMatRef::new(x_t.symbolic(), &self.wx_tvalues[..]);
let stack = MemStack::new(&mut self.scratch);
let xtwxmut = SparseColMatMut::new(xtwx_symbolic, xtwxvalues);
sparse_sparse_matmul_numeric(
xtwxmut,
Accum::Replace,
wx_t_ref,
wx_ref,
1.0,
&self.info,
self.par,
stack,
);
}
}
#[inline]
pub(crate) fn accumulate_outer_upper(
acc: &mut Array2<f64>,
x_t: SparseColMatRef<'_, usize, f64>,
weights: &Array1<f64>,
rows: std::ops::Range<usize>,
) {
assert_eq!(acc.nrows(), acc.ncols());
let p = acc.ncols();
let acc_data = acc
.as_slice_mut()
.expect("dense XᵀWX accumulator is row-major and contiguous");
for i in rows {
let w_i = weights[i];
if w_i == 0.0 {
continue;
}
let cols = x_t.row_idx_of_col_raw(i);
let vals = x_t.val_of_col(i);
let nnz_i = cols.len();
for jj in 0..nnz_i {
let j = cols[jj].unbound();
let wvj = w_i * vals[jj];
let row = &mut acc_data[j * p..j * p + p];
for kk in jj..nnz_i {
let k = cols[kk].unbound();
row[k] += wvj * vals[kk];
}
}
}
}
fn dense_design_from_csr(
x_design_csr: &SparseRowMat<usize, f64>,
) -> Result<Array2<f64>, EstimationError> {
let n = x_design_csr.nrows();
let p = x_design_csr.ncols();
let mut x_dense = Array2::<f64>::zeros((n, p));
let xview = x_design_csr.as_ref();
for i in 0..n {
let vals = xview.val_of_row(i);
let cols = xview.col_idx_of_row_raw(i);
if cols.len() != vals.len() {
crate::bail_invalid_estim!(
"sparse row structure mismatch: column/value lengths differ"
);
}
for (idx, &col) in cols.iter().enumerate() {
x_dense[[i, col.unbound()]] = vals[idx];
}
}
Ok(x_dense)
}
pub(super) fn build_firth_design_factor_sparse(
x_design_csr: &SparseRowMat<usize, f64>,
observation_weights: ArrayView1<f64>,
) -> Result<FirthDesignFactor, EstimationError> {
let x_dense = dense_design_from_csr(x_design_csr)?;
FirthDenseOperator::build_design_factor_with_observation_weights(
&x_dense,
Some(observation_weights),
)
}
pub(super) fn build_firth_design_factor_dense(
x_design: ArrayView2<f64>,
observation_weights: ArrayView1<f64>,
) -> Result<FirthDesignFactor, EstimationError> {
FirthDenseOperator::build_design_factor_with_observation_weights(
&x_design.to_owned(),
Some(observation_weights),
)
}
pub(super) fn jeffreys_pirls_diagnostics_and_hessian_from_factor(
factor: &FirthDesignFactor,
link: &InverseLink,
eta: ArrayView1<f64>,
) -> Result<(Array1<f64>, f64, Array1<f64>, Array2<f64>), EstimationError> {
let op = FirthDenseOperator::build_from_design_factor(factor, link, &eta.to_owned())?;
let hat_diag = &op.w * &op.h_diag;
let mut score_shift = Array1::<f64>::zeros(op.w.len());
for i in 0..op.w.len() {
if op.w[i] > 0.0 {
score_shift[i] = 0.5 * (op.w1[i] / op.w[i]) * op.h_diag[i];
}
}
let diag_term = gam_linalg::faer_ndarray::fast_xt_diag_x(
&op.x_dense,
&(&op.w2 * &op.h_diag),
);
let bpb = gam_linalg::faer_ndarray::fast_atb(&op.b_base, &op.p_b_base);
let mut hphi = 0.5 * (diag_term - bpb);
gam_linalg::matrix::symmetrize_in_place(&mut hphi);
if !hphi.iter().all(|value| value.is_finite()) {
crate::bail_invalid_estim!("Firth/Jeffreys coefficient Hessian is non-finite");
}
Ok((hat_diag, op.jeffreys_logdet(), score_shift, hphi))
}
pub(crate) fn ensure_positive_definitewithridge(
hess: &mut Array2<f64>,
label: &str,
) -> Result<f64, EstimationError> {
let ridge = if FIXED_STABILIZATION_RIDGE > 0.0 {
FIXED_STABILIZATION_RIDGE
} else {
0.0
};
if !hess.iter().all(|value| value.is_finite()) {
crate::bail_invalid_estim!(
"{label}: assembled Hessian contains non-finite entries; refusing to factor"
);
}
if hess.cholesky(Side::Lower).is_ok() {
return Ok(0.0);
}
if ridge > 0.0 {
for i in 0..hess.nrows() {
hess[[i, i]] += ridge;
}
if hess.cholesky(Side::Lower).is_ok() {
log::debug!("{} stabilized with fixed ridge {:.1e}.", label, ridge);
return Ok(ridge);
}
}
if let Ok((evals, _)) = hess.eigh(Side::Lower) {
let min_eig = evals.iter().fold(f64::INFINITY, |a, &b| a.min(b));
return Err(EstimationError::HessianNotPositiveDefinite {
min_eigenvalue: min_eig,
});
}
Err(EstimationError::HessianNotPositiveDefinite {
min_eigenvalue: f64::NEG_INFINITY,
})
}
pub(super) fn solve_direction_with_dense_factor(
factor: &FaerSymmetricFactor,
gradient: &Array1<f64>,
direction_out: &mut Array1<f64>,
) {
if direction_out.len() != gradient.len() {
*direction_out = Array1::zeros(gradient.len());
}
direction_out.assign(gradient);
let mut rhsview = array1_to_col_matmut(direction_out);
factor.solve_in_place(rhsview.as_mut());
direction_out.mapv_inplace(|v| -v);
}
pub(super) fn solve_newton_direction_dense(
hessian: &Array2<f64>,
gradient: &Array1<f64>,
direction_out: &mut Array1<f64>,
) -> Result<(), EstimationError> {
let dense_solve_start = std::time::Instant::now();
let p = hessian.nrows();
if direction_out.len() != gradient.len() {
*direction_out = Array1::zeros(gradient.len());
}
if gam_gpu::cuda_selected()
.map_err(|error| EstimationError::InvalidInput(error.to_string()))?
{
let rhs = Array2::from_shape_vec((p, 1), gradient.to_vec()).map_err(|e| {
EstimationError::InvalidInput(format!("CUDA PIRLS RHS layout failed: {e}"))
})?;
let solved = crate::gpu::pirls_gpu::cholesky_solve_only_gpu(hessian.view(), rhs.view())
.map_err(EstimationError::InvalidInput)?;
direction_out.assign(&solved.column(0));
direction_out.mapv_inplace(|v| -v);
if array_is_finite(direction_out) {
log::info!(
"[STAGE] PIRLS dense newton solve backend=CUDA p={} flops~{} elapsed={:.3}s route=\"cuSOLVER potrf/potrs\"",
p,
(p as u64).saturating_mul((p as u64).saturating_mul(p as u64)) / 3,
dense_solve_start.elapsed().as_secs_f64(),
);
return Ok(());
}
}
let cpu_route = String::from("CPU stable solver");
let factor = StableSolver::new()
.factorize(hessian)
.map_err(EstimationError::LinearSystemSolveFailed)?;
solve_direction_with_dense_factor(&factor, gradient, direction_out);
let validation_residual = {
let h_delta = hessian.dot(direction_out);
h_delta
.iter()
.zip(gradient.iter())
.map(|(h, g)| (h + g).abs())
.fold(0.0_f64, f64::max)
};
let g_inf = gradient.iter().map(|v| v.abs()).fold(0.0_f64, f64::max);
let rel = validation_residual / (1.0 + g_inf);
if !rel.is_finite() || rel > 1.0e-3 {
return Err(EstimationError::InvalidInput(format!(
"PIRLS Newton direction failed its unperturbed linear-system certificate: relative residual {rel:.3e} exceeds 1e-3"
)));
}
if array_is_finite(direction_out) {
log::info!(
"[STAGE] PIRLS dense newton solve backend=CPU p={} flops~{} elapsed={:.3}s route=\"{}\"",
p,
(p as u64).saturating_mul((p as u64).saturating_mul(p as u64)) / 3,
dense_solve_start.elapsed().as_secs_f64(),
cpu_route,
);
return Ok(());
}
Err(EstimationError::LinearSystemSolveFailed(
FaerLinalgError::FactorizationFailed {
context: "PIRLS dense newton solve exhausted",
},
))
}
pub(super) fn solve_newton_direction_from_root_with_firth_hessian(
root: &Array2<f64>,
root_residual: &Array1<f64>,
firth_hessian: Option<&Array2<f64>>,
direction_out: &mut Array1<f64>,
) -> Result<f64, EstimationError> {
let p = root.ncols();
if root.nrows() < p || root_residual.len() != root.nrows() {
crate::bail_invalid_estim!(
"PIRLS square-root solve dimension mismatch: root={}x{}, residual={}",
root.nrows(),
p,
root_residual.len()
);
}
let (q, r) = root
.qr()
.map_err(EstimationError::LinearSystemSolveFailed)?;
if r.nrows() != p || r.ncols() != p {
crate::bail_invalid_estim!(
"PIRLS square-root QR produced non-square R={}x{} for p={p}",
r.nrows(),
r.ncols()
);
}
let projected_residual = q.t().dot(root_residual);
if direction_out.len() != p {
*direction_out = Array1::zeros(p);
}
for reverse in 0..p {
let i = p - 1 - reverse;
let mut value = -projected_residual[i];
for k in (i + 1)..p {
value -= r[[i, k]] * direction_out[k];
}
let diagonal = r[[i, i]];
if !(diagonal.is_finite() && diagonal != 0.0) {
return Err(EstimationError::ModelIsIllConditioned {
condition_number: f64::INFINITY,
});
}
direction_out[i] = value / diagonal;
}
let mut least_squares_residual = root.dot(direction_out);
least_squares_residual += root_residual;
let normal_residual = root.t().dot(&least_squares_residual);
let residual_inf = inf_norm(normal_residual.iter().copied());
let root_inf = root
.rows()
.into_iter()
.map(|row| row.iter().map(|value| value.abs()).sum::<f64>())
.fold(0.0_f64, f64::max);
let root_transpose_inf = root
.columns()
.into_iter()
.map(|column| column.iter().map(|value| value.abs()).sum::<f64>())
.fold(0.0_f64, f64::max);
let direction_inf = inf_norm(direction_out.iter().copied());
let root_residual_inf = inf_norm(root_residual.iter().copied());
let scale = root_transpose_inf * (root_inf * direction_inf + root_residual_inf);
let backward_error = if scale > 0.0 {
residual_inf / scale
} else {
residual_inf
};
let tolerance = 256.0 * f64::EPSILON * root.nrows().max(p) as f64;
if !backward_error.is_finite() || backward_error > tolerance {
crate::bail_invalid_estim!(
"PIRLS square-root Newton direction failed its backward-error certificate: \
error {backward_error:.3e} exceeds {tolerance:.3e}"
);
}
if !array_is_finite(direction_out) {
crate::bail_invalid_estim!("PIRLS square-root Newton direction is non-finite");
}
if let Some(hphi) = firth_hessian {
let fisher_direction = direction_out.clone();
let lower = r.t().to_owned();
correct_fisher_direction_for_firth_hessian_from_root_factor(
&lower,
hphi,
&fisher_direction,
direction_out,
)?;
}
log::info!(
"[STAGE] PIRLS dense newton solve backend=CPU p={} rows={} route=\"Householder QR of PSD root\" backward_error={:.3e} damped_decrement_sq={:.3e}",
p,
root.nrows(),
backward_error,
projected_residual.dot(&projected_residual),
);
Ok(projected_residual.dot(&projected_residual))
}
pub(super) struct TallSkinnyQrLeastSquares {
p: usize,
pending_root: Array2<f64>,
pending_residual: Array1<f64>,
pending_rows: usize,
triangular_root: Option<Array2<f64>>,
projected_residual: Array1<f64>,
total_rows: usize,
root_row_sum_max: f64,
root_column_sums: Array1<f64>,
residual_inf: f64,
}
impl TallSkinnyQrLeastSquares {
pub(super) fn new(p: usize) -> Result<Self, EstimationError> {
if p == 0 {
crate::bail_invalid_estim!("tall-skinny QR requires at least one coefficient");
}
Ok(Self {
p,
pending_root: Array2::zeros((p, p).f()),
pending_residual: Array1::zeros(p),
pending_rows: 0,
triangular_root: None,
projected_residual: Array1::zeros(p),
total_rows: 0,
root_row_sum_max: 0.0,
root_column_sums: Array1::zeros(p),
residual_inf: 0.0,
})
}
pub(super) fn push_row(
&mut self,
root_row: ndarray::ArrayView1<'_, f64>,
root_residual: f64,
) -> Result<(), EstimationError> {
if root_row.len() != self.p {
crate::bail_invalid_estim!(
"tall-skinny QR row width {} does not match p={}",
root_row.len(),
self.p
);
}
if !root_residual.is_finite() || root_row.iter().any(|value| !value.is_finite()) {
crate::bail_invalid_estim!("tall-skinny QR received a non-finite augmented row");
}
let row_sum = root_row.iter().map(|value| value.abs()).sum::<f64>();
self.root_row_sum_max = self.root_row_sum_max.max(row_sum);
for (sum, value) in self.root_column_sums.iter_mut().zip(root_row.iter()) {
*sum += value.abs();
}
self.residual_inf = self.residual_inf.max(root_residual.abs());
self.pending_root
.row_mut(self.pending_rows)
.assign(&root_row);
self.pending_residual[self.pending_rows] = root_residual;
self.pending_rows += 1;
self.total_rows += 1;
if self.pending_rows == self.p {
self.flush()?;
}
Ok(())
}
fn flush(&mut self) -> Result<(), EstimationError> {
if self.pending_rows == 0 {
return Ok(());
}
let carried_rows = usize::from(self.triangular_root.is_some()) * self.p;
let rows = carried_rows + self.pending_rows;
if rows < self.p {
return Ok(());
}
let mut root = Array2::<f64>::zeros((rows, self.p).f());
let mut residual = Array1::<f64>::zeros(rows);
if let Some(previous) = self.triangular_root.as_ref() {
root.slice_mut(ndarray::s![..self.p, ..]).assign(previous);
residual
.slice_mut(ndarray::s![..self.p])
.assign(&self.projected_residual);
}
let start = carried_rows;
let end = start + self.pending_rows;
root.slice_mut(ndarray::s![start..end, ..])
.assign(&self.pending_root.slice(ndarray::s![..self.pending_rows, ..]));
residual
.slice_mut(ndarray::s![start..end])
.assign(&self.pending_residual.slice(ndarray::s![..self.pending_rows]));
let (q, r) = root
.qr()
.map_err(EstimationError::LinearSystemSolveFailed)?;
if r.dim() != (self.p, self.p) {
crate::bail_invalid_estim!(
"tall-skinny QR produced R={}x{} for p={}",
r.nrows(),
r.ncols(),
self.p
);
}
self.projected_residual = q.t().dot(&residual);
self.triangular_root = Some(r);
self.pending_root.fill(0.0);
self.pending_residual.fill(0.0);
self.pending_rows = 0;
Ok(())
}
pub(super) fn solve(
mut self,
firth_hessian: Option<&Array2<f64>>,
direction_out: &mut Array1<f64>,
) -> Result<f64, EstimationError> {
self.flush()?;
let r = self.triangular_root.ok_or_else(|| {
EstimationError::InvalidInput(format!(
"tall-skinny QR has only {} rows for p={}",
self.total_rows, self.p
))
})?;
if direction_out.len() != self.p {
*direction_out = Array1::zeros(self.p);
}
for reverse in 0..self.p {
let i = self.p - 1 - reverse;
let mut value = -self.projected_residual[i];
for k in (i + 1)..self.p {
value -= r[[i, k]] * direction_out[k];
}
let diagonal = r[[i, i]];
if !(diagonal.is_finite() && diagonal != 0.0) {
return Err(EstimationError::ModelIsIllConditioned {
condition_number: f64::INFINITY,
});
}
direction_out[i] = value / diagonal;
}
let mut compact_residual = r.dot(direction_out);
compact_residual += &self.projected_residual;
let normal_residual = r.t().dot(&compact_residual);
let residual_inf = inf_norm(normal_residual.iter().copied());
let root_transpose_inf = inf_norm(self.root_column_sums.iter().copied());
let direction_inf = inf_norm(direction_out.iter().copied());
let scale =
root_transpose_inf * (self.root_row_sum_max * direction_inf + self.residual_inf);
let backward_error = if scale > 0.0 {
residual_inf / scale
} else {
residual_inf
};
let tolerance = 256.0 * f64::EPSILON * self.total_rows.max(self.p) as f64;
if !backward_error.is_finite() || backward_error > tolerance {
crate::bail_invalid_estim!(
"PIRLS tall-skinny square-root direction failed its backward-error certificate: \
error {backward_error:.3e} exceeds {tolerance:.3e}"
);
}
if !array_is_finite(direction_out) {
crate::bail_invalid_estim!("PIRLS tall-skinny square-root direction is non-finite");
}
let decrement = self.projected_residual.dot(&self.projected_residual);
if let Some(hphi) = firth_hessian {
let fisher_direction = direction_out.clone();
let lower = r.t().to_owned();
correct_fisher_direction_for_firth_hessian_from_root_factor(
&lower,
hphi,
&fisher_direction,
direction_out,
)?;
}
log::info!(
"[STAGE] PIRLS tall-skinny newton solve backend=CPU p={} rows={} route=\"blocked Householder QR of sparse PSD root\" backward_error={:.3e} damped_decrement_sq={:.3e}",
self.p,
self.total_rows,
backward_error,
decrement,
);
Ok(decrement)
}
}
fn correct_fisher_direction_for_firth_hessian_from_root_factor(
fisher_root_lower: &Array2<f64>,
firth_hessian: &Array2<f64>,
fisher_direction: &Array1<f64>,
direction_out: &mut Array1<f64>,
) -> Result<(), EstimationError> {
let p = fisher_root_lower.nrows();
if fisher_root_lower.ncols() != p
|| firth_hessian.dim() != (p, p)
|| fisher_direction.len() != p
{
crate::bail_invalid_estim!(
"Firth congruence solve dimension mismatch: root={}x{}, Hphi={}x{}, direction={}",
fisher_root_lower.nrows(),
fisher_root_lower.ncols(),
firth_hessian.nrows(),
firth_hessian.ncols(),
fisher_direction.len()
);
}
let left_whitened = gam_linalg::triangular::forward_substitution_lower_matrix(
fisher_root_lower,
firth_hessian,
);
let whitened_transpose = gam_linalg::triangular::forward_substitution_lower_matrix(
fisher_root_lower,
&left_whitened.t().to_owned(),
);
let mut congruence = Array2::<f64>::eye(p) - whitened_transpose.t().to_owned();
gam_linalg::matrix::symmetrize_in_place(&mut congruence);
let congruence_factor = congruence
.cholesky(Side::Lower)
.map_err(EstimationError::LinearSystemSolveFailed)?;
let transformed_rhs = fisher_root_lower.t().dot(fisher_direction);
let transformed_direction = congruence_factor.solvevec(&transformed_rhs);
let direction = gam_linalg::triangular::back_substitution_lower_transpose(
fisher_root_lower,
&transformed_direction,
);
let residual = &congruence.dot(&transformed_direction) - &transformed_rhs;
let residual_inf = inf_norm(residual.iter().copied());
let congruence_inf = congruence
.rows()
.into_iter()
.map(|row| row.iter().map(|value| value.abs()).sum::<f64>())
.fold(0.0_f64, f64::max);
let transformed_direction_inf = inf_norm(transformed_direction.iter().copied());
let transformed_rhs_inf = inf_norm(transformed_rhs.iter().copied());
let scale = congruence_inf * transformed_direction_inf + transformed_rhs_inf;
let backward_error = if scale > 0.0 {
residual_inf / scale
} else {
residual_inf
};
let tolerance = 256.0 * f64::EPSILON * p.max(1) as f64;
if !backward_error.is_finite() || backward_error > tolerance {
crate::bail_invalid_estim!(
"Firth congruence Newton direction failed its backward-error certificate: error {backward_error:.3e} exceeds {tolerance:.3e}"
);
}
if !array_is_finite(&direction) {
crate::bail_invalid_estim!("Firth congruence Newton direction is non-finite");
}
if direction_out.len() != p {
*direction_out = Array1::zeros(p);
}
direction_out.assign(&direction);
Ok(())
}
#[cfg(test)]
mod square_root_solve_tests {
use super::*;
use ndarray::array;
fn solve_newton_direction_from_root(
root: &Array2<f64>,
root_residual: &Array1<f64>,
direction_out: &mut Array1<f64>,
) -> Result<f64, EstimationError> {
solve_newton_direction_from_root_with_firth_hessian(
root,
root_residual,
None,
direction_out,
)
}
#[test]
fn qr_root_solve_preserves_a_weak_rotated_direction() {
let root = array![[1.0e5, 1.0e5], [1.0, -1.0], [1.0e-3, 0.0], [0.0, 1.0e-3]];
let expected = array![1.0, -1.0];
let residual = -root.dot(&expected);
let mut actual = Array1::<f64>::zeros(2);
let decrement = solve_newton_direction_from_root(&root, &residual, &mut actual)
.expect("square-root solve");
assert!((actual[0] - expected[0]).abs() < 1.0e-9);
assert!((actual[1] - expected[1]).abs() < 1.0e-9);
assert!(decrement.is_finite() && decrement > 0.0);
}
#[test]
fn qr_root_solve_does_not_form_a_cancelling_normal_rhs() {
let root = array![[1.0e8, 1.0e8], [1.0, -1.0], [0.0, 1.0]];
let expected = array![0.25, -0.25];
let orthogonal_residual = array![1.0e-8, -1.0, -2.0];
let residual = -root.dot(&expected) + &orthogonal_residual;
let mut actual = Array1::<f64>::zeros(2);
solve_newton_direction_from_root(&root, &residual, &mut actual)
.expect("least-squares root solve");
assert!((actual[0] - expected[0]).abs() < 1.0e-8);
assert!((actual[1] - expected[1]).abs() < 1.0e-8);
let mut stationary_direction = Array1::<f64>::zeros(2);
let stationary_decrement = solve_newton_direction_from_root(
&root,
&orthogonal_residual,
&mut stationary_direction,
)
.expect("stationary least-squares root certificate");
assert!(stationary_direction.dot(&stationary_direction) <= 1.0e-28);
assert!(stationary_decrement <= 1.0e-28);
}
#[test]
fn tall_skinny_qr_matches_dense_qr_for_a_stiff_root() {
let root = array![
[1.0e8, 1.0e8],
[1.0, -1.0],
[0.0, 1.0],
[2.0, 0.0],
[0.0, 3.0]
];
let expected = array![0.25, -0.25];
let residual = -root.dot(&expected);
let mut dense_direction = Array1::<f64>::zeros(2);
let dense_decrement =
solve_newton_direction_from_root(&root, &residual, &mut dense_direction)
.expect("dense square-root solve");
let mut blocked = TallSkinnyQrLeastSquares::new(2).expect("blocked QR");
for i in 0..root.nrows() {
blocked
.push_row(root.row(i), residual[i])
.expect("append augmented row");
}
let mut blocked_direction = Array1::<f64>::zeros(2);
let blocked_decrement = blocked
.solve(None, &mut blocked_direction)
.expect("blocked square-root solve");
for (&blocked_value, &dense_value) in
blocked_direction.iter().zip(dense_direction.iter())
{
assert!((blocked_value - dense_value).abs() < 1.0e-10);
}
assert!((blocked_decrement - dense_decrement).abs() < 1.0e-8);
}
#[test]
fn firth_congruence_preserves_the_root_direction_without_reforming_the_score() {
let root = array![[1.0e3, 1.0e3], [1.0, -1.0], [0.0, 1.0]];
let fisher_hessian = root.t().dot(&root);
let exact_root_direction = array![0.25, -0.25];
let root_residual = -root.dot(&exact_root_direction);
let firth_hessian = array![[0.20, 0.03], [0.03, 0.15]];
let mut corrected = Array1::<f64>::zeros(2);
solve_newton_direction_from_root_with_firth_hessian(
&root,
&root_residual,
Some(&firth_hessian),
&mut corrected,
)
.expect("exact Firth congruence solve");
let true_hessian = &fisher_hessian - &firth_hessian;
let residual = true_hessian.dot(&corrected) - fisher_hessian.dot(&exact_root_direction);
assert!(inf_norm(residual.iter().copied()) < 1.0e-8);
assert!(array_is_finite(&corrected));
}
}
pub fn solve_newton_direction_implicit<F>(
apply_xtwx: F,
xtwx_diag: ArrayView1<'_, f64>,
dense_penalties: &[(f64, &Array2<f64>)],
op_penalties: &[(f64, &dyn gam_terms::analytic_penalties::PenaltyOp)],
gradient: &Array1<f64>,
direction_out: &mut Array1<f64>,
ridge: f64,
rel_tol: f64,
max_iter: usize,
) -> Result<(), EstimationError>
where
F: Fn(&Array1<f64>) -> Array1<f64>,
{
let p = gradient.len();
if xtwx_diag.len() != p {
crate::bail_invalid_estim!(
"solve_newton_direction_implicit: xtwx_diag length {} != gradient length {}",
xtwx_diag.len(),
p
);
}
for (_, s) in dense_penalties.iter() {
if s.nrows() != p || s.ncols() != p {
crate::bail_invalid_estim!(
"solve_newton_direction_implicit: dense penalty dim {}×{} != p={}",
s.nrows(),
s.ncols(),
p
);
}
}
for (_, op) in op_penalties.iter() {
if op.dim() != p {
crate::bail_invalid_estim!(
"solve_newton_direction_implicit: op penalty dim {} != p={}",
op.dim(),
p
);
}
}
if direction_out.len() != p {
*direction_out = Array1::zeros(p);
}
let pcg_start = std::time::Instant::now();
let mut precond_diag = xtwx_diag.to_owned();
if ridge > 0.0 {
precond_diag.mapv_inplace(|d| d + ridge);
}
for (lambda, s) in dense_penalties.iter() {
if *lambda == 0.0 {
continue;
}
for i in 0..p {
precond_diag[i] += *lambda * s[[i, i]];
}
}
for (lambda, op) in op_penalties.iter() {
if *lambda == 0.0 {
continue;
}
let d = op.diag();
for i in 0..p {
precond_diag[i] += *lambda * d[i];
}
}
let apply_h = |v: &Array1<f64>| -> Array1<f64> {
let mut hv = apply_xtwx(v);
if ridge > 0.0 {
hv.zip_mut_with(v, |h, &x| *h += ridge * x);
}
for (lambda, s) in dense_penalties.iter() {
if *lambda == 0.0 {
continue;
}
let sv = fast_av(s, v);
hv.scaled_add(*lambda, &sv);
}
for (lambda, op) in op_penalties.iter() {
if *lambda == 0.0 {
continue;
}
let mut sv = Array1::<f64>::zeros(p);
op.matvec(v.view(), sv.view_mut());
hv.scaled_add(*lambda, &sv);
}
hv
};
let solution =
gam_linalg::utils::solve_spd_pcg(apply_h, gradient, &precond_diag, rel_tol, max_iter)
.ok_or(EstimationError::LinearSystemSolveFailed(
FaerLinalgError::FactorizationFailed {
context: "PIRLS implicit PCG solve exhausted",
},
))?;
direction_out.assign(&solution);
direction_out.mapv_inplace(|v| -v);
if !array_is_finite(direction_out) {
return Err(EstimationError::LinearSystemSolveFailed(
FaerLinalgError::FactorizationFailed {
context: "PIRLS implicit PCG non-finite direction",
},
));
}
log::info!(
"[STAGE] PIRLS implicit (PCG) newton solve p={} dense_pens={} op_pens={} elapsed={:.3}s",
p,
dense_penalties.len(),
op_penalties.len(),
pcg_start.elapsed().as_secs_f64(),
);
Ok(())
}
pub(super) fn project_coefficients_to_lower_bounds(
beta: &mut Array1<f64>,
lower_bounds: &Array1<f64>,
) {
for i in 0..beta.len() {
let lb = lower_bounds[i];
if lb.is_finite() && beta[i] < lb {
beta[i] = lb;
}
}
}
pub(crate) const ACTIVE_BOUND_REL_TOL: f64 = 1e-6;
pub(crate) const ACTIVE_BOUND_ABS_TOL: f64 = 1e-10;
pub(super) fn projected_gradient_norm(
gradient: &Array1<f64>,
beta: &Array1<f64>,
lower_bounds: Option<&Array1<f64>>,
) -> f64 {
let Some(lb) = lower_bounds else {
return gradient.dot(gradient).sqrt();
};
let mut sum_sq = 0.0;
for i in 0..gradient.len() {
let g = gradient[i];
if lb[i].is_finite() && g > 0.0 {
let slack = beta[i] - lb[i];
let scale = beta[i].abs().max(lb[i].abs()).max(1.0);
let tol = ACTIVE_BOUND_REL_TOL * scale + ACTIVE_BOUND_ABS_TOL;
if slack < tol {
continue;
}
}
sum_sq += g * g;
}
sum_sq.sqrt()
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) enum PirlsSoftAccept {
NearStationaryPlateau,
BoundarySaturation,
RelativeBandPlateau,
}
#[derive(Clone, Copy, Debug)]
pub(super) enum SoftAcceptProgress {
Realized { dev_change: f64 },
Predicted {
predicted_reduction: f64,
current_penalized: f64,
},
}
#[inline]
pub(super) fn pirls_soft_acceptance(
state: &WorkingState,
projected_grad: f64,
progress: SoftAcceptProgress,
max_abs_eta: f64,
progress_tol: f64,
kkt_tol: f64,
) -> Option<PirlsSoftAccept> {
let objective_scale = state.deviance.abs() + state.penalty_term.abs();
let scaled_dev_tol = progress_tol * objective_scale;
let near_stationary_plateau = match progress {
SoftAcceptProgress::Realized { dev_change } => {
state.near_stationary_kkt(projected_grad, kkt_tol) && dev_change.abs() < scaled_dev_tol
}
SoftAcceptProgress::Predicted {
predicted_reduction,
current_penalized,
} => {
let reduction_noise_floor = current_penalized.abs() * 1e-12;
state.near_stationary_kkt(projected_grad, kkt_tol)
&& predicted_reduction.abs() <= reduction_noise_floor
}
};
if near_stationary_plateau {
return Some(PirlsSoftAccept::NearStationaryPlateau);
}
let dev_change = match progress {
SoftAcceptProgress::Realized { dev_change } => dev_change,
SoftAcceptProgress::Predicted { .. } => return None,
};
if max_abs_eta >= PIRLS_ETA_ABS_CAP * (1.0 - 1e-12) && dev_change.abs() < scaled_dev_tol {
return Some(PirlsSoftAccept::BoundarySaturation);
}
if state.relative_gradient_norm(projected_grad) <= progress_tol.max(1e-6)
&& dev_change.abs() < scaled_dev_tol * 0.1
&& dev_change >= 0.0
{
return Some(PirlsSoftAccept::RelativeBandPlateau);
}
None
}
pub(super) fn constrained_stationarity_norm(
gradient: &Array1<f64>,
beta: &Array1<f64>,
lower_bounds: Option<&Array1<f64>>,
linear_constraints: Option<&LinearInequalityConstraints>,
) -> f64 {
if let Some(constraints) = linear_constraints {
let kkt = compute_constraint_kkt_diagnostics(beta, gradient, constraints);
return kkt
.primal_feasibility
.max(kkt.dual_feasibility)
.max(kkt.complementarity)
.max(kkt.stationarity);
}
projected_gradient_norm(gradient, beta, lower_bounds)
}
pub(crate) fn count_dense_upper_nnz(matrix: &Array2<f64>, tol: f64) -> usize {
let p = matrix.nrows().min(matrix.ncols());
let mut nnz = 0usize;
for col in 0..p {
for row in 0..=col {
if matrix[[row, col]].abs() > tol {
nnz += 1;
}
}
}
nnz
}
pub(crate) fn estimate_sparse_native_decision(
workspace: &mut PirlsWorkspace,
x_original: &DesignMatrix,
s_lambda: &Array2<f64>,
coefficient_lower_bounds: Option<&Array1<f64>>,
linear_constraints_original: Option<&LinearInequalityConstraints>,
) -> SparsePirlsDecision {
let p = x_original.ncols();
let nnz_s_lambda = count_dense_upper_nnz(s_lambda, 1e-12);
let dense_reject = |reason: &'static str, nnz_x: usize| SparsePirlsDecision {
path: PirlsLinearSolvePath::DenseTransformed,
reason,
p,
nnz_x,
nnz_xtwx_symbolic: None,
nnz_s_lambda,
nnz_h_est: None,
density_h_est: None,
};
let has_finite_lower_bounds = coefficient_lower_bounds
.map(|lb| lb.iter().any(|bound| bound.is_finite()))
.unwrap_or(false);
if has_finite_lower_bounds || linear_constraints_original.is_some() {
return dense_reject("constraints_present", 0);
}
let x_sparse = if let Some(sparse) = x_original.as_sparse() {
sparse
} else {
let row_chunk_start = std::time::Instant::now();
let n = x_original.nrows();
let chunk = row_chunk_for_byte_budget(n, x_original.ncols());
let mut nnz: usize = 0;
let mut chunks_processed = 0usize;
if chunk > 0 && n > 0 {
let mut start = 0;
while start < n {
let end = (start + chunk).min(n);
chunks_processed += 1;
match x_original.try_row_chunk(start..end) {
Ok(rows) => {
nnz = nnz.saturating_add(rows.iter().filter(|v| v.abs() > 1e-12).count());
}
Err(_) => {
nnz = nnz.saturating_add((end - start).saturating_mul(x_original.ncols()));
}
}
start = end;
}
}
log::info!(
"[STAGE] PIRLS row-chunk generation chunks={} n={} p={} nnz={} elapsed={:.3}s",
chunks_processed,
n,
x_original.ncols(),
nnz,
row_chunk_start.elapsed().as_secs_f64(),
);
return dense_reject("design_not_sparse", nnz);
};
let nnz_x = x_sparse.val().len();
match workspace.sparse_penalized_system_stats(x_sparse, s_lambda) {
Ok(stats) => SparsePirlsDecision {
path: if stats.density_upper <= SPARSE_NATIVE_MAX_H_DENSITY {
PirlsLinearSolvePath::SparseNative
} else {
PirlsLinearSolvePath::DenseTransformed
},
reason: if stats.density_upper <= SPARSE_NATIVE_MAX_H_DENSITY {
"sparse_native_eligible"
} else {
"penalized_hessian_too_dense"
},
p,
nnz_x,
nnz_xtwx_symbolic: Some(stats.nnz_xtwx_symbolic),
nnz_s_lambda: stats.nnz_s_lambda_upper,
nnz_h_est: Some(stats.nnz_h_upper),
density_h_est: Some(stats.density_upper),
},
Err(_) => dense_reject("sparse_stats_failed", nnz_x),
}
}
pub(super) fn should_use_sparse_native_pirls(
workspace: &mut PirlsWorkspace,
x_original: &DesignMatrix,
s_lambda: &Array2<f64>,
coefficient_lower_bounds: Option<&Array1<f64>>,
linear_constraints_original: Option<&LinearInequalityConstraints>,
) -> SparsePirlsDecision {
estimate_sparse_native_decision(
workspace,
x_original,
s_lambda,
coefficient_lower_bounds,
linear_constraints_original,
)
}
pub(super) fn ensure_sparse_positive_definitewithridge<F>(
mut assemble: F,
) -> Result<
(
SparseColMat<usize, f64>,
gam_linalg::sparse_exact::SparseExactFactor,
f64,
),
EstimationError,
>
where
F: FnMut(f64) -> Result<SparseColMat<usize, f64>, EstimationError>,
{
let h0 = assemble(0.0)?;
if let Ok(factor) = factorize_sparse_spd(&h0) {
return Ok((h0, factor, 0.0));
}
let h_eps = assemble(FIXED_STABILIZATION_RIDGE)?;
if let Ok(factor) = factorize_sparse_spd(&h_eps) {
return Ok((h_eps, factor, FIXED_STABILIZATION_RIDGE));
}
let (gershgorin_min, diag_scale) = gershgorin_min_eig_lower_bound(&h_eps);
let scale = diag_scale.max(1.0);
let margin = FIXED_STABILIZATION_RIDGE * scale;
let direct_ridge = (margin - gershgorin_min).max(FIXED_STABILIZATION_RIDGE);
log::warn!(
"sparse penalized Hessian is not positive definite (Gershgorin λ_min ≥ {:.3e}, \
diag scale {:.3e}); regularizing curvature with direct ridge {:.3e}. Exported \
curvature/SEs are stabilized, not exact — investigate rank-deficiency or weight \
underflow in the Hessian assembly.",
gershgorin_min,
scale,
direct_ridge,
);
for ridge in [direct_ridge, direct_ridge * 2.0] {
let h = assemble(ridge)?;
if let Ok(factor) = factorize_sparse_spd(&h) {
return Ok((h, factor, ridge));
}
}
Err(EstimationError::HessianNotPositiveDefinite {
min_eigenvalue: gershgorin_min,
})
}
pub(crate) fn gershgorin_min_eig_lower_bound(h: &SparseColMat<usize, f64>) -> (f64, f64) {
let n = h.ncols();
let mut diag = vec![0.0_f64; n];
let mut radius = vec![0.0_f64; n];
let (symbolic, values) = h.parts();
let col_ptr = symbolic.col_ptr();
let row_idx = symbolic.row_idx();
for col in 0..n {
let start = col_ptr[col];
let end = col_ptr[col + 1];
for idx in start..end {
let row = row_idx[idx];
let value = values[idx];
if row == col {
diag[col] += value;
} else {
let a = value.abs();
radius[row] += a;
radius[col] += a;
}
}
}
let mut min_bound = f64::INFINITY;
let mut diag_scale = 0.0_f64;
for i in 0..n {
min_bound = min_bound.min(diag[i] - radius[i]);
diag_scale = diag_scale.max(diag[i].abs());
}
if !min_bound.is_finite() {
min_bound = f64::NEG_INFINITY;
}
(min_bound, diag_scale)
}
pub(crate) fn solve_subsystem_direction(
h_sub: ndarray::ArrayView2<f64>,
g_sub: ndarray::ArrayView1<f64>,
out: &mut Array1<f64>,
) -> Result<(), EstimationError> {
let n = g_sub.len();
if out.len() != n {
*out = Array1::zeros(n);
}
let factor = StableSolver::new()
.factorize_any(&h_sub)
.map_err(EstimationError::LinearSystemSolveFailed)?;
out.assign(&g_sub);
let mut rhs = array1_to_col_matmut(out);
factor.solve_in_place(rhs.as_mut());
out.mapv_inplace(|value| -value);
if array_is_finite(out) {
Ok(())
} else {
Err(EstimationError::InvalidInput(
"PIRLS constrained subsystem solve produced a non-finite direction".to_string(),
))
}
}
pub(super) fn linear_constraints_from_lower_bounds(
lower_bounds: &Array1<f64>,
) -> Option<LinearInequalityConstraints> {
LinearInequalityConstraints::from_per_coordinate_lower_bounds(lower_bounds)
}
pub(super) fn compute_constraint_kkt_diagnostics(
beta: &Array1<f64>,
gradient: &Array1<f64>,
constraints: &LinearInequalityConstraints,
) -> ConstraintKktDiagnostics {
active_set::compute_constraint_kkt_diagnostics(beta, gradient, constraints)
}
pub(super) fn select_active_set_release(
gradient: &Array1<f64>,
hd: &Array1<f64>,
active_idx: &[usize],
use_blands: bool,
) -> Option<usize> {
if use_blands {
for &i in active_idx {
let lambda_i = gradient[i] + hd[i];
let scale = gradient[i].abs().max(hd[i].abs()).max(1.0);
let tol = 64.0 * f64::EPSILON * scale;
if lambda_i < -tol {
return Some(i);
}
}
None
} else {
let mut worst = 0.0_f64;
let mut idx = None;
for &i in active_idx {
let lambda_i = gradient[i] + hd[i];
let tol = 64.0 * f64::EPSILON * gradient[i].abs().max(hd[i].abs()).max(1.0);
if lambda_i < -tol && lambda_i < worst {
worst = lambda_i;
idx = Some(i);
}
}
idx
}
}
pub fn solve_newton_directionwith_lower_bounds(
hessian: &Array2<f64>,
gradient: &Array1<f64>,
beta: &Array1<f64>,
lower_bounds: &Array1<f64>,
direction_out: &mut Array1<f64>,
active_hint: Option<&mut Vec<usize>>,
) -> Result<(), EstimationError> {
let p = gradient.len();
if lower_bounds.len() != p || beta.len() != p {
crate::bail_invalid_estim!(
"lower-bound size mismatch: beta={}, gradient={}, bounds={}",
beta.len(),
gradient.len(),
lower_bounds.len()
);
}
if direction_out.len() != p {
*direction_out = Array1::zeros(p);
}
direction_out.fill(0.0);
let has_active_hint = active_hint
.as_ref()
.map(|hint| !hint.is_empty())
.unwrap_or(false);
if !has_active_hint && solve_newton_direction_dense(hessian, gradient, direction_out).is_ok() {
let mut feasible = true;
for i in 0..p {
let lb = lower_bounds[i];
if lb.is_finite() && beta[i] + direction_out[i] < lb {
feasible = false;
break;
}
}
if feasible {
return Ok(());
}
}
let mut active = vec![false; p];
if let Some(hint) = active_hint.as_ref() {
for &idx in hint.iter() {
if idx < p {
active[idx] = true;
}
}
}
for i in 0..p {
let lb = lower_bounds[i];
if lb.is_finite() && gradient[i] > 0.0 {
let scale = beta[i].abs().max(lb.abs()).max(1.0);
let tol = ACTIVE_BOUND_REL_TOL * scale + ACTIVE_BOUND_ABS_TOL;
if beta[i] <= lb + tol {
active[i] = true;
}
}
}
const BLANDS_RULE_GRACE: usize = 2;
let blands_threshold = BLANDS_RULE_GRACE * (p + 1);
let max_iters = 8 * (p + 1);
let mut d_free = Array1::<f64>::zeros(p);
let mut h_ff_buf = Array2::<f64>::zeros((p, p));
let mut g_f_buf = Array1::<f64>::zeros(p);
for it in 0..max_iters {
let use_blands = it >= blands_threshold;
let free_idx: Vec<usize> = (0..p).filter(|&i| !active[i]).collect();
let active_idx: Vec<usize> = (0..p).filter(|&i| active[i]).collect();
direction_out.fill(0.0);
for &i in &active_idx {
let lb = lower_bounds[i];
if lb.is_finite() {
direction_out[i] = lb - beta[i];
}
}
if free_idx.is_empty() {
let hd = fast_av(hessian, direction_out);
if let Some(idx) = select_active_set_release(gradient, &hd, &active_idx, use_blands) {
active[idx] = false;
continue;
}
if let Some(hint) = active_hint {
hint.clear();
hint.extend((0..p).filter(|&i| active[i]));
}
return Ok(());
}
let n_free = free_idx.len();
{
let mut h_ff = h_ff_buf.slice_mut(ndarray::s![..n_free, ..n_free]);
let mut g_f = g_f_buf.slice_mut(ndarray::s![..n_free]);
for (ii, &i) in free_idx.iter().enumerate() {
let mut gi = gradient[i];
for &j in &active_idx {
gi += hessian[[i, j]] * direction_out[j];
}
g_f[ii] = gi;
for (jj, &j) in free_idx.iter().enumerate() {
h_ff[[ii, jj]] = hessian[[i, j]];
}
}
}
solve_subsystem_direction(
h_ff_buf.slice(ndarray::s![..n_free, ..n_free]),
g_f_buf.slice(ndarray::s![..n_free]),
&mut d_free,
)?;
for (ii, &i) in free_idx.iter().enumerate() {
direction_out[i] = d_free[ii];
}
let mut hit_idx: Option<usize> = None;
let mut best_alpha = 1.0_f64;
for &i in &free_idx {
let lb = lower_bounds[i];
if !lb.is_finite() {
continue;
}
let slack = beta[i] - lb;
let di = direction_out[i];
if let Some(alpha_i) = boundary_hit_step_fraction(slack, di, best_alpha) {
best_alpha = alpha_i;
hit_idx = Some(i);
}
}
if let Some(i_hit) = hit_idx {
for i in 0..p {
direction_out[i] *= best_alpha;
}
active[i_hit] = true;
continue;
}
let hd = fast_av(hessian, direction_out);
if let Some(idx) = select_active_set_release(gradient, &hd, &active_idx, use_blands) {
active[idx] = false;
continue;
}
if let Some(hint) = active_hint {
hint.clear();
hint.extend((0..p).filter(|&i| active[i]));
}
return Ok(());
}
Err(EstimationError::InvalidInput(format!(
"lower-bound active-set QP did not reach a consistent primal/dual KKT set in {max_iters} pivots"
)))
}
pub(super) fn solve_newton_directionwith_linear_constraints(
hessian: &Array2<f64>,
gradient: &Array1<f64>,
beta: &Array1<f64>,
constraints: &LinearInequalityConstraints,
direction_out: &mut Array1<f64>,
active_hint: Option<&mut Vec<usize>>,
) -> Result<(), EstimationError> {
active_set::solve_newton_direction_with_linear_constraints(
hessian,
gradient,
beta,
constraints,
direction_out,
active_hint,
)
}