Skip to main content

sparse_ir_core/
matrix.rs

1//! Small column-major dense containers for internal numerics.
2//!
3//! [`Mat`] and [`Mat3`] hold generic scalars (including `Df64`) for
4//! the SVE, polynomial, and kernel-matrix code paths. They are plain
5//! column-major `Vec<T>` buffers with infallible indexing: the first index
6//! varies fastest in memory, matching Fortran, BLAS, and tenferro.
7//!
8//! Public APIs that expose `f64`/`Complex<f64>` matrices use
9//! [`crate::Matrix`] (`TypedTensor<T, Rank<2>>`); [`Mat::into_typed`] and
10//! [`Mat::from_typed`] move between the two without copying.
11
12use std::ops::{Index, IndexMut};
13
14use num_traits::Zero;
15use tenferro_tensor::{Rank, TensorScalar, TypedTensor};
16
17/// Column-major dense matrix with infallible indexing.
18#[derive(Clone, Debug, PartialEq)]
19pub struct Mat<T> {
20    data: Vec<T>,
21    shape: (usize, usize),
22}
23
24impl<T> Mat<T> {
25    /// Wrap a column-major buffer.
26    ///
27    /// # Panics
28    /// Panics if `data.len() != nrows * ncols`; callers construct the buffer
29    /// with the matching size.
30    pub fn from_vec_col_major(shape: [usize; 2], data: Vec<T>) -> Self {
31        assert_eq!(
32            data.len(),
33            shape[0] * shape[1],
34            "Mat::from_vec_col_major: buffer length does not match shape"
35        );
36        Self {
37            data,
38            shape: (shape[0], shape[1]),
39        }
40    }
41
42    /// Build a matrix by evaluating `f(&[i, j])` for every element.
43    pub fn from_fn<F: FnMut(&[usize]) -> T>(shape: [usize; 2], mut f: F) -> Self {
44        let (m, n) = (shape[0], shape[1]);
45        let mut data = Vec::with_capacity(m * n);
46        for j in 0..n {
47            for i in 0..m {
48                data.push(f(&[i, j]));
49            }
50        }
51        Self {
52            data,
53            shape: (m, n),
54        }
55    }
56
57    /// Shape as `(nrows, ncols)`.
58    #[inline]
59    pub fn shape(&self) -> &(usize, usize) {
60        &self.shape
61    }
62
63    /// Shape as `[nrows, ncols]`, the form [`Mat::from_fn`] takes.
64    #[inline]
65    pub fn dims(&self) -> [usize; 2] {
66        [self.shape.0, self.shape.1]
67    }
68
69    /// Extent of axis `axis` (0 = rows, 1 = columns).
70    #[inline]
71    pub fn dim(&self, axis: usize) -> usize {
72        match axis {
73            0 => self.shape.0,
74            1 => self.shape.1,
75            _ => panic!("Mat::dim: axis {axis} out of range for a matrix"),
76        }
77    }
78
79    #[inline]
80    pub fn nrows(&self) -> usize {
81        self.shape.0
82    }
83
84    #[inline]
85    pub fn ncols(&self) -> usize {
86        self.shape.1
87    }
88
89    #[inline]
90    pub fn len(&self) -> usize {
91        self.data.len()
92    }
93
94    #[inline]
95    pub fn is_empty(&self) -> bool {
96        self.data.is_empty()
97    }
98
99    /// Column-major storage.
100    #[inline]
101    pub fn as_slice(&self) -> &[T] {
102        &self.data
103    }
104
105    /// Mutable column-major storage.
106    #[inline]
107    pub fn as_mut_slice(&mut self) -> &mut [T] {
108        &mut self.data
109    }
110
111    /// Consume into the column-major buffer.
112    pub fn into_vec(self) -> Vec<T> {
113        self.data
114    }
115
116    /// Contiguous column `j`.
117    #[inline]
118    pub fn col(&self, j: usize) -> &[T] {
119        let m = self.shape.0;
120        &self.data[j * m..(j + 1) * m]
121    }
122
123    /// Mutable contiguous column `j`.
124    #[inline]
125    pub fn col_mut(&mut self, j: usize) -> &mut [T] {
126        let m = self.shape.0;
127        &mut self.data[j * m..(j + 1) * m]
128    }
129
130    /// Iterate over elements in column-major order.
131    pub fn iter(&self) -> std::slice::Iter<'_, T> {
132        self.data.iter()
133    }
134
135    /// Iterate mutably over elements in column-major order.
136    pub fn iter_mut(&mut self) -> std::slice::IterMut<'_, T> {
137        self.data.iter_mut()
138    }
139
140    /// Apply `f` element-wise.
141    pub fn map<U, F: FnMut(&T) -> U>(&self, f: F) -> Mat<U> {
142        Mat {
143            data: self.data.iter().map(f).collect(),
144            shape: self.shape,
145        }
146    }
147}
148
149impl<T: Clone> Mat<T> {
150    /// Matrix filled with `elem`.
151    pub fn from_elem(shape: [usize; 2], elem: T) -> Self {
152        Self {
153            data: vec![elem; shape[0] * shape[1]],
154            shape: (shape[0], shape[1]),
155        }
156    }
157
158    /// Transposed copy.
159    pub fn transpose(&self) -> Self {
160        let (m, n) = self.shape;
161        Self::from_fn([n, m], |idx| self.data[idx[1] + m * idx[0]].clone())
162    }
163
164    /// Copy of the leading `ncols` columns.
165    pub fn leading_cols(&self, ncols: usize) -> Self {
166        assert!(ncols <= self.shape.1, "Mat::leading_cols: too many columns");
167        Self {
168            data: self.data[..self.shape.0 * ncols].to_vec(),
169            shape: (self.shape.0, ncols),
170        }
171    }
172}
173
174impl<T: Clone> Mat<T> {
175    /// Build from a list of rows (row-literal order, column-major storage).
176    ///
177    /// # Panics
178    /// Panics if the rows have different lengths.
179    pub fn from_rows(rows: Vec<Vec<T>>) -> Self {
180        let m = rows.len();
181        let n = rows.first().map_or(0, Vec::len);
182        assert!(
183            rows.iter().all(|r| r.len() == n),
184            "Mat::from_rows: ragged rows"
185        );
186        Self::from_fn([m, n], |idx| rows[idx[0]][idx[1]].clone())
187    }
188}
189
190/// Row-literal matrix constructor: `mat![[a, b], [c, d]]`.
191#[allow(unused_macros)]
192#[doc(hidden)]
193#[macro_export]
194macro_rules! mat {
195    ($([$($x:expr),* $(,)?]),+ $(,)?) => {
196        $crate::matrix::Mat::from_rows(vec![$(vec![$($x),*]),+])
197    };
198}
199#[allow(unused_imports)]
200#[doc(hidden)]
201pub use crate::mat;
202
203impl<T: Clone + Zero> Mat<T> {
204    /// Zero matrix.
205    pub fn zeros(shape: [usize; 2]) -> Self {
206        Self::from_elem(shape, T::zero())
207    }
208}
209
210impl<T: TensorScalar> Mat<T> {
211    /// Move into a tenferro rank-2 tensor without copying.
212    pub fn into_typed(self) -> TypedTensor<T, Rank<2>> {
213        let (m, n) = self.shape;
214        TypedTensor::from_vec_col_major([m, n], self.data)
215            .expect("Mat invariant: buffer length equals nrows * ncols")
216    }
217
218    /// Copy a host-resident tenferro matrix.
219    ///
220    /// # Errors
221    /// Returns an error when the tensor is not host-resident compact
222    /// column-major storage.
223    pub fn from_typed(t: &TypedTensor<T, Rank<2>>) -> Result<Self, tenferro_tensor::Error> {
224        let view = t.host_col_major_view()?;
225        let shape = *view.shape();
226        Ok(Self {
227            data: view.as_slice().to_vec(),
228            shape: (shape[0], shape[1]),
229        })
230    }
231}
232
233impl<T> Index<[usize; 2]> for Mat<T> {
234    type Output = T;
235    #[inline(always)]
236    fn index(&self, idx: [usize; 2]) -> &T {
237        debug_assert!(idx[0] < self.shape.0 && idx[1] < self.shape.1);
238        &self.data[idx[0] + self.shape.0 * idx[1]]
239    }
240}
241
242impl<T> IndexMut<[usize; 2]> for Mat<T> {
243    #[inline(always)]
244    fn index_mut(&mut self, idx: [usize; 2]) -> &mut T {
245        debug_assert!(idx[0] < self.shape.0 && idx[1] < self.shape.1);
246        &mut self.data[idx[0] + self.shape.0 * idx[1]]
247    }
248}
249
250impl<T> Index<&[usize]> for Mat<T> {
251    type Output = T;
252    #[inline(always)]
253    fn index(&self, idx: &[usize]) -> &T {
254        &self[[idx[0], idx[1]]]
255    }
256}
257
258/// Column-major dense rank-3 array with infallible indexing.
259#[derive(Clone, Debug, PartialEq)]
260pub struct Mat3<T> {
261    data: Vec<T>,
262    shape: (usize, usize, usize),
263}
264
265impl<T> Mat3<T> {
266    /// Shape as `(n0, n1, n2)`.
267    #[inline]
268    pub fn shape(&self) -> &(usize, usize, usize) {
269        &self.shape
270    }
271
272    /// Column-major storage.
273    #[inline]
274    pub fn as_slice(&self) -> &[T] {
275        &self.data
276    }
277
278    /// Consume into the column-major buffer.
279    pub fn into_vec(self) -> Vec<T> {
280        self.data
281    }
282
283    /// Wrap a column-major buffer.
284    ///
285    /// # Panics
286    /// Panics if the buffer length does not match the shape.
287    pub fn from_vec_col_major(shape: [usize; 3], data: Vec<T>) -> Self {
288        assert_eq!(
289            data.len(),
290            shape[0] * shape[1] * shape[2],
291            "Mat3::from_vec_col_major: buffer length does not match shape"
292        );
293        Self {
294            data,
295            shape: (shape[0], shape[1], shape[2]),
296        }
297    }
298}
299
300impl<T: Clone> Mat3<T> {
301    /// Array filled with `elem`.
302    pub fn from_elem(shape: [usize; 3], elem: T) -> Self {
303        Self {
304            data: vec![elem; shape[0] * shape[1] * shape[2]],
305            shape: (shape[0], shape[1], shape[2]),
306        }
307    }
308}
309
310impl<T: Clone + Zero> Mat3<T> {
311    /// Zero array.
312    pub fn zeros(shape: [usize; 3]) -> Self {
313        Self::from_elem(shape, T::zero())
314    }
315}
316
317impl<T> Index<[usize; 3]> for Mat3<T> {
318    type Output = T;
319    #[inline(always)]
320    fn index(&self, idx: [usize; 3]) -> &T {
321        let (a, b, _) = self.shape;
322        &self.data[idx[0] + a * (idx[1] + b * idx[2])]
323    }
324}
325
326impl<T> IndexMut<[usize; 3]> for Mat3<T> {
327    #[inline(always)]
328    fn index_mut(&mut self, idx: [usize; 3]) -> &mut T {
329        let (a, b, _) = self.shape;
330        &mut self.data[idx[0] + a * (idx[1] + b * idx[2])]
331    }
332}
333
334#[cfg(test)]
335mod tests {
336    use super::*;
337
338    #[test]
339    fn col_major_layout_and_transpose() {
340        let a = Mat::from_fn([2, 3], |idx| (10 * idx[0] + idx[1]) as f64);
341        assert_eq!(a.as_slice(), &[0.0, 10.0, 1.0, 11.0, 2.0, 12.0]);
342        assert_eq!(a[[1, 2]], 12.0);
343        assert_eq!(a.col(1), &[1.0, 11.0]);
344        let t = a.transpose();
345        assert_eq!(*t.shape(), (3, 2));
346        assert_eq!(t[[2, 1]], 12.0);
347    }
348
349    #[test]
350    fn typed_roundtrip_is_zero_copy_layout() {
351        let a = Mat::from_fn([2, 2], |idx| (idx[0] + 2 * idx[1]) as f64);
352        let t = a.clone().into_typed();
353        assert_eq!(t.shape(), &[2, 2]);
354        assert_eq!(*t.get2(1, 0).unwrap(), 1.0);
355        assert_eq!(Mat::from_typed(&t).unwrap(), a);
356    }
357
358    #[test]
359    fn rank3_indexing() {
360        let mut a = Mat3::<f64>::zeros([2, 3, 4]);
361        a[[1, 2, 3]] = 5.0;
362        assert_eq!(a.as_slice()[1 + 2 * (2 + 3 * 3)], 5.0);
363    }
364}