Skip to main content

datarust/
matrix.rs

1//! Dense and sparse matrix containers used throughout the crate.
2
3use crate::error::{DatarustError, Result};
4
5/// Row-major dense matrix of `f64` backed by a single contiguous `Vec<f64>`.
6///
7/// The flat layout (one allocation, stride-1 row traversal) keeps every numeric
8/// hot loop cache-friendly and auto-vectorizable — a substantial speedup over
9/// the previous `Vec<Vec<f64>>` representation on large dense inputs.
10#[derive(Debug, Clone, PartialEq)]
11pub struct Matrix {
12    /// Flat row-major data, length `rows * cols`.
13    data: Vec<f64>,
14    rows: usize,
15    cols: usize,
16}
17
18impl Matrix {
19    /// Creates a matrix from a nested vector, validating a rectangular shape.
20    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    /// Creates a matrix from a nested vector of rows.
49    pub fn from_rows(rows: Vec<Vec<f64>>) -> Result<Self> {
50        Self::new(rows)
51    }
52
53    /// Creates a matrix from row-major flat data of the given shape.
54    ///
55    /// With the flat internal representation this stores the buffer directly,
56    /// with no per-row chunking or copy.
57    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    /// Creates a matrix filled with zeros of the given shape.
81    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    /// Creates an `n` by `n` identity matrix.
93    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    /// Returns the number of rows.
109    #[inline]
110    pub fn nrows(&self) -> usize {
111        self.rows
112    }
113
114    /// Returns the number of columns.
115    #[inline]
116    pub fn ncols(&self) -> usize {
117        self.cols
118    }
119
120    /// Returns the underlying flat row-major data as a slice.
121    ///
122    /// Elements are laid out so that element `(i, j)` is at index `i * ncols + j`.
123    /// Use this for cache-friendly, auto-vectorizable numeric loops in preference
124    /// to per-element [`get`](Self::get) calls.
125    #[inline]
126    pub fn as_slice(&self) -> &[f64] {
127        &self.data
128    }
129
130    /// Returns the underlying flat row-major data as a mutable slice.
131    #[inline]
132    pub fn as_mut_slice(&mut self) -> &mut [f64] {
133        &mut self.data
134    }
135
136    /// Returns the element at row `i`, column `j`.
137    ///
138    /// # Panics
139    ///
140    /// Panics if `i >= nrows()` or `j >= ncols()`. Use [`checked_get`](Self::checked_get)
141    /// for a safe alternative that returns `None` on out-of-bounds access.
142    #[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        // SAFETY: debug_assert above + the struct invariant (data.len() == rows*cols).
153        unsafe { *self.data.get_unchecked(i * self.cols + j) }
154    }
155
156    /// Returns `Some(element)` at row `i`, column `j`, or `None` if the indices
157    /// are out of bounds.
158    #[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    /// Sets the element at row `i`, column `j`.
168    #[inline]
169    pub fn set(&mut self, i: usize, j: usize, v: f64) {
170        self.data[i * self.cols + j] = v;
171    }
172
173    /// Returns `Ok(())` if the matrix contains no NaN values.
174    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    /// Returns the row at index `i` as a contiguous slice.
189    #[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    /// Returns column `j` as a new vector.
196    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    /// Iterates over the rows as contiguous slices.
203    pub fn iter_rows(&self) -> impl Iterator<Item = &[f64]> {
204        self.data.chunks_exact(self.cols)
205    }
206
207    /// Returns the rows as a nested `Vec<Vec<f64>>`.
208    ///
209    /// This allocates and copies — prefer [`as_slice`](Self::as_slice) or
210    /// [`iter_rows`](Self::iter_rows) in hot loops. Kept for callers that
211    /// need the nested shape (e.g. passing to `stats` functions that still
212    /// take `&[Vec<f64>]`).
213    #[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    /// Consumes the matrix and returns the rows as a nested `Vec<Vec<f64>>`.
222    #[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    /// Returns the transpose of the matrix.
229    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    /// Multiplies two matrices and returns the product.
246    #[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            // Tuned pure-Rust GEMM: C(m×n) = 1.0*A(m×k)·B(k×n) + 0.0*C, row-major.
262            // dgemm takes isize strides.
263            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            // Indexed inner loops retained for cache-friendly row-major access.
286            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    /// Returns the mean of each column.
309    pub fn column_mean(&self) -> Vec<f64> {
310        crate::stats::column_mean_flat(&self.data, self.rows, self.cols)
311    }
312
313    /// Creates a matrix from a vector of columns.
314    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    /// Select a subset of columns by index (0-based).
342    ///
343    /// ```rust
344    /// use datarust::Matrix;
345    ///
346    /// let m = Matrix::new(vec![
347    ///     vec![1.0, 2.0, 3.0],
348    ///     vec![4.0, 5.0, 6.0],
349    /// ])?;
350    /// let sub = m.select_columns(&[0, 2])?;
351    /// assert_eq!(sub.ncols(), 2);
352    /// assert_eq!(sub.get(0, 0), 1.0);
353    /// assert_eq!(sub.get(0, 1), 3.0);
354    /// # Ok::<_, Box<dyn std::error::Error>>(())
355    /// ```
356    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    /// Select a subset of rows by index (0-based).
385    ///
386    /// ```rust
387    /// use datarust::Matrix;
388    ///
389    /// let m = Matrix::new(vec![
390    ///     vec![10.0],
391    ///     vec![20.0],
392    ///     vec![30.0],
393    /// ])?;
394    /// let sub = m.select_rows(&[0, 2])?;
395    /// assert_eq!(sub.nrows(), 2);
396    /// assert_eq!(sub.get(0, 0), 10.0);
397    /// assert_eq!(sub.get(1, 0), 30.0);
398    /// # Ok::<_, Box<dyn std::error::Error>>(())
399    /// ```
400    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        // Preserve the nested `{"data":[[...]]}` wire format even though the
434        // internal storage is now flat row-major.
435        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/// Row-major matrix of strings used by the categorical encoders.
458#[derive(Debug, Clone, PartialEq)]
459pub struct StrMatrix {
460    pub(crate) data: Vec<Vec<String>>,
461}
462
463impl StrMatrix {
464    /// Creates a string matrix from a nested vector, validating a rectangular shape.
465    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    /// Creates a single-column string matrix from an iterator of values.
485    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    /// Creates a string matrix from an iterator of rows.
498    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    /// Returns the number of rows.
511    #[inline]
512    pub fn nrows(&self) -> usize {
513        self.data.len()
514    }
515
516    /// Returns the number of columns.
517    #[inline]
518    pub fn ncols(&self) -> usize {
519        self.data[0].len()
520    }
521
522    /// Returns the string at row `i`, column `j`.
523    ///
524    /// # Panics
525    ///
526    /// Panics if `i >= nrows()` or `j >= ncols()`. Use
527    /// [`checked_get`](Self::checked_get) for a safe alternative.
528    #[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    /// Returns `Some(element)` at row `i`, column `j`, or `None` if the indices
542    /// are out of bounds.
543    #[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    /// Returns column `j` as a new vector of strings.
549    pub fn column(&self, j: usize) -> Vec<String> {
550        self.data.iter().map(|r| r[j].clone()).collect()
551    }
552
553    /// Returns the row at index `i` as a slice.
554    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    /// Fallibly construct a [`Matrix`] from a nested vector.
591    ///
592    /// Returns an error if the rows are empty or jagged. This replaces the
593    /// previous panicking `From` impl to keep validation consistent with
594    /// [`Matrix::new`].
595    fn try_from(data: Vec<Vec<f64>>) -> Result<Self> {
596        Matrix::new(data)
597    }
598}
599
600/// Compressed Sparse Row (CSR) matrix for memory-efficient storage of
601/// mostly-zero 2-D data, mirroring scipy.sparse `csr_matrix`.
602///
603/// Three arrays define the non-zero entries:
604/// - `indptr` (length `nrows + 1`): row `i` occupies
605///   `indptr[i]..indptr[i+1]` in `indices`/`data`.
606/// - `indices`: column index of each non-zero.
607/// - `data`: value of each non-zero.
608#[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    /// Build a CSR matrix from raw CSR arrays.
620    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    /// Build from `(row, col, value)` triplets. Zero-valued triplets are
666    /// dropped automatically. Within a row, entries are sorted by column.
667    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    /// Create an all-zeros sparse matrix with the given shape.
715    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    /// Returns the number of rows.
729    #[inline]
730    pub fn nrows(&self) -> usize {
731        self.nrows
732    }
733
734    /// Returns the number of columns.
735    #[inline]
736    pub fn ncols(&self) -> usize {
737        self.ncols
738    }
739
740    /// Number of stored non-zero entries.
741    #[inline]
742    pub fn nnz(&self) -> usize {
743        self.data.len()
744    }
745
746    /// Density (fraction of non-zero entries).
747    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    /// Get element at `(i, j)`. Returns 0.0 if not stored.
756    pub fn get(&self, i: usize, j: usize) -> f64 {
757        let start = self.indptr[i];
758        let end = self.indptr[i + 1];
759        // Binary search within the row's column indices.
760        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    /// Returns `Some(element)` at row `i`, column `j`, or `None` if `i` is
768    /// out of bounds or the value is not stored (equivalent to 0.0).
769    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    /// Iterate over `(col, value)` non-zero entries in row `i`.
780    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    /// Convert to a dense [`Matrix`].
790    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        // [[1*5+2*7, 1*6+2*8],[3*5+4*7, 3*6+4*8]] = [[19,22],[43,50]]
906        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}