1use 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 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 pub fn rows(&self) -> usize {
26 self.rows
27 }
28
29 pub fn cols(&self) -> usize {
31 self.cols
32 }
33
34 fn get_index(&self, row: usize, col: usize) -> usize {
36 row * self.cols + col
37 }
38
39 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 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 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 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 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 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 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 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 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 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 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 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}