use crate::linear_constraints::LinearInequalityConstraints;
use ndarray::{Array1, Array2, ArrayView1};
use rayon::iter::{IntoParallelIterator, ParallelIterator};
use std::sync::Arc;
pub const PRIMAL_FEASIBILITY_TOL: f64 = 1e-8;
enum RowMetrics<'a> {
Tiled {
norms: &'a [f64],
tile: usize,
bounds: Option<&'a [f64]>,
},
Carrier(&'a ConstraintSet),
}
impl RowMetrics<'_> {
#[inline]
fn read(&self, row: usize) -> Result<(f64, f64), String> {
match self {
RowMetrics::Tiled {
norms,
tile,
bounds,
} => {
let norm = norms[row % tile];
let bound = bounds.map_or(0.0, |values| values[row]);
Ok((norm, bound))
}
RowMetrics::Carrier(set) => Ok((set.row_norm(row)?, set.bound(row)?)),
}
}
}
#[derive(Clone, Debug)]
enum SweepTerminal {
RowUnavailable(String),
Undecidable { norm: f64, bound: f64, value: f64 },
VacuousRowWithPositiveBound,
}
struct ScaledViolationSweep {
terminal: Option<(usize, SweepTerminal)>,
worst: f64,
worst_row: Option<usize>,
}
impl ScaledViolationSweep {
fn none() -> Self {
Self {
terminal: None,
worst: 0.0,
worst_row: None,
}
}
fn record_terminal(&mut self, row: usize, terminal: SweepTerminal) {
let keep = match self.terminal {
Some((seen, _)) => row < seen,
None => true,
};
if keep {
self.terminal = Some((row, terminal));
}
}
fn record_violation(&mut self, row: usize, violation: f64) {
if violation > self.worst {
self.worst = violation;
self.worst_row = Some(row);
}
}
fn merge(mut self, other: Self) -> Self {
if let Some((row, terminal)) = other.terminal {
self.record_terminal(row, terminal);
}
let take_other = match (other.worst > self.worst, other.worst == self.worst) {
(true, _) => true,
(false, true) => match (other.worst_row, self.worst_row) {
(Some(candidate), Some(held)) => candidate < held,
(Some(_), None) => true,
_ => false,
},
_ => false,
};
if take_other {
self.worst = other.worst;
self.worst_row = other.worst_row;
}
self
}
fn verdict(self) -> Result<(f64, Option<usize>), String> {
match self.terminal {
Some((row, SweepTerminal::RowUnavailable(error))) => Err(format!(
"ConstraintSet::max_scaled_violation: row {row} has no readable norm or \
bound: {error}"
)),
Some((
row,
SweepTerminal::Undecidable {
norm,
bound,
value,
},
)) => Err(format!(
"ConstraintSet::max_scaled_violation: row {row} cannot be decided \
(row norm {norm:.3e}, bound {bound:.3e}, value {value:.3e}); \
feasibility of a non-finite iterate is undefined and every \
comparison in the sweep is false for NaN, so the row cannot \
be skipped (gam#2721)"
)),
Some((row, SweepTerminal::VacuousRowWithPositiveBound)) => {
Ok((f64::INFINITY, Some(row)))
}
None => Ok((self.worst, self.worst_row)),
}
}
}
pub fn feasibility_quantities_are_finite(quantities: &[f64]) -> bool {
quantities.iter().all(|q| q.is_finite())
}
pub(crate) fn contract_feasible_step_over_rows<B, N>(
values: &Array1<f64>,
directional: &Array1<f64>,
bound: B,
row_norm: N,
) -> Result<ContractFeasibleStep, ContractFeasibleStepError>
where
B: Fn(usize) -> Result<f64, String>,
N: Fn(usize) -> Result<f64, String>,
{
let tol = PRIMAL_FEASIBILITY_TOL;
let mut limit = ContractFeasibleStep::UNLIMITED;
for row in 0..values.len() {
let norm = row_norm(row).map_err(ContractFeasibleStepError::Carrier)?;
let bound = bound(row).map_err(ContractFeasibleStepError::Carrier)?;
if !feasibility_quantities_are_finite(&[norm, bound, values[row], directional[row]]) {
return Err(ContractFeasibleStepError::NonFinite {
row,
scaled_slack: (values[row] - bound) / norm,
scaled_drift: directional[row] / norm,
});
}
if !(norm.is_finite() && norm > 0.0) {
if bound > 0.0 {
return Err(ContractFeasibleStepError::InfeasibleIterate {
row,
scaled_slack: f64::NEG_INFINITY,
});
}
continue;
}
let slack = (values[row] - bound) / norm;
let drift = directional[row] / norm;
if slack < -tol {
return Err(ContractFeasibleStepError::InfeasibleIterate {
row,
scaled_slack: slack,
});
}
if drift >= 0.0 {
continue;
}
if slack + drift >= -tol {
continue;
}
let fraction = (slack.max(0.0) / -drift).clamp(0.0, 1.0);
if fraction < limit.fraction {
limit = ContractFeasibleStep {
fraction,
blocking_row: Some(row),
blocking_scaled_slack: slack,
blocking_scaled_drift: drift,
};
}
}
Ok(limit)
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct ContractFeasibleStep {
pub fraction: f64,
pub blocking_row: Option<usize>,
pub blocking_scaled_slack: f64,
pub blocking_scaled_drift: f64,
}
impl ContractFeasibleStep {
pub const UNLIMITED: Self = Self {
fraction: 1.0,
blocking_row: None,
blocking_scaled_slack: f64::INFINITY,
blocking_scaled_drift: 0.0,
};
}
#[derive(Clone, Debug, PartialEq)]
pub enum ContractFeasibleStepError {
Dimension {
beta: usize,
direction: usize,
expected: usize,
},
InfeasibleIterate { row: usize, scaled_slack: f64 },
NonFinite {
row: usize,
scaled_slack: f64,
scaled_drift: f64,
},
Carrier(String),
}
impl std::fmt::Display for ContractFeasibleStepError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ContractFeasibleStepError::Dimension {
beta,
direction,
expected,
} => write!(
f,
"constraint step dimension mismatch: beta={beta}, direction={direction}, constraints={expected}"
),
ContractFeasibleStepError::InfeasibleIterate { row, scaled_slack } => write!(
f,
"current iterate violates constraint row {row}: scaled slack={scaled_slack:.3e} \
below the primal-feasibility contract {PRIMAL_FEASIBILITY_TOL:.3e}"
),
ContractFeasibleStepError::NonFinite {
row,
scaled_slack,
scaled_drift,
} => write!(
f,
"constraint row {row} has a non-finite ratio test: scaled slack={scaled_slack:.3e}, \
scaled drift={scaled_drift:.3e}"
),
ContractFeasibleStepError::Carrier(reason) => write!(f, "{reason}"),
}
}
}
#[derive(Clone, Debug)]
pub struct KhatriRaoConeConstraints {
factor: Arc<Array2<f64>>,
factor_row_norms: Array1<f64>,
coupled_rows: Vec<usize>,
p_left: usize,
bounds: Option<Array1<f64>>,
}
impl KhatriRaoConeConstraints {
pub fn new(
factor: Arc<Array2<f64>>,
coupled_rows: Vec<usize>,
p_left: usize,
) -> Result<Self, String> {
if factor.nrows() == 0 || factor.ncols() == 0 {
return Err("KhatriRaoConeConstraints: factor must be non-empty".to_string());
}
if factor.iter().any(|v| !v.is_finite()) {
return Err("KhatriRaoConeConstraints: factor must be finite".to_string());
}
if coupled_rows.is_empty() {
return Err(
"KhatriRaoConeConstraints: at least one coupled coefficient row is required"
.to_string(),
);
}
let mut seen = vec![false; p_left];
for &k in &coupled_rows {
if k >= p_left {
return Err(format!(
"KhatriRaoConeConstraints: coupled row {k} out of range (p_left = {p_left})"
));
}
if seen[k] {
return Err(format!(
"KhatriRaoConeConstraints: coupled row {k} is duplicated"
));
}
seen[k] = true;
}
let factor_row_norms =
Array1::from_iter(factor.rows().into_iter().map(|row| row.dot(&row).sqrt()));
Ok(Self {
factor,
factor_row_norms,
coupled_rows,
p_left,
bounds: None,
})
}
pub fn factor(&self) -> &Array2<f64> {
self.factor.as_ref()
}
pub fn coupled_rows(&self) -> &[usize] {
&self.coupled_rows
}
pub fn p_left(&self) -> usize {
self.p_left
}
pub fn single_coupled_slot(&self, slot: usize) -> Result<Self, String> {
if slot >= self.coupled_rows.len() {
return Err(format!(
"KhatriRaoConeConstraints: coupled slot {slot} out of range ({} slots)",
self.coupled_rows.len()
));
}
let n = self.factor.nrows();
let bounds = self
.bounds
.as_ref()
.map(|all| all.slice(ndarray::s![slot * n..(slot + 1) * n]).to_owned());
Ok(Self {
factor: Arc::clone(&self.factor),
factor_row_norms: self.factor_row_norms.clone(),
coupled_rows: vec![0],
p_left: 1,
bounds,
})
}
pub fn nrows(&self) -> usize {
self.coupled_rows.len() * self.factor.nrows()
}
pub fn ncols(&self) -> usize {
self.p_left * self.factor.ncols()
}
#[inline]
fn split_row_id(&self, row: usize) -> Result<(usize, usize), String> {
let n = self.factor.nrows();
let slot = row / n;
if slot >= self.coupled_rows.len() {
return Err(format!(
"KhatriRaoConeConstraints: row id {row} out of range ({} rows)",
self.nrows()
));
}
Ok((slot, row % n))
}
pub fn values(&self, beta: ArrayView1<'_, f64>) -> Result<Array1<f64>, String> {
let p_cov = self.factor.ncols();
if beta.len() != self.ncols() {
return Err(format!(
"KhatriRaoConeConstraints: beta length {} != {}",
beta.len(),
self.ncols()
));
}
let n = self.factor.nrows();
let slots = self.coupled_rows.len();
let mut blocks = Array2::<f64>::zeros((p_cov, slots));
for (slot, &k) in self.coupled_rows.iter().enumerate() {
blocks
.column_mut(slot)
.assign(&beta.slice(ndarray::s![k * p_cov..(k + 1) * p_cov]));
}
let alpha = self.factor.dot(&blocks);
let mut out = Array1::<f64>::zeros(self.nrows());
for slot in 0..slots {
out.slice_mut(ndarray::s![slot * n..(slot + 1) * n])
.assign(&alpha.column(slot));
}
Ok(out)
}
pub fn row_norm(&self, row: usize) -> Result<f64, String> {
let (_, i) = self.split_row_id(row)?;
Ok(self.factor_row_norms[i])
}
pub fn row_column_support(&self, row: usize) -> Result<Vec<usize>, String> {
let (slot, i) = self.split_row_id(row)?;
let p_cov = self.factor.ncols();
let base = self.coupled_rows[slot] * p_cov;
Ok((0..p_cov)
.filter(|&j| self.factor[[i, j]] != 0.0)
.map(|j| base + j)
.collect())
}
pub fn bound(&self, row: usize) -> Result<f64, String> {
self.split_row_id(row)?;
Ok(self.bounds.as_ref().map_or(0.0, |bounds| bounds[row]))
}
pub(crate) fn row_norms_slice(&self) -> &[f64] {
self.factor_row_norms
.as_slice()
.expect("factor row norms are contiguous")
}
pub(crate) fn tile_rows(&self) -> usize {
self.factor.nrows()
}
pub(crate) fn bounds_slice(&self) -> Option<&[f64]> {
self.bounds
.as_ref()
.map(|bounds| bounds.as_slice().expect("bounds are contiguous"))
}
pub fn gather_rows(&self, rows: &[usize]) -> Result<LinearInequalityConstraints, String> {
let p_cov = self.factor.ncols();
let mut a = Array2::<f64>::zeros((rows.len(), self.ncols()));
let mut b = Array1::<f64>::zeros(rows.len());
for (out_row, &row) in rows.iter().enumerate() {
let (slot, i) = self.split_row_id(row)?;
let k = self.coupled_rows[slot];
a.row_mut(out_row)
.slice_mut(ndarray::s![k * p_cov..(k + 1) * p_cov])
.assign(&self.factor.row(i));
b[out_row] = self.bound(row)?;
}
LinearInequalityConstraints::new(a, b)
}
pub fn to_dense(&self) -> Result<LinearInequalityConstraints, String> {
let all: Vec<usize> = (0..self.nrows()).collect();
self.gather_rows(&all)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ConstraintRowId(pub usize);
impl ConstraintRowId {
#[inline]
pub fn index(self) -> usize {
self.0
}
}
#[derive(Clone, Debug)]
pub struct PlacedConstraintBlock {
pub col_start: usize,
pub set: ConstraintSet,
}
#[derive(Clone, Debug)]
pub enum ConstraintSet {
Dense(LinearInequalityConstraints),
KhatriRaoCone(KhatriRaoConeConstraints),
BlockDiagonal {
blocks: Vec<PlacedConstraintBlock>,
total_cols: usize,
},
}
impl ConstraintSet {
pub fn block_diagonal(
blocks: Vec<PlacedConstraintBlock>,
total_cols: usize,
) -> Result<Self, String> {
let mut ranges: Vec<(usize, usize)> = Vec::with_capacity(blocks.len());
for block in &blocks {
let end = block.col_start + block.set.ncols();
if end > total_cols {
return Err(format!(
"ConstraintSet::block_diagonal: block columns {}..{} exceed joint width {}",
block.col_start, end, total_cols
));
}
ranges.push((block.col_start, end));
}
ranges.sort_unstable();
for pair in ranges.windows(2) {
if pair[1].0 < pair[0].1 {
return Err(format!(
"ConstraintSet::block_diagonal: overlapping column ranges {:?} and {:?}",
pair[0], pair[1]
));
}
}
Ok(ConstraintSet::BlockDiagonal { blocks, total_cols })
}
fn block_for_row<'a>(
blocks: &'a [PlacedConstraintBlock],
row: usize,
) -> Result<(&'a PlacedConstraintBlock, usize), String> {
let mut offset = 0usize;
for block in blocks {
let rows = block.set.nrows();
if row < offset + rows {
return Ok((block, row - offset));
}
offset += rows;
}
Err(format!(
"ConstraintSet: row {row} out of range ({offset} rows)"
))
}
pub fn nrows(&self) -> usize {
match self {
ConstraintSet::Dense(dense) => dense.a.nrows(),
ConstraintSet::KhatriRaoCone(cone) => cone.nrows(),
ConstraintSet::BlockDiagonal { blocks, .. } => {
blocks.iter().map(|block| block.set.nrows()).sum()
}
}
}
pub fn ncols(&self) -> usize {
match self {
ConstraintSet::Dense(dense) => dense.a.ncols(),
ConstraintSet::KhatriRaoCone(cone) => cone.ncols(),
ConstraintSet::BlockDiagonal { total_cols, .. } => *total_cols,
}
}
pub fn values(&self, beta: ArrayView1<'_, f64>) -> Result<Array1<f64>, String> {
match self {
ConstraintSet::Dense(dense) => {
if beta.len() != dense.a.ncols() {
return Err(format!(
"ConstraintSet: beta length {} != {}",
beta.len(),
dense.a.ncols()
));
}
Ok(dense.a.dot(&beta))
}
ConstraintSet::KhatriRaoCone(cone) => cone.values(beta),
ConstraintSet::BlockDiagonal { blocks, total_cols } => {
if beta.len() != *total_cols {
return Err(format!(
"ConstraintSet: beta length {} != {}",
beta.len(),
total_cols
));
}
let mut out = Array1::<f64>::zeros(self.nrows());
let mut offset = 0usize;
for block in blocks {
let width = block.set.ncols();
let local = beta.slice(ndarray::s![block.col_start..block.col_start + width]);
let values = block.set.values(local)?;
let rows = values.len();
out.slice_mut(ndarray::s![offset..offset + rows])
.assign(&values);
offset += rows;
}
Ok(out)
}
}
}
pub fn bound(&self, row: usize) -> Result<f64, String> {
match self {
ConstraintSet::Dense(dense) => dense.b.get(row).copied().ok_or_else(|| {
format!(
"ConstraintSet: row {row} out of range ({} rows)",
dense.b.len()
)
}),
ConstraintSet::KhatriRaoCone(cone) => cone.bound(row),
ConstraintSet::BlockDiagonal { blocks, .. } => {
let (block, local) = Self::block_for_row(blocks, row)?;
block.set.bound(local)
}
}
}
pub fn row_norm(&self, row: usize) -> Result<f64, String> {
match self {
ConstraintSet::Dense(dense) => {
if row >= dense.a.nrows() {
return Err(format!(
"ConstraintSet: row {row} out of range ({} rows)",
dense.a.nrows()
));
}
let r = dense.a.row(row);
Ok(r.dot(&r).sqrt())
}
ConstraintSet::KhatriRaoCone(cone) => cone.row_norm(row),
ConstraintSet::BlockDiagonal { blocks, .. } => {
let (block, local) = Self::block_for_row(blocks, row)?;
block.set.row_norm(local)
}
}
}
pub fn shifted_to_delta(&self, beta: ArrayView1<'_, f64>) -> Result<Self, String> {
let values = self.values(beta)?;
match self {
ConstraintSet::Dense(dense) => Ok(ConstraintSet::Dense(
LinearInequalityConstraints::new(dense.a.clone(), &dense.b - &values)?,
)),
ConstraintSet::KhatriRaoCone(cone) => {
let mut shifted = cone.clone();
let base = shifted
.bounds
.take()
.unwrap_or_else(|| Array1::zeros(values.len()));
shifted.bounds = Some(&base - &values);
Ok(ConstraintSet::KhatriRaoCone(shifted))
}
ConstraintSet::BlockDiagonal { blocks, total_cols } => {
let mut shifted_blocks = Vec::with_capacity(blocks.len());
for block in blocks {
let width = block.set.ncols();
let local = beta.slice(ndarray::s![block.col_start..block.col_start + width]);
shifted_blocks.push(PlacedConstraintBlock {
col_start: block.col_start,
set: block.set.shifted_to_delta(local)?,
});
}
Ok(ConstraintSet::BlockDiagonal {
blocks: shifted_blocks,
total_cols: *total_cols,
})
}
}
}
pub fn max_scaled_violation(
&self,
beta: ArrayView1<'_, f64>,
) -> Result<(f64, Option<usize>), String> {
let values = self.values(beta)?;
let metrics = match self {
ConstraintSet::KhatriRaoCone(cone) => RowMetrics::Tiled {
norms: cone.row_norms_slice(),
tile: cone.tile_rows(),
bounds: cone.bounds_slice(),
},
_ => RowMetrics::Carrier(self),
};
let sweep = (0..values.len())
.into_par_iter()
.fold(ScaledViolationSweep::none, |mut sweep, row| {
let value = values[row];
let (norm, bound) = match metrics.read(row) {
Ok(pair) => pair,
Err(error) => {
sweep.record_terminal(row, SweepTerminal::RowUnavailable(error));
return sweep;
}
};
if !feasibility_quantities_are_finite(&[norm, bound, value]) {
sweep.record_terminal(
row,
SweepTerminal::Undecidable {
norm,
bound,
value,
},
);
return sweep;
}
if norm <= 0.0 {
if bound > 0.0 {
sweep.record_terminal(row, SweepTerminal::VacuousRowWithPositiveBound);
}
return sweep;
}
sweep.record_violation(row, (bound - value) / norm);
sweep
})
.reduce(ScaledViolationSweep::none, ScaledViolationSweep::merge);
sweep.verdict()
}
pub fn max_feasible_step(
&self,
beta: ArrayView1<'_, f64>,
delta: ArrayView1<'_, f64>,
skip_rows: &[usize],
) -> Result<(f64, Option<usize>), String> {
let values = self.values(beta)?;
let directional = self.values(delta)?;
let mut skip = vec![false; values.len()];
for &row in skip_rows {
if row < skip.len() {
skip[row] = true;
}
}
let mut step = 1.0_f64;
let mut blocking = None;
for row in 0..values.len() {
if skip[row] {
continue;
}
let norm = self.row_norm(row)?;
let bound = self.bound(row)?;
let value = values[row];
let rate = directional[row];
if !feasibility_quantities_are_finite(&[norm, bound, value, rate]) {
return Err(format!(
"ConstraintSet::max_feasible_step: row {row} cannot be decided \
(row norm {norm:.3e}, bound {bound:.3e}, value {value:.3e}, \
drift {rate:.3e}); every comparison in the ratio test is false \
for NaN, so skipping the row would report the whole step \
feasible (gam#2721)"
));
}
if norm <= 0.0 {
continue;
}
if rate >= 0.0 {
continue;
}
let t = (value - bound) / (-rate);
if t < step {
step = t.max(0.0);
blocking = Some(row);
}
}
Ok((step, blocking))
}
pub fn gather_rows(&self, rows: &[usize]) -> Result<LinearInequalityConstraints, String> {
match self {
ConstraintSet::Dense(dense) => {
let mut a = Array2::<f64>::zeros((rows.len(), dense.a.ncols()));
let mut b = Array1::<f64>::zeros(rows.len());
for (out_row, &row) in rows.iter().enumerate() {
if row >= dense.a.nrows() {
return Err(format!(
"ConstraintSet: row {row} out of range ({} rows)",
dense.a.nrows()
));
}
a.row_mut(out_row).assign(&dense.a.row(row));
b[out_row] = dense.b[row];
}
LinearInequalityConstraints::new(a, b)
}
ConstraintSet::KhatriRaoCone(cone) => cone.gather_rows(rows),
ConstraintSet::BlockDiagonal { blocks, total_cols } => {
let mut a = Array2::<f64>::zeros((rows.len(), *total_cols));
let mut b = Array1::<f64>::zeros(rows.len());
for (out_row, &row) in rows.iter().enumerate() {
let (block, local) = Self::block_for_row(blocks, row)?;
let gathered = block.set.gather_rows(&[local])?;
a.row_mut(out_row)
.slice_mut(ndarray::s![
block.col_start..block.col_start + block.set.ncols()
])
.assign(&gathered.a.row(0));
b[out_row] = gathered.b[0];
}
LinearInequalityConstraints::new(a, b)
}
}
}
pub fn to_dense(&self) -> Result<LinearInequalityConstraints, String> {
match self {
ConstraintSet::Dense(dense) => Ok(dense.clone()),
_ => {
let all: Vec<usize> = (0..self.nrows()).collect();
self.gather_rows(&all)
}
}
}
}
impl From<LinearInequalityConstraints> for ConstraintSet {
fn from(dense: LinearInequalityConstraints) -> Self {
ConstraintSet::Dense(dense)
}
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::array;
fn cone_fixture() -> KhatriRaoConeConstraints {
let psi = array![[1.0_f64, 0.5], [2.0, -1.0], [0.0, 3.0]];
KhatriRaoConeConstraints::new(Arc::new(psi), vec![1, 2], 3).expect("cone fixture")
}
fn beta_fixture() -> Array1<f64> {
array![9.0_f64, -4.0, 1.0, 2.0, 0.5, -0.25]
}
#[test]
fn cone_values_match_dense_system() {
let cone = cone_fixture();
let set = ConstraintSet::KhatriRaoCone(cone.clone());
let dense = ConstraintSet::Dense(cone.to_dense().expect("dense"));
let beta = beta_fixture();
let via_cone = set.values(beta.view()).expect("cone values");
let via_dense = dense.values(beta.view()).expect("dense values");
assert_eq!(via_cone.len(), 6);
for (a, b) in via_cone.iter().zip(via_dense.iter()) {
assert!((a - b).abs() < 1e-14, "cone/dense mismatch: {a} vs {b}");
}
assert!((via_cone[1] - 0.0).abs() < 1e-15);
}
#[test]
fn cone_row_norms_are_factor_row_norms_for_every_slot() {
let cone = cone_fixture();
let set = ConstraintSet::KhatriRaoCone(cone);
let expected = [(1.0_f64 + 0.25).sqrt(), (4.0_f64 + 1.0).sqrt(), 3.0_f64];
for slot in 0..2 {
for i in 0..3 {
let norm = set.row_norm(slot * 3 + i).expect("norm");
assert!((norm - expected[i]).abs() < 1e-15);
}
}
}
#[test]
fn max_scaled_violation_agrees_with_canonicalized_dense() {
let cone = cone_fixture();
let set = ConstraintSet::KhatriRaoCone(cone.clone());
let beta = beta_fixture();
let (violation, row) = set.max_scaled_violation(beta.view()).expect("violation");
let dense = cone
.to_dense()
.expect("dense")
.canonicalized()
.expect("canon");
let values = dense.a.dot(&beta);
let mut worst = 0.0_f64;
let mut worst_row = None;
for r in 0..values.len() {
let v = dense.b[r] - values[r];
if v > worst {
worst = v;
worst_row = Some(r);
}
}
assert!((violation - worst).abs() < 1e-14);
assert_eq!(row, worst_row);
assert!(violation > 0.0, "fixture must have a violated row");
}
#[test]
fn max_feasible_step_matches_scalar_ratio_test() {
let cone = cone_fixture();
let set = ConstraintSet::KhatriRaoCone(cone);
let beta = array![0.0_f64, 0.0, 1.0, 0.1, 1.0, 0.1];
let delta = array![0.0_f64, 0.0, 0.0, -1.0, 0.0, 0.0];
let (step, blocking) = set
.max_feasible_step(beta.view(), delta.view(), &[])
.expect("step");
assert!((step - 0.1).abs() < 1e-14, "expected 0.1, got {step}");
assert_eq!(blocking, Some(2));
let (step_skipped, blocking_skipped) = set
.max_feasible_step(beta.view(), delta.view(), &[2])
.expect("step skipped");
assert!((step_skipped - 1.0).abs() < 1e-14);
assert_eq!(blocking_skipped, None);
}
#[test]
fn gather_rows_places_factor_rows_in_the_coupled_slot() {
let cone = cone_fixture();
let gathered = cone.gather_rows(&[4]).expect("gather");
assert_eq!(gathered.a.nrows(), 1);
assert_eq!(gathered.a.ncols(), 6);
let expected = [0.0, 0.0, 0.0, 0.0, 2.0, -1.0];
for (j, &e) in expected.iter().enumerate() {
assert_eq!(gathered.a[[0, j]], e);
}
assert_eq!(gathered.b[0], 0.0);
}
#[test]
fn constructor_rejects_bad_coupled_rows() {
let psi = array![[1.0_f64, 0.0], [0.0, 1.0]];
assert!(KhatriRaoConeConstraints::new(Arc::new(psi.clone()), vec![3], 3).is_err());
assert!(KhatriRaoConeConstraints::new(Arc::new(psi.clone()), vec![1, 1], 3).is_err());
assert!(KhatriRaoConeConstraints::new(Arc::new(psi), vec![], 3).is_err());
}
#[test]
fn shifted_to_delta_matches_dense_shift() {
let cone = cone_fixture();
let set = ConstraintSet::KhatriRaoCone(cone);
let beta = beta_fixture();
let shifted = set.shifted_to_delta(beta.view()).expect("shift");
let dense = set.to_dense().expect("dense");
let expected_b = &dense.b - &dense.a.dot(&beta);
for row in 0..set.nrows() {
assert!(
(shifted.bound(row).expect("bound") - expected_b[row]).abs() < 1e-14,
"shifted bound mismatch at row {row}"
);
}
let zero = Array1::<f64>::zeros(set.ncols());
let (viol_delta, row_delta) = shifted
.max_scaled_violation(zero.view())
.expect("delta violation");
let (viol_orig, row_orig) = set.max_scaled_violation(beta.view()).expect("violation");
assert!((viol_delta - viol_orig).abs() < 1e-14);
assert_eq!(row_delta, row_orig);
}
#[test]
fn block_diagonal_composes_ids_bounds_and_values() {
let dense = LinearInequalityConstraints::new(
array![[1.0_f64, 0.0], [0.0, -2.0]],
array![0.5_f64, -1.0],
)
.expect("dense block");
let cone = cone_fixture();
let joint = ConstraintSet::block_diagonal(
vec![
PlacedConstraintBlock {
col_start: 0,
set: ConstraintSet::Dense(dense.clone()),
},
PlacedConstraintBlock {
col_start: 2,
set: ConstraintSet::KhatriRaoCone(cone.clone()),
},
],
8,
)
.expect("joint");
assert_eq!(joint.nrows(), 2 + 6);
assert_eq!(joint.ncols(), 8);
let mut beta = Array1::<f64>::zeros(8);
beta[0] = 2.0;
beta[1] = 1.0;
beta.slice_mut(ndarray::s![2..8]).assign(&beta_fixture());
let values = joint.values(beta.view()).expect("values");
assert!((values[0] - 2.0).abs() < 1e-15);
assert!((values[1] + 2.0).abs() < 1e-15);
let cone_values = cone.values(beta_fixture().view()).expect("cone values");
for (idx, &cv) in cone_values.iter().enumerate() {
assert!((values[2 + idx] - cv).abs() < 1e-15);
}
assert_eq!(joint.bound(0).expect("b0"), 0.5);
assert_eq!(joint.bound(2).expect("b2"), 0.0);
let gathered = joint.gather_rows(&[3]).expect("gather");
assert_eq!(gathered.a.ncols(), 8);
assert_eq!(gathered.a[[0, 4]], 2.0);
assert_eq!(gathered.a[[0, 5]], -1.0);
assert!(
ConstraintSet::block_diagonal(
vec![
PlacedConstraintBlock {
col_start: 0,
set: ConstraintSet::Dense(dense.clone()),
},
PlacedConstraintBlock {
col_start: 1,
set: ConstraintSet::Dense(dense),
},
],
8,
)
.is_err()
);
}
#[test]
fn zero_factor_rows_are_vacuous_not_violations() {
let psi = array![[0.0_f64, 0.0], [1.0, 1.0]];
let cone = KhatriRaoConeConstraints::new(Arc::new(psi), vec![1], 2).expect("cone");
let set = ConstraintSet::KhatriRaoCone(cone);
let beta = array![0.0_f64, 0.0, -5.0, 4.0];
let (violation, row) = set.max_scaled_violation(beta.view()).expect("violation");
assert_eq!(row, Some(1));
assert!((violation - 1.0 / 2.0_f64.sqrt()).abs() < 1e-14);
}
}