From b6ad90a40761d27d04ed5447dd0fd2f31adb786e Mon Sep 17 00:00:00 2001 From: hachem Date: Mon, 16 Sep 2024 13:27:22 +0200 Subject: Basic Vector Implementation --- libmu/src/maths/matrix.rs | 6 +- libmu/src/maths/mod.rs | 2 + libmu/src/maths/vector.rs | 246 ++++++++++++++++++++++++++++++++++++++++++++++ mu/src/main.rs | 19 +--- 4 files changed, 252 insertions(+), 21 deletions(-) create mode 100644 libmu/src/maths/vector.rs diff --git a/libmu/src/maths/matrix.rs b/libmu/src/maths/matrix.rs index 8a5465c..7662652 100644 --- a/libmu/src/maths/matrix.rs +++ b/libmu/src/maths/matrix.rs @@ -2,9 +2,9 @@ use super::Float; use core::{fmt, ops}; pub struct Matrix { - data: Vec, - rows: usize, - cols: usize, + pub data: Vec, + pub rows: usize, + pub cols: usize, } impl Matrix { diff --git a/libmu/src/maths/mod.rs b/libmu/src/maths/mod.rs index 3d144d7..651fcf2 100644 --- a/libmu/src/maths/mod.rs +++ b/libmu/src/maths/mod.rs @@ -1,7 +1,9 @@ pub mod complex; pub mod matrix; pub mod numeric_types; +pub mod vector; pub use complex::*; pub use matrix::*; pub use numeric_types::*; +pub use vector::*; diff --git a/libmu/src/maths/vector.rs b/libmu/src/maths/vector.rs new file mode 100644 index 0000000..1504013 --- /dev/null +++ b/libmu/src/maths/vector.rs @@ -0,0 +1,246 @@ +// TODO(Hachem): matrix multiplication with vector + +use super::{Float, Matrix}; +use core::ops; + +pub trait Vector { + fn new(data: Vec) -> Self; + fn get(&self, index: usize) -> T; + fn set(&mut self, index: usize, value: T); + fn size(&self) -> usize; + + fn dot(&self, other: &Self) -> T; + fn norm(&self) -> T; + + fn max(&self) -> T; + fn min(&self) -> T; + fn sum(&self) -> T; +} + +pub trait VectorMatrix { + fn from_matrix(matrix: &Matrix) -> Self; + fn to_matrix(&self) -> Matrix; +} + +pub struct VectorImpl(Vec); +pub type RowVector = VectorImpl; +pub type ColumnVector = VectorImpl; + +impl ColumnVector { + pub fn transpose(&self) -> RowVector { + RowVector::new(self.0.clone()) + } +} + +impl RowVector { + pub fn transpose(&self) -> ColumnVector { + ColumnVector::new(self.0.clone()) + } +} + +impl Vector for VectorImpl { + fn new(data: Vec) -> Self { + Self(data) + } + + fn get(&self, index: usize) -> T { + self.0[index] + } + + fn set(&mut self, index: usize, value: T) { + self.0[index] = value; + } + + fn size(&self) -> usize { + self.0.len() + } + + fn dot(&self, other: &Self) -> T { + self.0 + .iter() + .zip(other.0.iter()) + .map(|(a, b)| *a * *b) + .fold(T::zero(), |acc, x| acc + x) + } + + fn norm(&self) -> T { + self.0 + .iter() + .map(|x| *x * *x) + .fold(T::zero(), |acc, x| acc + x) + .sqrt() + } + + fn max(&self) -> T { + *self + .0 + .iter() + .max_by(|a, b| a.partial_cmp(b).unwrap()) + .unwrap_or(&T::zero()) + } + + fn min(&self) -> T { + *self + .0 + .iter() + .min_by(|a, b| a.partial_cmp(b).unwrap()) + .unwrap_or(&T::zero()) + } + + fn sum(&self) -> T { + self.0.iter().fold(T::zero(), |acc, x| acc + *x) + } +} + +impl VectorImpl { + pub fn add_to(&self, other: &Self) -> Option> { + if self.size() != other.size() { + return None; + } + + let mut result = VectorImpl::new(vec![T::zero(); ROWS * COLS]); + + for i in 0..self.size() { + let sum = self.get(i) + other.get(i); + result.set(i, sum); + } + + Some(result) + } + + pub fn subtract(&self, other: &Self) -> Option> { + if self.size() != other.size() { + return None; + } + + let mut result = VectorImpl::new(vec![T::zero(); ROWS * COLS]); + + for i in 0..self.size() { + let sum = self.get(i) - other.get(i); + result.set(i, sum); + } + + Some(result) + } + + pub fn scale(&self, scalar: T) -> VectorImpl { + let mut result = VectorImpl::new(vec![T::zero(); ROWS * COLS]); + + for i in 0..self.size() { + let product = self.get(i) * scalar; + result.set(i, product); + } + + result + } +} + +impl ops::Index + for VectorImpl +{ + type Output = T; + + fn index(&self, index: usize) -> &Self::Output { + &self.0[index] + } +} + +impl ops::IndexMut + for VectorImpl +{ + fn index_mut(&mut self, index: usize) -> &mut Self::Output { + &mut self.0[index] + } +} + +impl ops::Add<&VectorImpl> + for VectorImpl +{ + type Output = Option>; + + fn add(self, other: &VectorImpl) -> Self::Output { + self.add_to(other) + } +} + +impl ops::Sub<&VectorImpl> + for VectorImpl +{ + type Output = Option>; + + fn sub(self, other: &VectorImpl) -> Self::Output { + self.subtract(other) + } +} + +impl ops::Mul for VectorImpl { + type Output = VectorImpl; + + fn mul(self, scalar: T) -> Self::Output { + self.scale(scalar) + } +} + +impl ops::Div for VectorImpl { + type Output = VectorImpl; + + fn div(self, scalar: T) -> Self::Output { + self.scale(T::one() / scalar) + } +} + +impl ops::AddAssign> + for VectorImpl +{ + fn add_assign(&mut self, other: VectorImpl) { + if let Some(result) = self.add_to(&other) { + *self = result; + } + } +} + +impl ops::SubAssign> + for VectorImpl +{ + fn sub_assign(&mut self, other: VectorImpl) { + if let Some(result) = self.subtract(&other) { + *self = result; + } + } +} + +impl ops::MulAssign + for VectorImpl +{ + fn mul_assign(&mut self, scalar: T) { + *self = self.scale(scalar); + } +} + +impl ops::DivAssign + for VectorImpl +{ + fn div_assign(&mut self, scalar: T) { + *self = self.scale(T::one() / scalar); + } +} + +impl VectorMatrix for RowVector { + fn to_matrix(&self) -> Matrix { + Matrix::new(1, self.size(), self.0.clone()) + } + + fn from_matrix(matrix: &Matrix) -> Self { + Self::new(matrix.data.clone()) + } +} + +impl VectorMatrix for ColumnVector { + fn to_matrix(&self) -> Matrix { + Matrix::new(self.size(), 1, self.0.clone()) + } + + fn from_matrix(matrix: &Matrix) -> Self { + Self::new(matrix.data.clone()) + } +} diff --git a/mu/src/main.rs b/mu/src/main.rs index e3f23b2..f328e4d 100644 --- a/mu/src/main.rs +++ b/mu/src/main.rs @@ -1,18 +1 @@ -use libmu::Matrix; - -fn main() { - let mat1 = Matrix::::new(2, 2, vec![1.0, 2.0, 3.0, 4.0]); - let mat2 = Matrix::::new(2, 2, vec![5.0, 6.0, 7.0, 8.0]); - - let dot_product = mat1.dot(&mat2); - let kronecker_product = mat1.kronecker(&mat2); - - println!("Matrix 1:\n{}", mat1); - println!("Matrix 2:\n{}", mat2); - - if let Some(dp) = dot_product { - println!("Dot Product Matrix:\n{}", dp); - } - - println!("Kronecker Product Matrix:\n{}", kronecker_product); -} +fn main() {} -- cgit v1.3