use crate::api::context::Context;
use crate::api::expr::{Ex, ExprType};
use crate::base::combinatorics::factorial;
use crate::base::errors::SymplexError;
use crate::base::numeric::Q;
use crate::domains::decompositions::{
Diagonalization, Hessenberg, JordanForm, Lu, RankDecomposition,
};
use num_bigint::BigInt;
use num_rational::Rational64;
use std::fmt;
use std::sync::{Arc, OnceLock};
use tracing::{debug, trace, warn};
pub use crate::domains::exact_matrix::{
ExactMatrix, ExactScalar, LLL_DEFAULT_DELTA, QMatrix, ZMatrix,
};
pub use crate::output::codegen::{CodegenOptions, MathBackend, Precision};
use crate::domains::exact_matrix::PERMANENT_MAX_DIM;
#[derive(Clone)]
pub struct Matrix {
rows: Vec<Vec<Ex>>,
nrows: usize,
ncols: usize,
exact: OnceLock<Option<Arc<QMatrix>>>,
}
impl PartialEq for Matrix {
fn eq(&self, other: &Self) -> bool {
self.nrows == other.nrows && self.ncols == other.ncols && self.rows == other.rows
}
}
impl Eq for Matrix {}
fn invalid(operation: &'static str, reason: impl Into<String>) -> SymplexError {
SymplexError::invalid_argument(operation, reason)
}
fn failed(operation: &'static str, reason: impl Into<String>) -> SymplexError {
SymplexError::computation_failed(operation, reason)
}
pub(crate) fn reop(e: SymplexError, operation: &'static str) -> SymplexError {
match e {
SymplexError::ComputationFailed { reason, .. } => failed(operation, reason),
SymplexError::InvalidArgument { reason, .. } => invalid(operation, reason),
other => other,
}
}
pub(crate) fn ex_is_zero(e: &Ex) -> Option<bool> {
if e.is_zero_structural() {
return Some(true);
}
if let Some(b) = e.is_zero() {
return Some(b);
}
let s = e.eval().simplify();
if s.is_zero_structural() {
return Some(true);
}
if let Some(b) = s.is_zero() {
return Some(b);
}
if let Ok(z) = s.eval_complex64() {
if z.re == 0.0 && z.im == 0.0 {
return Some(true);
}
if z.re.abs() > 1e-12 || z.im.abs() > 1e-12 {
return Some(false);
}
return None;
}
if !has_radical(&s) {
let t = s.together();
if t != s {
let t = t.simplify();
if t.is_zero_structural() {
return Some(true);
}
if let Some(b) = t.is_zero() {
return Some(b);
}
}
}
None
}
fn has_radical(e: &Ex) -> bool {
use crate::base::node::ExprNode;
let inner = e.inner.read();
let arena = &inner.arena;
let mut stack = vec![e.raw_id()];
let mut seen = rustc_hash::FxHashSet::default();
while let Some(id) = stack.pop() {
if !seen.insert(id) {
continue;
}
if let ExprNode::Pow(_, exp) = arena.node(id)
&& let Some(r) = arena.as_num(*exp)
&& !r.is_integer()
{
return true;
}
stack.extend(arena.node(id).children());
}
false
}
fn fix_trig_parity(e: &Ex) -> Ex {
use crate::base::node::ExprNode;
use num_traits::Signed;
let targets: Vec<(crate::base::node::ExprId, bool, Q)> = {
let inner = e.inner.read();
let arena = &inner.arena;
let mut out = Vec::new();
let mut stack = vec![e.raw_id()];
let mut seen = rustc_hash::FxHashSet::default();
while let Some(id) = stack.pop() {
if !seen.insert(id) {
continue;
}
let (is_sin, arg) = match arena.node(id) {
ExprNode::Sin(a) => (true, *a),
ExprNode::Cos(a) => (false, *a),
_ => {
stack.extend(arena.node(id).children());
continue;
}
};
if let Some(r) = arena.as_num(arg)
&& r.is_negative()
{
out.push((id, is_sin, r.clone()));
}
}
out
};
if targets.is_empty() {
return e.clone();
}
let ctx = e.context();
let map: rustc_hash::FxHashMap<_, Ex> = targets
.into_iter()
.map(|(id, is_sin, r)| {
let pos = ctx.from_ratio(-r);
(id, if is_sin { -pos.sin() } else { pos.cos() })
})
.collect();
e.replace(|v| map.get(&v.id()).cloned())
}
fn has_root_of(e: &Ex) -> bool {
use crate::base::node::ExprNode;
let inner = e.inner.read();
let arena = &inner.arena;
let mut stack = vec![e.raw_id()];
let mut seen = rustc_hash::FxHashSet::default();
while let Some(id) = stack.pop() {
if !seen.insert(id) {
continue;
}
if matches!(arena.node(id), ExprNode::RootOf(..)) {
return true;
}
stack.extend(arena.node(id).children());
}
false
}
pub(crate) fn ex_is_positive(e: &Ex) -> Option<bool> {
if let Some(b) = e.is_positive() {
return Some(b);
}
let s = e.eval().simplify();
if let Some(b) = s.is_positive() {
return Some(b);
}
if let Ok(v) = s.eval_f64() {
return Some(v > 0.0);
}
None
}
pub(crate) fn ex_is_nonnegative(e: &Ex) -> Option<bool> {
if let Some(b) = e.is_nonnegative() {
return Some(b);
}
let s = e.eval().simplify();
if let Some(b) = s.is_nonnegative() {
return Some(b);
}
if let Ok(v) = s.eval_f64() {
return Some(v >= 0.0);
}
None
}
pub(crate) fn all3(iter: impl IntoIterator<Item = Option<bool>>) -> Option<bool> {
let mut unknown = false;
for v in iter {
match v {
Some(false) => return Some(false),
None => unknown = true,
Some(true) => {}
}
}
if unknown { None } else { Some(true) }
}
pub(crate) fn sqrt_rationalized(e: &Ex) -> Ex {
if let Some(r) = e.as_rational()
&& r.numer().sign() == num_bigint::Sign::Plus
{
let ctx = e.context();
let p = ctx.from_bigint(r.numer().clone());
let q = ctx.from_bigint(r.denom().clone());
return &(&p * &q).sqrt() / &q;
}
e.sqrt()
}
pub use crate::base::config::EXPRESSION_BUDGET;
pub(crate) fn tree_size_capped(e: &Ex, cap: usize) -> usize {
let inner = e.inner.read();
let arena = &inner.arena;
let mut stack = vec![e.raw_id()];
let mut count = 0usize;
while let Some(id) = stack.pop() {
count += 1;
if count > cap {
break;
}
stack.extend(arena.node(id).children());
}
count
}
pub(crate) fn budget_check<'a>(
entries: impl IntoIterator<Item = &'a Ex>,
operation: &'static str,
) -> Result<(), SymplexError> {
let mut total = 0usize;
for e in entries {
total += tree_size_capped(e, EXPRESSION_BUDGET - total);
if total > EXPRESSION_BUDGET {
return Err(failed(
operation,
format!(
"expression swell: intermediate result exceeds the budget of \
{EXPRESSION_BUDGET} expression nodes; simplify the input or \
substitute numeric values first"
),
));
}
}
Ok(())
}
impl Matrix {
pub fn new(rows: Vec<Vec<Ex>>) -> Result<Self, SymplexError> {
if rows.is_empty() {
return Err(invalid("Matrix::new", "must have at least one row"));
}
let ncols = rows[0].len();
if ncols == 0 {
return Err(invalid("Matrix::new", "must have at least one column"));
}
for (i, row) in rows.iter().enumerate() {
if row.len() != ncols {
return Err(invalid(
"Matrix::new",
format!("row {i} has length {} but expected {ncols}", row.len()),
));
}
}
let nrows = rows.len();
Ok(Matrix::with_shape(rows, nrows, ncols))
}
#[inline]
fn with_shape(rows: Vec<Vec<Ex>>, nrows: usize, ncols: usize) -> Self {
Matrix {
rows,
nrows,
ncols,
exact: OnceLock::new(),
}
}
#[inline]
pub(crate) fn from_rows_unchecked(rows: Vec<Vec<Ex>>) -> Self {
let nrows = rows.len();
let ncols = rows.first().map_or(0, Vec::len);
debug_assert!(nrows > 0 && ncols > 0);
debug_assert!(rows.iter().all(|r| r.len() == ncols));
Matrix::with_shape(rows, nrows, ncols)
}
pub fn zeros(ctx: &Context, n: usize, m: usize) -> Self {
assert!(n > 0 && m > 0, "Matrix::zeros: dimensions must be positive");
let zero = ctx.zero();
let rows = (0..n)
.map(|_| (0..m).map(|_| zero.clone()).collect())
.collect();
Matrix::with_shape(rows, n, m)
}
pub fn identity(ctx: &Context, n: usize) -> Self {
assert!(n > 0, "Matrix::identity: dimension must be positive");
let one = ctx.one();
let zero = ctx.zero();
let rows = (0..n)
.map(|i| {
(0..n)
.map(|j| if i == j { one.clone() } else { zero.clone() })
.collect()
})
.collect();
Matrix::with_shape(rows, n, n)
}
pub fn from_fn(n: usize, m: usize, mut f: impl FnMut(usize, usize) -> Ex) -> Self {
assert!(
n > 0 && m > 0,
"Matrix::from_fn: dimensions must be positive"
);
let rows = (0..n).map(|i| (0..m).map(|j| f(i, j)).collect()).collect();
Matrix::with_shape(rows, n, m)
}
pub fn row_vector(elems: Vec<Ex>) -> Self {
assert!(
!elems.is_empty(),
"Matrix::row_vector: need at least one element"
);
let ncols = elems.len();
Matrix::with_shape(vec![elems], 1, ncols)
}
pub fn col_vector(elems: Vec<Ex>) -> Self {
assert!(
!elems.is_empty(),
"Matrix::col_vector: need at least one element"
);
let nrows = elems.len();
let rows = elems.into_iter().map(|e| vec![e]).collect();
Matrix::with_shape(rows, nrows, 1)
}
pub fn diag(entries: &[Ex]) -> Matrix {
let n = entries.len();
assert!(n > 0, "Matrix::diag: entries must be non-empty");
let zero = entries[0].context().zero();
let rows: Vec<Vec<Ex>> = (0..n)
.map(|i| {
(0..n)
.map(|j| {
if i == j {
entries[i].clone()
} else {
zero.clone()
}
})
.collect()
})
.collect();
Matrix::from_rows_unchecked(rows)
}
pub fn from_i64(ctx: &Context, rows: &[&[i64]]) -> Result<Matrix, SymplexError> {
let data: Vec<Vec<Ex>> = rows
.iter()
.map(|row| row.iter().map(|&v| ctx.int(v)).collect())
.collect();
Matrix::new(data)
}
pub fn block_diag(blocks: &[&Matrix]) -> Result<Matrix, SymplexError> {
if blocks.is_empty() {
return Err(invalid("block_diag", "need at least one block"));
}
let nrows: usize = blocks.iter().map(|b| b.nrows).sum();
let ncols: usize = blocks.iter().map(|b| b.ncols).sum();
let zero = blocks[0].ctx_zero();
let mut rows: Vec<Vec<Ex>> = vec![vec![zero; ncols]; nrows];
let (mut r0, mut c0) = (0, 0);
for b in blocks {
for i in 0..b.nrows {
for j in 0..b.ncols {
rows[r0 + i][c0 + j] = b.rows[i][j].clone();
}
}
r0 += b.nrows;
c0 += b.ncols;
}
Ok(Matrix::with_shape(rows, nrows, ncols))
}
pub fn companion(coeffs: &[Ex]) -> Result<Matrix, SymplexError> {
let Some(first) = coeffs.first() else {
return Err(invalid(
"companion",
"need at least one coefficient (a monic polynomial of degree ≥ 1)",
));
};
let n = coeffs.len();
let ctx = first.context();
let zero = ctx.zero();
let one = ctx.one();
let mut rows: Vec<Vec<Ex>> = vec![vec![zero; n]; n];
for (i, c) in coeffs.iter().enumerate() {
rows[i][n - 1] = -c;
if i + 1 < n {
rows[i + 1][i] = one.clone();
}
}
Ok(Matrix::from_rows_unchecked(rows))
}
pub fn jordan_block(eigenvalue: &Ex, size: usize) -> Result<Matrix, SymplexError> {
if size == 0 {
return Err(invalid("jordan_block", "size must be positive"));
}
let ctx = eigenvalue.context();
let zero = ctx.zero();
let one = ctx.one();
Ok(Matrix::from_fn(size, size, |i, j| {
if i == j {
eigenvalue.clone()
} else if j == i + 1 {
one.clone()
} else {
zero.clone()
}
}))
}
}
impl TryFrom<Vec<Vec<Ex>>> for Matrix {
type Error = SymplexError;
fn try_from(rows: Vec<Vec<Ex>>) -> Result<Self, Self::Error> {
Matrix::new(rows)
}
}
impl Matrix {
#[inline]
pub fn nrows(&self) -> usize {
self.nrows
}
#[inline]
pub fn ncols(&self) -> usize {
self.ncols
}
#[inline]
pub fn shape(&self) -> (usize, usize) {
(self.nrows, self.ncols)
}
#[inline]
pub fn is_square(&self) -> bool {
self.nrows == self.ncols
}
#[inline]
pub fn get(&self, i: usize, j: usize) -> &Ex {
assert!(
i < self.nrows && j < self.ncols,
"Index ({i}, {j}) out of bounds for {}×{} matrix",
self.nrows,
self.ncols
);
&self.rows[i][j]
}
#[inline]
pub fn try_get(&self, i: usize, j: usize) -> Option<&Ex> {
if i < self.nrows && j < self.ncols {
Some(&self.rows[i][j])
} else {
None
}
}
#[inline]
pub fn get_mut(&mut self, i: usize, j: usize) -> &mut Ex {
assert!(
i < self.nrows && j < self.ncols,
"Index ({i}, {j}) out of bounds for {}×{} matrix",
self.nrows,
self.ncols
);
self.exact.take();
&mut self.rows[i][j]
}
#[inline]
pub fn set(&mut self, i: usize, j: usize, value: Ex) {
*self.get_mut(i, j) = value;
}
#[inline]
pub fn row(&self, i: usize) -> &[Ex] {
assert!(
i < self.nrows,
"Row index {i} out of bounds for {}×{} matrix",
self.nrows,
self.ncols
);
&self.rows[i]
}
pub fn col(&self, j: usize) -> Vec<Ex> {
assert!(
j < self.ncols,
"Column index {j} out of bounds for {}×{} matrix",
self.nrows,
self.ncols
);
self.rows.iter().map(|r| r[j].clone()).collect()
}
pub fn diagonal(&self) -> Vec<Ex> {
(0..self.nrows.min(self.ncols))
.map(|i| self.rows[i][i].clone())
.collect()
}
pub fn submatrix(&self, rows: std::ops::Range<usize>, cols: std::ops::Range<usize>) -> Matrix {
assert!(
rows.start < rows.end && rows.end <= self.nrows,
"submatrix: row range {rows:?} invalid for {} rows",
self.nrows
);
assert!(
cols.start < cols.end && cols.end <= self.ncols,
"submatrix: column range {cols:?} invalid for {} columns",
self.ncols
);
let data: Vec<Vec<Ex>> = rows.map(|i| self.rows[i][cols.clone()].to_vec()).collect();
Matrix::from_rows_unchecked(data)
}
pub fn iter(&self) -> impl Iterator<Item = &Ex> + '_ {
self.rows.iter().flatten()
}
pub fn to_vec(&self) -> Vec<Vec<Ex>> {
self.rows.clone()
}
pub fn eval_f64(&self) -> Result<Vec<Vec<f64>>, SymplexError> {
self.rows
.iter()
.map(|r| r.iter().map(Ex::eval_f64).collect())
.collect()
}
pub fn equals(&self, other: &Matrix) -> Option<bool> {
if self.shape() != other.shape() {
return Some(false);
}
all3(
self.iter()
.zip(other.iter())
.map(|(a, b)| ex_is_zero(&(a - b))),
)
}
pub(crate) fn ctx_zero(&self) -> Ex {
let elem = &self.rows[0][0];
let zero_id = elem.inner.read().arena.zero();
elem.wrap(zero_id)
}
pub(crate) fn ctx_one(&self) -> Ex {
let elem = &self.rows[0][0];
let one_id = elem.inner.read().arena.one();
elem.wrap(one_id)
}
pub fn context(&self) -> Context {
self.rows[0][0].context()
}
#[inline]
pub(crate) fn ctx(&self) -> Context {
self.context()
}
fn require_square(&self, operation: &'static str) -> Result<(), SymplexError> {
if self.nrows != self.ncols {
return Err(invalid(
operation,
format!(
"requires a square matrix, got {}×{}",
self.nrows, self.ncols
),
));
}
Ok(())
}
fn check_budget(&self, operation: &'static str) -> Result<(), SymplexError> {
budget_check(self.iter(), operation)
}
pub(crate) fn as_qmatrix(&self) -> Option<&QMatrix> {
self.exact
.get_or_init(|| {
let rows = self.to_rational_rows()?;
QMatrix::new(rows).ok().map(Arc::new)
})
.as_deref()
}
#[doc(hidden)]
pub fn is_exact_cached(&self) -> bool {
self.exact.get().is_some()
}
fn contains(&self, sym: &Ex) -> bool {
self.iter().any(|e| e.contains(sym))
}
pub(crate) fn fresh_symbol(&self, base: &str) -> Ex {
let ctx = self.ctx();
let mut k = 0usize;
loop {
let name = if k == 0 {
format!("__{base}")
} else {
format!("__{base}_{k}")
};
let sym = ctx.symbol(&name);
if !self.contains(&sym) {
return sym;
}
k += 1;
}
}
fn hide_dummy(&self, e: &Ex, dummy: &Ex) -> Ex {
if !e.contains(dummy) {
return e.clone();
}
let lam = self.ctx().symbol("λ");
if self.contains(&lam) {
return e.clone();
}
e.subs(dummy, &lam)
}
}
impl Matrix {
pub fn transpose(&self) -> Matrix {
let rows: Vec<Vec<Ex>> = (0..self.ncols)
.map(|j| (0..self.nrows).map(|i| self.rows[i][j].clone()).collect())
.collect();
Matrix::with_shape(rows, self.ncols, self.nrows)
}
pub fn adjoint(&self) -> Matrix {
self.transpose().map(|e| e.conjugate())
}
pub fn add(&self, other: &Matrix) -> Result<Matrix, SymplexError> {
if self.shape() != other.shape() {
return Err(invalid(
"add",
format!(
"cannot add matrices with shapes {:?} and {:?}",
self.shape(),
other.shape()
),
));
}
Ok(self.zip_with(other, |a, b| a + b))
}
pub fn sub(&self, other: &Matrix) -> Result<Matrix, SymplexError> {
if self.shape() != other.shape() {
return Err(invalid(
"sub",
format!(
"cannot subtract matrices with shapes {:?} and {:?}",
self.shape(),
other.shape()
),
));
}
Ok(self.zip_with(other, |a, b| a - b))
}
pub fn hadamard(&self, other: &Matrix) -> Result<Matrix, SymplexError> {
if self.shape() != other.shape() {
return Err(invalid(
"hadamard",
format!(
"cannot multiply element-wise matrices with shapes {:?} and {:?}",
self.shape(),
other.shape()
),
));
}
Ok(self.zip_with(other, |a, b| a * b))
}
fn zip_with(&self, other: &Matrix, f: impl Fn(&Ex, &Ex) -> Ex) -> Matrix {
let rows = (0..self.nrows)
.map(|i| {
(0..self.ncols)
.map(|j| f(&self.rows[i][j], &other.rows[i][j]))
.collect()
})
.collect();
Matrix::with_shape(rows, self.nrows, self.ncols)
}
pub fn scale(&self, scalar: &Ex) -> Matrix {
self.map(|elem| scalar * elem)
}
pub fn matmul(&self, other: &Matrix) -> Result<Matrix, SymplexError> {
if self.ncols != other.nrows {
return Err(invalid(
"matmul",
format!(
"cannot multiply {}×{} by {}×{} matrices",
self.nrows, self.ncols, other.nrows, other.ncols
),
));
}
if self.rows[0][0].ctx_id() == other.rows[0][0].ctx_id()
&& let (Some(qa), Some(qb)) = (self.as_qmatrix(), other.as_qmatrix())
{
let product = qa.matmul(qb).map_err(|e| reop(e, "matmul"))?;
return Ok(product.to_matrix(&self.ctx()));
}
let p = self.ncols;
let rows: Vec<Vec<Ex>> = (0..self.nrows)
.map(|i| {
(0..other.ncols)
.map(|j| {
let mut acc: Ex = &self.rows[i][0] * &other.rows[0][j];
for k in 1..p {
let term = &self.rows[i][k] * &other.rows[k][j];
acc += term;
}
acc
})
.collect()
})
.collect();
Ok(Matrix::with_shape(rows, self.nrows, other.ncols))
}
pub fn trace(&self) -> Result<Ex, SymplexError> {
self.require_square("trace")?;
let rational: Option<Q> = (0..self.nrows)
.try_fold(<Q as num_traits::Zero>::zero(), |acc, i| {
self.rows[i][i].as_rational().map(|v| acc + v)
});
if let Some(sum) = rational {
return Ok(self.ctx().from_ratio(sum));
}
let mut acc = self.rows[0][0].clone();
for i in 1..self.nrows {
acc += &self.rows[i][i];
}
Ok(acc)
}
pub fn det(&self) -> Result<Ex, SymplexError> {
self.require_square("det")?;
self.check_budget("det")?;
let n = self.nrows;
if n > 3
&& let Some(q) = self.as_qmatrix()
{
let d = q.det().map_err(|e| reop(e, "det"))?;
return Ok(self.ctx().from_ratio(d));
}
Ok(match n {
1 => self.rows[0][0].clone(),
2 => {
let a = &self.rows[0][0];
let b = &self.rows[0][1];
let c = &self.rows[1][0];
let d = &self.rows[1][1];
&(a * d) - &(b * c)
}
3 => det_cofactor_inner(&self.rows),
_ => {
let coeffs = self.berkowitz_monic("det")?;
let c0 = coeffs[n].clone();
if n % 2 == 1 { -c0 } else { c0 }
}
})
}
pub fn permanent(&self) -> Result<Ex, SymplexError> {
self.require_square("permanent")?;
let n = self.nrows;
if n > PERMANENT_MAX_DIM {
return Err(failed(
"permanent",
format!(
"needs 2^{n} terms for a {n}×{n} matrix; the limit is \
{PERMANENT_MAX_DIM}×{PERMANENT_MAX_DIM}"
),
));
}
if let Some(q) = self.as_qmatrix() {
let p = q.permanent().map_err(|e| reop(e, "permanent"))?;
return Ok(self.ctx().from_ratio(p));
}
self.check_budget("permanent")?;
let zero = self.ctx_zero();
let size = 1usize << n;
let mut memo: Vec<Ex> = vec![zero.clone(); size];
memo[0] = self.ctx_one();
for s in 1..size {
let row = s.count_ones() as usize - 1;
let mut acc: Option<Ex> = None;
for j in 0..n {
if (s >> j) & 1 == 0 {
continue;
}
let a = &self.rows[row][j];
let sub = &memo[s & !(1 << j)];
if a.is_zero_structural() || sub.is_zero_structural() {
continue;
}
let term = a * sub;
acc = Some(match acc {
None => term,
Some(x) => x + term,
});
}
let val = acc.unwrap_or_else(|| zero.clone());
budget_check(std::iter::once(&val), "permanent")?;
memo[s] = val;
}
let result = memo[size - 1].expand();
budget_check(std::iter::once(&result), "permanent")?;
Ok(result)
}
fn berkowitz_monic(&self, operation: &'static str) -> Result<Vec<Ex>, SymplexError> {
if let Some(q) = self.as_qmatrix() {
let ctx = self.ctx();
let coeffs = q.berkowitz_monic().map_err(|e| reop(e, operation))?;
return Ok(coeffs.into_iter().map(|c| ctx.from_ratio(c)).collect());
}
let n = self.nrows;
let one = self.ctx_one();
let zero = self.ctx_zero();
let mut vec: Vec<Ex> = vec![one.clone(), -&self.rows[n - 1][n - 1]];
for k in 2..=n {
let s = n - k; let a = &self.rows[s][s];
let r: Vec<&Ex> = (s + 1..n).map(|j| &self.rows[s][j]).collect();
let mut c: Vec<Ex> = (s + 1..n).map(|i| self.rows[i][s].clone()).collect();
let mut diags: Vec<Ex> = Vec::with_capacity(k + 1);
diags.push(one.clone());
diags.push(-a);
for step in 0..(k - 1) {
if step > 0 {
let next: Vec<Ex> = (s + 1..n)
.map(|i| {
let mut acc = zero.clone();
for (idx, j) in (s + 1..n).enumerate() {
acc += &self.rows[i][j] * &c[idx];
}
acc.expand()
})
.collect();
c = next;
}
let mut rc = zero.clone();
for (ri, ci) in r.iter().zip(c.iter()) {
rc += *ri * ci;
}
diags.push((-rc).expand());
}
let mut next_vec: Vec<Ex> = Vec::with_capacity(k + 1);
for i in 0..=k {
let mut acc = zero.clone();
for (j, v) in vec.iter().enumerate().take(k) {
if j <= i {
acc += &diags[i - j] * v;
}
}
next_vec.push(acc.expand());
}
vec = next_vec;
budget_check(vec.iter().chain(c.iter()), operation)?;
}
Ok(vec)
}
pub fn map(&self, mut f: impl FnMut(&Ex) -> Ex) -> Matrix {
let rows = self
.rows
.iter()
.map(|row| row.iter().map(&mut f).collect())
.collect();
Matrix::with_shape(rows, self.nrows, self.ncols)
}
pub fn map_indexed(&self, mut f: impl FnMut(usize, usize, &Ex) -> Ex) -> Matrix {
let rows = self
.rows
.iter()
.enumerate()
.map(|(i, row)| row.iter().enumerate().map(|(j, e)| f(i, j, e)).collect())
.collect();
Matrix::with_shape(rows, self.nrows, self.ncols)
}
pub fn minor_matrix(&self, row: usize, col: usize) -> Result<Matrix, SymplexError> {
self.require_square("minor_matrix")?;
if self.nrows <= 1 {
return Err(invalid("minor_matrix", "requires matrix dimension > 1"));
}
if row >= self.nrows || col >= self.ncols {
return Err(invalid(
"minor_matrix",
format!(
"index ({row}, {col}) out of range for {}×{} matrix",
self.nrows, self.ncols
),
));
}
let rows: Vec<Vec<Ex>> = self
.rows
.iter()
.enumerate()
.filter(|(r, _)| *r != row)
.map(|(_, row_data)| {
row_data
.iter()
.enumerate()
.filter(|(c, _)| *c != col)
.map(|(_, v)| v.clone())
.collect()
})
.collect();
Ok(Matrix::from_rows_unchecked(rows))
}
pub fn minor(&self, row: usize, col: usize) -> Result<Ex, SymplexError> {
self.minor_matrix(row, col)?.det()
}
pub fn cofactor(&self, row: usize, col: usize) -> Result<Ex, SymplexError> {
let minor_det = self.minor(row, col)?;
Ok(if (row + col).is_multiple_of(2) {
minor_det
} else {
-minor_det
})
}
pub fn adjugate(&self) -> Result<Matrix, SymplexError> {
self.require_square("adjugate")?;
let n = self.nrows;
if n == 1 {
return Ok(Matrix::from_rows_unchecked(vec![vec![self.ctx_one()]]));
}
let mut rows = Vec::with_capacity(n);
for j in 0..n {
let mut row = Vec::with_capacity(n);
for i in 0..n {
row.push(self.cofactor(i, j)?);
}
rows.push(row);
}
Ok(Matrix::from_rows_unchecked(rows))
}
pub fn inv(&self) -> Result<Matrix, SymplexError> {
self.require_square("inv")?;
if let Some(q) = self.as_qmatrix() {
return match q.inv() {
Ok(inv) => Ok(inv.to_matrix(&self.ctx())),
Err(SymplexError::ComputationFailed { .. }) => {
Err(failed("inv", "matrix is singular (determinant is zero)"))
}
Err(e) => Err(reop(e, "inv")),
};
}
self.check_budget("inv")?;
let d = self.det().map_err(|e| reop(e, "inv"))?;
if ex_is_zero(&d) == Some(true) {
return Err(failed("inv", "matrix is singular (determinant is zero)"));
}
let n = self.nrows;
if n == 1 {
let one_over_det = &self.ctx_one() / &d;
return Ok(Matrix::from_rows_unchecked(vec![vec![one_over_det]]));
}
let adj = self.adjugate().map_err(|e| reop(e, "inv"))?;
budget_check(adj.iter().chain(std::iter::once(&d)), "inv")?;
let one_over_det = &self.ctx_one() / &d;
Ok(adj.scale(&one_over_det))
}
pub fn solve(&self, b: &Matrix) -> Result<Matrix, SymplexError> {
let n = self.nrows;
self.require_square("solve")?;
if b.nrows != n {
return Err(invalid(
"solve",
format!(
"row count mismatch: A is {}×{}, b has {} rows",
n, self.ncols, b.nrows
),
));
}
if let (Some(qa), Some(qb)) = (self.as_qmatrix(), b.as_qmatrix()) {
return qa.solve(qb).map(|x| x.to_matrix(&self.ctx()));
}
let augmented = Matrix::hstack(&[self, b])?;
augmented.check_budget("solve")?;
let (rref_mat, pivots) = augmented.rref();
rref_mat.check_budget("solve")?;
if pivots.len() != n || pivots.iter().any(|&p| p >= n) {
return Err(failed(
"solve",
"matrix is singular; no unique solution exists",
));
}
let b_cols = b.ncols;
let sol_rows: Vec<Vec<Ex>> = (0..n)
.map(|i| {
(n..(n + b_cols))
.map(|j| rref_mat.rows[i][j].eval())
.collect()
})
.collect();
Ok(Matrix::from_rows_unchecked(sol_rows))
}
pub fn solve_least_squares(&self, b: &Matrix) -> Result<Matrix, SymplexError> {
if b.nrows != self.nrows {
return Err(invalid(
"solve_least_squares",
format!(
"row count mismatch: A is {}×{}, b has {} rows",
self.nrows, self.ncols, b.nrows
),
));
}
let at = self.transpose();
let ata = at.matmul(self)?;
let atb = at.matmul(b)?;
ata.solve(&atb).map_err(|_| {
failed(
"solve_least_squares",
"AᵀA is singular (A does not have full column rank)",
)
})
}
pub fn char_poly_coeffs(&self) -> Result<Vec<Ex>, SymplexError> {
self.require_square("char_poly_coeffs")?;
self.check_budget("char_poly_coeffs")?;
let n = self.nrows;
let monic = self.berkowitz_monic("char_poly_coeffs")?; let sign_flip = n % 2 == 1;
Ok((0..=n)
.map(|k| {
let c = monic[n - k].clone();
if sign_flip { -c } else { c }
})
.collect())
}
pub fn char_poly(&self, var: &Ex) -> Result<Ex, SymplexError> {
let coeffs = self.char_poly_coeffs()?;
let mut acc = coeffs[0].clone();
for (k, c) in coeffs.iter().enumerate().skip(1) {
if c.is_zero_structural() {
continue;
}
acc += c * &var.powi(k as i64);
}
Ok(acc.expand())
}
pub fn eigenvals_with_multiplicity(&self) -> Result<Vec<(Ex, usize)>, SymplexError> {
self.require_square("eigenvals")?;
let lam = self.fresh_symbol("lambda");
let pairs = self.eigen_pairs(&lam)?;
Ok(pairs
.into_iter()
.map(|(v, m)| (self.hide_dummy(&v, &lam), m))
.collect())
}
fn eigen_pairs(&self, lam: &Ex) -> Result<Vec<(Ex, usize)>, SymplexError> {
let n = self.nrows;
let coeffs = self.char_poly_coeffs()?;
let cp = {
let mut acc = coeffs[0].clone();
for (k, c) in coeffs.iter().enumerate().skip(1) {
if !c.is_zero_structural() {
acc += c * &lam.powi(k as i64);
}
}
acc.expand()
};
let symbolic = coeffs.iter().any(|c| c.expr_type() != ExprType::Number);
let mut pairs = if n <= 2 && symbolic {
low_degree_roots(&coeffs)
} else {
eigvals_with_multiplicity(&cp, lam)
};
let mut total: usize = pairs.iter().map(|(_, m)| *m).sum();
if total < n && n <= 2 {
pairs = low_degree_roots(&coeffs);
total = pairs.iter().map(|(_, m)| *m).sum();
}
if pairs.is_empty() {
return Err(failed(
"eigenvals",
"could not solve the characteristic polynomial; \
symbolic matrices larger than 2×2 need a factorable characteristic polynomial",
));
}
if total < n {
warn!(
"eigenvals: found algebraic multiplicity {} for a {}×{} matrix — \
characteristic polynomial may have factors beyond solver capability",
total, n, n
);
}
Ok(pairs)
}
pub fn eigenvals(&self) -> Result<Vec<Ex>, SymplexError> {
let pairs = self.eigenvals_with_multiplicity()?;
let mut out = Vec::new();
for (v, m) in pairs {
for _ in 0..m {
out.push(v.clone());
}
}
Ok(out)
}
pub fn eigenvects(&self) -> Result<Vec<(Ex, usize, Vec<Matrix>)>, SymplexError> {
self.require_square("eigenvects")?;
let n = self.nrows;
debug!(n, "eigenvects: computing for {}×{} matrix", n, n);
let lam = self.fresh_symbol("lambda");
let eigen_pairs = self.eigen_pairs(&lam)?;
let eye = Matrix::identity(&self.ctx(), n);
let mut result = Vec::new();
for (eigenval, alg_mult) in &eigen_pairs {
let a_minus_lambda_i = self.sub(&eye.scale(eigenval))?;
let vecs = a_minus_lambda_i.nullspace_semantic();
trace!(
alg_mult,
geom_mult = vecs.len(),
"eigenvects: eigenvalue has alg_mult={}, geom_mult={}",
alg_mult,
vecs.len()
);
let vecs = vecs
.into_iter()
.map(|v| v.map(|e| self.hide_dummy(e, &lam)))
.collect();
result.push((self.hide_dummy(eigenval, &lam), *alg_mult, vecs));
}
Ok(result)
}
pub fn is_diagonalizable(&self) -> Option<bool> {
if !self.is_square() {
return Some(false);
}
let eigvs = self.eigenvects().ok()?;
let total_alg: usize = eigvs.iter().map(|(_, m, _)| *m).sum();
if total_alg != self.nrows {
return None;
}
if eigvs.iter().all(|(_, m, v)| v.len() == *m) {
return Some(true);
}
if self.iter().all(Ex::is_constant) {
Some(false)
} else {
None
}
}
pub fn diagonalize(&self) -> Result<Diagonalization<Matrix>, SymplexError> {
self.require_square("diagonalize")?;
self.check_budget("diagonalize")?;
debug!(
"diagonalize: attempting for {}×{} matrix",
self.nrows, self.ncols
);
let eigvs = self.eigenvects().map_err(|e| reop(e, "diagonalize"))?;
let n = self.nrows;
let mut total_vecs = 0usize;
for (_, alg_mult, vecs) in &eigvs {
if vecs.len() != *alg_mult {
return Err(failed(
"diagonalize",
"matrix is not diagonalizable: geometric multiplicity \
does not equal algebraic multiplicity for all eigenvalues",
));
}
total_vecs += vecs.len();
}
if total_vecs != n {
return Err(failed(
"diagonalize",
"matrix is not diagonalizable: insufficient eigenvectors found",
));
}
let mut p_cols: Vec<&Matrix> = Vec::new();
let mut diag_entries: Vec<Ex> = Vec::new();
for (eigenval, _, vecs) in &eigvs {
for v in vecs {
p_cols.push(v);
diag_entries.push(eigenval.clone());
}
}
let p = Matrix::hstack(&p_cols)?;
p.check_budget("diagonalize")?;
let d = Matrix::diag(&diag_entries);
Ok(Diagonalization { p, d })
}
pub fn jordan_form(&self) -> Result<JordanForm<Matrix>, SymplexError> {
self.require_square("jordan_form")?;
self.check_budget("jordan_form")?;
let n = self.nrows;
debug!(n, "jordan_form: computing for {}×{} matrix", n, n);
let eye = Matrix::identity(&self.ctx(), n);
let eigvs = self.eigenvects().map_err(|e| reop(e, "jordan_form"))?;
let total_vecs: usize = eigvs.iter().map(|(_, _, v)| v.len()).sum();
let all_match = eigvs.iter().all(|(_, m, v)| v.len() == *m);
if all_match && total_vecs == n {
let Diagonalization { p, d } = self.diagonalize()?;
return Ok(JordanForm { p, j: d });
}
let total_alg: usize = eigvs.iter().map(|(_, m, _)| *m).sum();
if total_alg != n {
return Err(failed(
"jordan_form",
format!(
"eigenvalue solver found algebraic multiplicity sum {} but matrix is {}×{}",
total_alg, n, n
),
));
}
let mut jordan_rows: Vec<Vec<Ex>> = Vec::new();
let mut basis_cols: Vec<Matrix> = Vec::new();
for (eigenval, alg_mult, _) in &eigvs {
let a_minus_lambda = self.sub(&eye.scale(eigenval))?;
let mut chain: Vec<usize> = vec![0];
let mut power = a_minus_lambda.clone();
loop {
let nullity = n - power.rank_semantic();
let last = *chain.last().unwrap_or(&0);
if nullity == last || nullity >= *alg_mult {
if nullity > last {
chain.push(nullity);
}
break;
}
chain.push(nullity);
power = power.matmul(&a_minus_lambda)?;
power.check_budget("jordan_form")?;
}
let max_block_size = chain.len() - 1;
let mut block_counts: Vec<usize> = Vec::new();
for i in 0..max_block_size {
let diff_curr = chain[i + 1] - chain[i];
let diff_next = if i + 2 < chain.len() {
chain[i + 2] - chain[i + 1]
} else {
0
};
block_counts.push(diff_curr - diff_next);
}
let mut blocks: Vec<(usize, usize)> = block_counts
.iter()
.enumerate()
.filter(|(_, c)| **c > 0)
.map(|(i, c)| (i + 1, *c))
.collect();
blocks.sort_by_key(|b| std::cmp::Reverse(b.0));
let mut eig_basis: Vec<Matrix> = Vec::new();
for (block_size, count) in &blocks {
for _ in 0..*count {
let null_big = jordan_null_power(&a_minus_lambda, *block_size)?;
let null_small = if *block_size > 1 {
jordan_null_power(&a_minus_lambda, block_size - 1)?
} else {
Vec::new()
};
let exclude: Vec<&Matrix> = null_small.iter().chain(eig_basis.iter()).collect();
let Some(vec) = pick_independent_vec(&null_big, &exclude)? else {
return Err(failed(
"jordan_form",
"could not find independent generalized eigenvector",
));
};
let mut chain_vecs: Vec<Matrix> = Vec::with_capacity(*block_size);
for i in (0..*block_size).rev() {
chain_vecs.push(if i == 0 {
vec.clone()
} else {
matrix_pow_vec(&a_minus_lambda, &vec, i)?
});
}
eig_basis.extend(chain_vecs.iter().cloned());
basis_cols.extend(chain_vecs);
let bs = *block_size;
let col_offset = jordan_rows.len();
for row_idx in 0..bs {
let mut row = vec![self.ctx_zero(); n];
row[col_offset + row_idx] = eigenval.clone();
if row_idx + 1 < bs {
row[col_offset + row_idx + 1] = self.ctx_one();
}
jordan_rows.push(row);
}
}
}
}
if jordan_rows.len() != n || basis_cols.len() != n {
return Err(failed(
"jordan_form",
format!(
"internal error: expected {} basis vectors, got {}",
n,
basis_cols.len()
),
));
}
let j = Matrix::from_rows_unchecked(jordan_rows);
let col_refs: Vec<&Matrix> = basis_cols.iter().collect();
let p = Matrix::hstack(&col_refs)?;
p.check_budget("jordan_form")?;
Ok(JordanForm { p, j })
}
pub fn powi(&self, n: u32) -> Result<Matrix, SymplexError> {
self.require_square("powi")?;
if n == 0 {
return Ok(Matrix::identity(&self.ctx(), self.nrows));
}
if n == 1 {
return Ok(self.clone());
}
let mut result = Matrix::identity(&self.ctx(), self.nrows);
let mut base = self.clone();
let mut exp = n;
while exp > 0 {
if exp % 2 == 1 {
result = result.matmul(&base)?;
}
exp /= 2;
if exp > 0 {
base = base.matmul(&base)?;
}
}
Ok(result)
}
pub fn kronecker(&self, other: &Matrix) -> Matrix {
let mut rows = Vec::with_capacity(self.nrows * other.nrows);
for i in 0..self.nrows {
for k in 0..other.nrows {
let mut row = Vec::with_capacity(self.ncols * other.ncols);
for j in 0..self.ncols {
for l in 0..other.ncols {
row.push(&self.rows[i][j] * &other.rows[k][l]);
}
}
rows.push(row);
}
}
Matrix::from_rows_unchecked(rows)
}
pub fn exp_series(&self, order: usize) -> Result<Matrix, SymplexError> {
self.require_square("exp_series")?;
let n = self.nrows;
let ctx = self.ctx();
let mut result = Matrix::identity(&ctx, n);
let mut term = Matrix::identity(&ctx, n);
for k in 1..=order {
term = term.matmul(self)?;
let inv_k = ctx.rational(1, k as i64);
term = term.scale(&inv_k);
result = result.add(&term)?;
}
Ok(result)
}
pub fn matrix_exp(&self) -> Result<Matrix, SymplexError> {
self.matrix_exp_impl(None)
}
pub fn matrix_exp_t(&self, t: &Ex) -> Result<Matrix, SymplexError> {
self.matrix_exp_impl(Some(t))
}
fn matrix_exp_impl(&self, t: Option<&Ex>) -> Result<Matrix, SymplexError> {
self.require_square("matrix_exp")?;
debug!(
"matrix_exp: computing for {}×{} matrix",
self.nrows, self.ncols
);
let n = self.nrows;
let ctx = self.ctx();
let JordanForm { p, j } = self.jordan_form().map_err(|e| match e {
SymplexError::ComputationFailed { reason, .. }
if reason.starts_with("expression swell") =>
{
failed("matrix_exp", reason)
}
e => failed(
"matrix_exp",
format!(
"Jordan form unavailable ({e}); use exp_series(order) for a truncated approximation"
),
),
})?;
let mut exp_j_rows: Vec<Vec<Ex>> = vec![vec![self.ctx_zero(); n]; n];
let mut col = 0;
while col < n {
let lambda = j.rows[col][col].clone();
let mut block_size = 1;
while col + block_size < n {
let superdiag = &j.rows[col + block_size - 1][col + block_size];
let diag_next = &j.rows[col + block_size][col + block_size];
if !superdiag.is_one_structural() || diag_next != &lambda {
break;
}
block_size += 1;
}
trace!(col, block_size, "matrix_exp: processing Jordan block");
let exponent = match t {
Some(t) => &lambda * t,
None => lambda.clone(),
};
let exp_lambda = exponent.exp();
for i in 0..block_size {
for jj in i..block_size {
let d = jj - i;
let factorial_val = ctx.from_bigint(factorial(d as u64));
let mut entry = &exp_lambda / &factorial_val;
if d > 0
&& let Some(t) = t
{
entry *= t.powi(d as i64);
}
exp_j_rows[col + i][col + jj] = entry;
}
}
col += block_size;
}
let exp_j = Matrix::from_rows_unchecked(exp_j_rows);
let p_inv = p.inv().map_err(|e| match e {
SymplexError::ComputationFailed { reason, .. }
if reason.starts_with("expression swell") =>
{
failed("matrix_exp", reason)
}
_ => failed(
"matrix_exp",
"eigenvector matrix is singular (internal inconsistency)",
),
})?;
let result = p.matmul(&exp_j)?.matmul(&p_inv)?;
result.check_budget("matrix_exp")?;
let i_unit = ctx.i_unit();
Ok(result.map(|e| {
if has_root_of(e) {
return e.eval();
}
let s = e.simplify();
if s.contains(&i_unit) {
fix_trig_parity(&s.rewrite_as_trig())
.expand()
.eval()
.simplify()
} else {
s
}
}))
}
pub fn pinv(&self) -> Result<Matrix, SymplexError> {
let ctx = self.ctx();
if let Some(q) = self.as_qmatrix() {
return q
.pinv()
.map(|p| p.to_matrix(&ctx))
.map_err(|e| reop(e, "pinv"));
}
self.check_budget("pinv")?;
let inv_err = |what: &'static str| {
move |e: SymplexError| match e {
SymplexError::ComputationFailed { reason, .. }
if reason.starts_with("expression swell") =>
{
failed("pinv", reason)
}
_ => failed(
"pinv",
format!(
"{what} is singular: the columns are linearly dependent in a way the \
structural pivot search did not detect; simplify the entries first"
),
),
}
};
let (r, pivots) = self.rref();
if pivots.is_empty() {
return Ok(Matrix::zeros(&ctx, self.ncols, self.nrows));
}
let at = self.transpose();
if pivots.len() == self.ncols {
let ata = at.matmul(self)?;
let ata_inv = ata.inv().map_err(inv_err("AᵀA"))?;
return ata_inv.matmul(&at);
}
let c = self.select_cols(&pivots)?;
let f = r.select_rows(&(0..pivots.len()).collect::<Vec<_>>())?;
let ft = f.transpose();
let ct = c.transpose();
let fft_inv = f.matmul(&ft)?.inv().map_err(inv_err("FFᵀ"))?;
let ctc_inv = ct.matmul(&c)?.inv().map_err(inv_err("CᵀC"))?;
let result = ft.matmul(&fft_inv)?.matmul(&ctc_inv)?.matmul(&ct)?;
result.check_budget("pinv")?;
Ok(result)
}
pub fn rank_decomposition(&self) -> Result<RankDecomposition<Matrix>, SymplexError> {
let ctx = self.ctx();
if let Some(q) = self.as_qmatrix() {
let RankDecomposition { c, f } = q
.rank_decomposition()
.map_err(|e| reop(e, "rank_decomposition"))?;
return Ok(RankDecomposition {
c: c.to_matrix(&ctx),
f: f.to_matrix(&ctx),
});
}
let (r, pivots) = self.rref();
if pivots.is_empty() {
return Err(failed(
"rank_decomposition",
"matrix is zero (rank 0); the factors C (m×0) and F (0×n) would be empty",
));
}
let c = self.select_cols(&pivots)?;
let f = r.select_rows(&(0..pivots.len()).collect::<Vec<_>>())?;
Ok(RankDecomposition { c, f })
}
pub fn singular_values(&self) -> Result<Vec<Ex>, SymplexError> {
let ctx = self.ctx();
let (m, n) = self.shape();
let gram = if let Some(q) = self.as_qmatrix() {
let qt = q.transpose();
let g = if m >= n { qt.matmul(q) } else { q.matmul(&qt) };
g.map_err(|e| reop(e, "singular_values"))?.to_matrix(&ctx)
} else {
let at = self.transpose();
if m >= n {
at.matmul(self)?
} else {
self.matmul(&at)?
}
};
let dim = gram.nrows;
let eig = if gram.is_diagonal() == Some(true) {
gram.diagonal()
} else {
gram.eigenvals().map_err(|e| reop(e, "singular_values"))?
};
if eig.len() < dim {
return Err(failed(
"singular_values",
format!(
"found only {} of the {dim} eigenvalues of the Gram matrix AᵀA",
eig.len()
),
));
}
let mut vals: Vec<Ex> = eig.iter().map(|l| l.sqrt().eval()).collect();
vals.resize_with(vals.len().max(n), || ctx.zero());
sort_descending_numeric(&mut vals);
Ok(vals)
}
pub fn condition_number(&self) -> Result<Ex, SymplexError> {
let sv = self
.singular_values()
.map_err(|e| reop(e, "condition_number"))?;
let ctx = self.ctx();
let numeric = sv.iter().all(|v| v.eval_f64().is_ok());
let (max, min) = match (sv.first(), sv.last()) {
(Some(first), Some(last)) if numeric => (first.clone(), last.clone()),
_ => (
Ex::max_of(&ctx, sv.iter().cloned()),
Ex::min_of(&ctx, sv.iter().cloned()),
),
};
if ex_is_zero(&min) == Some(true) {
return Err(failed(
"condition_number",
"matrix is singular (smallest singular value is 0), condition number is infinite",
));
}
Ok(&max / &min)
}
pub fn hessenberg(&self) -> Result<Hessenberg<Matrix>, SymplexError> {
self.require_square("hessenberg")?;
let ctx = self.ctx();
if let Some(q) = self.as_qmatrix() {
let Hessenberg { h, p } = q.hessenberg().map_err(|e| reop(e, "hessenberg"))?;
return Ok(Hessenberg {
h: h.to_matrix(&ctx),
p: p.to_matrix(&ctx),
});
}
self.check_budget("hessenberg")?;
let n = self.nrows;
let zero = self.ctx_zero();
let one = self.ctx_one();
let mut h: Vec<Vec<Ex>> = self.rows.clone();
let mut p: Vec<Vec<Ex>> = (0..n)
.map(|i| {
(0..n)
.map(|j| if i == j { one.clone() } else { zero.clone() })
.collect()
})
.collect();
let tidy = |e: Ex| {
if e.is_constant() {
e.eval()
} else {
e.simplify()
}
};
for k in 0..n.saturating_sub(2) {
let mut pivot = None;
let mut fallback = None;
for (i, row) in h.iter().enumerate().skip(k + 1) {
match ex_is_zero(&row[k]) {
Some(false) => {
pivot = Some(i);
break;
}
None if fallback.is_none() => fallback = Some(i),
_ => {}
}
}
let Some(piv) = pivot.or(fallback) else {
for row in h.iter_mut().skip(k + 1) {
row[k] = zero.clone();
}
continue;
};
if piv != k + 1 {
h.swap(k + 1, piv);
for row in h.iter_mut() {
row.swap(k + 1, piv);
}
for row in p.iter_mut() {
row.swap(k + 1, piv);
}
}
let pv = h[k + 1][k].clone();
for j in (k + 2)..n {
if ex_is_zero(&h[j][k]) == Some(true) {
h[j][k] = zero.clone();
continue;
}
let f = tidy(&h[j][k] / &pv);
let pivot_row = h[k + 1].clone();
for (c, pr) in pivot_row.iter().enumerate() {
if c == k {
h[j][c] = zero.clone();
} else if !pr.is_zero_structural() {
let v = tidy(&h[j][c] - &(&f * pr));
h[j][c] = v;
}
}
for r in 0..n {
if !h[r][j].is_zero_structural() {
let v = tidy(&h[r][k + 1] + &(&f * &h[r][j]));
h[r][k + 1] = v;
}
if !p[r][j].is_zero_structural() {
let v = tidy(&p[r][k + 1] + &(&f * &p[r][j]));
p[r][k + 1] = v;
}
}
}
budget_check(h.iter().flatten().chain(p.iter().flatten()), "hessenberg")?;
}
Ok(Hessenberg {
h: Matrix::from_rows_unchecked(h),
p: Matrix::from_rows_unchecked(p),
})
}
}
fn sort_descending_numeric(vals: &mut [Ex]) {
let keys: Option<Vec<f64>> = vals.iter().map(|v| v.eval_f64().ok()).collect();
let Some(keys) = keys else { return };
let mut order: Vec<usize> = (0..vals.len()).collect();
order.sort_by(|&a, &b| keys[b].total_cmp(&keys[a]));
let sorted: Vec<Ex> = order.iter().map(|&i| vals[i].clone()).collect();
vals.clone_from_slice(&sorted);
}
fn det_cofactor_inner(m: &[Vec<Ex>]) -> Ex {
let n = m.len();
if n == 1 {
return m[0][0].clone();
}
if n == 2 {
let ad = &m[0][0] * &m[1][1];
let bc = &m[0][1] * &m[1][0];
return ad - bc;
}
let mut result: Option<Ex> = None;
for j in 0..n {
let minor: Vec<Vec<Ex>> = (1..n)
.map(|row| {
(0..n)
.filter(|&col| col != j)
.map(|col| m[row][col].clone())
.collect()
})
.collect();
let cofactor = det_cofactor_inner(&minor);
let term = &m[0][j] * &cofactor;
result = Some(match result {
None => {
if j % 2 == 0 {
term
} else {
-term
}
}
Some(acc) => {
if j % 2 == 0 {
acc + term
} else {
acc - term
}
}
});
}
result.unwrap_or_else(|| m[0][0].clone())
}
fn low_degree_roots(coeffs: &[Ex]) -> Vec<(Ex, usize)> {
match coeffs.len() {
2 => {
let root = (-&coeffs[0] / &coeffs[1]).eval().simplify();
vec![(root, 1)]
}
3 => {
let (c0, c1, c2) = (&coeffs[0], &coeffs[1], &coeffs[2]);
let ctx = c0.context();
let two = ctx.int(2);
let disc = (&c1.powi(2) - &(&(&ctx.int(4) * c2) * c0)).expand();
let denom = &two * c2;
if ex_is_zero(&disc) == Some(true) {
let root = (-c1 / &denom).eval();
return vec![(root, 2)];
}
let sq = quadratic_sqrt(&disc);
let r1 = (&(-c1 + &sq) / &denom).eval();
let r2 = (&(-c1 - &sq) / &denom).eval();
vec![(r1, 1), (r2, 1)]
}
_ => Vec::new(),
}
}
fn quadratic_sqrt(disc: &Ex) -> Ex {
use crate::base::node::ExprNode;
let ctx = disc.context();
let strip_abs = |e: &Ex| {
e.replace(|v| match v.node() {
ExprNode::Abs(inner) => Some(e.wrap(*inner)),
_ => None,
})
};
for (base, factor) in [(disc.clone(), ctx.one()), ((-disc).eval(), ctx.i_unit())] {
let root = strip_abs(&base.sqrt().simplify());
if has_radical(&root) {
continue;
}
let check = (&root.powi(2).expand() - &base.expand()).expand();
if check.is_zero_structural() {
return (&factor * &root).eval();
}
}
disc.sqrt()
}
pub fn jacobian(funcs: &[&Ex], vars: &[&Ex]) -> Matrix {
assert!(!funcs.is_empty(), "jacobian: funcs must be non-empty");
assert!(!vars.is_empty(), "jacobian: vars must be non-empty");
let rows: Vec<Vec<Ex>> = funcs
.iter()
.map(|fi| vars.iter().map(|vj| fi.diff(vj)).collect())
.collect();
Matrix::from_rows_unchecked(rows)
}
impl Matrix {
pub fn diff(&self, var: &Ex) -> Matrix {
self.map(|elem| elem.diff(var))
}
pub fn integrate(&self, var: &Ex) -> Matrix {
self.map(|elem| elem.integrate(var))
}
pub fn subs(&self, old: &Ex, new: &Ex) -> Matrix {
self.map(|elem| elem.subs(old, new))
}
pub fn eval(&self) -> Matrix {
self.map(|elem| elem.eval())
}
pub fn expand(&self) -> Matrix {
self.map(|elem| elem.expand())
}
pub fn simplify(&self) -> Matrix {
self.map(|elem| elem.simplify())
}
}
impl Matrix {
pub fn to_rust_fn(&self, name: &str, params: &[&str]) -> Result<String, SymplexError> {
self.to_rust_fn_with_options(name, params, &CodegenOptions::default())
}
pub fn to_rust_fn_with_options(
&self,
name: &str,
params: &[&str],
options: &CodegenOptions,
) -> Result<String, SymplexError> {
let first = &self.rows[0][0];
let mut guard = first.inner.write();
let arena = &mut guard.arena;
let entry_ids: Vec<crate::base::node::ExprId> = self
.rows
.iter()
.flat_map(|row| row.iter().map(|e| e.raw_id()))
.collect();
crate::output::codegen::matrix_to_rust_fn(
arena, &entry_ids, self.nrows, self.ncols, name, params, options,
)
}
}
impl Matrix {
pub fn to_latex(&self) -> String {
let mut s = String::from(r"\begin{bmatrix} ");
for i in 0..self.nrows {
if i > 0 {
s.push_str(r" \\ ");
}
for j in 0..self.ncols {
if j > 0 {
s.push_str(" & ");
}
s.push_str(&self.rows[i][j].to_latex());
}
}
s.push_str(r" \end{bmatrix}");
s
}
}
#[allow(clippy::needless_range_loop)]
impl Matrix {
pub fn lu(&self) -> Result<Lu<Matrix>, SymplexError> {
self.require_square("lu")?;
let n = self.nrows;
if let Some(q) = self.as_qmatrix() {
let ctx = self.ctx();
let Lu { l, u, perm } = q.lu().map_err(|e| reop(e, "lu"))?;
return Ok(Lu {
l: l.to_matrix(&ctx),
u: u.to_matrix(&ctx),
perm,
});
}
let mut perm: Vec<usize> = (0..n).collect();
let mut u: Vec<Vec<Ex>> = self.rows.clone();
let one = self.ctx_one();
let zero = self.ctx_zero();
let mut l: Vec<Vec<Ex>> = (0..n)
.map(|i| {
(0..n)
.map(|j| if i == j { one.clone() } else { zero.clone() })
.collect()
})
.collect();
for k in 0..n {
let mut pivot_row = None;
for i in k..n {
if !u[i][k].is_zero_structural() {
pivot_row = Some(i);
break;
}
}
let Some(pivot_row) = pivot_row else {
return Err(failed("lu", "matrix is singular (zero pivot column)"));
};
if pivot_row != k {
u.swap(k, pivot_row);
perm.swap(k, pivot_row);
for j in 0..k {
let tmp = l[k][j].clone();
l[k][j] = l[pivot_row][j].clone();
l[pivot_row][j] = tmp;
}
}
for i in (k + 1)..n {
if u[i][k].is_zero_structural() {
continue;
}
let factor = &u[i][k] / &u[k][k];
l[i][k] = factor.clone();
u[i][k] = zero.clone();
for j in (k + 1)..n {
let term = &factor * &u[k][j];
u[i][j] = &u[i][j] - &term;
}
}
}
Ok(Lu {
l: Matrix::from_rows_unchecked(l),
u: Matrix::from_rows_unchecked(u),
perm,
})
}
pub fn rref(&self) -> (Matrix, Vec<usize>) {
self.rref_by(&|e: &Ex| e.is_zero_structural())
}
fn rref_by(&self, is_zero: &dyn Fn(&Ex) -> bool) -> (Matrix, Vec<usize>) {
if let Some(q) = self.as_qmatrix() {
let (r, pivots) = q.rref();
return (r.to_matrix(&self.ctx()), pivots);
}
let nrows = self.nrows;
let ncols = self.ncols;
let mut rows: Vec<Vec<Ex>> = self.rows.clone();
let mut pivots = Vec::new();
let mut pivot_row = 0;
for col in 0..ncols {
if pivot_row >= nrows {
break;
}
let mut found = None;
for i in pivot_row..nrows {
if !is_zero(&rows[i][col]) {
found = Some(i);
break;
}
}
let Some(found) = found else { continue };
if found != pivot_row {
rows.swap(pivot_row, found);
}
let pivot_val = rows[pivot_row][col].clone();
for j in 0..ncols {
rows[pivot_row][j] = &rows[pivot_row][j] / &pivot_val;
}
for i in 0..nrows {
if i == pivot_row {
continue;
}
if !is_zero(&rows[i][col]) {
let factor = rows[i][col].clone();
for j in 0..ncols {
let term = &factor * &rows[pivot_row][j];
rows[i][j] = &rows[i][j] - &term;
}
}
}
pivots.push(col);
pivot_row += 1;
}
(Matrix::with_shape(rows, nrows, ncols), pivots)
}
pub fn rank(&self) -> usize {
let (_, pivots) = self.rref();
pivots.len()
}
fn rank_semantic(&self) -> usize {
self.rref_by(&eigen_zero_test).1.len()
}
pub fn nullspace(&self) -> Vec<Matrix> {
self.nullspace_by(&|e: &Ex| e.is_zero_structural())
}
fn nullspace_semantic(&self) -> Vec<Matrix> {
self.nullspace_by(&eigen_zero_test)
}
fn nullspace_by(&self, is_zero: &dyn Fn(&Ex) -> bool) -> Vec<Matrix> {
let (rref_mat, pivots) = self.rref_by(is_zero);
let n = self.ncols;
let pivot_set: std::collections::HashSet<usize> = pivots.iter().copied().collect();
let free_vars: Vec<usize> = (0..n).filter(|c| !pivot_set.contains(c)).collect();
let mut basis = Vec::new();
for &free_col in &free_vars {
let mut entries = vec![self.ctx_zero(); n];
entries[free_col] = self.ctx_one();
for (pivot_idx, &pivot_col) in pivots.iter().enumerate() {
entries[pivot_col] = -&rref_mat.rows[pivot_idx][free_col];
}
basis.push(Matrix::col_vector(entries));
}
basis
}
pub fn columnspace(&self) -> Vec<Matrix> {
let (_, pivots) = self.rref();
pivots
.iter()
.map(|&col| Matrix::col_vector(self.col(col)))
.collect()
}
pub fn rowspace(&self) -> Vec<Matrix> {
let (rref_mat, pivots) = self.rref();
(0..pivots.len())
.map(|i| Matrix::row_vector(rref_mat.rows[i].clone()))
.collect()
}
pub fn left_nullspace(&self) -> Vec<Matrix> {
self.transpose().nullspace()
}
}
impl Matrix {
pub fn norm_frobenius(&self) -> Ex {
let mut sum = self.ctx_zero();
for elem in self.iter() {
sum += &(elem * elem);
}
sum.sqrt()
}
pub fn norm(&self) -> Ex {
self.norm_frobenius()
}
pub fn hstack(matrices: &[&Matrix]) -> Result<Matrix, SymplexError> {
if matrices.is_empty() {
return Err(invalid("hstack", "need at least one matrix"));
}
let nrows = matrices[0].nrows;
for (idx, m) in matrices.iter().enumerate() {
if m.nrows != nrows {
return Err(invalid(
"hstack",
format!("matrix {idx} has {} rows, expected {nrows}", m.nrows),
));
}
}
let ncols: usize = matrices.iter().map(|m| m.ncols).sum();
let rows: Vec<Vec<Ex>> = (0..nrows)
.map(|i| {
matrices
.iter()
.flat_map(|m| m.rows[i].iter().cloned())
.collect()
})
.collect();
Ok(Matrix::with_shape(rows, nrows, ncols))
}
pub fn vstack(matrices: &[&Matrix]) -> Result<Matrix, SymplexError> {
if matrices.is_empty() {
return Err(invalid("vstack", "need at least one matrix"));
}
let ncols = matrices[0].ncols;
for (idx, m) in matrices.iter().enumerate() {
if m.ncols != ncols {
return Err(invalid(
"vstack",
format!("matrix {idx} has {} cols, expected {ncols}", m.ncols),
));
}
}
let nrows: usize = matrices.iter().map(|m| m.nrows).sum();
let rows: Vec<Vec<Ex>> = matrices
.iter()
.flat_map(|m| m.rows.iter().cloned())
.collect();
Ok(Matrix::with_shape(rows, nrows, ncols))
}
pub fn vec(&self) -> Matrix {
let mut entries = Vec::with_capacity(self.nrows * self.ncols);
for j in 0..self.ncols {
for i in 0..self.nrows {
entries.push(self.rows[i][j].clone());
}
}
Matrix::col_vector(entries)
}
}
impl Matrix {
fn check_indices(
operation: &'static str,
axis: &str,
idx: &[usize],
bound: usize,
) -> Result<(), SymplexError> {
if idx.is_empty() {
return Err(invalid(
operation,
format!("{axis} selection must contain at least one index"),
));
}
if let Some(&bad) = idx.iter().find(|&&k| k >= bound) {
return Err(invalid(
operation,
format!("{axis} index {bad} out of range for {bound} {axis}s"),
));
}
Ok(())
}
pub fn extract(&self, rows: &[usize], cols: &[usize]) -> Result<Matrix, SymplexError> {
Self::check_indices("extract", "row", rows, self.nrows)?;
Self::check_indices("extract", "column", cols, self.ncols)?;
let data: Vec<Vec<Ex>> = rows
.iter()
.map(|&i| cols.iter().map(|&j| self.rows[i][j].clone()).collect())
.collect();
Ok(Matrix::from_rows_unchecked(data))
}
pub fn select_rows(&self, rows: &[usize]) -> Result<Matrix, SymplexError> {
Self::check_indices("select_rows", "row", rows, self.nrows)?;
let data: Vec<Vec<Ex>> = rows.iter().map(|&i| self.rows[i].clone()).collect();
Ok(Matrix::from_rows_unchecked(data))
}
pub fn select_cols(&self, cols: &[usize]) -> Result<Matrix, SymplexError> {
Self::check_indices("select_cols", "column", cols, self.ncols)?;
let data: Vec<Vec<Ex>> = self
.rows
.iter()
.map(|r| cols.iter().map(|&j| r[j].clone()).collect())
.collect();
Ok(Matrix::from_rows_unchecked(data))
}
pub fn delete_row(&self, i: usize) -> Result<Matrix, SymplexError> {
if i >= self.nrows {
return Err(invalid(
"delete_row",
format!("row index {i} out of range for {} rows", self.nrows),
));
}
if self.nrows == 1 {
return Err(invalid(
"delete_row",
"cannot delete the only row of a matrix",
));
}
let data: Vec<Vec<Ex>> = self
.rows
.iter()
.enumerate()
.filter(|&(k, _)| k != i)
.map(|(_, r)| r.clone())
.collect();
Ok(Matrix::from_rows_unchecked(data))
}
pub fn delete_col(&self, j: usize) -> Result<Matrix, SymplexError> {
if j >= self.ncols {
return Err(invalid(
"delete_col",
format!("column index {j} out of range for {} columns", self.ncols),
));
}
if self.ncols == 1 {
return Err(invalid(
"delete_col",
"cannot delete the only column of a matrix",
));
}
let data: Vec<Vec<Ex>> = self
.rows
.iter()
.map(|r| {
r.iter()
.enumerate()
.filter(|&(k, _)| k != j)
.map(|(_, e)| e.clone())
.collect()
})
.collect();
Ok(Matrix::from_rows_unchecked(data))
}
pub fn row_del(&self, i: usize) -> Result<Matrix, SymplexError> {
self.delete_row(i).map_err(|e| reop(e, "row_del"))
}
pub fn col_del(&self, j: usize) -> Result<Matrix, SymplexError> {
self.delete_col(j).map_err(|e| reop(e, "col_del"))
}
pub fn row_insert(&self, pos: usize, rows: &Matrix) -> Result<Matrix, SymplexError> {
if pos > self.nrows {
return Err(invalid(
"row_insert",
format!(
"position {pos} out of range for {} rows (use {} to append)",
self.nrows, self.nrows
),
));
}
if rows.ncols != self.ncols {
return Err(invalid(
"row_insert",
format!(
"inserted rows have {} columns, expected {}",
rows.ncols, self.ncols
),
));
}
let mut data: Vec<Vec<Ex>> = Vec::with_capacity(self.nrows + rows.nrows);
data.extend(self.rows[..pos].iter().cloned());
data.extend(rows.rows.iter().cloned());
data.extend(self.rows[pos..].iter().cloned());
Ok(Matrix::with_shape(
data,
self.nrows + rows.nrows,
self.ncols,
))
}
pub fn col_insert(&self, pos: usize, cols: &Matrix) -> Result<Matrix, SymplexError> {
if pos > self.ncols {
return Err(invalid(
"col_insert",
format!(
"position {pos} out of range for {} columns (use {} to append)",
self.ncols, self.ncols
),
));
}
if cols.nrows != self.nrows {
return Err(invalid(
"col_insert",
format!(
"inserted columns have {} rows, expected {}",
cols.nrows, self.nrows
),
));
}
let ncols = self.ncols + cols.ncols;
let data: Vec<Vec<Ex>> = self
.rows
.iter()
.zip(&cols.rows)
.map(|(r, c)| {
let mut row = Vec::with_capacity(ncols);
row.extend(r[..pos].iter().cloned());
row.extend(c.iter().cloned());
row.extend(r[pos..].iter().cloned());
row
})
.collect();
Ok(Matrix::with_shape(data, self.nrows, ncols))
}
fn check_permutation(
operation: &'static str,
axis: &str,
perm: &[usize],
bound: usize,
) -> Result<(), SymplexError> {
if perm.len() != bound {
return Err(invalid(
operation,
format!(
"permutation has {} entries, expected one per {axis} ({bound})",
perm.len()
),
));
}
let mut seen = vec![false; bound];
for &k in perm {
if k >= bound {
return Err(invalid(
operation,
format!("{axis} index {k} out of range for {bound} {axis}s"),
));
}
if seen[k] {
return Err(invalid(
operation,
format!("{axis} index {k} appears twice; not a permutation"),
));
}
seen[k] = true;
}
Ok(())
}
pub fn permute_rows(&self, perm: &[usize]) -> Result<Matrix, SymplexError> {
Self::check_permutation("permute_rows", "row", perm, self.nrows)?;
self.select_rows(perm).map_err(|e| reop(e, "permute_rows"))
}
pub fn permute_cols(&self, perm: &[usize]) -> Result<Matrix, SymplexError> {
Self::check_permutation("permute_cols", "column", perm, self.ncols)?;
self.select_cols(perm).map_err(|e| reop(e, "permute_cols"))
}
pub fn is_integer_matrix(&self) -> Option<bool> {
let mut unknown = false;
for e in self.iter() {
match e.as_rational() {
Some(r) if r.is_integer() => {}
Some(_) => return Some(false),
None => unknown = true,
}
}
if unknown { None } else { Some(true) }
}
pub fn to_rational_rows(&self) -> Option<Vec<Vec<Q>>> {
self.rows
.iter()
.map(|r| r.iter().map(Ex::as_rational).collect())
.collect()
}
pub fn to_bigint_rows(&self) -> Option<Vec<Vec<BigInt>>> {
self.rows
.iter()
.map(|r| r.iter().map(Ex::as_bigint).collect())
.collect()
}
pub fn from_ratio(ctx: &Context, rows: &[Vec<Q>]) -> Result<Matrix, SymplexError> {
let data: Vec<Vec<Ex>> = rows
.iter()
.map(|r| r.iter().map(|q| ctx.from_ratio(q.clone())).collect())
.collect();
Matrix::new(data)
}
pub fn from_bigint(ctx: &Context, rows: &[Vec<BigInt>]) -> Result<Matrix, SymplexError> {
let data: Vec<Vec<Ex>> = rows
.iter()
.map(|r| r.iter().map(|n| ctx.from_bigint(n.clone())).collect())
.collect();
Matrix::new(data)
}
pub fn from_f64_rows(ctx: &Context, rows: &[Vec<f64>]) -> Result<Matrix, SymplexError> {
let mut data: Vec<Vec<Ex>> = Vec::with_capacity(rows.len());
for r in rows {
let mut out = Vec::with_capacity(r.len());
for &v in r {
out.push(ctx.from_f64(v).map_err(|e| reop(e, "from_f64_rows"))?);
}
data.push(out);
}
Matrix::new(data)
}
pub fn subs_map(&self, replacements: &[(&Ex, &Ex)]) -> Matrix {
self.map(|elem| elem.subs_map(replacements))
}
pub fn nnz(&self) -> usize {
self.iter().filter(|e| !e.is_zero_structural()).count()
}
}
impl Matrix {
pub fn hermite_normal_form(&self) -> Result<Matrix, SymplexError> {
crate::domains::normalforms::hermite_normal_form(self)
}
pub fn smith_normal_form(&self) -> Result<Matrix, SymplexError> {
crate::domains::normalforms::smith_normal_form(self)
}
pub fn integer_nullspace(&self) -> Result<Vec<Matrix>, SymplexError> {
crate::domains::normalforms::integer_nullspace(self)
}
pub fn inv_mod(&self, m: u64) -> Result<Matrix, SymplexError> {
let z = ZMatrix::try_from(self).map_err(|e| reop(e, "inv_mod"))?;
let r = z
.inv_mod(&BigInt::from(m))
.map_err(|e| reop(e, "inv_mod"))?;
Ok(r.to_matrix(&self.ctx()))
}
pub fn lll(&self, delta: Rational64) -> Result<Matrix, SymplexError> {
crate::domains::normalforms::lll(self, delta)
}
pub fn lll_default(&self) -> Result<Matrix, SymplexError> {
crate::domains::normalforms::lll(self, LLL_DEFAULT_DELTA)
}
}
fn eigvals_with_multiplicity(char_poly: &Ex, var: &Ex) -> Vec<(Ex, usize)> {
if let Some(pairs) = eigvals_via_poly_factor(char_poly, var)
&& !pairs.is_empty()
{
debug!(
count = pairs.len(),
"eigvals_with_multiplicity: used Poly::factor_over_z path"
);
return pairs;
}
debug!("eigvals_with_multiplicity: falling back to derivative-based detection");
eigvals_via_derivative(char_poly, var)
}
fn eigvals_via_poly_factor(char_poly: &Ex, var: &Ex) -> Option<Vec<(Ex, usize)>> {
let inner = char_poly.inner.read();
let arena = &inner.arena;
let poly = crate::poly::polybridge::expr_to_poly(arena, char_poly.raw_id(), var.raw_id())?;
let (_content, factors) = poly.factor_over_z();
if factors.is_empty() {
return None;
}
drop(inner);
let mut eigen_pairs: Vec<(Ex, usize)> = Vec::new();
for (factor, mult) in &factors {
let degree = factor.degree().unwrap_or(0);
if degree == 0 {
continue;
}
if degree == 1 {
let coeffs = factor.coeffs();
let a = &coeffs[1];
let b = &coeffs[0];
let root_val = -(b / a);
let root_numer: Result<i64, _> = root_val.numer().clone().try_into();
let root_denom: Result<i64, _> = root_val.denom().clone().try_into();
if let (Ok(n), Ok(d)) = (root_numer, root_denom) {
eigen_pairs.push((char_poly.context().rational(n, d), *mult as usize));
continue;
}
}
let mut write_inner = char_poly.inner.write();
let factor_expr =
crate::poly::polybridge::poly_to_expr(&mut write_inner.arena, factor, var.raw_id());
drop(write_inner);
let factor_ex = char_poly.wrap(factor_expr);
let roots = if (3..=4).contains(°ree) && !radical_form_is_compact(factor) {
debug!(
degree,
"eigvals_via_poly_factor: irreducible factor without compact radicals → RootOf"
);
(0..degree).map(|i| root_of(&factor_ex, i)).collect()
} else {
factor_ex.solve_or_empty(var)
};
if roots.is_empty() {
warn!(
"eigenvals: irreducible factor of degree {} yielded no roots — \
eigenvalues from this factor are missing",
degree
);
}
trace!(
degree,
root_count = roots.len(),
mult,
"eigvals_via_poly_factor: degree-{} factor, {} roots, mult {}",
degree,
roots.len(),
mult
);
for r in roots {
eigen_pairs.push((r, *mult as usize));
}
}
Some(eigen_pairs)
}
fn root_of(poly: &Ex, index: usize) -> Ex {
let mut inner = poly.inner.write();
let idx = inner.arena.int(index as i64);
let id = inner
.arena
.intern(crate::base::node::ExprNode::RootOf(poly.raw_id(), idx));
drop(inner);
poly.wrap(id)
}
fn radical_form_is_compact(f: &crate::poly::dense::Poly) -> bool {
use num_bigint::BigInt;
let Some(n) = f.degree() else { return true };
let c = |i: usize| f.coeff(i);
let k = |v: i64| num_rational::Ratio::from_integer(BigInt::from(v));
match n {
3 => &(&c(3) * &c(1)) * &k(3) == &c(2) * &c(2),
4 => {
let (a, b, cc, d) = (c(4), c(3), c(2), c(1));
let b3 = &(&b * &b) * &b;
let abc = &(&(&a * &b) * &cc) * &k(4);
let aad = &(&(&a * &a) * &d) * &k(8);
&(&b3 - &abc) + &aad == k(0)
}
_ => true,
}
}
fn eigvals_via_derivative(char_poly: &Ex, var: &Ex) -> Vec<(Ex, usize)> {
let all_roots = char_poly.solve_or_empty(var);
if all_roots.is_empty() {
return Vec::new();
}
let mut unique_roots: Vec<Ex> = Vec::new();
for root in &all_roots {
if !unique_roots.iter().any(|r| r == root) {
unique_roots.push(root.clone());
}
}
let mut eigen_pairs: Vec<(Ex, usize)> = Vec::new();
for root in &unique_roots {
let mut mult = 0usize;
let mut current = char_poly.clone();
for k in 0..20 {
let val = current.subs(var, root).eval().simplify();
if !val.is_zero_structural() {
mult = k;
break;
}
if k < 19 {
current = current.diff(var);
}
}
let mult = if mult == 0 { 1 } else { mult };
eigen_pairs.push((root.clone(), mult));
}
eigen_pairs
}
fn eigen_zero_test(e: &Ex) -> bool {
if e.is_zero_structural() {
return true;
}
if e.expr_type() == ExprType::Number {
return false;
}
match ex_is_zero(e) {
Some(b) => b,
None => {
matches!(e.eval_complex64(), Ok(z) if z.re.abs() < 1e-10 && z.im.abs() < 1e-10)
}
}
}
fn jordan_null_power(a_minus_lambda: &Matrix, power: usize) -> Result<Vec<Matrix>, SymplexError> {
if power == 0 {
return Ok(Vec::new());
}
let mut m = a_minus_lambda.clone();
for _ in 1..power {
m = m.matmul(a_minus_lambda)?;
}
Ok(m.nullspace_semantic())
}
fn matrix_pow_vec(
a_minus_lambda: &Matrix,
vec: &Matrix,
power: usize,
) -> Result<Matrix, SymplexError> {
let mut result = vec.clone();
for _ in 0..power {
result = a_minus_lambda.matmul(&result)?;
}
Ok(result)
}
fn pick_independent_vec(
candidates: &[Matrix],
exclude: &[&Matrix],
) -> Result<Option<Matrix>, SymplexError> {
if candidates.is_empty() {
return Ok(None);
}
if exclude.is_empty() {
return Ok(Some(candidates[0].clone()));
}
let base_rank = Matrix::hstack(exclude)?.rank_semantic();
for candidate in candidates {
let mut cols: Vec<&Matrix> = exclude.to_vec();
cols.push(candidate);
let combined = Matrix::hstack(&cols)?;
if combined.rank_semantic() > base_rank {
return Ok(Some(candidate.clone()));
}
}
Ok(None)
}
pub fn cross(a: &Matrix, b: &Matrix) -> Matrix {
assert!(
a.nrows() == 3 && a.ncols() == 1,
"cross: first argument must be a 3×1 column vector, got {}×{}",
a.nrows(),
a.ncols()
);
assert!(
b.nrows() == 3 && b.ncols() == 1,
"cross: second argument must be a 3×1 column vector, got {}×{}",
b.nrows(),
b.ncols()
);
let (a0, a1, a2) = (a.get(0, 0), a.get(1, 0), a.get(2, 0));
let (b0, b1, b2) = (b.get(0, 0), b.get(1, 0), b.get(2, 0));
Matrix::col_vector(vec![
a1 * b2 - a2 * b1,
a2 * b0 - a0 * b2,
a0 * b1 - a1 * b0,
])
}
pub fn dot(a: &Matrix, b: &Matrix) -> Ex {
assert_eq!(a.ncols(), 1, "dot: first argument must be a column vector");
assert_eq!(b.ncols(), 1, "dot: second argument must be a column vector");
assert_eq!(
a.nrows(),
b.nrows(),
"dot: vectors must have the same length ({} vs {})",
a.nrows(),
b.nrows()
);
let mut sum = a.get(0, 0) * b.get(0, 0);
for i in 1..a.nrows() {
sum += &(a.get(i, 0) * b.get(i, 0));
}
sum
}
impl fmt::Display for Matrix {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
if self.nrows == 1 {
write!(f, "[[")?;
for (j, elem) in self.rows[0].iter().enumerate() {
if j > 0 {
write!(f, ", ")?;
}
write!(f, "{elem}")?;
}
return write!(f, "]]");
}
let cells: Vec<Vec<String>> = self
.rows
.iter()
.map(|r| r.iter().map(ToString::to_string).collect())
.collect();
let widths: Vec<usize> = (0..self.ncols)
.map(|j| {
cells
.iter()
.map(|r| r[j].chars().count())
.max()
.unwrap_or(0)
})
.collect();
writeln!(f, "[")?;
for (i, row) in cells.iter().enumerate() {
write!(f, " [")?;
for (j, cell) in row.iter().enumerate() {
if j > 0 {
write!(f, ", ")?;
}
let pad = widths[j] - cell.chars().count();
write!(f, "{}{}", " ".repeat(pad), cell)?;
}
write!(f, "]")?;
if i + 1 < self.nrows {
writeln!(f, ",")?;
} else {
writeln!(f)?;
}
}
write!(f, "]")
}
}
impl fmt::Debug for Matrix {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "Matrix({}×{}, [", self.nrows, self.ncols)?;
for (i, row) in self.rows.iter().enumerate() {
if i > 0 {
write!(f, ", ")?;
}
write!(f, "[")?;
for (j, elem) in row.iter().enumerate() {
if j > 0 {
write!(f, ", ")?;
}
write!(f, "{elem}")?;
}
write!(f, "]")?;
}
write!(f, "])")
}
}
impl std::ops::Index<(usize, usize)> for Matrix {
type Output = Ex;
#[inline]
fn index(&self, (i, j): (usize, usize)) -> &Ex {
self.get(i, j)
}
}
impl std::ops::IndexMut<(usize, usize)> for Matrix {
#[inline]
fn index_mut(&mut self, (i, j): (usize, usize)) -> &mut Ex {
self.get_mut(i, j)
}
}
macro_rules! matrix_binop {
($trait:ident, $method:ident, $inner:ident, $msg:literal) => {
impl std::ops::$trait<&Matrix> for &Matrix {
type Output = Matrix;
fn $method(self, rhs: &Matrix) -> Matrix {
match self.$inner(rhs) {
Ok(m) => m,
Err(e) => panic!(concat!($msg, ": {}"), e),
}
}
}
impl std::ops::$trait<Matrix> for Matrix {
type Output = Matrix;
fn $method(self, rhs: Matrix) -> Matrix {
std::ops::$trait::$method(&self, &rhs)
}
}
impl std::ops::$trait<&Matrix> for Matrix {
type Output = Matrix;
fn $method(self, rhs: &Matrix) -> Matrix {
std::ops::$trait::$method(&self, rhs)
}
}
impl std::ops::$trait<Matrix> for &Matrix {
type Output = Matrix;
fn $method(self, rhs: Matrix) -> Matrix {
std::ops::$trait::$method(self, &rhs)
}
}
};
}
matrix_binop!(Add, add, add, "Matrix + Matrix");
matrix_binop!(Sub, sub, sub, "Matrix - Matrix");
matrix_binop!(Mul, mul, matmul, "Matrix * Matrix");
impl std::ops::Neg for &Matrix {
type Output = Matrix;
fn neg(self) -> Matrix {
self.map(|e| -e)
}
}
impl std::ops::Neg for Matrix {
type Output = Matrix;
fn neg(self) -> Matrix {
-&self
}
}
impl std::ops::Mul<&Ex> for &Matrix {
type Output = Matrix;
fn mul(self, rhs: &Ex) -> Matrix {
self.scale(rhs)
}
}
impl std::ops::Mul<Ex> for &Matrix {
type Output = Matrix;
fn mul(self, rhs: Ex) -> Matrix {
self.scale(&rhs)
}
}
impl std::ops::Mul<&Ex> for Matrix {
type Output = Matrix;
fn mul(self, rhs: &Ex) -> Matrix {
self.scale(rhs)
}
}
impl std::ops::Mul<Ex> for Matrix {
type Output = Matrix;
fn mul(self, rhs: Ex) -> Matrix {
self.scale(&rhs)
}
}
impl std::ops::Mul<&Matrix> for &Ex {
type Output = Matrix;
fn mul(self, rhs: &Matrix) -> Matrix {
rhs.scale(self)
}
}
impl std::ops::Mul<Matrix> for &Ex {
type Output = Matrix;
fn mul(self, rhs: Matrix) -> Matrix {
rhs.scale(self)
}
}
impl std::ops::Mul<&Matrix> for Ex {
type Output = Matrix;
fn mul(self, rhs: &Matrix) -> Matrix {
rhs.scale(&self)
}
}
impl std::ops::Mul<Matrix> for Ex {
type Output = Matrix;
fn mul(self, rhs: Matrix) -> Matrix {
rhs.scale(&self)
}
}
impl std::ops::Div<&Ex> for &Matrix {
type Output = Matrix;
fn div(self, rhs: &Ex) -> Matrix {
self.map(|e| e / rhs)
}
}
impl std::ops::Div<Ex> for &Matrix {
type Output = Matrix;
fn div(self, rhs: Ex) -> Matrix {
self / &rhs
}
}
impl std::ops::Div<&Ex> for Matrix {
type Output = Matrix;
fn div(self, rhs: &Ex) -> Matrix {
&self / rhs
}
}
impl std::ops::Div<Ex> for Matrix {
type Output = Matrix;
fn div(self, rhs: Ex) -> Matrix {
&self / &rhs
}
}
impl std::ops::Mul<i64> for &Matrix {
type Output = Matrix;
fn mul(self, rhs: i64) -> Matrix {
let s = self.ctx().int(rhs);
self.scale(&s)
}
}
impl std::ops::Mul<i64> for Matrix {
type Output = Matrix;
fn mul(self, rhs: i64) -> Matrix {
&self * rhs
}
}
impl std::ops::Mul<&Matrix> for i64 {
type Output = Matrix;
fn mul(self, rhs: &Matrix) -> Matrix {
rhs * self
}
}
impl std::ops::Mul<Matrix> for i64 {
type Output = Matrix;
fn mul(self, rhs: Matrix) -> Matrix {
&rhs * self
}
}
impl std::ops::Div<i64> for &Matrix {
type Output = Matrix;
fn div(self, rhs: i64) -> Matrix {
let s = self.ctx().int(rhs);
self / &s
}
}
impl std::ops::Div<i64> for Matrix {
type Output = Matrix;
fn div(self, rhs: i64) -> Matrix {
&self / rhs
}
}
#[cfg(test)]
mod tests {
use super::*;
fn tctx() -> Context {
Context::new()
}
fn assert_zero_matrix(m: &Matrix, label: &str) {
for i in 0..m.nrows() {
for j in 0..m.ncols() {
let d = m.get(i, j).expand().eval().simplify();
assert!(d.is_zero_structural(), "{label}: entry ({i},{j}) = {d} ≠ 0");
}
}
}
#[test]
fn identity_2x2() {
let ctx = tctx();
let m = Matrix::identity(&ctx, 2);
assert_eq!(m.nrows(), 2);
assert_eq!(m.ncols(), 2);
assert_eq!(format!("{}", m.get(0, 0)), "1");
assert_eq!(format!("{}", m.get(0, 1)), "0");
assert_eq!(format!("{}", m.get(1, 0)), "0");
assert_eq!(format!("{}", m.get(1, 1)), "1");
}
#[test]
fn identity_3x3() {
let ctx = tctx();
let m = Matrix::identity(&ctx, 3);
assert_eq!(m.shape(), (3, 3));
for i in 0..3 {
for j in 0..3 {
let expected = if i == j { "1" } else { "0" };
assert_eq!(format!("{}", m.get(i, j)), expected);
}
}
}
#[test]
fn zeros_and_from_fn() {
let ctx = tctx();
let z = Matrix::zeros(&ctx, 3, 3);
assert_eq!(format!("{}", z.get(1, 1)), "0");
let m = Matrix::from_fn(2, 2, |i, j| ctx.int((i * 2 + j + 1) as i64));
assert_eq!(format!("{}", m.get(0, 0)), "1");
assert_eq!(format!("{}", m.get(0, 1)), "2");
assert_eq!(format!("{}", m.get(1, 0)), "3");
assert_eq!(format!("{}", m.get(1, 1)), "4");
}
#[test]
fn row_and_col_vectors() {
let ctx = tctx();
let rv = Matrix::row_vector(vec![ctx.int(1), ctx.int(2), ctx.int(3)]);
assert_eq!(rv.shape(), (1, 3));
assert_eq!(format!("{}", rv.get(0, 1)), "2");
let cv = Matrix::col_vector(vec![ctx.int(10), ctx.int(20)]);
assert_eq!(cv.shape(), (2, 1));
assert_eq!(format!("{}", cv.get(1, 0)), "20");
}
#[test]
fn from_i64_and_try_from() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 2], &[3, 4]]).unwrap();
assert_eq!(m.shape(), (2, 2));
assert_eq!(m.get(1, 0), &ctx.int(3));
assert!(Matrix::from_i64(&ctx, &[&[1, 2], &[3]]).is_err());
let t: Matrix = vec![vec![ctx.int(7)]].try_into().unwrap();
assert_eq!(t.shape(), (1, 1));
let bad: Result<Matrix, _> = Vec::<Vec<Ex>>::new().try_into();
assert!(bad.is_err());
}
#[test]
fn block_diag_shapes() {
let ctx = tctx();
let a = Matrix::identity(&ctx, 2);
let b = Matrix::row_vector(vec![ctx.int(5), ctx.int(6), ctx.int(7)]);
let bd = Matrix::block_diag(&[&a, &b]).unwrap();
assert_eq!(bd.shape(), (3, 5));
assert_eq!(bd.get(2, 4), &ctx.int(7));
assert!(bd.get(0, 2).is_zero_structural());
assert!(Matrix::block_diag(&[]).is_err());
}
#[test]
fn get_mut_and_set_work() {
let ctx = tctx();
let mut m = Matrix::zeros(&ctx, 2, 2);
*m.get_mut(0, 1) = ctx.int(42);
m.set(1, 0, ctx.int(7));
m[(1, 1)] = ctx.int(9);
assert_eq!(format!("{}", m.get(0, 1)), "42");
assert_eq!(m[(1, 0)], ctx.int(7));
assert_eq!(m[(1, 1)], ctx.int(9));
}
#[test]
fn row_col_diagonal_accessors() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 2, 3], &[4, 5, 6]]).unwrap();
assert_eq!(m.row(0).len(), 3);
assert_eq!(m.col(2), vec![ctx.int(3), ctx.int(6)]);
assert_eq!(m.diagonal(), vec![ctx.int(1), ctx.int(5)]);
assert_eq!(m.iter().count(), 6);
assert_eq!(m.to_vec()[1][2], ctx.int(6));
assert_eq!(m.eval_f64().unwrap()[1], vec![4.0, 5.0, 6.0]);
}
#[test]
fn submatrix_and_minor_matrix() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 2, 3], &[4, 5, 6], &[7, 8, 9]]).unwrap();
let s = m.submatrix(0..2, 1..3);
assert_eq!(s, Matrix::from_i64(&ctx, &[&[2, 3], &[5, 6]]).unwrap());
let mm = m.minor_matrix(1, 1).unwrap();
assert_eq!(mm, Matrix::from_i64(&ctx, &[&[1, 3], &[7, 9]]).unwrap());
assert_eq!(m.minor(1, 1).unwrap(), ctx.int(-12));
assert!(m.minor_matrix(3, 0).is_err());
}
#[test]
fn try_get_works() {
let ctx = tctx();
let m = Matrix::zeros(&ctx, 2, 2);
assert!(m.try_get(0, 0).is_some());
assert!(m.try_get(1, 1).is_some());
assert!(m.try_get(2, 0).is_none());
assert!(m.try_get(0, 2).is_none());
}
#[test]
fn structural_equality() {
let ctx = tctx();
let a = Matrix::from_i64(&ctx, &[&[1, 2], &[3, 4]]).unwrap();
let b = Matrix::from_i64(&ctx, &[&[1, 2], &[3, 4]]).unwrap();
let c = Matrix::from_i64(&ctx, &[&[1, 2], &[3, 5]]).unwrap();
let d = Matrix::from_i64(&ctx, &[&[1, 2, 3, 4]]).unwrap();
assert_eq!(a, b);
assert_ne!(a, c);
assert_ne!(a, d);
assert_eq!(a.equals(&b), Some(true));
assert_eq!(a.equals(&c), Some(false));
assert_eq!(a.equals(&d), Some(false));
}
#[test]
fn transpose() {
let ctx = tctx();
let a = ctx.symbol("a");
let b = ctx.symbol("b");
let c = ctx.symbol("c");
let d = ctx.symbol("d");
let m = Matrix::new(vec![vec![a.clone(), b.clone()], vec![c.clone(), d.clone()]]).unwrap();
let t = m.transpose();
assert_eq!(t.get(0, 0), &a);
assert_eq!(t.get(0, 1), &c);
assert_eq!(t.get(1, 0), &b);
assert_eq!(t.get(1, 1), &d);
}
#[test]
fn transpose_non_square() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 2, 3], &[4, 5, 6]]).unwrap();
let t = m.transpose();
assert_eq!(t.shape(), (3, 2));
assert_eq!(format!("{}", t.get(2, 0)), "3");
assert_eq!(format!("{}", t.get(2, 1)), "6");
}
#[test]
fn adjoint_conjugates() {
let ctx = tctx();
let i = ctx.i_unit();
let m = Matrix::new(vec![
vec![ctx.int(1), &ctx.int(2) * &i],
vec![ctx.int(3), ctx.int(4)],
])
.unwrap();
let h = m.adjoint();
assert_eq!(h.get(1, 0), &(&ctx.int(-2) * &i));
assert_eq!(h.get(0, 1), &ctx.int(3));
}
#[test]
fn matrix_add() {
let ctx = tctx();
let m1 = Matrix::from_i64(&ctx, &[&[1, 2], &[3, 4]]).unwrap();
let m2 = Matrix::from_i64(&ctx, &[&[5, 6], &[7, 8]]).unwrap();
let sum = m1.add(&m2).unwrap();
assert_eq!(format!("{}", sum.get(0, 0)), "6");
assert_eq!(format!("{}", sum.get(1, 1)), "12");
assert_eq!(&m1 + &m2, sum);
}
#[test]
fn matrix_sub() {
let ctx = tctx();
let m1 = Matrix::new(vec![vec![ctx.int(10), ctx.int(20)]]).unwrap();
let m2 = Matrix::new(vec![vec![ctx.int(3), ctx.int(7)]]).unwrap();
let diff = m1.sub(&m2).unwrap();
assert_eq!(format!("{}", diff.get(0, 0)), "7");
assert_eq!(format!("{}", diff.get(0, 1)), "13");
}
#[test]
fn hadamard_product() {
let ctx = tctx();
let m1 = Matrix::from_i64(&ctx, &[&[1, 2], &[3, 4]]).unwrap();
let m2 = Matrix::from_i64(&ctx, &[&[5, 6], &[7, 8]]).unwrap();
let h = m1.hadamard(&m2).unwrap();
assert_eq!(h, Matrix::from_i64(&ctx, &[&[5, 12], &[21, 32]]).unwrap());
}
#[test]
fn scale() {
let ctx = tctx();
let m = Matrix::identity(&ctx, 2);
let two = ctx.int(2);
let scaled = m.scale(&two);
assert_eq!(format!("{}", scaled.get(0, 0)), "2");
assert_eq!(format!("{}", scaled.get(0, 1)), "0");
}
#[test]
fn scalar_operators_both_sides() {
let ctx = tctx();
let x = ctx.symbol("x");
let m = Matrix::from_i64(&ctx, &[&[1, 2]]).unwrap();
let a = &m * &x;
let b = &x * &m;
let c = m.clone() * x.clone();
let d = x.clone() * m.clone();
assert_eq!(a, b);
assert_eq!(a, c);
assert_eq!(a, d);
assert_eq!(a.get(0, 1), &(&ctx.int(2) * &x));
let e = &m * 3;
let f = 3 * &m;
assert_eq!(e, f);
assert_eq!(e.get(0, 1), &ctx.int(6));
let g = &m / 2;
assert_eq!(g.get(0, 0), &ctx.rational(1, 2));
let h = &m / &ctx.int(2);
assert_eq!(g, h);
let n = -&m;
assert_eq!(n.get(0, 1), &ctx.int(-2));
}
#[test]
fn matmul_2x2() {
let ctx = tctx();
let m = Matrix::identity(&ctx, 2);
let a = ctx.symbol("a");
let b = ctx.symbol("b");
let c = ctx.symbol("c");
let d = ctx.symbol("d");
let n = Matrix::new(vec![vec![a.clone(), b.clone()], vec![c.clone(), d.clone()]]).unwrap();
let result = m.matmul(&n).unwrap();
assert_eq!(result.get(0, 0), &a);
assert_eq!(result.get(1, 1), &d);
}
#[test]
fn matmul_non_square() {
let ctx = tctx();
let rv = Matrix::row_vector(vec![ctx.int(2), ctx.int(3)]);
let cv = Matrix::col_vector(vec![ctx.int(4), ctx.int(5)]);
let result = rv.matmul(&cv).unwrap();
assert_eq!(result.shape(), (1, 1));
assert_eq!(format!("{}", result.get(0, 0)), "23");
}
#[test]
fn matmul_numeric() {
let ctx = tctx();
let a = Matrix::from_i64(&ctx, &[&[1, 2], &[3, 4]]).unwrap();
let b = Matrix::from_i64(&ctx, &[&[5, 6], &[7, 8]]).unwrap();
let c = a.matmul(&b).unwrap();
assert_eq!(c, Matrix::from_i64(&ctx, &[&[19, 22], &[43, 50]]).unwrap());
}
#[test]
fn trace_2x2() {
let ctx = tctx();
let a = ctx.symbol("a");
let d = ctx.symbol("d");
let m = Matrix::new(vec![
vec![a.clone(), ctx.symbol("b")],
vec![ctx.symbol("c"), d.clone()],
])
.unwrap();
assert_eq!(m.trace().unwrap(), &a + &d);
}
#[test]
fn trace_numeric() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 2], &[3, 4]]).unwrap();
assert_eq!(format!("{}", m.trace().unwrap()), "5");
}
#[test]
fn det_1x1() {
let ctx = tctx();
let a = ctx.symbol("a");
let m = Matrix::new(vec![vec![a.clone()]]).unwrap();
assert_eq!(m.det().unwrap(), a);
}
#[test]
fn det_2x2() {
let ctx = tctx();
let a = ctx.symbol("a");
let b = ctx.symbol("b");
let c = ctx.symbol("c");
let d = ctx.symbol("d");
let m = Matrix::new(vec![vec![a.clone(), b.clone()], vec![c.clone(), d.clone()]]).unwrap();
assert_eq!(m.det().unwrap(), &(&a * &d) - &(&b * &c));
}
#[test]
fn det_2x2_numeric() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[3, 8], &[4, 6]]).unwrap();
assert_eq!(format!("{}", m.det().unwrap()), "-14");
}
#[test]
fn det_3x3_numeric() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 2, 3], &[4, 5, 6], &[7, 8, 9]]).unwrap();
assert_eq!(format!("{}", m.det().unwrap()), "0");
}
#[test]
fn det_3x3_nonsingular() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 0, 2], &[0, 1, 0], &[3, 0, 1]]).unwrap();
assert_eq!(format!("{}", m.det().unwrap()), "-5");
}
#[test]
fn det_4x4_symbolic_is_polynomial() {
let ctx = tctx();
let syms: Vec<Vec<Ex>> = (0..4)
.map(|i| (0..4).map(|j| ctx.symbol(&format!("a{i}{j}"))).collect())
.collect();
let m = Matrix::new(syms).unwrap();
let d = m.det().unwrap();
assert_eq!(d.term_count(), 24, "4×4 symbolic det has 24 terms: {d}");
let mut expected = ctx.zero();
for j in 0..4 {
expected += m.get(0, j) * &m.cofactor(0, j).unwrap();
}
assert!((&d - &expected).expand().is_zero_structural());
}
#[test]
fn det_4x4_mixed_symbolic_matches_bareiss_numeric_specialization() {
let ctx = tctx();
let x = ctx.symbol("x");
let m = Matrix::new(vec![
vec![x.clone(), ctx.int(2), ctx.int(0), ctx.int(1)],
vec![ctx.int(1), x.clone(), ctx.int(3), ctx.int(0)],
vec![ctx.int(0), ctx.int(1), x.clone(), ctx.int(2)],
vec![ctx.int(2), ctx.int(0), ctx.int(1), x.clone()],
])
.unwrap();
let d = m.det().unwrap();
assert!(d.is_polynomial(&x), "symbolic det must be polynomial: {d}");
for v in [-2i64, 0, 1, 3, 7] {
let numeric = m.subs(&x, &ctx.int(v)).det().unwrap().eval();
let via_sym = d.subs(&x, &ctx.int(v)).eval();
assert_eq!(numeric, via_sym, "det mismatch at x={v}");
}
}
#[test]
fn map_doubles() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 2], &[3, 4]]).unwrap();
let mut calls = 0;
let doubled = m.map(|e| {
calls += 1;
e * 2
});
assert_eq!(calls, 4);
assert_eq!(format!("{}", doubled.get(0, 0)), "2");
assert_eq!(format!("{}", doubled.get(1, 1)), "8");
let idx = m.map_indexed(|i, j, e| e + ctx.int((10 * i + j) as i64));
assert_eq!(idx.get(1, 1), &ctx.int(15));
}
#[test]
fn jacobian_test() {
let ctx = tctx();
let x = ctx.symbol("x");
let y = ctx.symbol("y");
let f1 = &x.powi(2) + &y;
let f2 = &x * &y;
let j = jacobian(&[&f1, &f2], &[&x, &y]);
assert_eq!(j.shape(), (2, 2));
assert_eq!(j.get(0, 0), &(&x * 2));
assert_eq!(j.get(0, 1), &ctx.int(1));
assert_eq!(j.get(1, 0), &y);
assert_eq!(j.get(1, 1), &x);
}
#[test]
fn matrix_diff_and_integrate() {
let ctx = tctx();
let x = ctx.symbol("x");
let m = Matrix::new(vec![vec![x.powi(2), x.sin()]]).unwrap();
let dm = m.diff(&x);
assert_eq!(dm.get(0, 0), &(&x * 2));
assert_eq!(dm.get(0, 1), &x.cos());
let im = m.integrate(&x);
assert_eq!(im.get(0, 0), &(&x.powi(3) / 3));
}
#[test]
fn matrix_subs() {
let ctx = tctx();
let x = ctx.symbol("x");
let m = Matrix::new(vec![vec![x.powi(2), x.clone()]]).unwrap();
let result = m.subs(&x, &ctx.int(3));
assert_eq!(format!("{}", result.get(0, 0)), "9");
assert_eq!(format!("{}", result.get(0, 1)), "3");
}
#[test]
fn matrix_eval() {
let ctx = tctx();
let m = Matrix::new(vec![vec![ctx.pi().cos(), ctx.int(2) + ctx.int(3)]]).unwrap();
let evaled = m.eval();
assert_eq!(format!("{}", evaled.get(0, 0)), "-1");
assert_eq!(format!("{}", evaled.get(0, 1)), "5");
}
#[test]
fn matrix_expand() {
let ctx = tctx();
let x = ctx.symbol("x");
let expr = (&x + ctx.int(1)).powi(2);
let m = Matrix::new(vec![vec![expr]]).unwrap();
let expanded = m.expand();
assert_eq!(expanded.get(0, 0), &(&x.powi(2) + &x * 2 + 1));
}
#[test]
fn char_poly_coeffs_2x2_and_3x3() {
let ctx = tctx();
let a = Matrix::from_i64(&ctx, &[&[1, 2], &[3, 4]]).unwrap();
assert_eq!(
a.char_poly_coeffs().unwrap(),
vec![ctx.int(-2), ctx.int(-5), ctx.int(1)]
);
let b = Matrix::diag(&[ctx.int(2), ctx.int(3), ctx.int(4)]);
assert_eq!(
b.char_poly_coeffs().unwrap(),
vec![ctx.int(24), ctx.int(-26), ctx.int(9), ctx.int(-1)]
);
}
#[test]
fn char_poly_symbolic_2x2() {
let ctx = tctx();
let (a, b, c, d) = (
ctx.symbol("a"),
ctx.symbol("b"),
ctx.symbol("c"),
ctx.symbol("d"),
);
let m = Matrix::new(vec![vec![a.clone(), b.clone()], vec![c.clone(), d.clone()]]).unwrap();
let coeffs = m.char_poly_coeffs().unwrap();
assert_eq!(coeffs.len(), 3);
assert!(
(&coeffs[0] - &(&a * &d - &b * &c))
.expand()
.is_zero_structural()
);
assert!((&coeffs[1] + &(&a + &d)).expand().is_zero_structural());
assert!(coeffs[2].is_one_structural());
}
#[test]
fn char_poly_5x5_numeric_is_polynomial_and_fast() {
let ctx = tctx();
let lam = ctx.symbol("lambda");
let m = Matrix::from_fn(5, 5, |i, j| {
ctx.int(((i * 7 + j * 3) % 5) as i64 + if i == j { 2 } else { 0 })
});
let start = std::time::Instant::now();
let cp = m.char_poly(&lam).unwrap();
assert!(start.elapsed().as_secs() < 5, "char_poly too slow");
assert_eq!(
cp.degree(&lam),
Some(5),
"cp must be a degree-5 polynomial: {cp}"
);
let p0 = cp.subs(&lam, &ctx.int(0)).eval();
assert_eq!(p0, m.det().unwrap().eval());
let coeffs = m.char_poly_coeffs().unwrap();
let mut acc = Matrix::zeros(&ctx, 5, 5);
for (k, c) in coeffs.iter().enumerate() {
acc = &acc + &(&m.powi(k as u32).unwrap() * c);
}
assert_zero_matrix(&acc.eval(), "Cayley–Hamilton 5×5");
}
#[test]
fn char_poly_5x5_symbolic_terminates_quickly() {
let ctx = tctx();
let syms: Vec<Vec<Ex>> = (0..5)
.map(|i| (0..5).map(|j| ctx.symbol(&format!("a{i}{j}"))).collect())
.collect();
let m = Matrix::new(syms).unwrap();
let start = std::time::Instant::now();
let coeffs = m.char_poly_coeffs().unwrap();
assert!(
start.elapsed().as_secs() < 20,
"5×5 symbolic char poly took {:?}",
start.elapsed()
);
assert_eq!(coeffs.len(), 6);
assert_eq!(coeffs[5], ctx.int(-1));
assert!(
(&coeffs[4] - &m.trace().unwrap())
.expand()
.is_zero_structural()
);
assert_eq!(coeffs[0].term_count(), 120);
assert!(
(&coeffs[0] - &m.det().unwrap())
.expand()
.is_zero_structural()
);
}
#[test]
fn eigenvects_2x2_distinct() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[2, 1], &[0, 3]]).unwrap();
let evs = m.eigenvects().expect("eigenvects should succeed");
assert_eq!(evs.len(), 2);
for (eigenval, mult, vecs) in &evs {
assert_eq!(*mult, 1);
assert_eq!(vecs.len(), 1);
let av = m.matmul(&vecs[0]).unwrap();
let lambda_v = vecs[0].scale(eigenval);
assert_zero_matrix(&(&av - &lambda_v), "A·v = λ·v");
}
}
#[test]
fn eigenvects_non_square_returns_error() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 2, 3], &[4, 5, 6]]).unwrap();
let result = m.eigenvects();
assert!(result.is_err());
let err_msg = format!("{}", result.unwrap_err());
assert!(err_msg.contains("square"), "{err_msg}");
}
#[test]
fn eigenvects_diagonal_matrix() {
let ctx = tctx();
let m = Matrix::diag(&[ctx.int(5), ctx.int(-3)]);
let evs = m.eigenvects().unwrap();
assert_eq!(evs.len(), 2);
let vals: Vec<_> = evs.iter().map(|(v, _, _)| v.clone()).collect();
assert!(vals.contains(&ctx.int(5)) && vals.contains(&ctx.int(-3)));
}
#[test]
fn eigenvals_repeat_with_multiplicity() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 1], &[0, 1]]).unwrap();
assert_eq!(m.eigenvals().unwrap(), vec![ctx.int(1), ctx.int(1)]);
assert_eq!(
m.eigenvals_with_multiplicity().unwrap(),
vec![(ctx.int(1), 2)]
);
}
#[test]
fn eigenvals_symbolic_2x2_quadratic_formula() {
let ctx = tctx();
let (a, b) = (ctx.symbol("a"), ctx.symbol("b"));
let m = Matrix::new(vec![vec![a.clone(), b.clone()], vec![b.clone(), a.clone()]]).unwrap();
let evs = m.eigenvals().unwrap();
assert_eq!(evs.len(), 2);
let lam = ctx.symbol("lambda");
let cp = m.char_poly(&lam).unwrap();
for ev in &evs {
assert!(!ev.to_string().contains("__lambda"), "dummy leaked: {ev}");
for (av, bv) in [(3i64, 2i64), (-1, 5), (7, -4)] {
let residual = cp
.subs(&lam, ev)
.subs(&a, &ctx.int(av))
.subs(&b, &ctx.int(bv))
.eval_f64()
.unwrap();
assert!(
residual.abs() < 1e-9,
"eigenvalue {ev} does not satisfy char poly at a={av}, b={bv}: {residual}"
);
}
}
let sum = (&evs[0] + &evs[1]).expand();
assert!((&sum - &m.trace().unwrap()).expand().is_zero_structural());
}
#[test]
fn eigen_dummy_never_leaks_even_for_rootof() {
let ctx = tctx();
let m = Matrix::new(vec![
vec![ctx.int(0), ctx.int(0), ctx.int(0), ctx.int(0), ctx.int(1)],
vec![ctx.int(1), ctx.int(0), ctx.int(0), ctx.int(0), ctx.int(1)],
vec![ctx.int(0), ctx.int(1), ctx.int(0), ctx.int(0), ctx.int(0)],
vec![ctx.int(0), ctx.int(0), ctx.int(1), ctx.int(0), ctx.int(0)],
vec![ctx.int(0), ctx.int(0), ctx.int(0), ctx.int(1), ctx.int(0)],
])
.unwrap();
let evs = m.eigenvals().unwrap();
assert_eq!(evs.len(), 5);
for ev in &evs {
let s = ev.to_string();
assert!(!s.contains("__"), "reserved dummy leaked into result: {s}");
assert!(s.contains("RootOf"), "expected RootOf, got {s}");
}
}
#[test]
fn eigen_dummy_avoids_user_symbol_named_lambda() {
let ctx = tctx();
let user = ctx.symbol("__lambda");
let m = Matrix::new(vec![
vec![user.clone(), ctx.int(1)],
vec![ctx.int(0), ctx.int(2)],
])
.unwrap();
let dummy = m.fresh_symbol("lambda");
assert_ne!(dummy, user);
assert_eq!(dummy.to_string(), "__lambda_1");
let one_by_one = Matrix::new(vec![vec![user.clone()]]).unwrap();
assert_eq!(one_by_one.eigenvals().unwrap(), vec![user.clone()]);
let evs = m.eigenvals().unwrap();
assert_eq!(evs.len(), 2);
let mut nums: Vec<f64> = evs
.iter()
.map(|e| e.subs(&user, &ctx.int(5)).eval_f64().unwrap())
.collect();
nums.sort_by(|a, b| a.partial_cmp(b).unwrap());
assert!(
(nums[0] - 2.0).abs() < 1e-12 && (nums[1] - 5.0).abs() < 1e-12,
"{nums:?}"
);
}
#[test]
fn diagonalize_upper_triangular() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[2, 1], &[0, 3]]).unwrap();
let Diagonalization { p, d } = m.diagonalize().unwrap();
assert_eq!(p.nrows(), 2);
assert_eq!(d.nrows(), 2);
let p_inv = p.inv().unwrap();
let reconstructed = p.matmul(&d).unwrap().matmul(&p_inv).unwrap();
assert_zero_matrix(&(&reconstructed - &m), "P·D·P⁻¹ = M");
}
#[test]
fn diagonalize_non_diagonalizable() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 1], &[0, 1]]).unwrap();
let result = m.diagonalize();
assert!(result.is_err());
let err_msg = format!("{}", result.unwrap_err());
assert!(err_msg.contains("not diagonalizable"), "{err_msg}");
assert_eq!(m.is_diagonalizable(), Some(false));
}
#[test]
fn is_diagonalizable_yes() {
let ctx = tctx();
let m = Matrix::diag(&[ctx.int(1), ctx.int(2), ctx.int(3)]);
assert_eq!(m.is_diagonalizable(), Some(true));
}
#[test]
fn jordan_form_diagonal() {
let ctx = tctx();
let m = Matrix::diag(&[ctx.int(1), ctx.int(2), ctx.int(3)]);
let JordanForm { p, j } = m.jordan_form().unwrap();
assert_eq!(j.nrows(), 3);
let p_inv = p.inv().unwrap();
let reconstructed = p.matmul(&j).unwrap().matmul(&p_inv).unwrap();
assert_zero_matrix(&(&reconstructed - &m), "P·J·P⁻¹ = M");
}
#[test]
fn jordan_form_defective_2x2() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 1], &[0, 1]]).unwrap();
let JordanForm { p, j } = m.jordan_form().unwrap();
assert_eq!(j.nrows(), 2);
assert!(j.get(0, 1).is_one_structural(), "J[0,1] should be 1");
let reconstructed = p.matmul(&j).unwrap().matmul(&p.inv().unwrap()).unwrap();
assert_zero_matrix(&(&reconstructed - &m), "P·J·P⁻¹ = M");
}
#[test]
fn jordan_form_upper_triangular_distinct() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[2, 1], &[0, 3]]).unwrap();
let JordanForm { p, j } = m.jordan_form().unwrap();
assert_eq!(j.nrows(), 2);
let reconstructed = p.matmul(&j).unwrap().matmul(&p.inv().unwrap()).unwrap();
assert_zero_matrix(&(&reconstructed - &m), "P·J·P⁻¹ = M");
}
#[test]
fn matrix_exp_identity() {
let ctx = tctx();
let m = Matrix::identity(&ctx, 2);
let result = m.matrix_exp().unwrap();
assert_eq!(result.get(0, 0), &ctx.e());
assert!(result.get(0, 1).is_zero_structural());
}
#[test]
fn matrix_exp_zero() {
let ctx = tctx();
let m = Matrix::zeros(&ctx, 2, 2);
let result = m.matrix_exp().unwrap();
assert_eq!(result, Matrix::identity(&ctx, 2));
}
#[test]
fn matrix_exp_t_nilpotent_and_rotation() {
let ctx = tctx();
let t = ctx.symbol("t");
let n = Matrix::from_i64(&ctx, &[&[0, 1, 0], &[0, 0, 1], &[0, 0, 0]]).unwrap();
let e = n.matrix_exp_t(&t).unwrap();
assert_eq!(e.get(0, 1), &t);
assert_eq!(e.get(0, 2), &(&t.powi(2) / 2));
assert!(e.get(1, 0).is_zero_structural());
let rot = Matrix::from_i64(&ctx, &[&[0, 1], &[-1, 0]]).unwrap();
let et = rot.matrix_exp_t(&t).unwrap();
assert_eq!(et.get(0, 0), &t.cos());
assert_eq!(et.get(0, 1), &t.sin());
assert_eq!(et.get(1, 0), &(-&t.sin()));
assert_eq!(et.get(1, 1), &t.cos());
let at = et.subs(&t, &ctx.rational(7, 10));
let expect = [[0.7f64.cos(), 0.7f64.sin()], [-0.7f64.sin(), 0.7f64.cos()]];
for (i, row) in expect.iter().enumerate() {
for (j, want) in row.iter().enumerate() {
let z = at.get(i, j).eval_complex64().unwrap();
assert!((z.re - want).abs() < 1e-12, "({i},{j}) re={}", z.re);
assert!(z.im.abs() < 1e-12, "({i},{j}) im={}", z.im);
}
}
}
#[test]
fn matrix_exp_non_square_returns_error() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 2, 3], &[4, 5, 6]]).unwrap();
assert!(m.matrix_exp().is_err());
}
#[test]
fn jordan_form_non_square_returns_error() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 2, 3], &[4, 5, 6]]).unwrap();
assert!(m.jordan_form().is_err());
}
#[test]
fn diagonalize_non_square_returns_error() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 2, 3], &[4, 5, 6]]).unwrap();
assert!(m.diagonalize().is_err());
assert_eq!(m.is_diagonalizable(), Some(false));
}
#[test]
fn lu_returns_result() {
let ctx = tctx();
let a = Matrix::from_i64(&ctx, &[&[4, 3], &[6, 3]]).unwrap();
let Lu { l, u, perm } = a.lu().unwrap();
assert_eq!(perm.len(), 2);
let pa = Matrix::new(perm.iter().map(|&i| a.row(i).to_vec()).collect()).unwrap();
assert_zero_matrix(&(&(&l * &u) - &pa), "L·U = P·A");
assert!(l.get(0, 1).is_zero_structural());
assert!(u.get(1, 0).is_zero_structural());
assert!(
Matrix::from_i64(&ctx, &[&[1, 2], &[2, 4]])
.unwrap()
.lu()
.is_err()
);
assert!(Matrix::from_i64(&ctx, &[&[1, 2, 3]]).unwrap().lu().is_err());
}
#[test]
fn inverse_numeric_via_gauss_jordan() {
let ctx = tctx();
let a = Matrix::from_i64(&ctx, &[&[2, 1, 0], &[1, 3, 1], &[0, 1, 4]]).unwrap();
let inv = a.inv().unwrap();
assert_eq!((&a * &inv).eval(), Matrix::identity(&ctx, 3));
}
#[test]
fn rowspace_and_left_nullspace() {
let ctx = tctx();
let a = Matrix::from_i64(&ctx, &[&[1, 2], &[2, 4], &[3, 6]]).unwrap();
assert_eq!(a.rowspace().len(), 1);
let ln = a.left_nullspace();
assert_eq!(ln.len(), 2);
for y in &ln {
let yt_a = y.transpose().matmul(&a).unwrap();
assert_zero_matrix(&yt_a.eval(), "yᵀA = 0");
}
}
#[test]
fn solve_least_squares_matches_normal_equations() {
let ctx = tctx();
let a = Matrix::from_i64(&ctx, &[&[1, 0], &[1, 1], &[1, 2]]).unwrap();
let b = Matrix::from_i64(&ctx, &[&[1], &[2], &[4]]).unwrap();
let x = a.solve_least_squares(&b).unwrap();
assert_eq!(x.get(0, 0), &ctx.rational(5, 6));
assert_eq!(x.get(1, 0), &ctx.rational(3, 2));
let rank_deficient = Matrix::from_i64(&ctx, &[&[1, 2], &[2, 4]]).unwrap();
assert!(
rank_deficient
.solve_least_squares(&Matrix::from_i64(&ctx, &[&[1], &[1]]).unwrap())
.is_err()
);
}
#[test]
fn vec_column_stacking() {
let ctx = tctx();
let a = Matrix::from_i64(&ctx, &[&[1, 2], &[3, 4]]).unwrap();
assert_eq!(
a.vec(),
Matrix::from_i64(&ctx, &[&[1], &[3], &[2], &[4]]).unwrap()
);
}
#[test]
fn matrix_display_aligned() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, -2], &[30, 4]]).unwrap();
assert_eq!(format!("{m}"), "[\n [ 1, -2],\n [30, 4]\n]");
assert_eq!(format!("{m:?}"), "Matrix(2×2, [[1, -2], [30, 4]])");
}
#[test]
fn row_vector_display() {
let ctx = tctx();
let m = Matrix::row_vector(vec![ctx.int(1), ctx.int(2), ctx.int(3)]);
assert_eq!(format!("{m}"), "[[1, 2, 3]]");
}
#[test]
fn add_mismatched_shapes_returns_err() {
let ctx = tctx();
let a = Matrix::zeros(&ctx, 2, 3);
let b = Matrix::zeros(&ctx, 3, 2);
let result = a.add(&b);
assert!(result.is_err());
let err_msg = format!("{}", result.unwrap_err());
assert!(err_msg.contains("add"), "{err_msg}");
}
#[test]
fn matmul_incompatible_returns_err() {
let ctx = tctx();
let a = Matrix::zeros(&ctx, 2, 3);
let b = Matrix::zeros(&ctx, 2, 3);
let result = a.matmul(&b);
assert!(result.is_err());
let err_msg = format!("{}", result.unwrap_err());
assert!(err_msg.contains("multiply"), "{err_msg}");
}
#[test]
fn trace_non_square_returns_err() {
let ctx = tctx();
let m = Matrix::zeros(&ctx, 2, 3);
let result = m.trace();
assert!(result.is_err());
assert!(format!("{}", result.unwrap_err()).contains("square"));
}
#[test]
fn det_non_square_returns_err() {
let ctx = tctx();
let m = Matrix::zeros(&ctx, 2, 3);
let result = m.det();
assert!(result.is_err());
assert!(format!("{}", result.unwrap_err()).contains("square"));
}
#[test]
#[should_panic(expected = "out of bounds")]
fn get_out_of_bounds_panics() {
let ctx = Context::new();
let m = Matrix::zeros(&ctx, 2, 2);
let _ = m.get(2, 0);
}
#[test]
#[should_panic(expected = "Matrix + Matrix")]
fn operator_add_mismatch_panics() {
let ctx = Context::new();
let _ = &Matrix::zeros(&ctx, 2, 2) + &Matrix::zeros(&ctx, 3, 3);
}
#[test]
fn new_empty_returns_error() {
let result = Matrix::new(vec![]);
assert!(result.is_err());
let err_msg = format!("{}", result.unwrap_err());
assert!(err_msg.contains("at least one row"), "{err_msg}");
}
#[test]
fn new_jagged_returns_error() {
let ctx = tctx();
let result = Matrix::new(vec![vec![ctx.int(1), ctx.int(2)], vec![ctx.int(3)]]);
assert!(result.is_err());
let err_msg = format!("{}", result.unwrap_err());
assert!(
err_msg.contains("row 1 has length 1 but expected 2"),
"{err_msg}"
);
}
#[test]
fn new_valid_succeeds() {
let ctx = tctx();
let m = Matrix::from_i64(&ctx, &[&[1, 2], &[3, 4]]).unwrap();
assert_eq!(m.shape(), (2, 2));
assert_eq!(format!("{}", m.get(0, 0)), "1");
assert_eq!(format!("{}", m.get(1, 1)), "4");
}
#[test]
fn three_valued_helpers() {
let ctx = tctx();
let x = ctx.symbol("x");
assert_eq!(ex_is_zero(&ctx.int(0)), Some(true));
assert_eq!(ex_is_zero(&ctx.int(2)), Some(false));
assert_eq!(ex_is_zero(&(ctx.int(2).sqrt() - 1)), Some(false));
assert_eq!(
ex_is_zero(&(&x.sin().powi(2) + &x.cos().powi(2) - 1)),
Some(true)
);
assert_eq!(ex_is_zero(&x), None);
assert_eq!(ex_is_positive(&(ctx.int(2).sqrt() - 1)), Some(true));
assert_eq!(
ex_is_positive(&(ctx.int(1) - ctx.int(2).sqrt())),
Some(false)
);
assert_eq!(ex_is_positive(&x), None);
assert_eq!(ex_is_nonnegative(&ctx.int(0)), Some(true));
assert_eq!(all3([Some(true), None]), None);
assert_eq!(all3([Some(true), None, Some(false)]), Some(false));
assert_eq!(all3([Some(true), Some(true)]), Some(true));
}
}