use std::collections::BTreeMap;
use crate::nl_reader::NlProblem;
use pounce_common::types::{Number, lower_bound_present, upper_bound_present};
const RUIZ_SWEEPS: usize = 10;
const SCALE_LO: Number = 1e-8;
const SCALE_HI: Number = 1e8;
#[derive(Debug, Clone)]
pub struct QuadRowCoef {
pub index: usize,
pub curvature: Number,
pub linear: Number,
pub rhs: Number,
}
pub fn quad_row_coefs(prob: &NlProblem) -> Vec<QuadRowCoef> {
let mut out = Vec::new();
for i in 0..prob.m {
let Some((hess, nl_lin, nl_const)) = prob.con_nonlinear[i].analyze_quadratic_full() else {
continue;
};
if hess.is_empty() {
continue; }
let mut row_sum: BTreeMap<usize, Number> = BTreeMap::new();
for (&(r, c), v) in &hess {
let a = v.abs();
*row_sum.entry(r).or_insert(0.0) += a;
if r != c {
*row_sum.entry(c).or_insert(0.0) += a;
}
}
let curvature = row_sum.values().fold(0.0_f64, |m, &v| m.max(v));
let linear = row_linear(prob, i, &nl_lin)
.values()
.fold(0.0_f64, |m, &v| m.max(v.abs()));
out.push(QuadRowCoef {
index: i,
curvature,
linear,
rhs: row_rhs(prob, i, nl_const),
});
}
out
}
fn row_linear(prob: &NlProblem, i: usize, nl_lin: &[(usize, Number)]) -> BTreeMap<usize, Number> {
let mut lin: BTreeMap<usize, Number> = BTreeMap::new();
for (var, coef) in &prob.con_linear[i] {
*lin.entry(*var).or_insert(0.0) += *coef;
}
for (var, coef) in nl_lin {
*lin.entry(*var).or_insert(0.0) += *coef;
}
lin
}
fn row_rhs(prob: &NlProblem, i: usize, nl_const: Number) -> Number {
let (lo, hi) = (prob.g_l[i], prob.g_u[i]);
let mut rhs = 0.0_f64;
if lower_bound_present(lo) {
rhs = rhs.max((lo - nl_const).abs());
}
if upper_bound_present(hi) {
rhs = rhs.max((hi - nl_const).abs());
}
rhs
}
#[derive(Debug, Clone)]
pub struct CurvatureScaling {
pub x: Vec<Number>,
pub g: Vec<Number>,
pub quadratic: bool,
}
struct QuadRow {
hess: BTreeMap<(usize, usize), Number>,
lin: BTreeMap<usize, Number>,
rhs: Number,
}
pub fn curvature_scaling(prob: &NlProblem) -> Option<CurvatureScaling> {
let n = prob.n;
let m = prob.m;
let (obj_hess, _obj_lin, _obj_const) = prob.obj_nonlinear.analyze_quadratic_full()?;
let mut rows: Vec<QuadRow> = Vec::with_capacity(m);
for i in 0..m {
let (hess, nl_lin, nl_const) = prob.con_nonlinear[i].analyze_quadratic_full()?;
rows.push(QuadRow {
hess,
lin: row_linear(prob, i, &nl_lin),
rhs: row_rhs(prob, i, nl_const),
});
}
let mut p_hat: BTreeMap<(usize, usize), Number> = BTreeMap::new();
let bump = |map: &mut BTreeMap<(usize, usize), Number>, key, v: Number| {
let slot = map.entry(key).or_insert(0.0);
if v > *slot {
*slot = v;
}
};
for (&k, v) in &obj_hess {
bump(&mut p_hat, k, v.abs());
}
let mut j_hat: Vec<BTreeMap<usize, Number>> = Vec::with_capacity(m);
for QuadRow { hess, lin, .. } in &rows {
for (&k, v) in hess {
bump(&mut p_hat, k, v.abs());
}
let mut row: BTreeMap<usize, Number> = BTreeMap::new();
for (&(r, c), v) in hess {
let a = v.abs();
let slot = row.entry(r).or_insert(0.0);
if a > *slot {
*slot = a;
}
let slot = row.entry(c).or_insert(0.0);
if a > *slot {
*slot = a;
}
}
for (&j, v) in lin {
let a = v.abs();
let slot = row.entry(j).or_insert(0.0);
if a > *slot {
*slot = a;
}
}
j_hat.push(row);
}
let dim = n + m;
let mut s = vec![1.0_f64; dim];
let mut rownorm = vec![0.0_f64; dim];
for _ in 0..RUIZ_SWEEPS {
rownorm.iter_mut().for_each(|v| *v = 0.0);
for (&(r, c), v) in &p_hat {
let x = (s[r] * v * s[c]).abs();
if x > rownorm[r] {
rownorm[r] = x;
}
if r != c && x > rownorm[c] {
rownorm[c] = x;
}
}
for (i, row) in j_hat.iter().enumerate() {
let ri = n + i;
for (&j, v) in row {
let x = (s[ri] * v * s[j]).abs();
if x > rownorm[ri] {
rownorm[ri] = x;
}
if x > rownorm[j] {
rownorm[j] = x;
}
}
}
for i in 0..dim {
if rownorm[i] > 0.0 {
s[i] /= rownorm[i].sqrt();
}
}
}
let d: Vec<Number> = s[..n].iter().map(|v| v.clamp(SCALE_LO, SCALE_HI)).collect();
let mut g = vec![1.0_f64; m];
for (i, QuadRow { hess, lin, rhs }) in rows.iter().enumerate() {
let mut row_sum: BTreeMap<usize, Number> = BTreeMap::new();
for (&(r, c), v) in hess {
let scaled = (d[r] * v * d[c]).abs();
*row_sum.entry(r).or_insert(0.0) += scaled;
if r != c {
*row_sum.entry(c).or_insert(0.0) += scaled;
}
}
let q_norm = row_sum.values().fold(0.0_f64, |m, &v| m.max(v));
let a_norm = lin
.iter()
.fold(0.0_f64, |m, (&j, v)| m.max((v * d[j]).abs()));
let scale = q_norm.max(a_norm).max(*rhs);
if scale > 0.0 {
g[i] = (1.0 / scale).clamp(SCALE_LO, SCALE_HI);
}
}
Some(CurvatureScaling {
x: d.iter().map(|v| 1.0 / v).collect(),
g,
quadratic: obj_hess
.values()
.chain(rows.iter().flat_map(|r| r.hess.values()))
.any(|v| *v != 0.0),
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::nl_reader::parse_nl_text;
const SPREAD_NL: &str = "\
g3 0 1 0
2 1 1 0 0
1 1
0 0
2 2 2
0 0 0 1
0 0 0 0 0
2 2
0 0
0 0 0 0 0
b
3
3
r
1 100000
C0
o54
2
o2
n0.5
o2
o2
n4.0
v0
v0
o2
n0.5
o2
o2
n2e-8
v1
v1
O0 0
o54
2
o2
n0.5
o2
o2
n2.0
v0
v0
o2
n0.5
o2
o2
n2e-8
v1
v1
k1
1
J0 2
0 0
1 0
G0 2
0 -1.0
1 -1.0
";
const NONLINEAR_NL: &str = "\
g3 0 1 0
1 1 1 0 0
1 0
0 0
1 1 1
0 0 0 1
0 0 0 0 0
1 1
0 0
0 0 0 0 0
b
3
r
1 2
C0
o44
v0
O0 0
n0
k0
J0 1
0 0
G0 1
0 1.0
";
#[test]
fn a_genuine_nonlinearity_is_declined_not_approximated() {
let prob = parse_nl_text(NONLINEAR_NL).expect("parse");
assert!(
curvature_scaling(&prob).is_none(),
"exp(x) has no constant Hessian; approximating it from the \
linear section would equilibrate against a fiction"
);
}
#[test]
fn a_quadratic_model_is_accepted() {
let prob = parse_nl_text(SPREAD_NL).expect("parse");
let sc = curvature_scaling(&prob).expect("degree ≤ 2 everywhere");
assert_eq!(sc.x.len(), 2);
assert_eq!(sc.g.len(), 1);
assert!(sc.x.iter().all(|v| v.is_finite() && *v > 0.0));
assert!(sc.g.iter().all(|v| v.is_finite() && *v > 0.0));
}
#[test]
fn x_factors_invert_d() {
let prob = parse_nl_text(SPREAD_NL).expect("parse");
let sc = curvature_scaling(&prob).expect("quadratic");
assert!(
sc.x[1] < sc.x[0],
"the small-coefficient variable must be shrunk, not grown: \
d = {:?}",
sc.x
);
assert!(
sc.x[0] / sc.x[1] > 1e2,
"expected a large ratio, got {:?}",
sc.x
);
}
#[test]
fn the_row_scale_normalizes_the_row() {
let prob = parse_nl_text(SPREAD_NL).expect("parse");
let sc = curvature_scaling(&prob).expect("quadratic");
let d: Vec<f64> = sc.x.iter().map(|v| 1.0 / v).collect();
let (hess, nl_lin, nl_const) = prob.con_nonlinear[0]
.analyze_quadratic_full()
.expect("quadratic row");
let mut row_sum: BTreeMap<usize, f64> = BTreeMap::new();
for (&(r, c), v) in &hess {
let s = (d[r] * v * d[c]).abs();
*row_sum.entry(r).or_insert(0.0) += s;
if r != c {
*row_sum.entry(c).or_insert(0.0) += s;
}
}
let q = row_sum.values().fold(0.0_f64, |m, &v| m.max(v));
let a = row_linear(&prob, 0, &nl_lin)
.iter()
.fold(0.0_f64, |m, (&j, v)| m.max((v * d[j]).abs()));
let b = row_rhs(&prob, 0, nl_const);
let scaled_max = q.max(a).max(b) * sc.g[0];
assert!(
(scaled_max - 1.0).abs() < 1e-12,
"scaled row magnitude should be exactly 1, got {scaled_max}"
);
}
#[test]
fn quad_row_coefs_reads_the_file() {
let prob = parse_nl_text(SPREAD_NL).expect("parse");
let rows = quad_row_coefs(&prob);
assert_eq!(rows.len(), 1);
let r = &rows[0];
assert_eq!(r.index, 0);
assert!((r.curvature - 4.0).abs() < 1e-12);
assert_eq!(r.linear, 0.0);
assert!((r.rhs - 1.0e5).abs() < 1e-9);
}
#[test]
fn quadratic_is_false_exactly_when_no_second_order_coefficient_exists() {
let lp = "\
g3 0 1 0
2 1 1 0 0
0 0
0 0
0 0 0
0 0 0 1
0 0 0 0 0
2 2
0 0
0 0 0 0 0
C0
n0
O0 0
n0
x2
0 0
1 0
r
2 1
b
2 0
2 0
k1
2
J0 2
0 1
1 1
G0 2
0 1
1 1
";
let prob = crate::nl_reader::parse_nl_text(lp).expect("parse LP");
let sc = curvature_scaling(&prob).expect("an LP is degree <= 2");
assert!(
!sc.quadratic,
"an LP has no `Q` to read; factors {:?} / {:?}",
sc.x, sc.g
);
let qp = lp.replace(
"O0 0
n0",
"O0 0
o5
v0
n2",
);
let prob = crate::nl_reader::parse_nl_text(&qp).expect("parse QP");
let sc = curvature_scaling(&prob).expect("x0^2 is degree 2");
assert!(sc.quadratic, "`x0^2` is a second-order coefficient");
}
}