use crate::error::{DatarustError, Result};
fn checked_element_count(rows: usize, cols: usize) -> Result<usize> {
rows.checked_mul(cols)
.ok_or_else(|| DatarustError::ShapeMismatch {
expected: "rows * cols within usize range".into(),
actual: format!("{} rows × {} cols overflows usize", rows, cols),
})
}
#[derive(Debug, Clone, PartialEq)]
pub struct Matrix {
data: Vec<f64>,
rows: usize,
cols: usize,
}
impl Matrix {
pub fn new(data: Vec<Vec<f64>>) -> Result<Self> {
if data.is_empty() {
return Err(DatarustError::EmptyInput("matrix has no rows".into()));
}
let cols = data[0].len();
if cols == 0 {
return Err(DatarustError::EmptyInput("matrix has no columns".into()));
}
for (i, row) in data.iter().enumerate() {
if row.len() != cols {
return Err(DatarustError::ShapeMismatch {
expected: format!("{} columns", cols),
actual: format!("{} columns at row {}", row.len(), i),
});
}
}
let rows = data.len();
let element_count = checked_element_count(rows, cols)?;
let mut flat = Vec::with_capacity(element_count);
for row in data {
flat.extend(row);
}
Ok(Self {
data: flat,
rows,
cols,
})
}
pub fn from_rows(rows: Vec<Vec<f64>>) -> Result<Self> {
Self::new(rows)
}
pub fn from_flat(rows: usize, cols: usize, flat: Vec<f64>) -> Result<Self> {
if rows == 0 || cols == 0 {
return Err(DatarustError::EmptyInput("zero dimension".into()));
}
let expected = checked_element_count(rows, cols)?;
if flat.len() != expected {
return Err(DatarustError::ShapeMismatch {
expected: format!("{} elements", expected),
actual: format!("{} elements", flat.len()),
});
}
Ok(Self {
data: flat,
rows,
cols,
})
}
pub fn zeros(rows: usize, cols: usize) -> Result<Self> {
if rows == 0 || cols == 0 {
return Err(DatarustError::EmptyInput("zero dimension".into()));
}
let element_count = checked_element_count(rows, cols)?;
Ok(Self {
data: vec![0.0; element_count],
rows,
cols,
})
}
pub fn identity(n: usize) -> Result<Self> {
if n == 0 {
return Err(DatarustError::EmptyInput("zero dimension".into()));
}
let element_count = checked_element_count(n, n)?;
let mut data = vec![0.0; element_count];
for i in 0..n {
data[i * n + i] = 1.0;
}
Ok(Self {
data,
rows: n,
cols: n,
})
}
#[inline]
pub fn nrows(&self) -> usize {
self.rows
}
#[inline]
pub fn ncols(&self) -> usize {
self.cols
}
#[inline]
pub fn as_slice(&self) -> &[f64] {
&self.data
}
#[inline]
pub fn as_mut_slice(&mut self) -> &mut [f64] {
&mut self.data
}
#[inline]
pub fn get(&self, i: usize, j: usize) -> f64 {
assert!(
i < self.rows && j < self.cols,
"Matrix::get: index ({}, {}) out of bounds for {}×{}",
i,
j,
self.rows,
self.cols
);
self.data[i * self.cols + j]
}
#[inline]
pub fn checked_get(&self, i: usize, j: usize) -> Option<f64> {
if i < self.rows && j < self.cols {
Some(self.data[i * self.cols + j])
} else {
None
}
}
#[inline]
pub fn set(&mut self, i: usize, j: usize, v: f64) {
self.data[i * self.cols + j] = v;
}
pub fn validate_no_nan(&self) -> Result<()> {
for (flat_idx, &v) in self.data.iter().enumerate() {
if v.is_nan() {
let i = flat_idx / self.cols;
let j = flat_idx % self.cols;
return Err(DatarustError::InvalidInput(format!(
"NaN value at position ({}, {})",
i, j
)));
}
}
Ok(())
}
pub fn validate_finite(&self) -> Result<()> {
for (flat_idx, &v) in self.data.iter().enumerate() {
if !v.is_finite() {
let i = flat_idx / self.cols;
let j = flat_idx % self.cols;
return Err(DatarustError::InvalidInput(format!(
"non-finite value {v} at position ({i}, {j})"
)));
}
}
Ok(())
}
pub fn validate_no_infinite(&self) -> Result<()> {
for (flat_idx, &v) in self.data.iter().enumerate() {
if v.is_infinite() {
let i = flat_idx / self.cols;
let j = flat_idx % self.cols;
return Err(DatarustError::InvalidInput(format!(
"infinite value {v} at position ({i}, {j})"
)));
}
}
Ok(())
}
#[inline]
pub fn row(&self, i: usize) -> &[f64] {
let start = i * self.cols;
&self.data[start..start + self.cols]
}
pub fn col(&self, j: usize) -> Vec<f64> {
(0..self.rows)
.map(|i| self.data[i * self.cols + j])
.collect()
}
pub fn iter_rows(&self) -> impl Iterator<Item = &[f64]> {
self.data.chunks_exact(self.cols)
}
#[doc(hidden)]
pub fn rows_ref(&self) -> Vec<Vec<f64>> {
self.data
.chunks_exact(self.cols)
.map(|chunk| chunk.to_vec())
.collect()
}
#[doc(hidden)]
pub fn into_rows(self) -> Vec<Vec<f64>> {
let cols = self.cols;
self.data.chunks(cols).map(|chunk| chunk.to_vec()).collect()
}
pub fn transpose(&self) -> Matrix {
let rows = self.rows;
let cols = self.cols;
let mut out = vec![0.0; rows * cols];
for i in 0..rows {
for j in 0..cols {
out[j * rows + i] = self.data[i * cols + j];
}
}
Matrix {
data: out,
rows: cols,
cols: rows,
}
}
#[allow(clippy::needless_range_loop)]
pub fn matmul(&self, other: &Matrix) -> Result<Matrix> {
if self.cols != other.rows {
return Err(DatarustError::ShapeMismatch {
expected: format!("second operand with {} rows", self.cols),
actual: format!("{} rows", other.rows),
});
}
let m = self.rows;
let k = self.cols;
let n = other.cols;
let element_count = checked_element_count(m, n)?;
let mut out = vec![0.0; element_count];
#[cfg(feature = "matrixmultiply")]
{
unsafe {
matrixmultiply::dgemm(
m,
k,
n,
1.0,
self.data.as_ptr(),
k as isize,
1,
other.data.as_ptr(),
n as isize,
1,
0.0,
out.as_mut_ptr(),
n as isize,
1,
);
}
}
#[cfg(not(feature = "matrixmultiply"))]
{
for i in 0..m {
let out_base = i * n;
let self_base = i * k;
for l in 0..k {
let a = self.data[self_base + l];
if a == 0.0 {
continue;
}
let other_base = l * n;
for j in 0..n {
out[out_base + j] += a * other.data[other_base + j];
}
}
}
}
Ok(Matrix {
data: out,
rows: m,
cols: n,
})
}
pub fn column_mean(&self) -> Vec<f64> {
crate::stats::column_mean_flat(&self.data, self.rows, self.cols)
}
pub fn from_columns(cols: Vec<Vec<f64>>) -> Result<Self> {
if cols.is_empty() || cols[0].is_empty() {
return Err(DatarustError::EmptyInput("no columns".into()));
}
let rows = cols[0].len();
for c in &cols {
if c.len() != rows {
return Err(DatarustError::ShapeMismatch {
expected: format!("{} rows", rows),
actual: format!("{} rows", c.len()),
});
}
}
let ncols = cols.len();
let element_count = checked_element_count(rows, ncols)?;
let mut data = vec![0.0; element_count];
for (j, col) in cols.iter().enumerate() {
for (i, &v) in col.iter().enumerate() {
data[i * ncols + j] = v;
}
}
Ok(Self {
data,
rows,
cols: ncols,
})
}
pub fn select_columns(&self, indices: &[usize]) -> Result<Self> {
if indices.is_empty() {
return Err(DatarustError::EmptyInput("no columns selected".into()));
}
let ncols = self.cols;
for &c in indices {
if c >= ncols {
return Err(DatarustError::InvalidInput(format!(
"column index {} out of range (ncols {})",
c, ncols
)));
}
}
let out_cols = indices.len();
let element_count = checked_element_count(self.rows, out_cols)?;
let mut out = Vec::with_capacity(element_count);
for i in 0..self.rows {
let base = i * ncols;
for &c in indices {
out.push(self.data[base + c]);
}
}
Ok(Self {
data: out,
rows: self.rows,
cols: out_cols,
})
}
pub fn select_rows(&self, indices: &[usize]) -> Result<Self> {
if indices.is_empty() {
return Err(DatarustError::EmptyInput("no rows selected".into()));
}
let nrows = self.rows;
for &r in indices {
if r >= nrows {
return Err(DatarustError::InvalidInput(format!(
"row index {} out of range (nrows {})",
r, nrows
)));
}
}
let element_count = checked_element_count(indices.len(), self.cols)?;
let mut out = Vec::with_capacity(element_count);
for &r in indices {
let start = r * self.cols;
out.extend_from_slice(&self.data[start..start + self.cols]);
}
Ok(Self {
data: out,
rows: indices.len(),
cols: self.cols,
})
}
}
#[cfg(feature = "serde")]
impl serde::Serialize for Matrix {
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
use serde::ser::SerializeStruct;
let rows: Vec<&[f64]> = self.data.chunks_exact(self.cols).collect();
let mut s = serializer.serialize_struct("Matrix", 1)?;
s.serialize_field("data", &rows)?;
s.end()
}
}
#[cfg(feature = "serde")]
impl<'de> serde::Deserialize<'de> for Matrix {
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(serde::Deserialize)]
struct Raw {
data: Vec<Vec<f64>>,
}
let raw = Raw::deserialize(deserializer)?;
Matrix::new(raw.data).map_err(serde::de::Error::custom)
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct StrMatrix {
pub(crate) data: Vec<Vec<String>>,
}
impl StrMatrix {
pub fn new(data: Vec<Vec<String>>) -> Result<Self> {
if data.is_empty() {
return Err(DatarustError::EmptyInput("matrix has no rows".into()));
}
let cols = data[0].len();
if cols == 0 {
return Err(DatarustError::EmptyInput("matrix has no columns".into()));
}
for (i, row) in data.iter().enumerate() {
if row.len() != cols {
return Err(DatarustError::ShapeMismatch {
expected: format!("{} columns", cols),
actual: format!("{} columns at row {}", row.len(), i),
});
}
}
Ok(Self { data })
}
pub fn from_column<I, S>(col: I) -> Result<Self>
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
let data: Vec<Vec<String>> = col.into_iter().map(|s| vec![s.into()]).collect::<Vec<_>>();
if data.is_empty() {
return Err(DatarustError::EmptyInput("column has no rows".into()));
}
Self::new(data)
}
pub fn from_strings<I, S>(rows: I) -> Result<Self>
where
I: IntoIterator<Item = Vec<S>>,
S: Into<String>,
{
let data: Vec<Vec<String>> = rows
.into_iter()
.map(|r| r.into_iter().map(|s| s.into()).collect())
.collect();
Self::new(data)
}
#[inline]
pub fn nrows(&self) -> usize {
self.data.len()
}
#[inline]
pub fn ncols(&self) -> usize {
self.data[0].len()
}
#[inline]
pub fn get(&self, i: usize, j: usize) -> &str {
debug_assert!(
i < self.data.len() && j < self.data[i].len(),
"StrMatrix::get: index ({}, {}) out of bounds for {}×{}",
i,
j,
self.data.len(),
self.data[0].len()
);
&self.data[i][j]
}
#[inline]
pub fn checked_get(&self, i: usize, j: usize) -> Option<&str> {
self.data.get(i)?.get(j).map(|s| s.as_str())
}
pub fn column(&self, j: usize) -> Vec<String> {
self.data.iter().map(|r| r[j].clone()).collect()
}
pub fn row(&self, i: usize) -> &[String] {
&self.data[i]
}
}
#[cfg(feature = "serde")]
impl serde::Serialize for StrMatrix {
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
use serde::ser::SerializeStruct;
let mut s = serializer.serialize_struct("StrMatrix", 1)?;
s.serialize_field("data", &self.data)?;
s.end()
}
}
#[cfg(feature = "serde")]
impl<'de> serde::Deserialize<'de> for StrMatrix {
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(serde::Deserialize)]
struct Raw {
data: Vec<Vec<String>>,
}
let raw = Raw::deserialize(deserializer)?;
StrMatrix::new(raw.data).map_err(serde::de::Error::custom)
}
}
impl TryFrom<Vec<Vec<f64>>> for Matrix {
type Error = DatarustError;
fn try_from(data: Vec<Vec<f64>>) -> Result<Self> {
Matrix::new(data)
}
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct SparseMatrix {
nrows: usize,
ncols: usize,
indptr: Vec<usize>,
indices: Vec<usize>,
data: Vec<f64>,
}
impl SparseMatrix {
pub fn new(
nrows: usize,
ncols: usize,
indptr: Vec<usize>,
indices: Vec<usize>,
data: Vec<f64>,
) -> Result<Self> {
if nrows == 0 || ncols == 0 {
return Err(DatarustError::EmptyInput("zero dimension".into()));
}
let expected_indptr_len =
nrows
.checked_add(1)
.ok_or_else(|| DatarustError::ShapeMismatch {
expected: "nrows + 1 within usize range".into(),
actual: format!("{} rows overflows the CSR indptr length", nrows),
})?;
if indptr.len() != expected_indptr_len {
return Err(DatarustError::ShapeMismatch {
expected: format!("{} indptr entries", expected_indptr_len),
actual: format!("{} indptr entries", indptr.len()),
});
}
if indices.len() != data.len() {
return Err(DatarustError::ShapeMismatch {
expected: format!("{} indices", data.len()),
actual: format!("{} indices", indices.len()),
});
}
let nnz = data.len();
if indptr[0] != 0 || indptr[nrows] != nnz {
return Err(DatarustError::InvalidInput(
"indptr must start at 0 and end at nnz".into(),
));
}
for row in 0..nrows {
let start = indptr[row];
let end = indptr[row + 1];
if start > end || end > nnz {
return Err(DatarustError::InvalidInput(format!(
"indptr must be non-decreasing and stay within nnz (row {})",
row
)));
}
let row_indices = &indices[start..end];
for &c in row_indices {
if c >= ncols {
return Err(DatarustError::InvalidInput(format!(
"column index {} out of range (ncols {})",
c, ncols
)));
}
}
if row_indices.windows(2).any(|pair| pair[0] >= pair[1]) {
return Err(DatarustError::InvalidInput(format!(
"column indices in row {} must be strictly increasing",
row
)));
}
}
Ok(Self {
nrows,
ncols,
indptr,
indices,
data,
})
}
pub fn from_triplets(
nrows: usize,
ncols: usize,
triplets: &[(usize, usize, f64)],
) -> Result<Self> {
if nrows == 0 || ncols == 0 {
return Err(DatarustError::EmptyInput("zero dimension".into()));
}
let indptr_len = nrows
.checked_add(1)
.ok_or_else(|| DatarustError::ShapeMismatch {
expected: "nrows + 1 within usize range".into(),
actual: format!("{} rows overflows the CSR indptr length", nrows),
})?;
let mut per_row: Vec<Vec<(usize, f64)>> = vec![vec![]; nrows];
for &(r, c, v) in triplets {
if r >= nrows {
return Err(DatarustError::InvalidInput(format!(
"row {} out of range (nrows {})",
r, nrows
)));
}
if c >= ncols {
return Err(DatarustError::InvalidInput(format!(
"col {} out of range (ncols {})",
c, ncols
)));
}
if v != 0.0 {
per_row[r].push((c, v));
}
}
let mut indptr = Vec::with_capacity(indptr_len);
let mut indices = Vec::new();
let mut data = Vec::new();
indptr.push(0);
for row_entries in &mut per_row {
row_entries.sort_by_key(|(c, _)| *c);
let mut position = 0;
while position < row_entries.len() {
let column = row_entries[position].0;
let mut value = 0.0;
while position < row_entries.len() && row_entries[position].0 == column {
value += row_entries[position].1;
position += 1;
}
if value != 0.0 {
indices.push(column);
data.push(value);
}
}
indptr.push(indices.len());
}
Ok(Self {
nrows,
ncols,
indptr,
indices,
data,
})
}
pub fn zeros(nrows: usize, ncols: usize) -> Result<Self> {
if nrows == 0 || ncols == 0 {
return Err(DatarustError::EmptyInput("zero dimension".into()));
}
let indptr_len = nrows
.checked_add(1)
.ok_or_else(|| DatarustError::ShapeMismatch {
expected: "nrows + 1 within usize range".into(),
actual: format!("{} rows overflows the CSR indptr length", nrows),
})?;
Ok(Self {
nrows,
ncols,
indptr: vec![0; indptr_len],
indices: vec![],
data: vec![],
})
}
#[inline]
pub fn nrows(&self) -> usize {
self.nrows
}
#[inline]
pub fn ncols(&self) -> usize {
self.ncols
}
#[inline]
pub fn nnz(&self) -> usize {
self.data.len()
}
pub fn density(&self) -> f64 {
let total = self.nrows.saturating_mul(self.ncols);
if total == 0 {
return 0.0;
}
self.nnz() as f64 / total as f64
}
pub fn get(&self, i: usize, j: usize) -> f64 {
assert!(
i < self.nrows && j < self.ncols,
"SparseMatrix::get: index ({}, {}) out of bounds for {}×{}",
i,
j,
self.nrows,
self.ncols
);
let start = self.indptr[i];
let end = self.indptr[i + 1];
let slice = &self.indices[start..end];
match slice.binary_search(&j) {
Ok(local) => self.data[start + local],
Err(_) => 0.0,
}
}
pub fn checked_get(&self, i: usize, j: usize) -> Option<f64> {
if i >= self.nrows || j >= self.ncols {
return None;
}
let start = *self.indptr.get(i)?;
let end = *self.indptr.get(i + 1)?;
let slice = self.indices.get(start..end)?;
match slice.binary_search(&j) {
Ok(local) => Some(self.data[start + local]),
Err(_) => Some(0.0),
}
}
pub fn row_nz(&self, i: usize) -> impl Iterator<Item = (usize, f64)> + '_ {
assert!(
i < self.nrows,
"SparseMatrix::row_nz: row {} out of bounds for {} rows",
i,
self.nrows
);
let start = self.indptr[i];
let end = self.indptr[i + 1];
self.indices[start..end]
.iter()
.zip(self.data[start..end].iter())
.map(|(&c, &v)| (c, v))
}
pub fn to_dense(&self) -> Result<Matrix> {
let mut rows = vec![vec![0.0; self.ncols]; self.nrows];
for (i, row) in rows.iter_mut().enumerate() {
for (c, v) in self.row_nz(i) {
row[c] = v;
}
}
Matrix::new(rows)
}
}
#[cfg(feature = "serde")]
impl<'de> serde::Deserialize<'de> for SparseMatrix {
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(serde::Deserialize)]
struct Raw {
nrows: usize,
ncols: usize,
indptr: Vec<usize>,
indices: Vec<usize>,
data: Vec<f64>,
}
let raw = Raw::deserialize(deserializer)?;
SparseMatrix::new(raw.nrows, raw.ncols, raw.indptr, raw.indices, raw.data)
.map_err(serde::de::Error::custom)
}
}
#[cfg(test)]
#[allow(dead_code)]
pub(crate) fn approx_eq_matrices(a: &Matrix, b: &Matrix, tol: f64) -> bool {
if a.nrows() != b.nrows() || a.ncols() != b.ncols() {
return false;
}
for i in 0..a.nrows() {
for j in 0..a.ncols() {
if (a.get(i, j) - b.get(i, j)).abs() > tol {
return false;
}
}
}
true
}
#[cfg(test)]
#[allow(dead_code)]
pub(crate) fn approx_eq_vecs(a: &[f64], b: &[f64], tol: f64) -> bool {
if a.len() != b.len() {
return false;
}
a.iter().zip(b.iter()).all(|(x, y)| (x - y).abs() <= tol)
}
#[cfg(test)]
#[macro_use]
mod assert_macros {
#[allow(unused_macros)]
macro_rules! assert_mat_eq {
($a:expr, $b:expr, $tol:expr) => {{
assert!(
$crate::matrix::approx_eq_matrices(&$a, &$b, $tol),
"matrices not equal within tolerance {}\n left: {:?}\nright: {:?}",
$tol,
$a.as_slice(),
$b.as_slice()
);
}};
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn new_valid() {
let m = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
assert_eq!(m.nrows(), 2);
assert_eq!(m.ncols(), 2);
}
#[test]
fn new_jagged_rejected() {
let err = Matrix::new(vec![vec![1.0, 2.0], vec![3.0]]).unwrap_err();
assert!(matches!(err, DatarustError::ShapeMismatch { .. }));
}
#[test]
fn new_empty_rejected() {
assert!(Matrix::new(vec![]).is_err());
assert!(Matrix::new(vec![vec![]]).is_err());
}
#[test]
fn from_flat() {
let m = Matrix::from_flat(2, 3, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]).unwrap();
assert_eq!(m.get(0, 2), 3.0);
assert_eq!(m.get(1, 0), 4.0);
assert_eq!(m.get(1, 2), 6.0);
}
#[test]
fn from_flat_bad_count() {
let err = Matrix::from_flat(2, 2, vec![1.0, 2.0, 3.0]).unwrap_err();
assert!(matches!(err, DatarustError::ShapeMismatch { .. }));
}
#[test]
fn zeros_and_identity() {
let z = Matrix::zeros(2, 3).unwrap();
assert_eq!(z.get(1, 2), 0.0);
let id = Matrix::identity(3).unwrap();
assert_eq!(id.get(0, 0), 1.0);
assert_eq!(id.get(0, 1), 0.0);
assert_eq!(id.get(1, 1), 1.0);
}
#[test]
fn transpose() {
let m = Matrix::new(vec![vec![1.0, 2.0, 3.0], vec![4.0, 5.0, 6.0]]).unwrap();
let t = m.transpose();
assert_eq!(t.nrows(), 3);
assert_eq!(t.ncols(), 2);
assert_eq!(t.get(2, 0), 3.0);
assert_eq!(t.get(1, 1), 5.0);
}
#[test]
fn matmul() {
let a = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
let b = Matrix::new(vec![vec![5.0, 6.0], vec![7.0, 8.0]]).unwrap();
let c = a.matmul(&b).unwrap();
assert_eq!(c.get(0, 0), 19.0);
assert_eq!(c.get(0, 1), 22.0);
assert_eq!(c.get(1, 0), 43.0);
assert_eq!(c.get(1, 1), 50.0);
}
#[test]
fn matmul_shape_mismatch() {
let a = Matrix::new(vec![vec![1.0, 2.0, 3.0]]).unwrap();
let b = Matrix::new(vec![vec![1.0, 2.0]]).unwrap();
assert!(a.matmul(&b).is_err());
}
#[test]
fn from_columns() {
let m = Matrix::from_columns(vec![vec![1.0, 2.0], vec![10.0, 20.0]]).unwrap();
assert_eq!(m.nrows(), 2);
assert_eq!(m.ncols(), 2);
assert_eq!(m.get(0, 0), 1.0);
assert_eq!(m.get(1, 1), 20.0);
}
#[test]
fn strmatrix_from_column() {
let s = StrMatrix::from_column(["a", "b", "a"]).unwrap();
assert_eq!(s.nrows(), 3);
assert_eq!(s.ncols(), 1);
assert_eq!(s.get(2, 0), "a");
}
#[test]
fn strmatrix_from_strings() {
let s = StrMatrix::from_strings(vec![vec!["x", "y"], vec!["x", "z"]]).unwrap();
assert_eq!(s.ncols(), 2);
assert_eq!(s.get(1, 1), "z");
}
#[test]
fn sparse_from_triplets_basic() {
let sp = SparseMatrix::from_triplets(
3,
4,
&[(0, 0, 1.0), (1, 2, 3.0), (2, 3, 5.0), (0, 3, 7.0)],
)
.unwrap();
assert_eq!(sp.nrows(), 3);
assert_eq!(sp.ncols(), 4);
assert_eq!(sp.nnz(), 4);
assert_eq!(sp.get(0, 0), 1.0);
assert_eq!(sp.get(1, 2), 3.0);
assert_eq!(sp.get(0, 3), 7.0);
assert_eq!(sp.get(1, 0), 0.0);
}
#[test]
fn sparse_zero_triplets_dropped() {
let sp =
SparseMatrix::from_triplets(2, 2, &[(0, 0, 0.0), (0, 1, 5.0), (1, 0, 0.0)]).unwrap();
assert_eq!(sp.nnz(), 1);
assert_eq!(sp.get(0, 1), 5.0);
}
#[test]
fn sparse_to_dense() {
let sp = SparseMatrix::from_triplets(2, 3, &[(0, 1, 2.0), (1, 0, 4.0)]).unwrap();
let dense = sp.to_dense().unwrap();
assert_eq!(dense.row(0), [0.0, 2.0, 0.0]);
assert_eq!(dense.row(1), [4.0, 0.0, 0.0]);
}
#[test]
fn sparse_zeros() {
let sp = SparseMatrix::zeros(2, 3).unwrap();
assert_eq!(sp.nnz(), 0);
assert_eq!(sp.density(), 0.0);
assert_eq!(sp.get(0, 1), 0.0);
}
#[test]
fn sparse_density() {
let sp = SparseMatrix::from_triplets(2, 4, &[(0, 0, 1.0), (1, 3, 1.0)]).unwrap();
assert!((sp.density() - 0.25).abs() < 1e-12);
}
#[test]
fn sparse_row_nz() {
let sp =
SparseMatrix::from_triplets(2, 3, &[(0, 1, 2.0), (0, 2, 9.0), (1, 0, 4.0)]).unwrap();
let row0: Vec<(usize, f64)> = sp.row_nz(0).collect();
assert_eq!(row0, vec![(1, 2.0), (2, 9.0)]);
let row1: Vec<(usize, f64)> = sp.row_nz(1).collect();
assert_eq!(row1, vec![(0, 4.0)]);
}
#[test]
fn sparse_bad_indptr_rejected() {
let err = SparseMatrix::new(2, 2, vec![0, 0, 5], vec![0], vec![1.0]).unwrap_err();
assert!(matches!(err, DatarustError::InvalidInput(_)));
}
#[test]
fn sparse_col_out_of_range_rejected() {
let err = SparseMatrix::from_triplets(2, 2, &[(0, 5, 1.0)]).unwrap_err();
assert!(matches!(err, DatarustError::InvalidInput(_)));
}
#[test]
fn select_columns_basic() {
let m = Matrix::new(vec![vec![1.0, 2.0, 3.0], vec![4.0, 5.0, 6.0]]).unwrap();
let sub = m.select_columns(&[0, 2]).unwrap();
assert_eq!(sub.ncols(), 2);
assert_eq!(sub.get(0, 0), 1.0);
assert_eq!(sub.get(0, 1), 3.0);
assert_eq!(sub.get(1, 0), 4.0);
assert_eq!(sub.get(1, 1), 6.0);
}
#[test]
fn select_columns_out_of_range() {
let m = Matrix::new(vec![vec![1.0, 2.0]]).unwrap();
assert!(m.select_columns(&[0, 5]).is_err());
}
#[test]
fn select_columns_empty() {
let m = Matrix::new(vec![vec![1.0]]).unwrap();
assert!(m.select_columns(&[]).is_err());
}
#[test]
fn select_rows_basic() {
let m = Matrix::new(vec![vec![10.0], vec![20.0], vec![30.0]]).unwrap();
let sub = m.select_rows(&[0, 2]).unwrap();
assert_eq!(sub.nrows(), 2);
assert_eq!(sub.get(0, 0), 10.0);
assert_eq!(sub.get(1, 0), 30.0);
}
#[test]
fn select_rows_out_of_range() {
let m = Matrix::new(vec![vec![1.0]]).unwrap();
assert!(m.select_rows(&[0, 5]).is_err());
}
#[test]
fn select_rows_empty() {
let m = Matrix::new(vec![vec![1.0]]).unwrap();
assert!(m.select_rows(&[]).is_err());
}
#[test]
fn select_columns_reordered() {
let m = Matrix::new(vec![vec![1.0, 2.0, 3.0], vec![4.0, 5.0, 6.0]]).unwrap();
let sub = m.select_columns(&[2, 0]).unwrap();
assert_eq!(sub.get(0, 0), 3.0);
assert_eq!(sub.get(0, 1), 1.0);
}
#[test]
fn select_rows_duplicates() {
let m = Matrix::new(vec![vec![10.0], vec![20.0]]).unwrap();
let sub = m.select_rows(&[0, 0, 1]).unwrap();
assert_eq!(sub.nrows(), 3);
assert_eq!(sub.get(0, 0), 10.0);
assert_eq!(sub.get(1, 0), 10.0);
assert_eq!(sub.get(2, 0), 20.0);
}
#[test]
fn dense_accessors_and_row_conversions_preserve_layout() {
let mut m = Matrix::from_rows(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
assert_eq!(m.as_slice(), &[1.0, 2.0, 3.0, 4.0]);
assert_eq!(m.checked_get(1, 0), Some(3.0));
assert_eq!(m.checked_get(2, 0), None);
assert_eq!(m.checked_get(0, 2), None);
m.as_mut_slice()[1] = 20.0;
m.set(1, 0, 30.0);
assert_eq!(m.row(0), [1.0, 20.0]);
assert_eq!(m.col(0), vec![1.0, 30.0]);
assert_eq!(
m.iter_rows().collect::<Vec<_>>(),
vec![&[1.0, 20.0][..], &[30.0, 4.0][..]]
);
assert_eq!(m.rows_ref(), vec![vec![1.0, 20.0], vec![30.0, 4.0]]);
assert_eq!(
m.clone().into_rows(),
vec![vec![1.0, 20.0], vec![30.0, 4.0]]
);
assert!(m.validate_no_nan().is_ok());
m.set(1, 1, f64::NAN);
assert!(matches!(
m.validate_no_nan(),
Err(DatarustError::InvalidInput(_))
));
}
#[test]
fn dense_constructor_edge_cases_are_rejected() {
assert!(Matrix::zeros(0, 1).is_err());
assert!(Matrix::identity(0).is_err());
assert!(Matrix::from_flat(0, 1, vec![]).is_err());
assert!(Matrix::from_flat(usize::MAX, 2, vec![]).is_err());
assert!(Matrix::zeros(usize::MAX, 2).is_err());
assert!(Matrix::identity(usize::MAX).is_err());
assert!(Matrix::from_columns(vec![]).is_err());
assert!(Matrix::from_columns(vec![vec![]]).is_err());
assert!(Matrix::from_columns(vec![vec![1.0], vec![2.0, 3.0]]).is_err());
assert!(Matrix::try_from(vec![vec![1.0], vec![]]).is_err());
}
#[test]
fn finite_validators_distinguish_nan_from_infinity() {
let finite = Matrix::new(vec![vec![1.0, -2.0]]).unwrap();
assert!(finite.validate_finite().is_ok());
assert!(finite.validate_no_infinite().is_ok());
let nan = Matrix::new(vec![vec![f64::NAN]]).unwrap();
assert!(nan.validate_finite().is_err());
assert!(nan.validate_no_infinite().is_ok());
let infinite = Matrix::new(vec![vec![f64::INFINITY]]).unwrap();
assert!(infinite.validate_finite().is_err());
assert!(infinite.validate_no_infinite().is_err());
}
#[test]
fn strmatrix_accessors_and_validation_cover_edge_cases() {
let strings = StrMatrix::new(vec![
vec!["a".into(), "b".into()],
vec!["c".into(), "d".into()],
])
.unwrap();
assert_eq!(strings.checked_get(1, 1), Some("d"));
assert_eq!(strings.checked_get(2, 0), None);
assert_eq!(strings.checked_get(0, 2), None);
assert_eq!(strings.column(1), vec!["b".to_string(), "d".to_string()]);
assert_eq!(strings.row(0), ["a".to_string(), "b".to_string()]);
assert!(StrMatrix::new(vec![]).is_err());
assert!(StrMatrix::new(vec![vec![]]).is_err());
assert!(StrMatrix::new(vec![vec!["a".into()], vec!["b".into(), "c".into()]]).is_err());
assert!(StrMatrix::from_column(Vec::<String>::new()).is_err());
}
#[test]
fn raw_sparse_construction_and_bounds_checks() {
let sparse = SparseMatrix::new(2, 3, vec![0, 1, 2], vec![0, 2], vec![1.0, 3.0]).unwrap();
assert_eq!(sparse.checked_get(0, 0), Some(1.0));
assert_eq!(sparse.checked_get(0, 2), Some(0.0));
assert_eq!(sparse.checked_get(2, 0), None);
assert_eq!(sparse.checked_get(0, 3), None);
assert_eq!(
sparse.to_dense().unwrap().rows_ref(),
vec![vec![1.0, 0.0, 0.0], vec![0.0, 0.0, 3.0]]
);
assert!(SparseMatrix::new(0, 1, vec![0], vec![], vec![]).is_err());
assert!(SparseMatrix::new(2, 2, vec![0, 0], vec![], vec![]).is_err());
assert!(SparseMatrix::new(1, 2, vec![0, 1], vec![], vec![1.0]).is_err());
assert!(SparseMatrix::new(1, 2, vec![1, 1], vec![0], vec![1.0]).is_err());
assert!(SparseMatrix::new(1, 2, vec![0, 1], vec![2], vec![1.0]).is_err());
assert!(SparseMatrix::new(2, 2, vec![0, 2, 1], vec![0], vec![1.0]).is_err());
assert!(SparseMatrix::new(1, 3, vec![0, 2], vec![2, 1], vec![1.0, 2.0]).is_err());
assert!(SparseMatrix::new(1, 3, vec![0, 2], vec![1, 1], vec![1.0, 2.0]).is_err());
assert!(SparseMatrix::from_triplets(0, 1, &[]).is_err());
assert!(SparseMatrix::from_triplets(1, 1, &[(1, 0, 1.0)]).is_err());
}
#[test]
fn sparse_duplicate_triplets_are_summed_and_zero_sums_are_dropped() {
let sparse = SparseMatrix::from_triplets(
1,
3,
&[(0, 1, 2.0), (0, 1, 3.0), (0, 2, 4.0), (0, 2, -4.0)],
)
.unwrap();
assert_eq!(sparse.nnz(), 1);
assert_eq!(sparse.get(0, 1), 5.0);
assert_eq!(sparse.get(0, 2), 0.0);
}
#[test]
#[should_panic(expected = "Matrix::get")]
fn dense_get_panics_on_out_of_bounds_access() {
Matrix::new(vec![vec![1.0]]).unwrap().get(1, 0);
}
#[test]
#[should_panic(expected = "SparseMatrix::get")]
fn sparse_get_panics_on_out_of_bounds_access() {
SparseMatrix::zeros(1, 1).unwrap().get(0, 1);
}
#[cfg(feature = "serde")]
#[test]
fn serde_round_trips_dense_and_string_matrices_and_rejects_jagged_data() {
let matrix = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
let encoded = serde_json::to_string(&matrix).unwrap();
assert_eq!(serde_json::from_str::<Matrix>(&encoded).unwrap(), matrix);
assert!(serde_json::from_str::<Matrix>(r#"{"data":[[1.0],[2.0,3.0]]}"#).is_err());
let strings = StrMatrix::from_strings(vec![vec!["north", "south"]]).unwrap();
let encoded = serde_json::to_string(&strings).unwrap();
assert_eq!(
serde_json::from_str::<StrMatrix>(&encoded).unwrap(),
strings
);
assert!(
serde_json::from_str::<StrMatrix>(r#"{"data":[["north"],["south","east"]]}"#).is_err()
);
let sparse = SparseMatrix::from_triplets(1, 2, &[(0, 1, 3.0)]).unwrap();
let encoded = serde_json::to_string(&sparse).unwrap();
assert_eq!(
serde_json::from_str::<SparseMatrix>(&encoded).unwrap(),
sparse
);
assert!(serde_json::from_str::<SparseMatrix>(
r#"{"nrows":1,"ncols":3,"indptr":[0,2],"indices":[2,1],"data":[1.0,2.0]}"#
)
.is_err());
}
}