Skip to main content

wassily_geometry/
matrix.rs

1//! A zero indexed row-major matrix backed by a `Vec`.
2//! Allows acces to elements of matrix A by A\[i\]\[j\].
3use num_traits::{zero, AsPrimitive, Float, One, Zero};
4use std::ops::{Add, Div, Index, IndexMut, Mul};
5
6#[derive(Debug)]
7pub struct Matrix<T> {
8    rows: usize,
9    cols: usize,
10    pub data: Vec<T>,
11}
12
13impl<T> Matrix<T> {
14    /// Create a new matrix with the given number of rows and columns.
15    pub fn new<U: AsPrimitive<usize>>(rows: U, cols: U, data: Vec<T>) -> Self {
16        assert_eq!(rows.as_() * cols.as_(), data.len());
17        Self {
18            rows: rows.as_(),
19            cols: cols.as_(),
20            data,
21        }
22    }
23
24    /// The number of rows in the matrix.
25    pub fn rows(&self) -> usize {
26        self.rows
27    }
28
29    /// The number of columns in the matrix.
30    pub fn cols(&self) -> usize {
31        self.cols
32    }
33
34    /// The index in the underyling data vector for the given row and column.
35    fn get_index(&self, row: usize, col: usize) -> usize {
36        row * self.cols + col
37    }
38
39    /// Create a new matrix using a function to generate the data.
40    pub fn generate<F, U: AsPrimitive<usize>>(rows: U, cols: U, generator: F) -> Self
41    where
42        F: Fn(usize, usize) -> T,
43        U: AsPrimitive<usize>,
44    {
45        let mut data: Vec<T> = vec![];
46        for r in 0..rows.as_() {
47            for c in 0..cols.as_() {
48                data.push(generator(r, c))
49            }
50        }
51        Matrix {
52            rows: rows.as_(),
53            cols: cols.as_(),
54            data,
55        }
56    }
57
58    /// Get a reference to the element at the given row and column.
59    pub fn get_ref(&self, row: usize, col: usize) -> Option<&T> {
60        if row < self.rows && col < self.cols {
61            Some(&self.data[self.get_index(row, col)])
62        } else {
63            None
64        }
65    }
66
67    /// Insert a value at the given row and column.
68    pub fn put(&mut self, row: usize, col: usize, item: T) -> bool {
69        if row >= self.rows || col >= self.cols {
70            false
71        } else {
72            let idx = self.get_index(row, col);
73            self.data[idx] = item;
74            true
75        }
76    }
77
78    /// Is this a valid row and column?
79    pub fn valid<U: Into<usize>>(&self, row: U, col: U) -> bool {
80        row.into() < self.rows() && col.into() < self.cols()
81    }
82}
83
84impl<T> Matrix<T>
85where
86    T: Clone + Copy,
87{
88    /// Create a new matrix with the given number of rows and columns filled with
89    /// a given value.
90    pub fn fill(rows: usize, cols: usize, datum: T) -> Self {
91        let data = vec![datum; rows * cols];
92        Self { rows, cols, data }
93    }
94
95    /// Return the element at the given row and column.
96    pub fn get(&self, row: usize, col: usize) -> Option<T> {
97        if row < self.rows && col < self.cols {
98            let idx = self.get_index(row, col);
99            Some(self.data[idx])
100        } else {
101            None
102        }
103    }
104
105    /// Insert a column at position n.
106    pub fn insert_col(&self, n: usize, column: Vec<T>) -> Self {
107        assert_eq!(column.len(), self.rows());
108        Matrix::generate(self.rows(), self.cols() + 1, |r, c| match c.cmp(&n) {
109            std::cmp::Ordering::Less => self[r][c],
110            std::cmp::Ordering::Equal => column[r],
111            std::cmp::Ordering::Greater => self[r][c - 1],
112        })
113    }
114
115    /// Insert a row at position n.
116    pub fn insert_row(&self, n: usize, row: Vec<T>) -> Self {
117        assert_eq!(row.len(), self.cols());
118        Matrix::generate(self.rows() + 1, self.cols(), |r, c| match r.cmp(&n) {
119            std::cmp::Ordering::Less => self[r][c],
120            std::cmp::Ordering::Equal => row[c],
121            std::cmp::Ordering::Greater => self[r - 1][c],
122        })
123    }
124
125    // Transpose the matrix.
126    pub fn transpose(&self) -> Self {
127        Matrix::generate(self.cols(), self.rows(), |r, c| self[c][r])
128    }
129}
130
131impl<T> Matrix<T>
132where
133    T: Zero + Clone,
134{
135    /// A matrix of all zeros.
136    pub fn zeros<U: AsPrimitive<usize>>(rows: U, cols: U) -> Self {
137        let data = vec![T::zero(); rows.as_() * cols.as_()];
138        Self {
139            rows: rows.as_(),
140            cols: cols.as_(),
141            data,
142        }
143    }
144}
145
146impl<T> Matrix<T>
147where
148    T: One + Clone,
149{
150    /// A matrix of all ones.
151    pub fn ones<U: AsPrimitive<usize>>(rows: U, cols: U) -> Self {
152        let data = vec![T::one(); rows.as_() * cols.as_()];
153        Self {
154            rows: rows.as_(),
155            cols: cols.as_(),
156            data,
157        }
158    }
159}
160
161impl<T> Matrix<T>
162where
163    T: Float,
164{
165    /// Convolves the matrix with the given kernel.
166    pub fn convolve(&self, kernel: &Matrix<T>) -> Matrix<T> {
167        let mut m: Matrix<T> = Matrix {
168            rows: self.rows,
169            cols: self.cols,
170            data: self.data.clone(),
171        };
172        let k = kernel.rows / 2;
173        for i in k..self.rows - k {
174            for j in k..self.cols - k {
175                let mut acc = T::zero();
176                for r in 0..kernel.rows {
177                    for c in 0..kernel.cols {
178                        acc = acc + self[i - k + r][j - k + c] * kernel[r][c];
179                    }
180                }
181                m[i][j] = acc;
182            }
183        }
184        m
185    }
186}
187
188impl<T> Index<usize> for Matrix<T> {
189    type Output = [T];
190    fn index(&self, index: usize) -> &Self::Output {
191        let start = index * self.cols;
192        &self.data[start..start + self.cols]
193    }
194}
195
196impl<T> IndexMut<usize> for Matrix<T> {
197    fn index_mut(&mut self, index: usize) -> &mut Self::Output {
198        let start = index * self.cols;
199        &mut self.data[start..start + self.cols]
200    }
201}
202
203impl<T> Mul<T> for &Matrix<T>
204where
205    T: Mul<Output = T> + Zero + Copy,
206{
207    type Output = Matrix<T>;
208
209    fn mul(self, rhs: T) -> Self::Output {
210        let mut m: Matrix<T> = Matrix::fill(self.rows(), self.cols(), zero());
211        for r in 0..self.rows() {
212            for c in 0..self.cols() {
213                m[r][c] = self[r][c] * rhs;
214            }
215        }
216        m
217    }
218}
219
220impl<T> Div<T> for &Matrix<T>
221where
222    T: Div<Output = T> + Zero + Copy,
223{
224    type Output = Matrix<T>;
225
226    fn div(self, rhs: T) -> Self::Output {
227        let mut m: Matrix<T> = Matrix::fill(self.rows(), self.cols(), zero());
228        for r in 0..self.rows() {
229            for c in 0..self.cols() {
230                m[r][c] = self[r][c] / rhs;
231            }
232        }
233        m
234    }
235}
236
237impl<T> Mul<&Vec<T>> for &Matrix<T>
238where
239    T: Add<Output = T> + Mul<Output = T> + Zero + Copy,
240{
241    type Output = Vec<T>;
242
243    fn mul(self, rhs: &Vec<T>) -> Self::Output {
244        assert_eq!(self.cols(), rhs.len());
245        let mut v: Vec<T> = vec![];
246        for r in 0..self.rows() {
247            v.push(
248                self[r]
249                    .iter()
250                    .zip(rhs)
251                    .fold(zero(), |accum: T, item| accum + *item.0 * *item.1),
252            );
253        }
254        v
255    }
256}
257
258impl<T> Mul<&Matrix<T>> for &Matrix<T>
259where
260    T: Add<Output = T> + Mul<Output = T> + Zero + Copy,
261{
262    type Output = Matrix<T>;
263
264    fn mul(self, rhs: &Matrix<T>) -> Self::Output {
265        assert_eq!(self.cols(), rhs.rows());
266        let mut m: Matrix<T> = Matrix::fill(self.rows(), rhs.cols(), zero());
267        for r in 0..self.rows() {
268            for c in 0..rhs.cols() {
269                let mut a = zero();
270                for i in 0..self.cols() {
271                    a = a + self[r][i] * rhs[i][c];
272                }
273                m[r][c] = a;
274            }
275        }
276        m
277    }
278}
279
280impl<T> PartialEq for Matrix<T>
281where
282    T: PartialEq,
283{
284    fn eq(&self, other: &Self) -> bool {
285        self.rows == other.rows && self.cols == other.cols && self.data == other.data
286    }
287}
288
289impl<T> Eq for Matrix<T> where T: Eq {}
290
291#[cfg(test)]
292mod tests {
293    use std::vec;
294
295    use super::*;
296    #[test]
297    fn gen_test() {
298        let m = Matrix::generate(2, 3, |i, j| (i, j));
299        assert_eq!(m.data, vec![(0, 0), (0, 1), (0, 2), (1, 0), (1, 1), (1, 2)]);
300    }
301
302    #[test]
303    fn get_test() {
304        let m = Matrix::generate(2, 3, |i, j| (i, j));
305        assert_eq!(m.get(1, 1), Some((1, 1)));
306        assert_eq!(m.get(2, 1), None);
307        assert_eq!(m.get(0, 3), None);
308    }
309
310    #[test]
311    fn get_ref_test() {
312        let m = Matrix::generate(2, 3, |i, j| (i, j));
313        assert_eq!(m.get_ref(1, 1), Some(&(1, 1)));
314        assert_eq!(m.get_ref(2, 1), None);
315        assert_eq!(m.get_ref(0, 3), None);
316    }
317
318    #[test]
319    fn put_test() {
320        let mut m = Matrix::generate(2, 3, |i, j| (i, j));
321        assert_eq!(m.put(1, 1, (5, 5)), true);
322        assert_eq!(m.get(1, 1), Some((5, 5)));
323    }
324
325    #[test]
326    fn fill_test() {
327        let m = Matrix::fill(1, 2, true);
328        assert_eq!(m.data, vec![true, true]);
329    }
330
331    #[test]
332    fn index_test() {
333        let m = Matrix::generate(2, 3, |i, j| (i, j));
334        assert_eq!(m[1][1], (1, 1));
335    }
336
337    #[test]
338    fn indexmut_test() {
339        let mut m = Matrix::generate(2, 3, |i, j| (i, j));
340        m[1][1] = (5, 5);
341        assert_eq!(m[1][1], (5, 5));
342    }
343
344    #[test]
345    fn convolve_test() {
346        let m = Matrix::<f32>::ones(5, 5);
347        let k = Matrix::<f32>::ones(3, 3);
348        let c = m.convolve(&k);
349        assert_eq!(c[0][0], 1.0);
350        assert_eq!(c[1][1], 9.0);
351    }
352
353    #[test]
354    fn mul_test() {
355        let m = Matrix::new(2, 2, vec![1, 2, 3, 4]);
356        assert_eq!(&m * &vec![5, 10], vec![25, 55]);
357    }
358
359    #[test]
360    fn mul_mat_test() {
361        let m1 = Matrix::new(3, 2, vec![1, 2, 3, 4, 5, 6]);
362        let m2 = Matrix::new(2, 2, vec![5, 10, 50, 100]);
363        assert_eq!((&m1 * &m2).data, vec![105, 210, 215, 430, 325, 650,]);
364    }
365
366    #[test]
367    fn mul_scalar_test() {
368        let m = Matrix::new(2, 2, vec![1, 2, 3, 4]);
369        assert_eq!((&m * 2).data, vec![2, 4, 6, 8]);
370    }
371
372    #[test]
373    fn insert_col_test() {
374        let m = Matrix::new(2, 2, vec![1, 2, 3, 4]);
375        let m1 = m.insert_col(1, vec![5, 5]);
376        assert_eq!(m1.data, vec![1, 5, 2, 3, 5, 4]);
377    }
378
379    #[test]
380    fn insert_row_test() {
381        let m = Matrix::new(2, 2, vec![1, 2, 3, 4]);
382        let m1 = m.insert_row(1, vec![5, 5]);
383        assert_eq!(m1.data, vec![1, 2, 5, 5, 3, 4]);
384    }
385
386    #[test]
387    fn transpose_test() {
388        let m = Matrix::new(2, 2, vec![1, 2, 3, 4]);
389        let m1 = m.transpose();
390        assert_eq!(m1.data, vec![1, 3, 2, 4]);
391    }
392}