Skip to main content

datarust/
matrix.rs

1//! Dense and sparse matrix containers used throughout the crate.
2
3use 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/// Row-major dense matrix of `f64` backed by a single contiguous `Vec<f64>`.
14///
15/// The flat layout (one allocation, stride-1 row traversal) keeps every numeric
16/// hot loop cache-friendly and auto-vectorizable — a substantial speedup over
17/// the previous `Vec<Vec<f64>>` representation on large dense inputs.
18///
19/// # Examples
20///
21/// ```rust,no_run
22/// use datarust::Matrix;
23///
24/// let m = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
25/// assert_eq!(m.nrows(), 2);
26/// assert_eq!(m.ncols(), 2);
27/// assert_eq!(m.get(0, 0), 1.0);
28/// ```
29#[derive(Debug, Clone, PartialEq)]
30pub struct Matrix {
31    /// Flat row-major data, length `rows * cols`.
32    data: Vec<f64>,
33    rows: usize,
34    cols: usize,
35}
36
37impl Matrix {
38    /// Creates a matrix from a nested vector, validating a rectangular shape.
39    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    /// Creates a matrix from a nested vector of rows.
69    pub fn from_rows(rows: Vec<Vec<f64>>) -> Result<Self> {
70        Self::new(rows)
71    }
72
73    /// Creates a matrix from row-major flat data of the given shape.
74    ///
75    /// With the flat internal representation this stores the buffer directly,
76    /// with no per-row chunking or copy.
77    ///
78    /// ```rust
79    /// use datarust::Matrix;
80    ///
81    /// let m = Matrix::from_flat(2, 3, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]).unwrap();
82    /// assert_eq!(m.nrows(), 2);
83    /// assert_eq!(m.ncols(), 3);
84    /// assert_eq!(m.get(1, 2), 6.0);
85    /// ```
86    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    /// Creates a matrix filled with zeros of the given shape.
105    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    /// Creates an `n` by `n` identity matrix.
118    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    /// Returns the number of rows.
135    #[inline]
136    pub fn nrows(&self) -> usize {
137        self.rows
138    }
139
140    /// Returns the number of columns.
141    #[inline]
142    pub fn ncols(&self) -> usize {
143        self.cols
144    }
145
146    /// Returns the underlying flat row-major data as a slice.
147    ///
148    /// Elements are laid out so that element `(i, j)` is at index `i * ncols + j`.
149    /// Use this for cache-friendly, auto-vectorizable numeric loops in preference
150    /// to per-element [`get`](Self::get) calls.
151    #[inline]
152    pub fn as_slice(&self) -> &[f64] {
153        &self.data
154    }
155
156    /// Returns the underlying flat row-major data as a mutable slice.
157    #[inline]
158    pub fn as_mut_slice(&mut self) -> &mut [f64] {
159        &mut self.data
160    }
161
162    /// Returns the element at row `i`, column `j`.
163    ///
164    /// # Panics
165    ///
166    /// Panics if `i >= nrows()` or `j >= ncols()`. Use [`checked_get`](Self::checked_get)
167    /// for a safe alternative that returns `None` on out-of-bounds access.
168    #[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    /// Returns `Some(element)` at row `i`, column `j`, or `None` if the indices
182    /// are out of bounds.
183    #[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    /// Sets the element at row `i`, column `j`.
193    #[inline]
194    pub fn set(&mut self, i: usize, j: usize, v: f64) {
195        self.data[i * self.cols + j] = v;
196    }
197
198    /// Returns `Ok(())` if the matrix contains no NaN values.
199    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    /// Returns `Ok(())` if every matrix element is finite.
214    ///
215    /// Unlike [`validate_no_nan`](Self::validate_no_nan), this also rejects
216    /// positive and negative infinity. Algorithms whose numerical kernels do
217    /// not define missing-value semantics should call this before fitting or
218    /// predicting.
219    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    /// Returns `Ok(())` if the matrix contains no positive or negative
233    /// infinity. `NaN` values are allowed for missing-value-aware algorithms
234    /// such as imputers.
235    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    /// Returns the row at index `i` as a contiguous slice.
249    ///
250    /// ```rust
251    /// use datarust::Matrix;
252    ///
253    /// let m = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
254    /// assert_eq!(m.row(0), &[1.0, 2.0]);
255    /// ```
256    #[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    /// Returns column `j` as a new vector.
263    ///
264    /// ```rust
265    /// use datarust::Matrix;
266    ///
267    /// let m = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
268    /// assert_eq!(m.col(1), vec![2.0, 4.0]);
269    /// ```
270    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    /// Iterates over the rows as contiguous slices.
277    pub fn iter_rows(&self) -> impl Iterator<Item = &[f64]> {
278        self.data.chunks_exact(self.cols)
279    }
280
281    /// Returns the rows as a nested `Vec<Vec<f64>>`.
282    ///
283    /// This allocates and copies — prefer [`as_slice`](Self::as_slice) or
284    /// [`iter_rows`](Self::iter_rows) in hot loops. Kept for callers that
285    /// need the nested shape (e.g. passing to `stats` functions that still
286    /// take `&[Vec<f64>]`).
287    #[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    /// Consumes the matrix and returns the rows as a nested `Vec<Vec<f64>>`.
296    #[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    /// Returns the transpose of the matrix.
303    ///
304    /// ```rust
305    /// use datarust::Matrix;
306    ///
307    /// let m = Matrix::new(vec![vec![1.0, 2.0, 3.0], vec![4.0, 5.0, 6.0]]).unwrap();
308    /// let t = m.transpose();
309    /// assert_eq!(t.nrows(), 3);
310    /// assert_eq!(t.ncols(), 2);
311    /// assert_eq!(t.get(0, 1), 4.0);
312    /// ```
313    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    /// Multiplies two matrices and returns the product.
330    ///
331    /// ```rust
332    /// use datarust::Matrix;
333    ///
334    /// let a = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
335    /// let b = Matrix::new(vec![vec![5.0, 6.0], vec![7.0, 8.0]]).unwrap();
336    /// let c = a.matmul(&b).unwrap();
337    /// assert_eq!(c.get(0, 0), 19.0); // 1*5 + 2*7
338    /// assert_eq!(c.get(0, 1), 22.0); // 1*6 + 2*8
339    /// ```
340    #[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            // Tuned pure-Rust GEMM: C(m×n) = 1.0*A(m×k)·B(k×n) + 0.0*C, row-major.
357            // dgemm takes isize strides.
358            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            // Indexed inner loops retained for cache-friendly row-major access.
381            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    /// Returns the mean of each column.
404    pub fn column_mean(&self) -> Vec<f64> {
405        crate::stats::column_mean_flat(&self.data, self.rows, self.cols)
406    }
407
408    /// Creates a matrix from a vector of columns.
409    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    /// Select a subset of columns by index (0-based).
438    ///
439    /// ```rust
440    /// use datarust::Matrix;
441    ///
442    /// let m = Matrix::new(vec![
443    ///     vec![1.0, 2.0, 3.0],
444    ///     vec![4.0, 5.0, 6.0],
445    /// ])?;
446    /// let sub = m.select_columns(&[0, 2])?;
447    /// assert_eq!(sub.ncols(), 2);
448    /// assert_eq!(sub.get(0, 0), 1.0);
449    /// assert_eq!(sub.get(0, 1), 3.0);
450    /// # Ok::<_, Box<dyn std::error::Error>>(())
451    /// ```
452    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    /// Select a subset of rows by index (0-based).
482    ///
483    /// ```rust
484    /// use datarust::Matrix;
485    ///
486    /// let m = Matrix::new(vec![
487    ///     vec![10.0],
488    ///     vec![20.0],
489    ///     vec![30.0],
490    /// ])?;
491    /// let sub = m.select_rows(&[0, 2])?;
492    /// assert_eq!(sub.nrows(), 2);
493    /// assert_eq!(sub.get(0, 0), 10.0);
494    /// assert_eq!(sub.get(1, 0), 30.0);
495    /// # Ok::<_, Box<dyn std::error::Error>>(())
496    /// ```
497    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        // Preserve the nested `{"data":[[...]]}` wire format even though the
532        // internal storage is now flat row-major.
533        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/// Row-major matrix of strings used by the categorical encoders.
556#[derive(Debug, Clone, PartialEq)]
557pub struct StrMatrix {
558    pub(crate) data: Vec<Vec<String>>,
559}
560
561impl StrMatrix {
562    /// Creates a string matrix from a nested vector, validating a rectangular shape.
563    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    /// Creates a single-column string matrix from an iterator of values.
583    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    /// Creates a string matrix from an iterator of rows.
596    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    /// Returns the number of rows.
609    #[inline]
610    pub fn nrows(&self) -> usize {
611        self.data.len()
612    }
613
614    /// Returns the number of columns.
615    #[inline]
616    pub fn ncols(&self) -> usize {
617        self.data[0].len()
618    }
619
620    /// Returns the string at row `i`, column `j`.
621    ///
622    /// # Panics
623    ///
624    /// Panics if `i >= nrows()` or `j >= ncols()`. Use
625    /// [`checked_get`](Self::checked_get) for a safe alternative.
626    #[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    /// Returns `Some(element)` at row `i`, column `j`, or `None` if the indices
640    /// are out of bounds.
641    #[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    /// Returns column `j` as a new vector of strings.
647    pub fn column(&self, j: usize) -> Vec<String> {
648        self.data.iter().map(|r| r[j].clone()).collect()
649    }
650
651    /// Returns column `j` as a vector of string slices borrowing from the matrix.
652    ///
653    /// Unlike [`column`](Self::column), this does not allocate or copy the cell
654    /// strings, which is faster when the column is only inspected.
655    pub fn column_refs(&self, j: usize) -> Vec<&str> {
656        self.data.iter().map(|r| r[j].as_str()).collect()
657    }
658
659    /// Returns the row at index `i` as a slice.
660    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    /// Fallibly construct a [`Matrix`] from a nested vector.
697    ///
698    /// Returns an error if the rows are empty or jagged. This replaces the
699    /// previous panicking `From` impl to keep validation consistent with
700    /// [`Matrix::new`].
701    fn try_from(data: Vec<Vec<f64>>) -> Result<Self> {
702        Matrix::new(data)
703    }
704}
705
706/// Compressed Sparse Row (CSR) matrix for memory-efficient storage of
707/// mostly-zero 2-D data, mirroring scipy.sparse `csr_matrix`.
708///
709/// Three arrays define the non-zero entries:
710/// - `indptr` (length `nrows + 1`): row `i` occupies
711///   `indptr[i]..indptr[i+1]` in `indices`/`data`.
712/// - `indices`: column index of each non-zero.
713/// - `data`: value of each non-zero.
714#[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    /// Build a CSR matrix from raw CSR arrays.
726    ///
727    /// `indptr` must be non-decreasing and stay within the stored-entry count.
728    /// Column indices must be strictly increasing within every row.
729    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    /// Build from `(row, col, value)` triplets. Zero-valued triplets are
799    /// dropped automatically. Within a row, entries are sorted by column and
800    /// duplicate coordinates are summed.
801    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    /// Create an all-zeros sparse matrix with the given shape.
864    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    /// Returns the number of rows.
884    #[inline]
885    pub fn nrows(&self) -> usize {
886        self.nrows
887    }
888
889    /// Returns the number of columns.
890    #[inline]
891    pub fn ncols(&self) -> usize {
892        self.ncols
893    }
894
895    /// Number of stored non-zero entries.
896    #[inline]
897    pub fn nnz(&self) -> usize {
898        self.data.len()
899    }
900
901    /// Density (fraction of non-zero entries).
902    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    /// Get element at `(i, j)`. Returns 0.0 if not stored.
911    ///
912    /// # Panics
913    ///
914    /// Panics if `i >= nrows()` or `j >= ncols()`. Use
915    /// [`checked_get`](Self::checked_get) for bounds-checked access.
916    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        // Binary search within the row's column indices.
928        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    /// Returns `Some(element)` at row `i`, column `j`, including `Some(0.0)`
936    /// when the value is not stored, or `None` if either index is out of bounds.
937    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    /// Iterate over `(col, value)` non-zero entries in row `i`.
951    ///
952    /// # Panics
953    ///
954    /// Panics if `i >= nrows()`.
955    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    /// Convert to a dense [`Matrix`].
971    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        // [[1*5+2*7, 1*6+2*8],[3*5+4*7, 3*6+4*8]] = [[19,22],[43,50]]
1108        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}