1use std::ops::{Index, IndexMut};
13
14use num_traits::Zero;
15use tenferro_tensor::{Rank, TensorScalar, TypedTensor};
16
17#[derive(Clone, Debug, PartialEq)]
19pub struct Mat<T> {
20 data: Vec<T>,
21 shape: (usize, usize),
22}
23
24impl<T> Mat<T> {
25 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 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 #[inline]
59 pub fn shape(&self) -> &(usize, usize) {
60 &self.shape
61 }
62
63 #[inline]
65 pub fn dims(&self) -> [usize; 2] {
66 [self.shape.0, self.shape.1]
67 }
68
69 #[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 #[inline]
101 pub fn as_slice(&self) -> &[T] {
102 &self.data
103 }
104
105 #[inline]
107 pub fn as_mut_slice(&mut self) -> &mut [T] {
108 &mut self.data
109 }
110
111 pub fn into_vec(self) -> Vec<T> {
113 self.data
114 }
115
116 #[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 #[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 pub fn iter(&self) -> std::slice::Iter<'_, T> {
132 self.data.iter()
133 }
134
135 pub fn iter_mut(&mut self) -> std::slice::IterMut<'_, T> {
137 self.data.iter_mut()
138 }
139
140 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 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 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 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 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#[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 pub fn zeros(shape: [usize; 2]) -> Self {
206 Self::from_elem(shape, T::zero())
207 }
208}
209
210impl<T: TensorScalar> Mat<T> {
211 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 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#[derive(Clone, Debug, PartialEq)]
260pub struct Mat3<T> {
261 data: Vec<T>,
262 shape: (usize, usize, usize),
263}
264
265impl<T> Mat3<T> {
266 #[inline]
268 pub fn shape(&self) -> &(usize, usize, usize) {
269 &self.shape
270 }
271
272 #[inline]
274 pub fn as_slice(&self) -> &[T] {
275 &self.data
276 }
277
278 pub fn into_vec(self) -> Vec<T> {
280 self.data
281 }
282
283 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 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 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}