From 5abda23c3ed21db3c0a4541b85b6e2e916ce569e Mon Sep 17 00:00:00 2001 From: hachem Date: Mon, 16 Sep 2024 20:53:46 +0200 Subject: Matrix and vector products + macros --- libmu/src/maths/vector.rs | 90 +++++++++++++++++++++++++++++++++++++++++++++-- 1 file changed, 88 insertions(+), 2 deletions(-) (limited to 'libmu/src/maths/vector.rs') diff --git a/libmu/src/maths/vector.rs b/libmu/src/maths/vector.rs index 1504013..682907e 100644 --- a/libmu/src/maths/vector.rs +++ b/libmu/src/maths/vector.rs @@ -1,8 +1,26 @@ -// TODO(Hachem): matrix multiplication with vector - use super::{Float, Matrix}; use core::ops; +#[macro_export] +macro_rules! row_vector { + ($($x:expr),*) => { + RowVector::new(vec![$($x),*]) + }; + ($($x:expr,)*) => { + RowVector::new(vec![$($x),*]) + }; +} + +#[macro_export] +macro_rules! column_vector { + ($($x:expr),*) => { + ColumnVector::new(vec![$($x),*]) + }; + ($($x:expr,)*) => { + ColumnVector::new(vec![$($x),*]) + }; +} + pub trait Vector { fn new(data: Vec) -> Self; fn get(&self, index: usize) -> T; @@ -27,12 +45,48 @@ pub type RowVector = VectorImpl; pub type ColumnVector = VectorImpl; impl ColumnVector { + pub fn mul_matrix(&self, matrix: &Matrix) -> Option> { + if matrix.cols != self.size() { + return None; + } + + let mut result = ColumnVector::new(vec![T::zero(); matrix.rows]); + + for i in 0..matrix.rows { + let mut sum = T::zero(); + for j in 0..matrix.cols { + sum = sum + (matrix.get(i, j) * self.get(j)); + } + result.set(i, sum); + } + + Some(result) + } + pub fn transpose(&self) -> RowVector { RowVector::new(self.0.clone()) } } impl RowVector { + pub fn mul_matrix(&self, matrix: &Matrix) -> Option> { + if self.size() != matrix.rows { + return None; + } + + let mut result = RowVector::new(vec![T::zero(); matrix.cols]); + + for j in 0..matrix.cols { + let mut sum = T::zero(); + for i in 0..matrix.rows { + sum = sum + (self.get(i) * matrix.get(i, j)); + } + result.set(j, sum); + } + + Some(result) + } + pub fn transpose(&self) -> ColumnVector { ColumnVector::new(self.0.clone()) } @@ -244,3 +298,35 @@ impl VectorMatrix for ColumnVector { Self::new(matrix.data.clone()) } } + +impl ops::Mul<&Matrix> for RowVector { + type Output = Option>; + + fn mul(self, matrix: &Matrix) -> Self::Output { + self.mul_matrix(matrix) + } +} + +impl ops::Mul<&Matrix> for ColumnVector { + type Output = Option>; + + fn mul(self, matrix: &Matrix) -> Self::Output { + self.mul_matrix(matrix) + } +} + +impl ops::MulAssign<&Matrix> for RowVector { + fn mul_assign(&mut self, matrix: &Matrix) { + if let Some(result) = self.mul_matrix(matrix) { + *self = result; + } + } +} + +impl ops::MulAssign<&Matrix> for ColumnVector { + fn mul_assign(&mut self, matrix: &Matrix) { + if let Some(result) = self.mul_matrix(matrix) { + *self = result; + } + } +} -- cgit v1.3