use pounce_common::types::Number;
use crate::backsolver::{BoundRow, SensBacksolver};
#[derive(Clone, Debug)]
pub struct WatchedRow {
pub coefficients: Vec<(usize, Number)>,
pub limit: Number,
pub base_value: Number,
pub active: Option<(usize, Number)>,
}
#[derive(Clone)]
pub struct RowLimitView<B> {
base: B,
n_x: usize,
rows: Vec<WatchedRow>,
base_dim: usize,
bound_rows: Vec<BoundRow>,
shifts: Vec<(usize, usize, Number)>,
}
impl<B: SensBacksolver> RowLimitView<B> {
pub fn new(base: B, n_x: usize, rows: Vec<WatchedRow>) -> Option<Self> {
let base_dim = base.dim();
if n_x > base_dim || base.natural_units_factor().is_some() {
return None;
}
let n_t = rows.len();
let mut bound_rows: Vec<BoundRow> = Vec::new();
if let Some(base_rows) = base.bound_rows() {
for b in base_rows {
if b.var_row >= n_x || b.row < n_x || b.row >= base_dim {
return None;
}
bound_rows.push(BoundRow {
row: b.row + n_t,
var_row: b.var_row,
lower: b.lower,
});
}
}
let mut shifts = Vec::new();
for (j, r) in rows.iter().enumerate() {
if r.coefficients.iter().any(|&(c, _)| c >= n_x) {
return None;
}
if let Some((row, mult)) = r.active {
if row < n_x || row >= base_dim {
return None;
}
let obs = n_x + j;
bound_rows.push(BoundRow {
row: row + n_t,
var_row: obs,
lower: false,
});
shifts.push((row + n_t, obs, mult));
}
}
Some(Self {
base,
n_x,
rows,
base_dim,
bound_rows,
shifts,
})
}
pub fn n_observers(&self) -> usize {
self.rows.len()
}
pub fn from_base(&self, base_row: usize) -> Option<usize> {
if base_row < self.n_x {
Some(base_row)
} else if base_row < self.base_dim {
Some(base_row + self.rows.len())
} else {
None
}
}
pub fn to_base(&self, row: usize) -> Option<usize> {
let n_t = self.rows.len();
if row < self.n_x {
Some(row)
} else if row < self.n_x + n_t {
None
} else if row < self.base_dim + n_t {
Some(row - n_t)
} else {
None
}
}
pub fn primal_box(
&self,
x_curr: &[Number],
lo: &[Number],
hi: &[Number],
) -> Option<(Vec<Number>, Vec<Number>, Vec<Number>)> {
if x_curr.len() != self.n_x || lo.len() != self.n_x || hi.len() != self.n_x {
return None;
}
let mut x = x_curr.to_vec();
let mut l = lo.to_vec();
let mut h = hi.to_vec();
for r in &self.rows {
x.push(r.base_value);
l.push(Number::NEG_INFINITY);
h.push(r.limit);
}
Some((x, l, h))
}
pub fn lift_rhs(&self, base_rhs: &[Number]) -> Option<Vec<Number>> {
if base_rhs.len() != self.base_dim {
return None;
}
let mut out = vec![0.0; self.dim()];
for (i, &v) in base_rhs.iter().enumerate() {
let row = self.from_base(i)?;
out[row] = v;
}
Some(out)
}
pub fn lift_step(&self, base_lhs: &[Number]) -> Option<Vec<Number>> {
if base_lhs.len() != self.base_dim {
return None;
}
let zeros = vec![0.0; self.rows.len()];
let mut out = vec![0.0; self.dim()];
self.unfold(base_lhs, &zeros, &zeros, &mut out);
Some(out)
}
pub fn all_bound_rows(&self) -> &[BoundRow] {
&self.bound_rows
}
pub fn lift_multipliers(
&self,
base: &[crate::boundcheck::BoundMultiplier],
) -> Option<Vec<crate::boundcheck::BoundMultiplier>> {
let mut out = Vec::with_capacity(base.len() + self.shifts.len());
for m in base {
out.push(crate::boundcheck::BoundMultiplier {
row: self.from_base(m.row)?,
base: m.base,
});
}
for &(row, _, mult) in &self.shifts {
out.push(crate::boundcheck::BoundMultiplier { row, base: mult });
}
Some(out)
}
fn fold(&self, rhs: &[Number]) -> (Vec<Number>, Vec<Number>, Vec<Number>) {
let n_t = self.rows.len();
let r_t = rhs[self.n_x..self.n_x + n_t].to_vec();
let r_mu = rhs[self.base_dim + n_t..].to_vec();
let mut base_rhs = vec![0.0; self.base_dim];
base_rhs[..self.n_x].copy_from_slice(&rhs[..self.n_x]);
base_rhs[self.n_x..].copy_from_slice(&rhs[self.n_x + n_t..self.base_dim + n_t]);
for (j, r) in self.rows.iter().enumerate() {
let v = r_t[j];
if v != 0.0 {
for &(c, coef) in &r.coefficients {
base_rhs[c] += coef * v;
}
}
}
(base_rhs, r_t, r_mu)
}
fn unfold(&self, base_lhs: &[Number], r_t: &[Number], r_mu: &[Number], lhs: &mut [Number]) {
let n_t = self.rows.len();
lhs[..self.n_x].copy_from_slice(&base_lhs[..self.n_x]);
lhs[self.n_x + n_t..self.base_dim + n_t].copy_from_slice(&base_lhs[self.n_x..]);
for (j, r) in self.rows.iter().enumerate() {
let gx: Number = r
.coefficients
.iter()
.map(|&(c, coef)| coef * base_lhs[c])
.sum();
lhs[self.n_x + j] = gx + r_mu[j];
lhs[self.base_dim + n_t + j] = r_t[j];
}
}
fn around<F>(&self, rhs: &[Number], lhs: &mut [Number], f: F) -> bool
where
F: FnOnce(&[Number], &mut [Number]) -> bool,
{
if rhs.len() != self.dim() || lhs.len() != self.dim() {
return false;
}
let (base_rhs, r_t, r_mu) = self.fold(rhs);
let mut base_lhs = vec![0.0; self.base_dim];
if !f(&base_rhs, &mut base_lhs) {
return false;
}
self.unfold(&base_lhs, &r_t, &r_mu, lhs);
true
}
fn released_in_base(&self, released: &[usize]) -> Option<Vec<usize>> {
released.iter().map(|&r| self.to_base(r)).collect()
}
}
impl<B: SensBacksolver> SensBacksolver for RowLimitView<B> {
fn dim(&self) -> usize {
self.base_dim + 2 * self.rows.len()
}
fn solve(&self, rhs: &[Number], lhs: &mut [Number]) -> bool {
self.around(rhs, lhs, |b, l| self.base.solve(b, l))
}
fn natural_units_factor(&self) -> Option<&[Number]> {
None
}
fn bound_rows(&self) -> Option<&[BoundRow]> {
Some(&self.bound_rows)
}
fn supports_release(&self) -> bool {
self.base.supports_release()
}
fn solve_released(&self, released: &[usize], rhs: &[Number], lhs: &mut [Number]) -> bool {
let Some(key) = self.released_in_base(released) else {
return false;
};
self.around(rhs, lhs, |b, l| self.base.solve_released(&key, b, l))
}
fn solve_released_step(&self, released: &[usize], rhs: &[Number], lhs: &mut [Number]) -> bool {
if rhs.len() != self.dim() {
return false;
}
let Some(key) = self.released_in_base(released) else {
return false;
};
let mut rhs = rhs.to_vec();
for &(row, obs, mult) in &self.shifts {
if !released.contains(&row) {
continue;
}
rhs[row] = 0.0;
rhs[obs] += mult;
}
self.around(&rhs, lhs, |b, l| self.base.solve_released_step(&key, b, l))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::backsolver::DenseLuBacksolver;
struct Base {
n: usize,
k: Vec<Number>,
rows: Vec<BoundRow>,
}
impl Base {
fn factor(&self, released: &[usize]) -> Option<DenseLuBacksolver> {
let mut k = self.k.clone();
for &r in released {
for j in 0..self.n {
k[r * self.n + j] = 0.0;
k[j * self.n + r] = 0.0;
}
k[r * self.n + r] = -1.0;
}
DenseLuBacksolver::from_dense(self.n, &k).ok()
}
}
impl SensBacksolver for Base {
fn dim(&self) -> usize {
self.n
}
fn solve(&self, rhs: &[Number], lhs: &mut [Number]) -> bool {
self.factor(&[]).is_some_and(|f| f.solve(rhs, lhs))
}
fn bound_rows(&self) -> Option<&[BoundRow]> {
Some(&self.rows)
}
fn supports_release(&self) -> bool {
true
}
fn solve_released(&self, released: &[usize], rhs: &[Number], lhs: &mut [Number]) -> bool {
self.factor(released).is_some_and(|f| f.solve(rhs, lhs))
}
fn solve_released_step(
&self,
released: &[usize],
rhs: &[Number],
lhs: &mut [Number],
) -> bool {
self.solve_released(released, rhs, lhs)
}
}
fn base(rows: Vec<BoundRow>) -> Base {
#[rustfmt::skip]
let k = vec![
2.0, 0.3, -0.4, 1.0, 0.0,
0.3, 1.7, 0.2, 0.0, 1.0,
-0.4, 0.2, 2.5, 1.0, 1.0,
1.0, 0.0, 1.0, 0.0, 0.0,
0.0, 1.0, 1.0, 0.0, 0.0,
];
Base { n: 5, k, rows }
}
fn watched(active: Option<(usize, Number)>, coefficients: Vec<(usize, Number)>) -> WatchedRow {
WatchedRow {
coefficients,
limit: 1.0,
base_value: 0.25,
active,
}
}
#[test]
fn the_triangular_solve_is_the_augmented_system() {
let n_x = 3;
let rows = vec![
watched(None, vec![(0, 1.0), (2, -0.5)]),
watched(None, vec![(1, 2.0)]),
];
let b = base(Vec::new());
let base_k = b.k.clone();
let base_dim = b.n;
let n_t = rows.len();
let view = RowLimitView::new(b, n_x, rows.clone()).expect("the view must build");
let dim = view.dim();
assert_eq!(dim, base_dim + 2 * n_t);
let t0 = n_x;
let mu0 = base_dim + n_t;
let mut a = vec![0.0; dim * dim];
let view_of = |i: usize| if i < n_x { i } else { i + n_t };
for i in 0..base_dim {
for j in 0..base_dim {
a[view_of(i) * dim + view_of(j)] = base_k[i * base_dim + j];
}
}
for (j, r) in rows.iter().enumerate() {
a[(t0 + j) * dim + mu0 + j] = 1.0;
a[(mu0 + j) * dim + t0 + j] = 1.0;
for &(c, coef) in &r.coefficients {
a[c * dim + mu0 + j] = -coef;
a[(mu0 + j) * dim + c] = -coef;
}
}
let dense =
DenseLuBacksolver::from_dense(dim, &a).expect("the augmented system is regular");
let rhs: Vec<Number> = (0..dim).map(|i| 0.7 - 0.31 * (i as Number)).collect();
let mut want = vec![0.0; dim];
assert!(dense.solve(&rhs, &mut want));
let mut got = vec![0.0; dim];
assert!(view.solve(&rhs, &mut got));
for i in 0..dim {
assert!(
(got[i] - want[i]).abs() < 1e-10,
"row {i}: view {got:?} vs dense {want:?}",
);
}
}
#[test]
fn lifting_a_step_agrees_with_solving_it() {
let n_x = 3;
let rows = vec![
watched(None, vec![(0, 1.0), (2, -0.5)]),
watched(Some((3, 0.7)), vec![(1, 2.0)]),
];
let b = base(Vec::new());
let base_dim = b.n;
let view = RowLimitView::new(b, n_x, rows).expect("the view must build");
let base_rhs: Vec<Number> = vec![0.4, -1.1, 0.9, 0.3, -0.6];
let lifted_rhs = view.lift_rhs(&base_rhs).expect("the rhs lifts");
let mut want = vec![0.0; view.dim()];
assert!(view.solve(&lifted_rhs, &mut want), "the view must solve");
let mut base_lhs = vec![0.0; base_dim];
assert!(
view.base.solve(&base_rhs, &mut base_lhs),
"the base must solve",
);
let got = view.lift_step(&base_lhs).expect("the step lifts");
for (i, (g, w)) in got.iter().zip(&want).enumerate() {
assert!(
(g - w).abs() < 1e-10,
"row {i}: lift_step {g} against solve {w}\n{got:?}\n{want:?}",
);
}
assert!(
got[n_x..n_x + 2].iter().any(|v| v.abs() > 1e-6),
"the observers read {:?}, so this test would pass on a \
zero-filling lift",
&got[n_x..n_x + 2],
);
}
#[test]
fn lift_step_refuses_a_vector_that_is_not_the_base() {
let rows = vec![watched(None, vec![(0, 1.0)])];
let view = RowLimitView::new(base(Vec::new()), 3, rows).expect("the view must build");
assert!(view.lift_step(&vec![0.0; view.dim()]).is_none());
assert!(view.lift_step(&[]).is_none());
assert!(view.lift_step(&vec![0.0; 5]).is_some());
}
#[test]
fn each_active_row_keeps_its_own_observer() {
let n_x = 3;
let (m0, m1) = (0.75, 0.2);
let rows = vec![
watched(Some((3, m0)), vec![(0, 1.0), (2, -0.5)]),
watched(Some((4, m1)), vec![(1, 2.0)]),
];
let n_t = rows.len();
let b = base(Vec::new());
let reference = base(Vec::new());
let view = RowLimitView::new(b, n_x, rows.clone()).expect("the view must build");
assert_eq!(
view.all_bound_rows(),
&[
BoundRow {
row: 3 + n_t,
var_row: n_x,
lower: false
},
BoundRow {
row: 4 + n_t,
var_row: n_x + 1,
lower: false
},
],
);
let lifted = view
.lift_multipliers(&[])
.expect("an empty base list still lifts");
assert_eq!(lifted.len(), 2);
for (got, want) in lifted.iter().zip([(3 + n_t, m0), (4 + n_t, m1)]) {
assert_eq!((got.row, got.base), want);
}
let released = vec![3 + n_t];
let mut got = vec![0.0; view.dim()];
assert!(view.solve_released_step(&released, &vec![0.0; view.dim()], &mut got));
let mut want_rhs = vec![0.0; 5];
for &(c, coef) in &rows[0].coefficients {
want_rhs[c] += coef * m0;
}
let mut want = vec![0.0; 5];
assert!(reference.solve_released(&[3], &want_rhs, &mut want));
for i in 0..n_x {
assert!(
(got[i] - want[i]).abs() < 1e-10,
"x row {i}: {got:?} vs {want:?}",
);
}
assert!(
got[..n_x].iter().any(|v| v.abs() > 1e-6),
"the shift must actually move something: {got:?}",
);
}
}