use ferrolearn_core::FerroError;
use ferrolearn_sparse::{CooMatrix, CscMatrix};
use ndarray::{Array1, array};
fn sample_a() -> CscMatrix<f64> {
CscMatrix::new(
3,
3,
vec![0, 2, 3, 5],
vec![0, 2, 1, 0, 2],
vec![1.0, 4.0, 3.0, 2.0, 5.0],
)
.unwrap()
}
fn sample_b() -> CscMatrix<f64> {
CscMatrix::new(
3,
3,
vec![0, 2, 4, 6],
vec![0, 2, 0, 1, 1, 2],
vec![1.0, 1.0, 2.0, 1.0, 3.0, 1.0],
)
.unwrap()
}
fn assert_dense_eq(d: &ndarray::Array2<f64>, expected: &[[f64; 3]; 3]) {
for (r, row) in expected.iter().enumerate() {
for (c, &v) in row.iter().enumerate() {
assert_eq!(d[[r, c]], v, "mismatch at ({r},{c})");
}
}
}
#[test]
fn csc_from_dense_to_dense_and_nnz_match_scipy() {
let dense = array![[1.0_f64, 0.0, 2.0], [0.0, 3.0, 0.0], [4.0, 0.0, 5.0]];
let m = CscMatrix::from_dense(&dense.view(), 0.0);
assert_eq!(m.nnz(), 5);
let expected = [[1.0, 0.0, 2.0], [0.0, 3.0, 0.0], [4.0, 0.0, 5.0]];
assert_dense_eq(&m.to_dense(), &expected);
}
#[test]
fn csc_from_dense_small_matches_scipy() {
let dense = array![[0.0_f64, 1.0], [2.0, 0.0]];
let m = CscMatrix::from_dense(&dense.view(), 0.0);
assert_eq!(m.nnz(), 2);
let d = m.to_dense();
assert_eq!(d[[0, 0]], 0.0);
assert_eq!(d[[0, 1]], 1.0);
assert_eq!(d[[1, 0]], 2.0);
assert_eq!(d[[1, 1]], 0.0);
}
#[test]
fn csc_from_coo_to_dense_matches_scipy() {
let mut coo: CooMatrix<f64> = CooMatrix::new(3, 3);
coo.push(0, 0, 1.0).unwrap();
coo.push(0, 2, 2.0).unwrap();
coo.push(1, 1, 3.0).unwrap();
coo.push(2, 0, 4.0).unwrap();
coo.push(2, 2, 5.0).unwrap();
let csc = CscMatrix::from_coo(&coo).unwrap();
assert_eq!(csc.nnz(), 5);
let expected = [[1.0, 0.0, 2.0], [0.0, 3.0, 0.0], [4.0, 0.0, 5.0]];
assert_dense_eq(&csc.to_dense(), &expected);
}
#[test]
fn csc_to_csr_roundtrip_matches_scipy() {
let a = sample_a();
let csr = a.to_csr();
let d = csr.to_dense();
let expected = [[1.0, 0.0, 2.0], [0.0, 3.0, 0.0], [4.0, 0.0, 5.0]];
for (r, row) in expected.iter().enumerate() {
for (c, &v) in row.iter().enumerate() {
assert_eq!(d[[r, c]], v, "tocsr mismatch at ({r},{c})");
}
}
}
#[test]
fn csc_mul_vec_matches_scipy() {
let a = sample_a();
let v = Array1::from(vec![1.0_f64, 2.0, 3.0]);
let r = a.mul_vec(&v).unwrap();
assert_eq!(r[0], 7.0);
assert_eq!(r[1], 6.0);
assert_eq!(r[2], 19.0);
}
#[test]
fn csc_add_self_matches_scipy() {
let a = sample_a();
let sum = a.add(&a).unwrap();
let expected = [[2.0, 0.0, 4.0], [0.0, 6.0, 0.0], [8.0, 0.0, 10.0]];
assert_dense_eq(&sum.to_dense(), &expected);
}
#[test]
fn csc_add_other_matches_scipy() {
let a = sample_a();
let b = sample_b();
let sum = a.add(&b).unwrap();
let expected = [[2.0, 2.0, 2.0], [0.0, 4.0, 3.0], [5.0, 0.0, 6.0]];
assert_dense_eq(&sum.to_dense(), &expected);
}
#[test]
fn csc_scalar_mul_matches_scipy() {
let expected = [[2.0, 0.0, 4.0], [0.0, 6.0, 0.0], [8.0, 0.0, 10.0]];
let a = sample_a();
let m2 = a.mul_scalar(2.0);
assert_dense_eq(&m2.to_dense(), &expected);
let mut a2 = sample_a();
a2.scale(2.0);
assert_dense_eq(&a2.to_dense(), &expected);
}
#[test]
fn csc_col_slice_matches_scipy() {
let a = sample_a();
let sliced = a.col_slice(0, 2).unwrap();
assert_eq!(sliced.n_rows(), 3);
assert_eq!(sliced.n_cols(), 2);
let d = sliced.to_dense();
let expected = [[1.0, 0.0], [0.0, 3.0], [4.0, 0.0]];
for (r, row) in expected.iter().enumerate() {
for (c, &v) in row.iter().enumerate() {
assert_eq!(d[[r, c]], v, "col_slice mismatch at ({r},{c})");
}
}
}
#[test]
fn csc_add_shape_mismatch_is_err() {
let a = sample_a();
let c = CscMatrix::<f64>::new(2, 3, vec![0, 0, 0, 0], vec![], vec![]).unwrap();
assert!(
a.add(&c).is_err(),
"shape-mismatched add must return Err (scipy raises ValueError)"
);
}
#[test]
fn csc_mul_vec_shape_mismatch_is_err() {
let a = sample_a();
let v = Array1::from(vec![1.0_f64, 2.0]);
assert!(
a.mul_vec(&v).is_err(),
"wrong-length matvec must return Err (scipy raises ValueError)"
);
}
#[test]
fn csc_geometry_matches_scipy() {
let a = sample_a();
assert_eq!((a.n_rows(), a.n_cols()), (3, 3));
assert_eq!(a.nnz(), 5);
}
fn sample_b_elmul() -> CscMatrix<f64> {
CscMatrix::new(
3,
3,
vec![0, 1, 3, 5],
vec![0, 0, 1, 1, 2],
vec![1.0, 1.0, 1.0, 1.0, 1.0],
)
.unwrap()
}
#[test]
fn csc_multiply_matches_scipy() {
let a = sample_a();
let b = sample_b_elmul();
let prod = a.multiply(&b).unwrap();
let expected = [[1.0, 0.0, 0.0], [0.0, 3.0, 0.0], [0.0, 0.0, 5.0]];
assert_dense_eq(&prod.to_dense(), &expected);
}
#[test]
fn csc_sub_matches_scipy() {
let a = sample_a();
let b = sample_b_elmul();
let diff = a.sub(&b).unwrap();
let expected = [[0.0, -1.0, 2.0], [0.0, 2.0, -1.0], [4.0, 0.0, 4.0]];
assert_dense_eq(&diff.to_dense(), &expected);
}
#[test]
fn csc_multiply_shape_mismatch_is_err() {
let a = sample_a();
let c = CscMatrix::<f64>::new(2, 3, vec![0, 0, 0, 0], vec![], vec![]).unwrap();
assert!(
a.multiply(&c).is_err(),
"shape-mismatched multiply must return Err (scipy raises ValueError)"
);
}
#[test]
fn csc_sub_shape_mismatch_is_err() {
let a = sample_a();
let c = CscMatrix::<f64>::new(2, 3, vec![0, 0, 0, 0], vec![], vec![]).unwrap();
assert!(
a.sub(&c).is_err(),
"shape-mismatched sub must return Err (scipy raises ValueError)"
);
}
fn sample_c() -> CscMatrix<f64> {
let dense = array![[1.0_f64, 2.0], [3.0, 4.0], [5.0, 6.0]];
CscMatrix::from_dense(&dense.view(), 0.0)
}
#[test]
fn csc_matmul_matches_scipy() {
let a = sample_a();
let b = sample_b_elmul();
let prod = a.matmul(&b).unwrap();
let expected = [[1.0, 1.0, 2.0], [0.0, 3.0, 3.0], [4.0, 4.0, 5.0]];
assert_dense_eq(&prod.to_dense(), &expected);
}
#[test]
fn csc_matmul_non_square() {
let a = sample_a();
let c = sample_c();
let prod = a.matmul(&c).unwrap();
assert_eq!(prod.n_rows(), 3);
assert_eq!(prod.n_cols(), 2);
let d = prod.to_dense();
let expected = [[11.0, 14.0], [9.0, 12.0], [29.0, 38.0]];
for (r, row) in expected.iter().enumerate() {
for (col, &v) in row.iter().enumerate() {
assert_eq!(d[[r, col]], v, "matmul mismatch at ({r},{col})");
}
}
}
#[test]
fn csc_matmul_shape_mismatch_is_err() {
let a = sample_a();
let d = CscMatrix::<f64>::new(2, 2, vec![0, 0, 0], vec![], vec![]).unwrap();
assert!(
a.matmul(&d).is_err(),
"inner-dimension-mismatched matmul must return Err (scipy raises ValueError)"
);
}
#[test]
fn csc_get_element_matches_scipy() {
let a = sample_a();
assert_eq!(a.get(1, 1).unwrap(), 3.0);
assert_eq!(a.get(0, 0).unwrap(), 1.0);
assert_eq!(a.get(0, 2).unwrap(), 2.0);
assert_eq!(a.get(2, 0).unwrap(), 4.0);
}
#[test]
fn csc_get_absent_is_zero() {
let a = sample_a();
assert_eq!(a.get(0, 1).unwrap(), 0.0);
}
#[test]
fn csc_get_out_of_bounds_is_err() {
let a = sample_a();
assert!(
matches!(
a.get(3, 0),
Err(ferrolearn_core::FerroError::InvalidParameter { .. })
),
"row index out of bounds must return Err(InvalidParameter)"
);
assert!(
matches!(
a.get(0, 3),
Err(ferrolearn_core::FerroError::InvalidParameter { .. })
),
"col index out of bounds must return Err(InvalidParameter)"
);
}
#[test]
fn csc_getrow_matches_scipy() {
let a = sample_a();
let r0 = a.getrow(0).unwrap();
assert_eq!(r0.n_rows(), 1);
assert_eq!(r0.n_cols(), 3);
let d0 = r0.to_dense();
assert_eq!(d0[[0, 0]], 1.0);
assert_eq!(d0[[0, 1]], 0.0);
assert_eq!(d0[[0, 2]], 2.0);
let r1 = a.getrow(1).unwrap();
let d1 = r1.to_dense();
assert_eq!(d1[[0, 0]], 0.0);
assert_eq!(d1[[0, 1]], 3.0);
assert_eq!(d1[[0, 2]], 0.0);
}
#[test]
fn csc_getcol_matches_scipy() {
let a = sample_a();
let c0 = a.getcol(0).unwrap();
assert_eq!(c0.n_rows(), 3);
assert_eq!(c0.n_cols(), 1);
let d0 = c0.to_dense();
assert_eq!(d0[[0, 0]], 1.0);
assert_eq!(d0[[1, 0]], 0.0);
assert_eq!(d0[[2, 0]], 4.0);
let c2 = a.getcol(2).unwrap();
let d2 = c2.to_dense();
assert_eq!(d2[[0, 0]], 2.0);
assert_eq!(d2[[1, 0]], 0.0);
assert_eq!(d2[[2, 0]], 5.0);
}
#[test]
fn csc_getrow_getcol_out_of_bounds_is_err() {
let a = sample_a();
assert!(
matches!(
a.getrow(3),
Err(ferrolearn_core::FerroError::InvalidParameter { .. })
),
"row index out of bounds must return Err(InvalidParameter)"
);
assert!(
matches!(
a.getcol(3),
Err(ferrolearn_core::FerroError::InvalidParameter { .. })
),
"col index out of bounds must return Err(InvalidParameter)"
);
}
#[test]
fn csc_shape_data_indices_indptr_match_scipy() {
let dense = array![[1.0_f64, 0.0, 2.0], [0.0, 3.0, 0.0], [4.0, 0.0, 5.0]];
let a = CscMatrix::from_dense(&dense.view(), 0.0);
assert_eq!(a.shape(), (3, 3));
assert_eq!(a.data(), &[1.0, 4.0, 3.0, 2.0, 5.0]);
assert_eq!(a.indices(), &[0, 2, 1, 0, 2]);
assert_eq!(a.indptr(), vec![0, 2, 3, 5]);
}
#[test]
fn csc_max_min_folds_implicit_zero() -> Result<(), FerroError> {
let neg = array![[-3.0_f64, 0.0, 0.0], [0.0, -1.0, 0.0], [0.0, 0.0, -5.0]];
let m_neg = CscMatrix::from_dense(&neg.view(), 0.0);
assert_eq!(m_neg.max(), 0.0);
assert_eq!(m_neg.min(), -5.0);
let pos = array![[3.0_f64, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 5.0]];
let m_pos = CscMatrix::from_dense(&pos.view(), 0.0);
assert_eq!(m_pos.max(), 5.0);
assert_eq!(m_pos.min(), 0.0);
Ok(())
}
#[test]
fn csc_astype_truncates() -> Result<(), FerroError> {
let dense = array![[3.7_f64, 0.0, 0.0], [0.0, -2.9, 0.0], [0.0, 0.0, 5.0]];
let m = CscMatrix::from_dense(&dense.view(), 0.0);
let cast: CscMatrix<i64> = m.astype(|&v| v as i64)?;
assert_eq!(cast.data(), &[3_i64, -2, 5]);
assert_eq!(cast.indptr(), m.indptr());
assert_eq!(cast.indices(), m.indices());
assert_eq!(cast.shape(), m.shape());
assert_eq!(cast.nnz(), 3);
Ok(())
}
#[test]
fn csc_copy_preserves_structure() -> Result<(), FerroError> {
let a = sample_a();
let c = a.copy();
assert_eq!(c.nnz(), a.nnz());
assert_eq!(c.data(), a.data());
assert_eq!(c.to_dense(), a.to_dense());
assert_eq!(a.nnz(), 5);
let expected = [[1.0, 0.0, 2.0], [0.0, 3.0, 0.0], [4.0, 0.0, 5.0]];
assert_dense_eq(&a.to_dense(), &expected);
Ok(())
}
#[test]
fn csc_eliminate_zeros_matches_scipy() -> Result<(), FerroError> {
let m = CscMatrix::new(
3,
3,
vec![0, 1, 2, 3],
vec![0, 1, 2],
vec![3.0_f64, 0.0, 5.0],
)?;
assert_eq!(m.nnz(), 3);
let pruned = m.eliminate_zeros()?;
assert_eq!(pruned.nnz(), 2);
assert_eq!(pruned.data(), &[3.0, 5.0]);
assert_eq!(pruned.indices(), &[0, 2]);
assert_eq!(pruned.indptr(), vec![0, 1, 1, 2]);
let expected = [[3.0, 0.0, 0.0], [0.0, 0.0, 0.0], [0.0, 0.0, 5.0]];
assert_dense_eq(&pruned.to_dense(), &expected);
Ok(())
}
#[test]
fn csc_power_matches_scipy() -> Result<(), FerroError> {
let dense = array![[2.0_f64, 0.0], [0.0, -3.0]];
let m = CscMatrix::from_dense(&dense.view(), 0.0);
let p2 = m.power(2.0)?;
assert_eq!(p2.data(), &[4.0, 9.0]);
let p3 = m.power(3.0)?;
assert_eq!(p3.data(), &[8.0, -27.0]);
Ok(())
}
#[test]
fn csc_argmax_matches_scipy() -> Result<(), FerroError> {
let dense = array![[1.0_f64, 0.0, 2.0], [0.0, 5.0, 0.0], [4.0, 0.0, 3.0]];
let a = CscMatrix::from_dense(&dense.view(), 0.0);
assert_eq!(a.argmax()?, 4);
assert_eq!(a.argmin()?, 1);
Ok(())
}
#[test]
fn csc_argmin_dense_all_negative() -> Result<(), FerroError> {
let dense = array![[-1.0_f64, -2.0], [-3.0, -4.0]];
let b = CscMatrix::from_dense(&dense.view(), 0.0);
assert_eq!(b.argmax()?, 0);
assert_eq!(b.argmin()?, 3);
Ok(())
}
#[test]
fn csc_argmax_implicit_zero() -> Result<(), FerroError> {
let dense = array![[-1.0_f64, 0.0], [-3.0, -4.0]];
let c = CscMatrix::from_dense(&dense.view(), 0.0);
assert_eq!(c.argmax()?, 1);
assert_eq!(c.argmin()?, 3);
Ok(())
}
#[test]
fn csc_argmax_all_zero() -> Result<(), FerroError> {
let dense = array![[0.0_f64, 0.0, 0.0], [0.0, 0.0, 0.0]];
let z = CscMatrix::from_dense(&dense.view(), 0.0);
assert_eq!(z.nnz(), 0);
assert_eq!(z.argmax()?, 0);
assert_eq!(z.argmin()?, 0);
Ok(())
}
#[test]
fn csc_argmax_ties_earliest() -> Result<(), FerroError> {
let t_dense = array![[5.0_f64, 5.0], [1.0, 5.0]];
let t = CscMatrix::from_dense(&t_dense.view(), 0.0);
assert_eq!(t.argmax()?, 0);
let p_dense = array![[2.0_f64, 0.0, 2.0]];
let p = CscMatrix::from_dense(&p_dense.view(), 0.0);
assert_eq!(p.argmax()?, 0);
assert_eq!(p.argmin()?, 1);
Ok(())
}
#[test]
fn csc_sort_indices_canonical() -> Result<(), FerroError> {
let a = sample_a();
let sorted = a.sort_indices()?;
assert_eq!(sorted.data(), &[1.0, 4.0, 3.0, 2.0, 5.0]);
assert_eq!(sorted.indices(), &[0, 2, 1, 0, 2]);
assert_eq!(sorted.indptr(), vec![0, 2, 3, 5]);
let expected = [[1.0, 0.0, 2.0], [0.0, 3.0, 0.0], [4.0, 0.0, 5.0]];
assert_dense_eq(&sorted.to_dense(), &expected);
Ok(())
}
#[test]
fn csc_sum_duplicates_canonical() -> Result<(), FerroError> {
let a = sample_a();
let coalesced = a.sum_duplicates()?;
assert_eq!(coalesced.data(), &[1.0, 4.0, 3.0, 2.0, 5.0]);
assert_eq!(coalesced.indices(), &[0, 2, 1, 0, 2]);
assert_eq!(coalesced.indptr(), vec![0, 2, 3, 5]);
let expected = [[1.0, 0.0, 2.0], [0.0, 3.0, 0.0], [4.0, 0.0, 5.0]];
assert_dense_eq(&coalesced.to_dense(), &expected);
Ok(())
}
#[test]
fn csc_sum_duplicates_preserves_zero_sum() -> Result<(), FerroError> {
let m = CscMatrix::new(
3,
3,
vec![0, 1, 2, 3],
vec![0, 1, 2],
vec![3.0_f64, 0.0, 5.0],
)?;
assert_eq!(m.nnz(), 3);
let coalesced = m.sum_duplicates()?;
assert_eq!(coalesced.nnz(), 3);
assert_eq!(coalesced.data(), &[3.0, 0.0, 5.0]);
assert_eq!(coalesced.indices(), &[0, 1, 2]);
assert_eq!(coalesced.indptr(), vec![0, 1, 2, 3]);
Ok(())
}