1use crate::error::{DatarustError, Result};
4
5fn checked_element_count(rows: usize, cols: usize) -> Result<usize> {
6 rows.checked_mul(cols)
7 .ok_or_else(|| DatarustError::ShapeMismatch {
8 expected: "rows * cols within usize range".into(),
9 actual: format!("{} rows × {} cols overflows usize", rows, cols),
10 })
11}
12
13#[derive(Debug, Clone, PartialEq)]
30pub struct Matrix {
31 data: Vec<f64>,
33 rows: usize,
34 cols: usize,
35}
36
37impl Matrix {
38 pub fn new(data: Vec<Vec<f64>>) -> Result<Self> {
40 if data.is_empty() {
41 return Err(DatarustError::EmptyInput("matrix has no rows".into()));
42 }
43 let cols = data[0].len();
44 if cols == 0 {
45 return Err(DatarustError::EmptyInput("matrix has no columns".into()));
46 }
47 for (i, row) in data.iter().enumerate() {
48 if row.len() != cols {
49 return Err(DatarustError::ShapeMismatch {
50 expected: format!("{} columns", cols),
51 actual: format!("{} columns at row {}", row.len(), i),
52 });
53 }
54 }
55 let rows = data.len();
56 let element_count = checked_element_count(rows, cols)?;
57 let mut flat = Vec::with_capacity(element_count);
58 for row in data {
59 flat.extend(row);
60 }
61 Ok(Self {
62 data: flat,
63 rows,
64 cols,
65 })
66 }
67
68 pub fn from_rows(rows: Vec<Vec<f64>>) -> Result<Self> {
70 Self::new(rows)
71 }
72
73 pub fn from_flat(rows: usize, cols: usize, flat: Vec<f64>) -> Result<Self> {
87 if rows == 0 || cols == 0 {
88 return Err(DatarustError::EmptyInput("zero dimension".into()));
89 }
90 let expected = checked_element_count(rows, cols)?;
91 if flat.len() != expected {
92 return Err(DatarustError::ShapeMismatch {
93 expected: format!("{} elements", expected),
94 actual: format!("{} elements", flat.len()),
95 });
96 }
97 Ok(Self {
98 data: flat,
99 rows,
100 cols,
101 })
102 }
103
104 pub fn zeros(rows: usize, cols: usize) -> Result<Self> {
106 if rows == 0 || cols == 0 {
107 return Err(DatarustError::EmptyInput("zero dimension".into()));
108 }
109 let element_count = checked_element_count(rows, cols)?;
110 Ok(Self {
111 data: vec![0.0; element_count],
112 rows,
113 cols,
114 })
115 }
116
117 pub fn identity(n: usize) -> Result<Self> {
119 if n == 0 {
120 return Err(DatarustError::EmptyInput("zero dimension".into()));
121 }
122 let element_count = checked_element_count(n, n)?;
123 let mut data = vec![0.0; element_count];
124 for i in 0..n {
125 data[i * n + i] = 1.0;
126 }
127 Ok(Self {
128 data,
129 rows: n,
130 cols: n,
131 })
132 }
133
134 #[inline]
136 pub fn nrows(&self) -> usize {
137 self.rows
138 }
139
140 #[inline]
142 pub fn ncols(&self) -> usize {
143 self.cols
144 }
145
146 #[inline]
152 pub fn as_slice(&self) -> &[f64] {
153 &self.data
154 }
155
156 #[inline]
158 pub fn as_mut_slice(&mut self) -> &mut [f64] {
159 &mut self.data
160 }
161
162 #[inline]
169 pub fn get(&self, i: usize, j: usize) -> f64 {
170 assert!(
171 i < self.rows && j < self.cols,
172 "Matrix::get: index ({}, {}) out of bounds for {}×{}",
173 i,
174 j,
175 self.rows,
176 self.cols
177 );
178 self.data[i * self.cols + j]
179 }
180
181 #[inline]
184 pub fn checked_get(&self, i: usize, j: usize) -> Option<f64> {
185 if i < self.rows && j < self.cols {
186 Some(self.data[i * self.cols + j])
187 } else {
188 None
189 }
190 }
191
192 #[inline]
194 pub fn set(&mut self, i: usize, j: usize, v: f64) {
195 self.data[i * self.cols + j] = v;
196 }
197
198 pub fn validate_no_nan(&self) -> Result<()> {
200 for (flat_idx, &v) in self.data.iter().enumerate() {
201 if v.is_nan() {
202 let i = flat_idx / self.cols;
203 let j = flat_idx % self.cols;
204 return Err(DatarustError::InvalidInput(format!(
205 "NaN value at position ({}, {})",
206 i, j
207 )));
208 }
209 }
210 Ok(())
211 }
212
213 pub fn validate_finite(&self) -> Result<()> {
220 for (flat_idx, &v) in self.data.iter().enumerate() {
221 if !v.is_finite() {
222 let i = flat_idx / self.cols;
223 let j = flat_idx % self.cols;
224 return Err(DatarustError::InvalidInput(format!(
225 "non-finite value {v} at position ({i}, {j})"
226 )));
227 }
228 }
229 Ok(())
230 }
231
232 pub fn validate_no_infinite(&self) -> Result<()> {
236 for (flat_idx, &v) in self.data.iter().enumerate() {
237 if v.is_infinite() {
238 let i = flat_idx / self.cols;
239 let j = flat_idx % self.cols;
240 return Err(DatarustError::InvalidInput(format!(
241 "infinite value {v} at position ({i}, {j})"
242 )));
243 }
244 }
245 Ok(())
246 }
247
248 #[inline]
257 pub fn row(&self, i: usize) -> &[f64] {
258 let start = i * self.cols;
259 &self.data[start..start + self.cols]
260 }
261
262 pub fn col(&self, j: usize) -> Vec<f64> {
271 (0..self.rows)
272 .map(|i| self.data[i * self.cols + j])
273 .collect()
274 }
275
276 pub fn iter_rows(&self) -> impl Iterator<Item = &[f64]> {
278 self.data.chunks_exact(self.cols)
279 }
280
281 #[doc(hidden)]
288 pub fn rows_ref(&self) -> Vec<Vec<f64>> {
289 self.data
290 .chunks_exact(self.cols)
291 .map(|chunk| chunk.to_vec())
292 .collect()
293 }
294
295 #[doc(hidden)]
297 pub fn into_rows(self) -> Vec<Vec<f64>> {
298 let cols = self.cols;
299 self.data.chunks(cols).map(|chunk| chunk.to_vec()).collect()
300 }
301
302 pub fn transpose(&self) -> Matrix {
314 let rows = self.rows;
315 let cols = self.cols;
316 let mut out = vec![0.0; rows * cols];
317 for i in 0..rows {
318 for j in 0..cols {
319 out[j * rows + i] = self.data[i * cols + j];
320 }
321 }
322 Matrix {
323 data: out,
324 rows: cols,
325 cols: rows,
326 }
327 }
328
329 #[allow(clippy::needless_range_loop)]
341 pub fn matmul(&self, other: &Matrix) -> Result<Matrix> {
342 if self.cols != other.rows {
343 return Err(DatarustError::ShapeMismatch {
344 expected: format!("second operand with {} rows", self.cols),
345 actual: format!("{} rows", other.rows),
346 });
347 }
348 let m = self.rows;
349 let k = self.cols;
350 let n = other.cols;
351 let element_count = checked_element_count(m, n)?;
352 let mut out = vec![0.0; element_count];
353
354 #[cfg(feature = "matrixmultiply")]
355 {
356 unsafe {
359 matrixmultiply::dgemm(
360 m,
361 k,
362 n,
363 1.0,
364 self.data.as_ptr(),
365 k as isize,
366 1,
367 other.data.as_ptr(),
368 n as isize,
369 1,
370 0.0,
371 out.as_mut_ptr(),
372 n as isize,
373 1,
374 );
375 }
376 }
377
378 #[cfg(not(feature = "matrixmultiply"))]
379 {
380 for i in 0..m {
382 let out_base = i * n;
383 let self_base = i * k;
384 for l in 0..k {
385 let a = self.data[self_base + l];
386 if a == 0.0 {
387 continue;
388 }
389 let other_base = l * n;
390 for j in 0..n {
391 out[out_base + j] += a * other.data[other_base + j];
392 }
393 }
394 }
395 }
396 Ok(Matrix {
397 data: out,
398 rows: m,
399 cols: n,
400 })
401 }
402
403 pub fn column_mean(&self) -> Vec<f64> {
405 crate::stats::column_mean_flat(&self.data, self.rows, self.cols)
406 }
407
408 pub fn from_columns(cols: Vec<Vec<f64>>) -> Result<Self> {
410 if cols.is_empty() || cols[0].is_empty() {
411 return Err(DatarustError::EmptyInput("no columns".into()));
412 }
413 let rows = cols[0].len();
414 for c in &cols {
415 if c.len() != rows {
416 return Err(DatarustError::ShapeMismatch {
417 expected: format!("{} rows", rows),
418 actual: format!("{} rows", c.len()),
419 });
420 }
421 }
422 let ncols = cols.len();
423 let element_count = checked_element_count(rows, ncols)?;
424 let mut data = vec![0.0; element_count];
425 for (j, col) in cols.iter().enumerate() {
426 for (i, &v) in col.iter().enumerate() {
427 data[i * ncols + j] = v;
428 }
429 }
430 Ok(Self {
431 data,
432 rows,
433 cols: ncols,
434 })
435 }
436
437 pub fn select_columns(&self, indices: &[usize]) -> Result<Self> {
453 if indices.is_empty() {
454 return Err(DatarustError::EmptyInput("no columns selected".into()));
455 }
456 let ncols = self.cols;
457 for &c in indices {
458 if c >= ncols {
459 return Err(DatarustError::InvalidInput(format!(
460 "column index {} out of range (ncols {})",
461 c, ncols
462 )));
463 }
464 }
465 let out_cols = indices.len();
466 let element_count = checked_element_count(self.rows, out_cols)?;
467 let mut out = Vec::with_capacity(element_count);
468 for i in 0..self.rows {
469 let base = i * ncols;
470 for &c in indices {
471 out.push(self.data[base + c]);
472 }
473 }
474 Ok(Self {
475 data: out,
476 rows: self.rows,
477 cols: out_cols,
478 })
479 }
480
481 pub fn select_rows(&self, indices: &[usize]) -> Result<Self> {
498 if indices.is_empty() {
499 return Err(DatarustError::EmptyInput("no rows selected".into()));
500 }
501 let nrows = self.rows;
502 for &r in indices {
503 if r >= nrows {
504 return Err(DatarustError::InvalidInput(format!(
505 "row index {} out of range (nrows {})",
506 r, nrows
507 )));
508 }
509 }
510 let element_count = checked_element_count(indices.len(), self.cols)?;
511 let mut out = Vec::with_capacity(element_count);
512 for &r in indices {
513 let start = r * self.cols;
514 out.extend_from_slice(&self.data[start..start + self.cols]);
515 }
516 Ok(Self {
517 data: out,
518 rows: indices.len(),
519 cols: self.cols,
520 })
521 }
522}
523
524#[cfg(feature = "serde")]
525impl serde::Serialize for Matrix {
526 fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
527 where
528 S: serde::Serializer,
529 {
530 use serde::ser::SerializeStruct;
531 let rows: Vec<&[f64]> = self.data.chunks_exact(self.cols).collect();
534 let mut s = serializer.serialize_struct("Matrix", 1)?;
535 s.serialize_field("data", &rows)?;
536 s.end()
537 }
538}
539
540#[cfg(feature = "serde")]
541impl<'de> serde::Deserialize<'de> for Matrix {
542 fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
543 where
544 D: serde::Deserializer<'de>,
545 {
546 #[derive(serde::Deserialize)]
547 struct Raw {
548 data: Vec<Vec<f64>>,
549 }
550 let raw = Raw::deserialize(deserializer)?;
551 Matrix::new(raw.data).map_err(serde::de::Error::custom)
552 }
553}
554
555#[derive(Debug, Clone, PartialEq)]
557pub struct StrMatrix {
558 pub(crate) data: Vec<Vec<String>>,
559}
560
561impl StrMatrix {
562 pub fn new(data: Vec<Vec<String>>) -> Result<Self> {
564 if data.is_empty() {
565 return Err(DatarustError::EmptyInput("matrix has no rows".into()));
566 }
567 let cols = data[0].len();
568 if cols == 0 {
569 return Err(DatarustError::EmptyInput("matrix has no columns".into()));
570 }
571 for (i, row) in data.iter().enumerate() {
572 if row.len() != cols {
573 return Err(DatarustError::ShapeMismatch {
574 expected: format!("{} columns", cols),
575 actual: format!("{} columns at row {}", row.len(), i),
576 });
577 }
578 }
579 Ok(Self { data })
580 }
581
582 pub fn from_column<I, S>(col: I) -> Result<Self>
584 where
585 I: IntoIterator<Item = S>,
586 S: Into<String>,
587 {
588 let data: Vec<Vec<String>> = col.into_iter().map(|s| vec![s.into()]).collect::<Vec<_>>();
589 if data.is_empty() {
590 return Err(DatarustError::EmptyInput("column has no rows".into()));
591 }
592 Self::new(data)
593 }
594
595 pub fn from_strings<I, S>(rows: I) -> Result<Self>
597 where
598 I: IntoIterator<Item = Vec<S>>,
599 S: Into<String>,
600 {
601 let data: Vec<Vec<String>> = rows
602 .into_iter()
603 .map(|r| r.into_iter().map(|s| s.into()).collect())
604 .collect();
605 Self::new(data)
606 }
607
608 #[inline]
610 pub fn nrows(&self) -> usize {
611 self.data.len()
612 }
613
614 #[inline]
616 pub fn ncols(&self) -> usize {
617 self.data[0].len()
618 }
619
620 #[inline]
627 pub fn get(&self, i: usize, j: usize) -> &str {
628 debug_assert!(
629 i < self.data.len() && j < self.data[i].len(),
630 "StrMatrix::get: index ({}, {}) out of bounds for {}×{}",
631 i,
632 j,
633 self.data.len(),
634 self.data[0].len()
635 );
636 &self.data[i][j]
637 }
638
639 #[inline]
642 pub fn checked_get(&self, i: usize, j: usize) -> Option<&str> {
643 self.data.get(i)?.get(j).map(|s| s.as_str())
644 }
645
646 pub fn column(&self, j: usize) -> Vec<String> {
648 self.data.iter().map(|r| r[j].clone()).collect()
649 }
650
651 pub fn column_refs(&self, j: usize) -> Vec<&str> {
656 self.data.iter().map(|r| r[j].as_str()).collect()
657 }
658
659 pub fn row(&self, i: usize) -> &[String] {
661 &self.data[i]
662 }
663}
664
665#[cfg(feature = "serde")]
666impl serde::Serialize for StrMatrix {
667 fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
668 where
669 S: serde::Serializer,
670 {
671 use serde::ser::SerializeStruct;
672 let mut s = serializer.serialize_struct("StrMatrix", 1)?;
673 s.serialize_field("data", &self.data)?;
674 s.end()
675 }
676}
677
678#[cfg(feature = "serde")]
679impl<'de> serde::Deserialize<'de> for StrMatrix {
680 fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
681 where
682 D: serde::Deserializer<'de>,
683 {
684 #[derive(serde::Deserialize)]
685 struct Raw {
686 data: Vec<Vec<String>>,
687 }
688 let raw = Raw::deserialize(deserializer)?;
689 StrMatrix::new(raw.data).map_err(serde::de::Error::custom)
690 }
691}
692
693impl TryFrom<Vec<Vec<f64>>> for Matrix {
694 type Error = DatarustError;
695
696 fn try_from(data: Vec<Vec<f64>>) -> Result<Self> {
702 Matrix::new(data)
703 }
704}
705
706#[derive(Debug, Clone, PartialEq)]
715#[cfg_attr(feature = "serde", derive(serde::Serialize))]
716pub struct SparseMatrix {
717 nrows: usize,
718 ncols: usize,
719 indptr: Vec<usize>,
720 indices: Vec<usize>,
721 data: Vec<f64>,
722}
723
724impl SparseMatrix {
725 pub fn new(
730 nrows: usize,
731 ncols: usize,
732 indptr: Vec<usize>,
733 indices: Vec<usize>,
734 data: Vec<f64>,
735 ) -> Result<Self> {
736 if nrows == 0 || ncols == 0 {
737 return Err(DatarustError::EmptyInput("zero dimension".into()));
738 }
739 let expected_indptr_len =
740 nrows
741 .checked_add(1)
742 .ok_or_else(|| DatarustError::ShapeMismatch {
743 expected: "nrows + 1 within usize range".into(),
744 actual: format!("{} rows overflows the CSR indptr length", nrows),
745 })?;
746 if indptr.len() != expected_indptr_len {
747 return Err(DatarustError::ShapeMismatch {
748 expected: format!("{} indptr entries", expected_indptr_len),
749 actual: format!("{} indptr entries", indptr.len()),
750 });
751 }
752 if indices.len() != data.len() {
753 return Err(DatarustError::ShapeMismatch {
754 expected: format!("{} indices", data.len()),
755 actual: format!("{} indices", indices.len()),
756 });
757 }
758 let nnz = data.len();
759 if indptr[0] != 0 || indptr[nrows] != nnz {
760 return Err(DatarustError::InvalidInput(
761 "indptr must start at 0 and end at nnz".into(),
762 ));
763 }
764 for row in 0..nrows {
765 let start = indptr[row];
766 let end = indptr[row + 1];
767 if start > end || end > nnz {
768 return Err(DatarustError::InvalidInput(format!(
769 "indptr must be non-decreasing and stay within nnz (row {})",
770 row
771 )));
772 }
773 let row_indices = &indices[start..end];
774 for &c in row_indices {
775 if c >= ncols {
776 return Err(DatarustError::InvalidInput(format!(
777 "column index {} out of range (ncols {})",
778 c, ncols
779 )));
780 }
781 }
782 if row_indices.windows(2).any(|pair| pair[0] >= pair[1]) {
783 return Err(DatarustError::InvalidInput(format!(
784 "column indices in row {} must be strictly increasing",
785 row
786 )));
787 }
788 }
789 Ok(Self {
790 nrows,
791 ncols,
792 indptr,
793 indices,
794 data,
795 })
796 }
797
798 pub fn from_triplets(
802 nrows: usize,
803 ncols: usize,
804 triplets: &[(usize, usize, f64)],
805 ) -> Result<Self> {
806 if nrows == 0 || ncols == 0 {
807 return Err(DatarustError::EmptyInput("zero dimension".into()));
808 }
809 let indptr_len = nrows
810 .checked_add(1)
811 .ok_or_else(|| DatarustError::ShapeMismatch {
812 expected: "nrows + 1 within usize range".into(),
813 actual: format!("{} rows overflows the CSR indptr length", nrows),
814 })?;
815 let mut per_row: Vec<Vec<(usize, f64)>> = vec![vec![]; nrows];
816 for &(r, c, v) in triplets {
817 if r >= nrows {
818 return Err(DatarustError::InvalidInput(format!(
819 "row {} out of range (nrows {})",
820 r, nrows
821 )));
822 }
823 if c >= ncols {
824 return Err(DatarustError::InvalidInput(format!(
825 "col {} out of range (ncols {})",
826 c, ncols
827 )));
828 }
829 if v != 0.0 {
830 per_row[r].push((c, v));
831 }
832 }
833 let mut indptr = Vec::with_capacity(indptr_len);
834 let mut indices = Vec::new();
835 let mut data = Vec::new();
836 indptr.push(0);
837 for row_entries in &mut per_row {
838 row_entries.sort_by_key(|(c, _)| *c);
839 let mut position = 0;
840 while position < row_entries.len() {
841 let column = row_entries[position].0;
842 let mut value = 0.0;
843 while position < row_entries.len() && row_entries[position].0 == column {
844 value += row_entries[position].1;
845 position += 1;
846 }
847 if value != 0.0 {
848 indices.push(column);
849 data.push(value);
850 }
851 }
852 indptr.push(indices.len());
853 }
854 Ok(Self {
855 nrows,
856 ncols,
857 indptr,
858 indices,
859 data,
860 })
861 }
862
863 pub fn zeros(nrows: usize, ncols: usize) -> Result<Self> {
865 if nrows == 0 || ncols == 0 {
866 return Err(DatarustError::EmptyInput("zero dimension".into()));
867 }
868 let indptr_len = nrows
869 .checked_add(1)
870 .ok_or_else(|| DatarustError::ShapeMismatch {
871 expected: "nrows + 1 within usize range".into(),
872 actual: format!("{} rows overflows the CSR indptr length", nrows),
873 })?;
874 Ok(Self {
875 nrows,
876 ncols,
877 indptr: vec![0; indptr_len],
878 indices: vec![],
879 data: vec![],
880 })
881 }
882
883 #[inline]
885 pub fn nrows(&self) -> usize {
886 self.nrows
887 }
888
889 #[inline]
891 pub fn ncols(&self) -> usize {
892 self.ncols
893 }
894
895 #[inline]
897 pub fn nnz(&self) -> usize {
898 self.data.len()
899 }
900
901 pub fn density(&self) -> f64 {
903 let total = self.nrows.saturating_mul(self.ncols);
904 if total == 0 {
905 return 0.0;
906 }
907 self.nnz() as f64 / total as f64
908 }
909
910 pub fn get(&self, i: usize, j: usize) -> f64 {
917 assert!(
918 i < self.nrows && j < self.ncols,
919 "SparseMatrix::get: index ({}, {}) out of bounds for {}×{}",
920 i,
921 j,
922 self.nrows,
923 self.ncols
924 );
925 let start = self.indptr[i];
926 let end = self.indptr[i + 1];
927 let slice = &self.indices[start..end];
929 match slice.binary_search(&j) {
930 Ok(local) => self.data[start + local],
931 Err(_) => 0.0,
932 }
933 }
934
935 pub fn checked_get(&self, i: usize, j: usize) -> Option<f64> {
938 if i >= self.nrows || j >= self.ncols {
939 return None;
940 }
941 let start = *self.indptr.get(i)?;
942 let end = *self.indptr.get(i + 1)?;
943 let slice = self.indices.get(start..end)?;
944 match slice.binary_search(&j) {
945 Ok(local) => Some(self.data[start + local]),
946 Err(_) => Some(0.0),
947 }
948 }
949
950 pub fn row_nz(&self, i: usize) -> impl Iterator<Item = (usize, f64)> + '_ {
956 assert!(
957 i < self.nrows,
958 "SparseMatrix::row_nz: row {} out of bounds for {} rows",
959 i,
960 self.nrows
961 );
962 let start = self.indptr[i];
963 let end = self.indptr[i + 1];
964 self.indices[start..end]
965 .iter()
966 .zip(self.data[start..end].iter())
967 .map(|(&c, &v)| (c, v))
968 }
969
970 pub fn to_dense(&self) -> Result<Matrix> {
972 let mut rows = vec![vec![0.0; self.ncols]; self.nrows];
973 for (i, row) in rows.iter_mut().enumerate() {
974 for (c, v) in self.row_nz(i) {
975 row[c] = v;
976 }
977 }
978 Matrix::new(rows)
979 }
980}
981
982#[cfg(feature = "serde")]
983impl<'de> serde::Deserialize<'de> for SparseMatrix {
984 fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
985 where
986 D: serde::Deserializer<'de>,
987 {
988 #[derive(serde::Deserialize)]
989 struct Raw {
990 nrows: usize,
991 ncols: usize,
992 indptr: Vec<usize>,
993 indices: Vec<usize>,
994 data: Vec<f64>,
995 }
996
997 let raw = Raw::deserialize(deserializer)?;
998 SparseMatrix::new(raw.nrows, raw.ncols, raw.indptr, raw.indices, raw.data)
999 .map_err(serde::de::Error::custom)
1000 }
1001}
1002
1003#[cfg(test)]
1004#[allow(dead_code)]
1005pub(crate) fn approx_eq_matrices(a: &Matrix, b: &Matrix, tol: f64) -> bool {
1006 if a.nrows() != b.nrows() || a.ncols() != b.ncols() {
1007 return false;
1008 }
1009 for i in 0..a.nrows() {
1010 for j in 0..a.ncols() {
1011 if (a.get(i, j) - b.get(i, j)).abs() > tol {
1012 return false;
1013 }
1014 }
1015 }
1016 true
1017}
1018
1019#[cfg(test)]
1020#[allow(dead_code)]
1021pub(crate) fn approx_eq_vecs(a: &[f64], b: &[f64], tol: f64) -> bool {
1022 if a.len() != b.len() {
1023 return false;
1024 }
1025 a.iter().zip(b.iter()).all(|(x, y)| (x - y).abs() <= tol)
1026}
1027
1028#[cfg(test)]
1029#[macro_use]
1030mod assert_macros {
1031 #[allow(unused_macros)]
1032 macro_rules! assert_mat_eq {
1033 ($a:expr, $b:expr, $tol:expr) => {{
1034 assert!(
1035 $crate::matrix::approx_eq_matrices(&$a, &$b, $tol),
1036 "matrices not equal within tolerance {}\n left: {:?}\nright: {:?}",
1037 $tol,
1038 $a.as_slice(),
1039 $b.as_slice()
1040 );
1041 }};
1042 }
1043}
1044
1045#[cfg(test)]
1046mod tests {
1047 use super::*;
1048
1049 #[test]
1050 fn new_valid() {
1051 let m = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
1052 assert_eq!(m.nrows(), 2);
1053 assert_eq!(m.ncols(), 2);
1054 }
1055
1056 #[test]
1057 fn new_jagged_rejected() {
1058 let err = Matrix::new(vec![vec![1.0, 2.0], vec![3.0]]).unwrap_err();
1059 assert!(matches!(err, DatarustError::ShapeMismatch { .. }));
1060 }
1061
1062 #[test]
1063 fn new_empty_rejected() {
1064 assert!(Matrix::new(vec![]).is_err());
1065 assert!(Matrix::new(vec![vec![]]).is_err());
1066 }
1067
1068 #[test]
1069 fn from_flat() {
1070 let m = Matrix::from_flat(2, 3, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]).unwrap();
1071 assert_eq!(m.get(0, 2), 3.0);
1072 assert_eq!(m.get(1, 0), 4.0);
1073 assert_eq!(m.get(1, 2), 6.0);
1074 }
1075
1076 #[test]
1077 fn from_flat_bad_count() {
1078 let err = Matrix::from_flat(2, 2, vec![1.0, 2.0, 3.0]).unwrap_err();
1079 assert!(matches!(err, DatarustError::ShapeMismatch { .. }));
1080 }
1081
1082 #[test]
1083 fn zeros_and_identity() {
1084 let z = Matrix::zeros(2, 3).unwrap();
1085 assert_eq!(z.get(1, 2), 0.0);
1086 let id = Matrix::identity(3).unwrap();
1087 assert_eq!(id.get(0, 0), 1.0);
1088 assert_eq!(id.get(0, 1), 0.0);
1089 assert_eq!(id.get(1, 1), 1.0);
1090 }
1091
1092 #[test]
1093 fn transpose() {
1094 let m = Matrix::new(vec![vec![1.0, 2.0, 3.0], vec![4.0, 5.0, 6.0]]).unwrap();
1095 let t = m.transpose();
1096 assert_eq!(t.nrows(), 3);
1097 assert_eq!(t.ncols(), 2);
1098 assert_eq!(t.get(2, 0), 3.0);
1099 assert_eq!(t.get(1, 1), 5.0);
1100 }
1101
1102 #[test]
1103 fn matmul() {
1104 let a = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
1105 let b = Matrix::new(vec![vec![5.0, 6.0], vec![7.0, 8.0]]).unwrap();
1106 let c = a.matmul(&b).unwrap();
1107 assert_eq!(c.get(0, 0), 19.0);
1109 assert_eq!(c.get(0, 1), 22.0);
1110 assert_eq!(c.get(1, 0), 43.0);
1111 assert_eq!(c.get(1, 1), 50.0);
1112 }
1113
1114 #[test]
1115 fn matmul_shape_mismatch() {
1116 let a = Matrix::new(vec![vec![1.0, 2.0, 3.0]]).unwrap();
1117 let b = Matrix::new(vec![vec![1.0, 2.0]]).unwrap();
1118 assert!(a.matmul(&b).is_err());
1119 }
1120
1121 #[test]
1122 fn from_columns() {
1123 let m = Matrix::from_columns(vec![vec![1.0, 2.0], vec![10.0, 20.0]]).unwrap();
1124 assert_eq!(m.nrows(), 2);
1125 assert_eq!(m.ncols(), 2);
1126 assert_eq!(m.get(0, 0), 1.0);
1127 assert_eq!(m.get(1, 1), 20.0);
1128 }
1129
1130 #[test]
1131 fn strmatrix_from_column() {
1132 let s = StrMatrix::from_column(["a", "b", "a"]).unwrap();
1133 assert_eq!(s.nrows(), 3);
1134 assert_eq!(s.ncols(), 1);
1135 assert_eq!(s.get(2, 0), "a");
1136 }
1137
1138 #[test]
1139 fn strmatrix_from_strings() {
1140 let s = StrMatrix::from_strings(vec![vec!["x", "y"], vec!["x", "z"]]).unwrap();
1141 assert_eq!(s.ncols(), 2);
1142 assert_eq!(s.get(1, 1), "z");
1143 }
1144
1145 #[test]
1146 fn sparse_from_triplets_basic() {
1147 let sp = SparseMatrix::from_triplets(
1148 3,
1149 4,
1150 &[(0, 0, 1.0), (1, 2, 3.0), (2, 3, 5.0), (0, 3, 7.0)],
1151 )
1152 .unwrap();
1153 assert_eq!(sp.nrows(), 3);
1154 assert_eq!(sp.ncols(), 4);
1155 assert_eq!(sp.nnz(), 4);
1156 assert_eq!(sp.get(0, 0), 1.0);
1157 assert_eq!(sp.get(1, 2), 3.0);
1158 assert_eq!(sp.get(0, 3), 7.0);
1159 assert_eq!(sp.get(1, 0), 0.0);
1160 }
1161
1162 #[test]
1163 fn sparse_zero_triplets_dropped() {
1164 let sp =
1165 SparseMatrix::from_triplets(2, 2, &[(0, 0, 0.0), (0, 1, 5.0), (1, 0, 0.0)]).unwrap();
1166 assert_eq!(sp.nnz(), 1);
1167 assert_eq!(sp.get(0, 1), 5.0);
1168 }
1169
1170 #[test]
1171 fn sparse_to_dense() {
1172 let sp = SparseMatrix::from_triplets(2, 3, &[(0, 1, 2.0), (1, 0, 4.0)]).unwrap();
1173 let dense = sp.to_dense().unwrap();
1174 assert_eq!(dense.row(0), [0.0, 2.0, 0.0]);
1175 assert_eq!(dense.row(1), [4.0, 0.0, 0.0]);
1176 }
1177
1178 #[test]
1179 fn sparse_zeros() {
1180 let sp = SparseMatrix::zeros(2, 3).unwrap();
1181 assert_eq!(sp.nnz(), 0);
1182 assert_eq!(sp.density(), 0.0);
1183 assert_eq!(sp.get(0, 1), 0.0);
1184 }
1185
1186 #[test]
1187 fn sparse_density() {
1188 let sp = SparseMatrix::from_triplets(2, 4, &[(0, 0, 1.0), (1, 3, 1.0)]).unwrap();
1189 assert!((sp.density() - 0.25).abs() < 1e-12);
1190 }
1191
1192 #[test]
1193 fn sparse_row_nz() {
1194 let sp =
1195 SparseMatrix::from_triplets(2, 3, &[(0, 1, 2.0), (0, 2, 9.0), (1, 0, 4.0)]).unwrap();
1196 let row0: Vec<(usize, f64)> = sp.row_nz(0).collect();
1197 assert_eq!(row0, vec![(1, 2.0), (2, 9.0)]);
1198 let row1: Vec<(usize, f64)> = sp.row_nz(1).collect();
1199 assert_eq!(row1, vec![(0, 4.0)]);
1200 }
1201
1202 #[test]
1203 fn sparse_bad_indptr_rejected() {
1204 let err = SparseMatrix::new(2, 2, vec![0, 0, 5], vec![0], vec![1.0]).unwrap_err();
1205 assert!(matches!(err, DatarustError::InvalidInput(_)));
1206 }
1207
1208 #[test]
1209 fn sparse_col_out_of_range_rejected() {
1210 let err = SparseMatrix::from_triplets(2, 2, &[(0, 5, 1.0)]).unwrap_err();
1211 assert!(matches!(err, DatarustError::InvalidInput(_)));
1212 }
1213
1214 #[test]
1215 fn select_columns_basic() {
1216 let m = Matrix::new(vec![vec![1.0, 2.0, 3.0], vec![4.0, 5.0, 6.0]]).unwrap();
1217 let sub = m.select_columns(&[0, 2]).unwrap();
1218 assert_eq!(sub.ncols(), 2);
1219 assert_eq!(sub.get(0, 0), 1.0);
1220 assert_eq!(sub.get(0, 1), 3.0);
1221 assert_eq!(sub.get(1, 0), 4.0);
1222 assert_eq!(sub.get(1, 1), 6.0);
1223 }
1224
1225 #[test]
1226 fn select_columns_out_of_range() {
1227 let m = Matrix::new(vec![vec![1.0, 2.0]]).unwrap();
1228 assert!(m.select_columns(&[0, 5]).is_err());
1229 }
1230
1231 #[test]
1232 fn select_columns_empty() {
1233 let m = Matrix::new(vec![vec![1.0]]).unwrap();
1234 assert!(m.select_columns(&[]).is_err());
1235 }
1236
1237 #[test]
1238 fn select_rows_basic() {
1239 let m = Matrix::new(vec![vec![10.0], vec![20.0], vec![30.0]]).unwrap();
1240 let sub = m.select_rows(&[0, 2]).unwrap();
1241 assert_eq!(sub.nrows(), 2);
1242 assert_eq!(sub.get(0, 0), 10.0);
1243 assert_eq!(sub.get(1, 0), 30.0);
1244 }
1245
1246 #[test]
1247 fn select_rows_out_of_range() {
1248 let m = Matrix::new(vec![vec![1.0]]).unwrap();
1249 assert!(m.select_rows(&[0, 5]).is_err());
1250 }
1251
1252 #[test]
1253 fn select_rows_empty() {
1254 let m = Matrix::new(vec![vec![1.0]]).unwrap();
1255 assert!(m.select_rows(&[]).is_err());
1256 }
1257
1258 #[test]
1259 fn select_columns_reordered() {
1260 let m = Matrix::new(vec![vec![1.0, 2.0, 3.0], vec![4.0, 5.0, 6.0]]).unwrap();
1261 let sub = m.select_columns(&[2, 0]).unwrap();
1262 assert_eq!(sub.get(0, 0), 3.0);
1263 assert_eq!(sub.get(0, 1), 1.0);
1264 }
1265
1266 #[test]
1267 fn select_rows_duplicates() {
1268 let m = Matrix::new(vec![vec![10.0], vec![20.0]]).unwrap();
1269 let sub = m.select_rows(&[0, 0, 1]).unwrap();
1270 assert_eq!(sub.nrows(), 3);
1271 assert_eq!(sub.get(0, 0), 10.0);
1272 assert_eq!(sub.get(1, 0), 10.0);
1273 assert_eq!(sub.get(2, 0), 20.0);
1274 }
1275
1276 #[test]
1277 fn dense_accessors_and_row_conversions_preserve_layout() {
1278 let mut m = Matrix::from_rows(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
1279
1280 assert_eq!(m.as_slice(), &[1.0, 2.0, 3.0, 4.0]);
1281 assert_eq!(m.checked_get(1, 0), Some(3.0));
1282 assert_eq!(m.checked_get(2, 0), None);
1283 assert_eq!(m.checked_get(0, 2), None);
1284 m.as_mut_slice()[1] = 20.0;
1285 m.set(1, 0, 30.0);
1286 assert_eq!(m.row(0), [1.0, 20.0]);
1287 assert_eq!(m.col(0), vec![1.0, 30.0]);
1288 assert_eq!(
1289 m.iter_rows().collect::<Vec<_>>(),
1290 vec![&[1.0, 20.0][..], &[30.0, 4.0][..]]
1291 );
1292 assert_eq!(m.rows_ref(), vec![vec![1.0, 20.0], vec![30.0, 4.0]]);
1293 assert_eq!(
1294 m.clone().into_rows(),
1295 vec![vec![1.0, 20.0], vec![30.0, 4.0]]
1296 );
1297 assert!(m.validate_no_nan().is_ok());
1298
1299 m.set(1, 1, f64::NAN);
1300 assert!(matches!(
1301 m.validate_no_nan(),
1302 Err(DatarustError::InvalidInput(_))
1303 ));
1304 }
1305
1306 #[test]
1307 fn dense_constructor_edge_cases_are_rejected() {
1308 assert!(Matrix::zeros(0, 1).is_err());
1309 assert!(Matrix::identity(0).is_err());
1310 assert!(Matrix::from_flat(0, 1, vec![]).is_err());
1311 assert!(Matrix::from_flat(usize::MAX, 2, vec![]).is_err());
1312 assert!(Matrix::zeros(usize::MAX, 2).is_err());
1313 assert!(Matrix::identity(usize::MAX).is_err());
1314 assert!(Matrix::from_columns(vec![]).is_err());
1315 assert!(Matrix::from_columns(vec![vec![]]).is_err());
1316 assert!(Matrix::from_columns(vec![vec![1.0], vec![2.0, 3.0]]).is_err());
1317 assert!(Matrix::try_from(vec![vec![1.0], vec![]]).is_err());
1318 }
1319
1320 #[test]
1321 fn finite_validators_distinguish_nan_from_infinity() {
1322 let finite = Matrix::new(vec![vec![1.0, -2.0]]).unwrap();
1323 assert!(finite.validate_finite().is_ok());
1324 assert!(finite.validate_no_infinite().is_ok());
1325
1326 let nan = Matrix::new(vec![vec![f64::NAN]]).unwrap();
1327 assert!(nan.validate_finite().is_err());
1328 assert!(nan.validate_no_infinite().is_ok());
1329
1330 let infinite = Matrix::new(vec![vec![f64::INFINITY]]).unwrap();
1331 assert!(infinite.validate_finite().is_err());
1332 assert!(infinite.validate_no_infinite().is_err());
1333 }
1334
1335 #[test]
1336 fn strmatrix_accessors_and_validation_cover_edge_cases() {
1337 let strings = StrMatrix::new(vec![
1338 vec!["a".into(), "b".into()],
1339 vec!["c".into(), "d".into()],
1340 ])
1341 .unwrap();
1342 assert_eq!(strings.checked_get(1, 1), Some("d"));
1343 assert_eq!(strings.checked_get(2, 0), None);
1344 assert_eq!(strings.checked_get(0, 2), None);
1345 assert_eq!(strings.column(1), vec!["b".to_string(), "d".to_string()]);
1346 assert_eq!(strings.row(0), ["a".to_string(), "b".to_string()]);
1347 assert!(StrMatrix::new(vec![]).is_err());
1348 assert!(StrMatrix::new(vec![vec![]]).is_err());
1349 assert!(StrMatrix::new(vec![vec!["a".into()], vec!["b".into(), "c".into()]]).is_err());
1350 assert!(StrMatrix::from_column(Vec::<String>::new()).is_err());
1351 }
1352
1353 #[test]
1354 fn raw_sparse_construction_and_bounds_checks() {
1355 let sparse = SparseMatrix::new(2, 3, vec![0, 1, 2], vec![0, 2], vec![1.0, 3.0]).unwrap();
1356 assert_eq!(sparse.checked_get(0, 0), Some(1.0));
1357 assert_eq!(sparse.checked_get(0, 2), Some(0.0));
1358 assert_eq!(sparse.checked_get(2, 0), None);
1359 assert_eq!(sparse.checked_get(0, 3), None);
1360 assert_eq!(
1361 sparse.to_dense().unwrap().rows_ref(),
1362 vec![vec![1.0, 0.0, 0.0], vec![0.0, 0.0, 3.0]]
1363 );
1364
1365 assert!(SparseMatrix::new(0, 1, vec![0], vec![], vec![]).is_err());
1366 assert!(SparseMatrix::new(2, 2, vec![0, 0], vec![], vec![]).is_err());
1367 assert!(SparseMatrix::new(1, 2, vec![0, 1], vec![], vec![1.0]).is_err());
1368 assert!(SparseMatrix::new(1, 2, vec![1, 1], vec![0], vec![1.0]).is_err());
1369 assert!(SparseMatrix::new(1, 2, vec![0, 1], vec![2], vec![1.0]).is_err());
1370 assert!(SparseMatrix::new(2, 2, vec![0, 2, 1], vec![0], vec![1.0]).is_err());
1371 assert!(SparseMatrix::new(1, 3, vec![0, 2], vec![2, 1], vec![1.0, 2.0]).is_err());
1372 assert!(SparseMatrix::new(1, 3, vec![0, 2], vec![1, 1], vec![1.0, 2.0]).is_err());
1373 assert!(SparseMatrix::from_triplets(0, 1, &[]).is_err());
1374 assert!(SparseMatrix::from_triplets(1, 1, &[(1, 0, 1.0)]).is_err());
1375 }
1376
1377 #[test]
1378 fn sparse_duplicate_triplets_are_summed_and_zero_sums_are_dropped() {
1379 let sparse = SparseMatrix::from_triplets(
1380 1,
1381 3,
1382 &[(0, 1, 2.0), (0, 1, 3.0), (0, 2, 4.0), (0, 2, -4.0)],
1383 )
1384 .unwrap();
1385 assert_eq!(sparse.nnz(), 1);
1386 assert_eq!(sparse.get(0, 1), 5.0);
1387 assert_eq!(sparse.get(0, 2), 0.0);
1388 }
1389
1390 #[test]
1391 #[should_panic(expected = "Matrix::get")]
1392 fn dense_get_panics_on_out_of_bounds_access() {
1393 Matrix::new(vec![vec![1.0]]).unwrap().get(1, 0);
1394 }
1395
1396 #[test]
1397 #[should_panic(expected = "SparseMatrix::get")]
1398 fn sparse_get_panics_on_out_of_bounds_access() {
1399 SparseMatrix::zeros(1, 1).unwrap().get(0, 1);
1400 }
1401
1402 #[cfg(feature = "serde")]
1403 #[test]
1404 fn serde_round_trips_dense_and_string_matrices_and_rejects_jagged_data() {
1405 let matrix = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
1406 let encoded = serde_json::to_string(&matrix).unwrap();
1407 assert_eq!(serde_json::from_str::<Matrix>(&encoded).unwrap(), matrix);
1408 assert!(serde_json::from_str::<Matrix>(r#"{"data":[[1.0],[2.0,3.0]]}"#).is_err());
1409
1410 let strings = StrMatrix::from_strings(vec![vec!["north", "south"]]).unwrap();
1411 let encoded = serde_json::to_string(&strings).unwrap();
1412 assert_eq!(
1413 serde_json::from_str::<StrMatrix>(&encoded).unwrap(),
1414 strings
1415 );
1416 assert!(
1417 serde_json::from_str::<StrMatrix>(r#"{"data":[["north"],["south","east"]]}"#).is_err()
1418 );
1419
1420 let sparse = SparseMatrix::from_triplets(1, 2, &[(0, 1, 3.0)]).unwrap();
1421 let encoded = serde_json::to_string(&sparse).unwrap();
1422 assert_eq!(
1423 serde_json::from_str::<SparseMatrix>(&encoded).unwrap(),
1424 sparse
1425 );
1426 assert!(serde_json::from_str::<SparseMatrix>(
1427 r#"{"nrows":1,"ncols":3,"indptr":[0,2],"indices":[2,1],"data":[1.0,2.0]}"#
1428 )
1429 .is_err());
1430 }
1431}