use crate::array::Array;
use crate::error::{NumRs2Error, Result};
use crate::kernels::{borrow, cast};
use num_traits::Float;
use scirs2_core::ndarray::{Array1, Array2, ArrayView1, ArrayView2};
use scirs2_linalg::error::LinalgError;
use std::fmt::Debug;
pub(crate) const SCIRS2_MIN_DIM: usize = 16;
pub(crate) const QR_MAX_DIM: usize = 96;
pub(crate) const SVD_MAX_DIM: usize = 256;
pub(crate) const SOLVE_MIN_DIM: usize = 4;
fn map_linalg_err(e: LinalgError) -> NumRs2Error {
match e {
LinalgError::SingularMatrixError(s) => {
NumRs2Error::InvalidOperation(format!("Singular matrix: {}", s))
}
LinalgError::DimensionError(s) | LinalgError::ShapeError(s) => {
NumRs2Error::DimensionMismatch(s)
}
LinalgError::NonPositiveDefiniteError(s) => {
NumRs2Error::InvalidOperation(format!("Matrix is not positive definite: {}", s))
}
LinalgError::ConvergenceError(s) => {
NumRs2Error::ComputationError(format!("Convergence failed: {}", s))
}
other => NumRs2Error::ComputationError(format!("Linear algebra error: {}", other)),
}
}
fn square_dim<T: Clone>(a: &Array<T>, min_dim: usize) -> Option<usize> {
square_dim_in_range(a, min_dim, usize::MAX)
}
fn square_dim_in_range<T: Clone>(a: &Array<T>, min_dim: usize, max_dim: usize) -> Option<usize> {
let shape = a.shape();
if shape.len() != 2 || shape[0] != shape[1] || shape[0] < min_dim || shape[0] >= max_dim {
return None;
}
Some(shape[0])
}
fn into_row_major_vec<F: Clone>(m: Array2<F>) -> Vec<F> {
if m.is_standard_layout() {
let len = m.len();
let (buf, offset) = m.into_raw_vec_and_offset();
let start = offset.unwrap_or(0);
if start == 0 && buf.len() == len {
return buf;
}
return buf[start..start + len].to_vec();
}
m.iter().cloned().collect()
}
macro_rules! concrete_converters {
($f:ty, $matrix_fn:ident, $vector_fn:ident, $vec_cast:path) => {
fn $matrix_fn<T: Clone + 'static>(m: Array2<$f>) -> Option<Array<T>> {
let (rows, cols) = (m.nrows(), m.ncols());
let data: Vec<T> = $vec_cast(into_row_major_vec(m))?;
Some(Array::from_vec_shape(data, &[rows, cols]).unwrap_or_else(|e| panic!("{e}")))
}
fn $vector_fn<T: Clone + 'static>(v: Array1<$f>) -> Option<Array<T>> {
let data: Vec<T> = $vec_cast(v.to_vec())?;
Some(Array::from_vec(data))
}
};
}
concrete_converters!(f64, matrix_from_f64, vector_from_f64, cast::vec_from_f64);
concrete_converters!(f32, matrix_from_f32, vector_from_f32, cast::vec_from_f32);
type SvdAttempt<T> = Option<Result<(Array<T>, Array<T>, Array<T>)>>;
pub(crate) fn try_scirs2_det<T>(a: &Array<T>) -> Option<Result<T>>
where
T: Float + Clone + Debug + 'static,
{
let n = square_dim(a, SCIRS2_MIN_DIM)?;
let op = borrow::operand(a);
if let Some(s) = cast::as_f64(&op) {
let view = ArrayView2::from_shape((n, n), s).ok()?;
let det = scirs2_linalg::det(&view, None).ok()?;
return Some(Ok(cast::f64_to(det)?));
}
if let Some(s) = cast::as_f32(&op) {
let view = ArrayView2::from_shape((n, n), s).ok()?;
let det = scirs2_linalg::det(&view, None).ok()?;
return Some(Ok(cast::f32_to(det)?));
}
None
}
pub(crate) fn try_scirs2_inv<T>(a: &Array<T>) -> Option<Result<Array<T>>>
where
T: Float + Clone + Debug + 'static,
{
let n = square_dim(a, SCIRS2_MIN_DIM)?;
let op = borrow::operand(a);
if let Some(s) = cast::as_f64(&op) {
let view = ArrayView2::from_shape((n, n), s).ok()?;
let inv = scirs2_linalg::inv(&view, None).ok()?;
return Some(Ok(matrix_from_f64(inv)?));
}
if let Some(s) = cast::as_f32(&op) {
let view = ArrayView2::from_shape((n, n), s).ok()?;
let inv = scirs2_linalg::inv(&view, None).ok()?;
return Some(Ok(matrix_from_f32(inv)?));
}
None
}
pub(crate) fn try_scirs2_solve<T>(a: &Array<T>, b: &Array<T>) -> Option<Result<Array<T>>>
where
T: Float + Clone + Debug + 'static,
{
let n = square_dim(a, SOLVE_MIN_DIM)?;
let b_shape = b.shape();
if b_shape.len() != 1 || b_shape[0] != n {
return None;
}
let a_op = borrow::operand(a);
let b_op = borrow::operand(b);
if let (Some(a_s), Some(b_s)) = (cast::as_f64(&a_op), cast::as_f64(&b_op)) {
let a_view = ArrayView2::from_shape((n, n), a_s).ok()?;
let b_view = ArrayView1::from_shape(n, b_s).ok()?;
return Some(finish_solve(
scirs2_linalg::solve(&a_view, &b_view, None).map(|x| vector_from_f64(x)),
));
}
if let (Some(a_s), Some(b_s)) = (cast::as_f32(&a_op), cast::as_f32(&b_op)) {
let a_view = ArrayView2::from_shape((n, n), a_s).ok()?;
let b_view = ArrayView1::from_shape(n, b_s).ok()?;
return Some(finish_solve(
scirs2_linalg::solve(&a_view, &b_view, None).map(|x| vector_from_f32(x)),
));
}
None
}
fn finish_solve<T>(r: std::result::Result<Option<Array<T>>, LinalgError>) -> Result<Array<T>> {
match r {
Ok(Some(arr)) => Ok(arr),
Ok(None) => Err(NumRs2Error::ConversionError(
"solve: solution is not representable in the array's element type".to_string(),
)),
Err(e) => Err(map_linalg_err(e)),
}
}
pub(crate) fn try_scirs2_svd<T>(a: &Array<T>) -> SvdAttempt<T>
where
T: Float + Clone + Debug + 'static,
{
let n = square_dim_in_range(a, SCIRS2_MIN_DIM, SVD_MAX_DIM)?;
let op = borrow::operand(a);
if let Some(s) = cast::as_f64(&op) {
let view = ArrayView2::from_shape((n, n), s).ok()?;
let (u, sv, vt) = scirs2_linalg::svd(&view, true, None).ok()?;
return Some(Ok((
matrix_from_f64(u)?,
vector_from_f64(sv)?,
matrix_from_f64(vt)?,
)));
}
if let Some(s) = cast::as_f32(&op) {
let view = ArrayView2::from_shape((n, n), s).ok()?;
let (u, sv, vt) = scirs2_linalg::svd(&view, true, None).ok()?;
return Some(Ok((
matrix_from_f32(u)?,
vector_from_f32(sv)?,
matrix_from_f32(vt)?,
)));
}
None
}
pub(crate) fn try_scirs2_eigh<T>(a: &Array<T>) -> Option<Result<(Array<T>, Array<T>)>>
where
T: Float + Clone + Debug + 'static,
{
let n = square_dim(a, SCIRS2_MIN_DIM)?;
let op = borrow::operand(a);
if let Some(s) = cast::as_f64(&op) {
if !is_exactly_symmetric(s, n) {
return None;
}
let view = ArrayView2::from_shape((n, n), s).ok()?;
let (vals, vecs) = scirs2_linalg::eigh(&view, None).ok()?;
let (vals, vecs) = sort_descending(vals, vecs);
return Some(Ok((vector_from_f64(vals)?, matrix_from_f64(vecs)?)));
}
if let Some(s) = cast::as_f32(&op) {
if !is_exactly_symmetric(s, n) {
return None;
}
let view = ArrayView2::from_shape((n, n), s).ok()?;
let (vals, vecs) = scirs2_linalg::eigh(&view, None).ok()?;
let (vals, vecs) = sort_descending(vals, vecs);
return Some(Ok((vector_from_f32(vals)?, matrix_from_f32(vecs)?)));
}
None
}
fn sort_descending<F: Float>(vals: Array1<F>, vecs: Array2<F>) -> (Array1<F>, Array2<F>) {
let n = vals.len();
let mut order: Vec<usize> = (0..n).collect();
order.sort_by(|&i, &j| {
vals[j]
.partial_cmp(&vals[i])
.unwrap_or(std::cmp::Ordering::Equal)
});
if order.iter().copied().eq(0..n) {
return (vals, vecs);
}
let sorted_vals = Array1::from_iter(order.iter().map(|&i| vals[i]));
let rows = vecs.nrows();
let mut sorted_vecs = Array2::<F>::zeros((rows, n));
for (new_col, &old_col) in order.iter().enumerate() {
for row in 0..rows {
sorted_vecs[[row, new_col]] = vecs[[row, old_col]];
}
}
(sorted_vals, sorted_vecs)
}
pub(crate) fn try_scirs2_eig<T>(a: &Array<T>) -> Option<Result<(Array<T>, Array<T>)>>
where
T: Float + Clone + Debug + 'static,
{
try_scirs2_eigh(a)
}
pub(crate) fn try_scirs2_cholesky<T>(a: &Array<T>) -> Option<Result<Array<T>>>
where
T: Float + Clone + Debug + 'static,
{
let n = square_dim(a, SCIRS2_MIN_DIM)?;
let op = borrow::operand(a);
if let Some(s) = cast::as_f64(&op) {
let view = ArrayView2::from_shape((n, n), s).ok()?;
let l = scirs2_linalg::cholesky(&view, None).ok()?;
return Some(Ok(matrix_from_f64(l)?));
}
if let Some(s) = cast::as_f32(&op) {
let view = ArrayView2::from_shape((n, n), s).ok()?;
let l = scirs2_linalg::cholesky(&view, None).ok()?;
return Some(Ok(matrix_from_f32(l)?));
}
None
}
pub(crate) fn try_scirs2_qr<T>(a: &Array<T>) -> Option<Result<(Array<T>, Array<T>)>>
where
T: Float + Clone + Debug + 'static,
{
let n = square_dim_in_range(a, SCIRS2_MIN_DIM, QR_MAX_DIM)?;
let op = borrow::operand(a);
if let Some(s) = cast::as_f64(&op) {
let view = ArrayView2::from_shape((n, n), s).ok()?;
let (q, r) = scirs2_linalg::qr(&view, None).ok()?;
return Some(Ok((matrix_from_f64(q)?, matrix_from_f64(r)?)));
}
if let Some(s) = cast::as_f32(&op) {
let view = ArrayView2::from_shape((n, n), s).ok()?;
let (q, r) = scirs2_linalg::qr(&view, None).ok()?;
return Some(Ok((matrix_from_f32(q)?, matrix_from_f32(r)?)));
}
None
}
fn is_exactly_symmetric<F: PartialEq>(buf: &[F], n: usize) -> bool {
for i in 0..n {
for j in (i + 1)..n {
if buf[i * n + j] != buf[j * n + i] {
return false;
}
}
}
true
}
#[cfg(test)]
mod tests {
use super::*;
fn identity(n: usize) -> Array<f64> {
let mut data = vec![0.0_f64; n * n];
for i in 0..n {
data[i * n + i] = 1.0;
}
Array::from_vec(data).reshape(&[n, n])
}
#[test]
fn small_matrices_are_not_eligible() {
for n in [1_usize, 2, 3, 4, 15] {
let a = identity(n);
assert!(try_scirs2_det(&a).is_none(), "det n={n}");
assert!(try_scirs2_inv(&a).is_none(), "inv n={n}");
assert!(try_scirs2_svd(&a).is_none(), "svd n={n}");
assert!(try_scirs2_qr(&a).is_none(), "qr n={n}");
assert!(try_scirs2_cholesky(&a).is_none(), "cholesky n={n}");
assert!(try_scirs2_eig(&a).is_none(), "eig n={n}");
}
}
#[test]
fn solve_gate_starts_at_four() {
for n in [1_usize, 2, 3] {
let a = identity(n);
let b = Array::from_vec(vec![1.0_f64; n]);
assert!(try_scirs2_solve(&a, &b).is_none(), "solve n={n}");
}
let a = identity(4);
let b = Array::from_vec(vec![1.0_f64; 4]);
assert!(try_scirs2_solve(&a, &b).is_some(), "solve n=4 must engage");
}
#[test]
fn large_matrices_are_not_eligible_for_qr_and_svd() {
let big = identity(QR_MAX_DIM);
assert!(try_scirs2_qr(&big).is_none(), "qr at QR_MAX_DIM");
assert!(
try_scirs2_qr(&identity(QR_MAX_DIM - 1)).is_some(),
"qr just below QR_MAX_DIM must still engage"
);
assert!(try_scirs2_det(&big).is_some(), "det has no upper bound");
assert!(try_scirs2_inv(&big).is_some(), "inv has no upper bound");
assert!(
try_scirs2_cholesky(&big).is_some(),
"cholesky has no upper bound"
);
assert!(try_scirs2_eig(&big).is_some(), "eig has no upper bound");
assert!(
try_scirs2_svd(&big).is_some(),
"svd's bound is higher than qr's"
);
let huge = identity(SVD_MAX_DIM);
assert!(try_scirs2_svd(&huge).is_none(), "svd at SVD_MAX_DIM");
assert!(
try_scirs2_svd(&identity(SVD_MAX_DIM - 1)).is_some(),
"svd just below SVD_MAX_DIM must still engage"
);
assert!(try_scirs2_det(&huge).is_some(), "det has no upper bound");
}
#[test]
fn non_square_input_is_not_eligible() {
let a = Array::from_vec(vec![1.0_f64; 16 * 20]).reshape(&[16, 20]);
assert!(try_scirs2_det(&a).is_none());
assert!(try_scirs2_inv(&a).is_none());
assert!(try_scirs2_svd(&a).is_none());
assert!(try_scirs2_qr(&a).is_none());
assert!(try_scirs2_cholesky(&a).is_none());
}
#[test]
fn solve_declines_mismatched_rhs() {
let a = identity(16);
let wrong_len = Array::from_vec(vec![1.0_f64; 15]);
assert!(try_scirs2_solve(&a, &wrong_len).is_none());
let wrong_rank = Array::from_vec(vec![1.0_f64; 16]).reshape(&[4, 4]);
assert!(try_scirs2_solve(&a, &wrong_rank).is_none());
}
#[test]
fn eig_declines_non_symmetric_input() {
let mut a = identity(16);
a.set(&[0, 1], 2.0).expect("in-bounds set");
assert!(
try_scirs2_eig(&a).is_none(),
"non-symmetric input must stay on the existing QR-iteration path"
);
}
#[test]
fn eig_accepts_symmetric_input() {
let mut a = identity(16);
a.set(&[0, 1], 2.0).expect("in-bounds set");
a.set(&[1, 0], 2.0).expect("in-bounds set");
assert!(try_scirs2_eig(&a).is_some());
}
#[test]
fn symmetry_test_is_exact() {
let n = 3;
let mut buf = vec![1.0_f64, 2.0, 3.0, 2.0, 4.0, 5.0, 3.0, 5.0, 6.0];
assert!(is_exactly_symmetric(&buf, n));
let nudged = f64::from_bits(buf[1].to_bits() + 1);
assert_ne!(nudged, buf[1], "next representable must actually differ");
buf[1] = nudged;
assert!(!is_exactly_symmetric(&buf, n));
}
#[test]
fn f32_operands_are_eligible_too() {
let mut data = vec![0.0_f32; 16 * 16];
for i in 0..16 {
data[i * 16 + i] = 2.0;
}
let a = Array::from_vec(data).reshape(&[16, 16]);
let det = try_scirs2_det(&a).expect("f32 must dispatch").expect("det");
assert!((det - 65536.0_f32).abs() < 1.0, "got {det}");
}
#[test]
fn sort_descending_permutes_columns_with_values() {
let vals = Array1::from_vec(vec![-3.0_f64, 2.0, 5.0]);
let vecs =
Array2::from_shape_vec((3, 3), vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0])
.expect("valid shape");
let (sv, svecs) = sort_descending(vals, vecs);
assert_eq!(sv.to_vec(), vec![5.0, 2.0, -3.0]);
assert_eq!(svecs.column(0).to_vec(), vec![3.0, 6.0, 9.0]);
assert_eq!(svecs.column(1).to_vec(), vec![2.0, 5.0, 8.0]);
assert_eq!(svecs.column(2).to_vec(), vec![1.0, 4.0, 7.0]);
}
#[test]
fn sort_descending_is_identity_when_already_ordered() {
let vals = Array1::from_vec(vec![5.0_f64, 2.0, -3.0]);
let vecs = Array2::from_shape_vec((3, 3), (0..9).map(|v| v as f64).collect())
.expect("valid shape");
let (sv, svecs) = sort_descending(vals.clone(), vecs.clone());
assert_eq!(sv, vals);
assert_eq!(svecs, vecs);
}
#[test]
fn sort_descending_is_algebraic_not_by_magnitude() {
let vals = Array1::from_vec(vec![1.0_f64, -9.0, 4.0]);
let vecs = Array2::<f64>::zeros((3, 3));
let (sv, _) = sort_descending(vals, vecs);
assert_eq!(sv.to_vec(), vec![4.0, 1.0, -9.0]);
}
#[test]
fn into_row_major_vec_handles_non_standard_layout() {
let m = Array2::from_shape_vec((2, 2), vec![1.0_f64, 2.0, 3.0, 4.0]).expect("shape");
let t = m.reversed_axes();
assert!(!t.is_standard_layout());
assert_eq!(into_row_major_vec(t), vec![1.0, 3.0, 2.0, 4.0]);
}
#[test]
fn into_row_major_vec_reuses_standard_layout_buffer() {
let m = Array2::from_shape_vec((2, 2), vec![1.0_f64, 2.0, 3.0, 4.0]).expect("shape");
assert_eq!(into_row_major_vec(m), vec![1.0, 2.0, 3.0, 4.0]);
}
#[test]
fn strided_operand_is_read_in_logical_order() {
let n = 16;
let mut data = vec![0.0_f64; n * n];
for i in 0..n {
for j in 0..n {
data[i * n + j] = if i == j { (i as f64) + 2.0 } else { 0.25 };
}
}
let a = Array::from_vec(data).reshape(&[n, n]);
let at = a.transpose();
let da = try_scirs2_det(&a).expect("eligible").expect("det ok");
let dt = try_scirs2_det(&at).expect("eligible").expect("det ok");
assert!((da - dt).abs() < 1e-6 * da.abs().max(1.0));
}
}