1use crate::error::{DatarustError, Result};
4
5#[derive(Debug, Clone, PartialEq)]
11pub struct Matrix {
12 data: Vec<f64>,
14 rows: usize,
15 cols: usize,
16}
17
18impl Matrix {
19 pub fn new(data: Vec<Vec<f64>>) -> Result<Self> {
21 if data.is_empty() {
22 return Err(DatarustError::EmptyInput("matrix has no rows".into()));
23 }
24 let cols = data[0].len();
25 if cols == 0 {
26 return Err(DatarustError::EmptyInput("matrix has no columns".into()));
27 }
28 for (i, row) in data.iter().enumerate() {
29 if row.len() != cols {
30 return Err(DatarustError::ShapeMismatch {
31 expected: format!("{} columns", cols),
32 actual: format!("{} columns at row {}", row.len(), i),
33 });
34 }
35 }
36 let rows = data.len();
37 let mut flat = Vec::with_capacity(rows * cols);
38 for row in data {
39 flat.extend(row);
40 }
41 Ok(Self {
42 data: flat,
43 rows,
44 cols,
45 })
46 }
47
48 pub fn from_rows(rows: Vec<Vec<f64>>) -> Result<Self> {
50 Self::new(rows)
51 }
52
53 pub fn from_flat(rows: usize, cols: usize, flat: Vec<f64>) -> Result<Self> {
58 if rows == 0 || cols == 0 {
59 return Err(DatarustError::EmptyInput("zero dimension".into()));
60 }
61 let expected = rows
62 .checked_mul(cols)
63 .ok_or_else(|| DatarustError::ShapeMismatch {
64 expected: "rows * cols within usize range".into(),
65 actual: format!("{} rows × {} cols overflows usize", rows, cols),
66 })?;
67 if flat.len() != expected {
68 return Err(DatarustError::ShapeMismatch {
69 expected: format!("{} elements", expected),
70 actual: format!("{} elements", flat.len()),
71 });
72 }
73 Ok(Self {
74 data: flat,
75 rows,
76 cols,
77 })
78 }
79
80 pub fn zeros(rows: usize, cols: usize) -> Result<Self> {
82 if rows == 0 || cols == 0 {
83 return Err(DatarustError::EmptyInput("zero dimension".into()));
84 }
85 Ok(Self {
86 data: vec![0.0; rows * cols],
87 rows,
88 cols,
89 })
90 }
91
92 pub fn identity(n: usize) -> Result<Self> {
94 if n == 0 {
95 return Err(DatarustError::EmptyInput("zero dimension".into()));
96 }
97 let mut data = vec![0.0; n * n];
98 for i in 0..n {
99 data[i * n + i] = 1.0;
100 }
101 Ok(Self {
102 data,
103 rows: n,
104 cols: n,
105 })
106 }
107
108 #[inline]
110 pub fn nrows(&self) -> usize {
111 self.rows
112 }
113
114 #[inline]
116 pub fn ncols(&self) -> usize {
117 self.cols
118 }
119
120 #[inline]
126 pub fn as_slice(&self) -> &[f64] {
127 &self.data
128 }
129
130 #[inline]
132 pub fn as_mut_slice(&mut self) -> &mut [f64] {
133 &mut self.data
134 }
135
136 #[inline]
143 pub fn get(&self, i: usize, j: usize) -> f64 {
144 debug_assert!(
145 i < self.rows && j < self.cols,
146 "Matrix::get: index ({}, {}) out of bounds for {}×{}",
147 i,
148 j,
149 self.rows,
150 self.cols
151 );
152 unsafe { *self.data.get_unchecked(i * self.cols + j) }
154 }
155
156 #[inline]
159 pub fn checked_get(&self, i: usize, j: usize) -> Option<f64> {
160 if i < self.rows && j < self.cols {
161 Some(self.data[i * self.cols + j])
162 } else {
163 None
164 }
165 }
166
167 #[inline]
169 pub fn set(&mut self, i: usize, j: usize, v: f64) {
170 self.data[i * self.cols + j] = v;
171 }
172
173 pub fn validate_no_nan(&self) -> Result<()> {
175 for (flat_idx, &v) in self.data.iter().enumerate() {
176 if v.is_nan() {
177 let i = flat_idx / self.cols;
178 let j = flat_idx % self.cols;
179 return Err(DatarustError::InvalidInput(format!(
180 "NaN value at position ({}, {})",
181 i, j
182 )));
183 }
184 }
185 Ok(())
186 }
187
188 #[inline]
190 pub fn row(&self, i: usize) -> &[f64] {
191 let start = i * self.cols;
192 &self.data[start..start + self.cols]
193 }
194
195 pub fn col(&self, j: usize) -> Vec<f64> {
197 (0..self.rows)
198 .map(|i| self.data[i * self.cols + j])
199 .collect()
200 }
201
202 pub fn iter_rows(&self) -> impl Iterator<Item = &[f64]> {
204 self.data.chunks_exact(self.cols)
205 }
206
207 #[doc(hidden)]
214 pub fn rows_ref(&self) -> Vec<Vec<f64>> {
215 self.data
216 .chunks_exact(self.cols)
217 .map(|chunk| chunk.to_vec())
218 .collect()
219 }
220
221 #[doc(hidden)]
223 pub fn into_rows(self) -> Vec<Vec<f64>> {
224 let cols = self.cols;
225 self.data.chunks(cols).map(|chunk| chunk.to_vec()).collect()
226 }
227
228 pub fn transpose(&self) -> Matrix {
230 let rows = self.rows;
231 let cols = self.cols;
232 let mut out = vec![0.0; rows * cols];
233 for i in 0..rows {
234 for j in 0..cols {
235 out[j * rows + i] = self.data[i * cols + j];
236 }
237 }
238 Matrix {
239 data: out,
240 rows: cols,
241 cols: rows,
242 }
243 }
244
245 #[allow(clippy::needless_range_loop)]
247 pub fn matmul(&self, other: &Matrix) -> Result<Matrix> {
248 if self.cols != other.rows {
249 return Err(DatarustError::ShapeMismatch {
250 expected: format!("second operand with {} rows", self.cols),
251 actual: format!("{} rows", other.rows),
252 });
253 }
254 let m = self.rows;
255 let k = self.cols;
256 let n = other.cols;
257 let mut out = vec![0.0; m * n];
258
259 #[cfg(feature = "matrixmultiply")]
260 {
261 unsafe {
264 matrixmultiply::dgemm(
265 m,
266 k,
267 n,
268 1.0,
269 self.data.as_ptr(),
270 k as isize,
271 1,
272 other.data.as_ptr(),
273 n as isize,
274 1,
275 0.0,
276 out.as_mut_ptr(),
277 n as isize,
278 1,
279 );
280 }
281 }
282
283 #[cfg(not(feature = "matrixmultiply"))]
284 {
285 for i in 0..m {
287 let out_base = i * n;
288 let self_base = i * k;
289 for l in 0..k {
290 let a = self.data[self_base + l];
291 if a == 0.0 {
292 continue;
293 }
294 let other_base = l * n;
295 for j in 0..n {
296 out[out_base + j] += a * other.data[other_base + j];
297 }
298 }
299 }
300 }
301 Ok(Matrix {
302 data: out,
303 rows: m,
304 cols: n,
305 })
306 }
307
308 pub fn column_mean(&self) -> Vec<f64> {
310 crate::stats::column_mean_flat(&self.data, self.rows, self.cols)
311 }
312
313 pub fn from_columns(cols: Vec<Vec<f64>>) -> Result<Self> {
315 if cols.is_empty() || cols[0].is_empty() {
316 return Err(DatarustError::EmptyInput("no columns".into()));
317 }
318 let rows = cols[0].len();
319 for c in &cols {
320 if c.len() != rows {
321 return Err(DatarustError::ShapeMismatch {
322 expected: format!("{} rows", rows),
323 actual: format!("{} rows", c.len()),
324 });
325 }
326 }
327 let ncols = cols.len();
328 let mut data = vec![0.0; rows * ncols];
329 for (j, col) in cols.iter().enumerate() {
330 for (i, &v) in col.iter().enumerate() {
331 data[i * ncols + j] = v;
332 }
333 }
334 Ok(Self {
335 data,
336 rows,
337 cols: ncols,
338 })
339 }
340
341 pub fn select_columns(&self, indices: &[usize]) -> Result<Self> {
357 if indices.is_empty() {
358 return Err(DatarustError::EmptyInput("no columns selected".into()));
359 }
360 let ncols = self.cols;
361 for &c in indices {
362 if c >= ncols {
363 return Err(DatarustError::InvalidInput(format!(
364 "column index {} out of range (ncols {})",
365 c, ncols
366 )));
367 }
368 }
369 let out_cols = indices.len();
370 let mut out = Vec::with_capacity(self.rows * out_cols);
371 for i in 0..self.rows {
372 let base = i * ncols;
373 for &c in indices {
374 out.push(self.data[base + c]);
375 }
376 }
377 Ok(Self {
378 data: out,
379 rows: self.rows,
380 cols: out_cols,
381 })
382 }
383
384 pub fn select_rows(&self, indices: &[usize]) -> Result<Self> {
401 if indices.is_empty() {
402 return Err(DatarustError::EmptyInput("no rows selected".into()));
403 }
404 let nrows = self.rows;
405 for &r in indices {
406 if r >= nrows {
407 return Err(DatarustError::InvalidInput(format!(
408 "row index {} out of range (nrows {})",
409 r, nrows
410 )));
411 }
412 }
413 let mut out = Vec::with_capacity(indices.len() * self.cols);
414 for &r in indices {
415 let start = r * self.cols;
416 out.extend_from_slice(&self.data[start..start + self.cols]);
417 }
418 Ok(Self {
419 data: out,
420 rows: indices.len(),
421 cols: self.cols,
422 })
423 }
424}
425
426#[cfg(feature = "serde")]
427impl serde::Serialize for Matrix {
428 fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
429 where
430 S: serde::Serializer,
431 {
432 use serde::ser::SerializeStruct;
433 let rows: Vec<&[f64]> = self.data.chunks_exact(self.cols).collect();
436 let mut s = serializer.serialize_struct("Matrix", 1)?;
437 s.serialize_field("data", &rows)?;
438 s.end()
439 }
440}
441
442#[cfg(feature = "serde")]
443impl<'de> serde::Deserialize<'de> for Matrix {
444 fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
445 where
446 D: serde::Deserializer<'de>,
447 {
448 #[derive(serde::Deserialize)]
449 struct Raw {
450 data: Vec<Vec<f64>>,
451 }
452 let raw = Raw::deserialize(deserializer)?;
453 Matrix::new(raw.data).map_err(serde::de::Error::custom)
454 }
455}
456
457#[derive(Debug, Clone, PartialEq)]
459pub struct StrMatrix {
460 pub(crate) data: Vec<Vec<String>>,
461}
462
463impl StrMatrix {
464 pub fn new(data: Vec<Vec<String>>) -> Result<Self> {
466 if data.is_empty() {
467 return Err(DatarustError::EmptyInput("matrix has no rows".into()));
468 }
469 let cols = data[0].len();
470 if cols == 0 {
471 return Err(DatarustError::EmptyInput("matrix has no columns".into()));
472 }
473 for (i, row) in data.iter().enumerate() {
474 if row.len() != cols {
475 return Err(DatarustError::ShapeMismatch {
476 expected: format!("{} columns", cols),
477 actual: format!("{} columns at row {}", row.len(), i),
478 });
479 }
480 }
481 Ok(Self { data })
482 }
483
484 pub fn from_column<I, S>(col: I) -> Result<Self>
486 where
487 I: IntoIterator<Item = S>,
488 S: Into<String>,
489 {
490 let data: Vec<Vec<String>> = col.into_iter().map(|s| vec![s.into()]).collect::<Vec<_>>();
491 if data.is_empty() {
492 return Err(DatarustError::EmptyInput("column has no rows".into()));
493 }
494 Self::new(data)
495 }
496
497 pub fn from_strings<I, S>(rows: I) -> Result<Self>
499 where
500 I: IntoIterator<Item = Vec<S>>,
501 S: Into<String>,
502 {
503 let data: Vec<Vec<String>> = rows
504 .into_iter()
505 .map(|r| r.into_iter().map(|s| s.into()).collect())
506 .collect();
507 Self::new(data)
508 }
509
510 #[inline]
512 pub fn nrows(&self) -> usize {
513 self.data.len()
514 }
515
516 #[inline]
518 pub fn ncols(&self) -> usize {
519 self.data[0].len()
520 }
521
522 #[inline]
529 pub fn get(&self, i: usize, j: usize) -> &str {
530 debug_assert!(
531 i < self.data.len() && j < self.data[i].len(),
532 "StrMatrix::get: index ({}, {}) out of bounds for {}×{}",
533 i,
534 j,
535 self.data.len(),
536 self.data[0].len()
537 );
538 &self.data[i][j]
539 }
540
541 #[inline]
544 pub fn checked_get(&self, i: usize, j: usize) -> Option<&str> {
545 self.data.get(i)?.get(j).map(|s| s.as_str())
546 }
547
548 pub fn column(&self, j: usize) -> Vec<String> {
550 self.data.iter().map(|r| r[j].clone()).collect()
551 }
552
553 pub fn row(&self, i: usize) -> &[String] {
555 &self.data[i]
556 }
557}
558
559#[cfg(feature = "serde")]
560impl serde::Serialize for StrMatrix {
561 fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
562 where
563 S: serde::Serializer,
564 {
565 use serde::ser::SerializeStruct;
566 let mut s = serializer.serialize_struct("StrMatrix", 1)?;
567 s.serialize_field("data", &self.data)?;
568 s.end()
569 }
570}
571
572#[cfg(feature = "serde")]
573impl<'de> serde::Deserialize<'de> for StrMatrix {
574 fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
575 where
576 D: serde::Deserializer<'de>,
577 {
578 #[derive(serde::Deserialize)]
579 struct Raw {
580 data: Vec<Vec<String>>,
581 }
582 let raw = Raw::deserialize(deserializer)?;
583 StrMatrix::new(raw.data).map_err(serde::de::Error::custom)
584 }
585}
586
587impl TryFrom<Vec<Vec<f64>>> for Matrix {
588 type Error = DatarustError;
589
590 fn try_from(data: Vec<Vec<f64>>) -> Result<Self> {
596 Matrix::new(data)
597 }
598}
599
600#[derive(Debug, Clone, PartialEq)]
609#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
610pub struct SparseMatrix {
611 nrows: usize,
612 ncols: usize,
613 indptr: Vec<usize>,
614 indices: Vec<usize>,
615 data: Vec<f64>,
616}
617
618impl SparseMatrix {
619 pub fn new(
621 nrows: usize,
622 ncols: usize,
623 indptr: Vec<usize>,
624 indices: Vec<usize>,
625 data: Vec<f64>,
626 ) -> Result<Self> {
627 if nrows == 0 || ncols == 0 {
628 return Err(DatarustError::EmptyInput("zero dimension".into()));
629 }
630 if indptr.len() != nrows + 1 {
631 return Err(DatarustError::ShapeMismatch {
632 expected: format!("{} indptr entries", nrows + 1),
633 actual: format!("{} indptr entries", indptr.len()),
634 });
635 }
636 if indices.len() != data.len() {
637 return Err(DatarustError::ShapeMismatch {
638 expected: format!("{} indices", data.len()),
639 actual: format!("{} indices", indices.len()),
640 });
641 }
642 let nnz = data.len();
643 if indptr[0] != 0 || indptr[nrows] != nnz {
644 return Err(DatarustError::InvalidInput(
645 "indptr must start at 0 and end at nnz".into(),
646 ));
647 }
648 for &c in &indices {
649 if c >= ncols {
650 return Err(DatarustError::InvalidInput(format!(
651 "column index {} out of range (ncols {})",
652 c, ncols
653 )));
654 }
655 }
656 Ok(Self {
657 nrows,
658 ncols,
659 indptr,
660 indices,
661 data,
662 })
663 }
664
665 pub fn from_triplets(
668 nrows: usize,
669 ncols: usize,
670 triplets: &[(usize, usize, f64)],
671 ) -> Result<Self> {
672 if nrows == 0 || ncols == 0 {
673 return Err(DatarustError::EmptyInput("zero dimension".into()));
674 }
675 let mut per_row: Vec<Vec<(usize, f64)>> = vec![vec![]; nrows];
676 for &(r, c, v) in triplets {
677 if r >= nrows {
678 return Err(DatarustError::InvalidInput(format!(
679 "row {} out of range (nrows {})",
680 r, nrows
681 )));
682 }
683 if c >= ncols {
684 return Err(DatarustError::InvalidInput(format!(
685 "col {} out of range (ncols {})",
686 c, ncols
687 )));
688 }
689 if v != 0.0 {
690 per_row[r].push((c, v));
691 }
692 }
693 let mut indptr = Vec::with_capacity(nrows + 1);
694 let mut indices = Vec::new();
695 let mut data = Vec::new();
696 indptr.push(0);
697 for row_entries in &mut per_row {
698 row_entries.sort_by_key(|(c, _)| *c);
699 for &(c, v) in row_entries.iter() {
700 indices.push(c);
701 data.push(v);
702 }
703 indptr.push(indices.len());
704 }
705 Ok(Self {
706 nrows,
707 ncols,
708 indptr,
709 indices,
710 data,
711 })
712 }
713
714 pub fn zeros(nrows: usize, ncols: usize) -> Result<Self> {
716 if nrows == 0 || ncols == 0 {
717 return Err(DatarustError::EmptyInput("zero dimension".into()));
718 }
719 Ok(Self {
720 nrows,
721 ncols,
722 indptr: vec![0; nrows + 1],
723 indices: vec![],
724 data: vec![],
725 })
726 }
727
728 #[inline]
730 pub fn nrows(&self) -> usize {
731 self.nrows
732 }
733
734 #[inline]
736 pub fn ncols(&self) -> usize {
737 self.ncols
738 }
739
740 #[inline]
742 pub fn nnz(&self) -> usize {
743 self.data.len()
744 }
745
746 pub fn density(&self) -> f64 {
748 let total = self.nrows.saturating_mul(self.ncols);
749 if total == 0 {
750 return 0.0;
751 }
752 self.nnz() as f64 / total as f64
753 }
754
755 pub fn get(&self, i: usize, j: usize) -> f64 {
757 let start = self.indptr[i];
758 let end = self.indptr[i + 1];
759 let slice = &self.indices[start..end];
761 match slice.binary_search(&j) {
762 Ok(local) => self.data[start + local],
763 Err(_) => 0.0,
764 }
765 }
766
767 pub fn checked_get(&self, i: usize, j: usize) -> Option<f64> {
770 let start = *self.indptr.get(i)?;
771 let end = *self.indptr.get(i + 1)?;
772 let slice = self.indices.get(start..end)?;
773 match slice.binary_search(&j) {
774 Ok(local) => Some(self.data[start + local]),
775 Err(_) => Some(0.0),
776 }
777 }
778
779 pub fn row_nz(&self, i: usize) -> impl Iterator<Item = (usize, f64)> + '_ {
781 let start = self.indptr[i];
782 let end = self.indptr[i + 1];
783 self.indices[start..end]
784 .iter()
785 .zip(self.data[start..end].iter())
786 .map(|(&c, &v)| (c, v))
787 }
788
789 pub fn to_dense(&self) -> Result<Matrix> {
791 let mut rows = vec![vec![0.0; self.ncols]; self.nrows];
792 for (i, row) in rows.iter_mut().enumerate() {
793 for (c, v) in self.row_nz(i) {
794 row[c] = v;
795 }
796 }
797 Matrix::new(rows)
798 }
799}
800
801#[cfg(test)]
802#[allow(dead_code)]
803pub(crate) fn approx_eq_matrices(a: &Matrix, b: &Matrix, tol: f64) -> bool {
804 if a.nrows() != b.nrows() || a.ncols() != b.ncols() {
805 return false;
806 }
807 for i in 0..a.nrows() {
808 for j in 0..a.ncols() {
809 if (a.get(i, j) - b.get(i, j)).abs() > tol {
810 return false;
811 }
812 }
813 }
814 true
815}
816
817#[cfg(test)]
818#[allow(dead_code)]
819pub(crate) fn approx_eq_vecs(a: &[f64], b: &[f64], tol: f64) -> bool {
820 if a.len() != b.len() {
821 return false;
822 }
823 a.iter().zip(b.iter()).all(|(x, y)| (x - y).abs() <= tol)
824}
825
826#[cfg(test)]
827#[macro_use]
828mod assert_macros {
829 #[allow(unused_macros)]
830 macro_rules! assert_mat_eq {
831 ($a:expr, $b:expr, $tol:expr) => {{
832 assert!(
833 $crate::matrix::approx_eq_matrices(&$a, &$b, $tol),
834 "matrices not equal within tolerance {}\n left: {:?}\nright: {:?}",
835 $tol,
836 $a.as_slice(),
837 $b.as_slice()
838 );
839 }};
840 }
841}
842
843#[cfg(test)]
844mod tests {
845 use super::*;
846
847 #[test]
848 fn new_valid() {
849 let m = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
850 assert_eq!(m.nrows(), 2);
851 assert_eq!(m.ncols(), 2);
852 }
853
854 #[test]
855 fn new_jagged_rejected() {
856 let err = Matrix::new(vec![vec![1.0, 2.0], vec![3.0]]).unwrap_err();
857 assert!(matches!(err, DatarustError::ShapeMismatch { .. }));
858 }
859
860 #[test]
861 fn new_empty_rejected() {
862 assert!(Matrix::new(vec![]).is_err());
863 assert!(Matrix::new(vec![vec![]]).is_err());
864 }
865
866 #[test]
867 fn from_flat() {
868 let m = Matrix::from_flat(2, 3, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]).unwrap();
869 assert_eq!(m.get(0, 2), 3.0);
870 assert_eq!(m.get(1, 0), 4.0);
871 assert_eq!(m.get(1, 2), 6.0);
872 }
873
874 #[test]
875 fn from_flat_bad_count() {
876 let err = Matrix::from_flat(2, 2, vec![1.0, 2.0, 3.0]).unwrap_err();
877 assert!(matches!(err, DatarustError::ShapeMismatch { .. }));
878 }
879
880 #[test]
881 fn zeros_and_identity() {
882 let z = Matrix::zeros(2, 3).unwrap();
883 assert_eq!(z.get(1, 2), 0.0);
884 let id = Matrix::identity(3).unwrap();
885 assert_eq!(id.get(0, 0), 1.0);
886 assert_eq!(id.get(0, 1), 0.0);
887 assert_eq!(id.get(1, 1), 1.0);
888 }
889
890 #[test]
891 fn transpose() {
892 let m = Matrix::new(vec![vec![1.0, 2.0, 3.0], vec![4.0, 5.0, 6.0]]).unwrap();
893 let t = m.transpose();
894 assert_eq!(t.nrows(), 3);
895 assert_eq!(t.ncols(), 2);
896 assert_eq!(t.get(2, 0), 3.0);
897 assert_eq!(t.get(1, 1), 5.0);
898 }
899
900 #[test]
901 fn matmul() {
902 let a = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
903 let b = Matrix::new(vec![vec![5.0, 6.0], vec![7.0, 8.0]]).unwrap();
904 let c = a.matmul(&b).unwrap();
905 assert_eq!(c.get(0, 0), 19.0);
907 assert_eq!(c.get(0, 1), 22.0);
908 assert_eq!(c.get(1, 0), 43.0);
909 assert_eq!(c.get(1, 1), 50.0);
910 }
911
912 #[test]
913 fn matmul_shape_mismatch() {
914 let a = Matrix::new(vec![vec![1.0, 2.0, 3.0]]).unwrap();
915 let b = Matrix::new(vec![vec![1.0, 2.0]]).unwrap();
916 assert!(a.matmul(&b).is_err());
917 }
918
919 #[test]
920 fn from_columns() {
921 let m = Matrix::from_columns(vec![vec![1.0, 2.0], vec![10.0, 20.0]]).unwrap();
922 assert_eq!(m.nrows(), 2);
923 assert_eq!(m.ncols(), 2);
924 assert_eq!(m.get(0, 0), 1.0);
925 assert_eq!(m.get(1, 1), 20.0);
926 }
927
928 #[test]
929 fn strmatrix_from_column() {
930 let s = StrMatrix::from_column(["a", "b", "a"]).unwrap();
931 assert_eq!(s.nrows(), 3);
932 assert_eq!(s.ncols(), 1);
933 assert_eq!(s.get(2, 0), "a");
934 }
935
936 #[test]
937 fn strmatrix_from_strings() {
938 let s = StrMatrix::from_strings(vec![vec!["x", "y"], vec!["x", "z"]]).unwrap();
939 assert_eq!(s.ncols(), 2);
940 assert_eq!(s.get(1, 1), "z");
941 }
942
943 #[test]
944 fn sparse_from_triplets_basic() {
945 let sp = SparseMatrix::from_triplets(
946 3,
947 4,
948 &[(0, 0, 1.0), (1, 2, 3.0), (2, 3, 5.0), (0, 3, 7.0)],
949 )
950 .unwrap();
951 assert_eq!(sp.nrows(), 3);
952 assert_eq!(sp.ncols(), 4);
953 assert_eq!(sp.nnz(), 4);
954 assert_eq!(sp.get(0, 0), 1.0);
955 assert_eq!(sp.get(1, 2), 3.0);
956 assert_eq!(sp.get(0, 3), 7.0);
957 assert_eq!(sp.get(1, 0), 0.0);
958 }
959
960 #[test]
961 fn sparse_zero_triplets_dropped() {
962 let sp =
963 SparseMatrix::from_triplets(2, 2, &[(0, 0, 0.0), (0, 1, 5.0), (1, 0, 0.0)]).unwrap();
964 assert_eq!(sp.nnz(), 1);
965 assert_eq!(sp.get(0, 1), 5.0);
966 }
967
968 #[test]
969 fn sparse_to_dense() {
970 let sp = SparseMatrix::from_triplets(2, 3, &[(0, 1, 2.0), (1, 0, 4.0)]).unwrap();
971 let dense = sp.to_dense().unwrap();
972 assert_eq!(dense.row(0), [0.0, 2.0, 0.0]);
973 assert_eq!(dense.row(1), [4.0, 0.0, 0.0]);
974 }
975
976 #[test]
977 fn sparse_zeros() {
978 let sp = SparseMatrix::zeros(2, 3).unwrap();
979 assert_eq!(sp.nnz(), 0);
980 assert_eq!(sp.density(), 0.0);
981 assert_eq!(sp.get(0, 1), 0.0);
982 }
983
984 #[test]
985 fn sparse_density() {
986 let sp = SparseMatrix::from_triplets(2, 4, &[(0, 0, 1.0), (1, 3, 1.0)]).unwrap();
987 assert!((sp.density() - 0.25).abs() < 1e-12);
988 }
989
990 #[test]
991 fn sparse_row_nz() {
992 let sp =
993 SparseMatrix::from_triplets(2, 3, &[(0, 1, 2.0), (0, 2, 9.0), (1, 0, 4.0)]).unwrap();
994 let row0: Vec<(usize, f64)> = sp.row_nz(0).collect();
995 assert_eq!(row0, vec![(1, 2.0), (2, 9.0)]);
996 let row1: Vec<(usize, f64)> = sp.row_nz(1).collect();
997 assert_eq!(row1, vec![(0, 4.0)]);
998 }
999
1000 #[test]
1001 fn sparse_bad_indptr_rejected() {
1002 let err = SparseMatrix::new(2, 2, vec![0, 0, 5], vec![0], vec![1.0]).unwrap_err();
1003 assert!(matches!(err, DatarustError::InvalidInput(_)));
1004 }
1005
1006 #[test]
1007 fn sparse_col_out_of_range_rejected() {
1008 let err = SparseMatrix::from_triplets(2, 2, &[(0, 5, 1.0)]).unwrap_err();
1009 assert!(matches!(err, DatarustError::InvalidInput(_)));
1010 }
1011
1012 #[test]
1013 fn select_columns_basic() {
1014 let m = Matrix::new(vec![vec![1.0, 2.0, 3.0], vec![4.0, 5.0, 6.0]]).unwrap();
1015 let sub = m.select_columns(&[0, 2]).unwrap();
1016 assert_eq!(sub.ncols(), 2);
1017 assert_eq!(sub.get(0, 0), 1.0);
1018 assert_eq!(sub.get(0, 1), 3.0);
1019 assert_eq!(sub.get(1, 0), 4.0);
1020 assert_eq!(sub.get(1, 1), 6.0);
1021 }
1022
1023 #[test]
1024 fn select_columns_out_of_range() {
1025 let m = Matrix::new(vec![vec![1.0, 2.0]]).unwrap();
1026 assert!(m.select_columns(&[0, 5]).is_err());
1027 }
1028
1029 #[test]
1030 fn select_columns_empty() {
1031 let m = Matrix::new(vec![vec![1.0]]).unwrap();
1032 assert!(m.select_columns(&[]).is_err());
1033 }
1034
1035 #[test]
1036 fn select_rows_basic() {
1037 let m = Matrix::new(vec![vec![10.0], vec![20.0], vec![30.0]]).unwrap();
1038 let sub = m.select_rows(&[0, 2]).unwrap();
1039 assert_eq!(sub.nrows(), 2);
1040 assert_eq!(sub.get(0, 0), 10.0);
1041 assert_eq!(sub.get(1, 0), 30.0);
1042 }
1043
1044 #[test]
1045 fn select_rows_out_of_range() {
1046 let m = Matrix::new(vec![vec![1.0]]).unwrap();
1047 assert!(m.select_rows(&[0, 5]).is_err());
1048 }
1049
1050 #[test]
1051 fn select_rows_empty() {
1052 let m = Matrix::new(vec![vec![1.0]]).unwrap();
1053 assert!(m.select_rows(&[]).is_err());
1054 }
1055
1056 #[test]
1057 fn select_columns_reordered() {
1058 let m = Matrix::new(vec![vec![1.0, 2.0, 3.0], vec![4.0, 5.0, 6.0]]).unwrap();
1059 let sub = m.select_columns(&[2, 0]).unwrap();
1060 assert_eq!(sub.get(0, 0), 3.0);
1061 assert_eq!(sub.get(0, 1), 1.0);
1062 }
1063
1064 #[test]
1065 fn select_rows_duplicates() {
1066 let m = Matrix::new(vec![vec![10.0], vec![20.0]]).unwrap();
1067 let sub = m.select_rows(&[0, 0, 1]).unwrap();
1068 assert_eq!(sub.nrows(), 3);
1069 assert_eq!(sub.get(0, 0), 10.0);
1070 assert_eq!(sub.get(1, 0), 10.0);
1071 assert_eq!(sub.get(2, 0), 20.0);
1072 }
1073
1074 #[test]
1075 fn dense_accessors_and_row_conversions_preserve_layout() {
1076 let mut m = Matrix::from_rows(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
1077
1078 assert_eq!(m.as_slice(), &[1.0, 2.0, 3.0, 4.0]);
1079 assert_eq!(m.checked_get(1, 0), Some(3.0));
1080 assert_eq!(m.checked_get(2, 0), None);
1081 assert_eq!(m.checked_get(0, 2), None);
1082 m.as_mut_slice()[1] = 20.0;
1083 m.set(1, 0, 30.0);
1084 assert_eq!(m.row(0), [1.0, 20.0]);
1085 assert_eq!(m.col(0), vec![1.0, 30.0]);
1086 assert_eq!(
1087 m.iter_rows().collect::<Vec<_>>(),
1088 vec![&[1.0, 20.0][..], &[30.0, 4.0][..]]
1089 );
1090 assert_eq!(m.rows_ref(), vec![vec![1.0, 20.0], vec![30.0, 4.0]]);
1091 assert_eq!(
1092 m.clone().into_rows(),
1093 vec![vec![1.0, 20.0], vec![30.0, 4.0]]
1094 );
1095 assert!(m.validate_no_nan().is_ok());
1096
1097 m.set(1, 1, f64::NAN);
1098 assert!(matches!(
1099 m.validate_no_nan(),
1100 Err(DatarustError::InvalidInput(_))
1101 ));
1102 }
1103
1104 #[test]
1105 fn dense_constructor_edge_cases_are_rejected() {
1106 assert!(Matrix::zeros(0, 1).is_err());
1107 assert!(Matrix::identity(0).is_err());
1108 assert!(Matrix::from_flat(0, 1, vec![]).is_err());
1109 assert!(Matrix::from_flat(usize::MAX, 2, vec![]).is_err());
1110 assert!(Matrix::from_columns(vec![]).is_err());
1111 assert!(Matrix::from_columns(vec![vec![]]).is_err());
1112 assert!(Matrix::from_columns(vec![vec![1.0], vec![2.0, 3.0]]).is_err());
1113 assert!(Matrix::try_from(vec![vec![1.0], vec![]]).is_err());
1114 }
1115
1116 #[test]
1117 fn strmatrix_accessors_and_validation_cover_edge_cases() {
1118 let strings = StrMatrix::new(vec![
1119 vec!["a".into(), "b".into()],
1120 vec!["c".into(), "d".into()],
1121 ])
1122 .unwrap();
1123 assert_eq!(strings.checked_get(1, 1), Some("d"));
1124 assert_eq!(strings.checked_get(2, 0), None);
1125 assert_eq!(strings.checked_get(0, 2), None);
1126 assert_eq!(strings.column(1), vec!["b".to_string(), "d".to_string()]);
1127 assert_eq!(strings.row(0), ["a".to_string(), "b".to_string()]);
1128 assert!(StrMatrix::new(vec![]).is_err());
1129 assert!(StrMatrix::new(vec![vec![]]).is_err());
1130 assert!(StrMatrix::new(vec![vec!["a".into()], vec!["b".into(), "c".into()]]).is_err());
1131 assert!(StrMatrix::from_column(Vec::<String>::new()).is_err());
1132 }
1133
1134 #[test]
1135 fn raw_sparse_construction_and_bounds_checks() {
1136 let sparse = SparseMatrix::new(2, 3, vec![0, 1, 2], vec![0, 2], vec![1.0, 3.0]).unwrap();
1137 assert_eq!(sparse.checked_get(0, 0), Some(1.0));
1138 assert_eq!(sparse.checked_get(0, 2), Some(0.0));
1139 assert_eq!(sparse.checked_get(2, 0), None);
1140 assert_eq!(
1141 sparse.to_dense().unwrap().rows_ref(),
1142 vec![vec![1.0, 0.0, 0.0], vec![0.0, 0.0, 3.0]]
1143 );
1144
1145 assert!(SparseMatrix::new(0, 1, vec![0], vec![], vec![]).is_err());
1146 assert!(SparseMatrix::new(2, 2, vec![0, 0], vec![], vec![]).is_err());
1147 assert!(SparseMatrix::new(1, 2, vec![0, 1], vec![], vec![1.0]).is_err());
1148 assert!(SparseMatrix::new(1, 2, vec![1, 1], vec![0], vec![1.0]).is_err());
1149 assert!(SparseMatrix::new(1, 2, vec![0, 1], vec![2], vec![1.0]).is_err());
1150 assert!(SparseMatrix::from_triplets(0, 1, &[]).is_err());
1151 assert!(SparseMatrix::from_triplets(1, 1, &[(1, 0, 1.0)]).is_err());
1152 }
1153
1154 #[cfg(feature = "serde")]
1155 #[test]
1156 fn serde_round_trips_dense_and_string_matrices_and_rejects_jagged_data() {
1157 let matrix = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
1158 let encoded = serde_json::to_string(&matrix).unwrap();
1159 assert_eq!(serde_json::from_str::<Matrix>(&encoded).unwrap(), matrix);
1160 assert!(serde_json::from_str::<Matrix>(r#"{"data":[[1.0],[2.0,3.0]]}"#).is_err());
1161
1162 let strings = StrMatrix::from_strings(vec![vec!["north", "south"]]).unwrap();
1163 let encoded = serde_json::to_string(&strings).unwrap();
1164 assert_eq!(
1165 serde_json::from_str::<StrMatrix>(&encoded).unwrap(),
1166 strings
1167 );
1168 assert!(
1169 serde_json::from_str::<StrMatrix>(r#"{"data":[["north"],["south","east"]]}"#).is_err()
1170 );
1171 }
1172}