diff options
| author | hachem <im@hachem.wtf> | 2024-09-16 09:41:54 +0200 |
|---|---|---|
| committer | hachem <im@hachem.wtf> | 2024-09-16 09:41:54 +0200 |
| commit | 6d750e1ca95f7c4dcf4423b9cebc30fdc471d083 (patch) | |
| tree | 45e3206a25556b3a4b21b3e873f85137fb467f49 | |
| parent | fea812469ad0a2099e100a4f232b121f5adeec79 (diff) | |
Matrix struct
| -rw-r--r-- | libmu/src/lib.rs | 2 | ||||
| -rw-r--r-- | libmu/src/maths/matrix.rs | 221 | ||||
| -rw-r--r-- | libmu/src/maths/mod.rs | 2 | ||||
| -rw-r--r-- | libmu/src/maths/numeric_types.rs | 3 | ||||
| -rw-r--r-- | mu/src/main.rs | 19 |
5 files changed, 242 insertions, 5 deletions
diff --git a/libmu/src/lib.rs b/libmu/src/lib.rs index 83fb75c..1dea3c3 100644 --- a/libmu/src/lib.rs +++ b/libmu/src/lib.rs @@ -1,5 +1,3 @@ -#![no_std] - mod maths; pub use maths::*; diff --git a/libmu/src/maths/matrix.rs b/libmu/src/maths/matrix.rs new file mode 100644 index 0000000..8a5465c --- /dev/null +++ b/libmu/src/maths/matrix.rs @@ -0,0 +1,221 @@ +use super::Float; +use core::{fmt, ops}; + +pub struct Matrix<T: Float> { + data: Vec<T>, + rows: usize, + cols: usize, +} + +impl<T: Float> Matrix<T> { + pub fn new(rows: usize, cols: usize, data: Vec<T>) -> Self { + Matrix { data, rows, cols } + } + + pub fn get(&self, row: usize, col: usize) -> T { + self.data[row * self.cols + col] + } + + pub fn set(&mut self, row: usize, col: usize, value: T) { + self.data[row * self.cols + col] = value; + } + + pub fn dot(&self, other: &Self) -> Option<Matrix<T>> { + if self.cols != other.rows { + return None; + } + + let mut result = Matrix::new( + self.rows, + other.cols, + vec![T::zero(); self.rows * other.cols], + ); + for i in 0..self.rows { + for j in 0..other.cols { + let mut sum = T::zero(); + for k in 0..self.cols { + sum = sum + (self.get(i, k) * other.get(k, j)); + } + result.set(i, j, sum); + } + } + Some(result) + } + + pub fn kronecker(&self, other: &Self) -> Matrix<T> { + let new_rows = self.rows * other.rows; + let new_cols = self.cols * other.cols; + + let mut result = Matrix::new(new_rows, new_cols, vec![T::zero(); new_rows * new_cols]); + + for i in 0..self.rows { + for j in 0..self.cols { + let self_val = self.get(i, j); + for k in 0..other.rows { + for l in 0..other.cols { + let result_row = i * other.rows + k; + let result_col = j * other.cols + l; + result.set(result_row, result_col, self_val * other.get(k, l)) + } + } + } + } + + result + } + + pub fn transpose(&self) -> Matrix<T> { + let mut result = Matrix::new(self.cols, self.rows, vec![T::zero(); self.cols * self.rows]); + + for i in 0..self.rows { + for j in 0..self.cols { + let value = self.get(i, j); + result.set(j, i, value); + } + } + + result + } + + pub fn add_to(&self, other: &Self) -> Option<Matrix<T>> { + if self.rows != other.rows || self.cols != other.cols { + return None; + } + + let mut result = Matrix::new(self.rows, self.cols, vec![T::zero(); self.rows * self.cols]); + + for i in 0..self.rows { + for j in 0..self.cols { + let sum = self.get(i, j) + other.get(i, j); + result.set(i, j, sum); + } + } + Some(result) + } + + pub fn subtract(&self, other: &Self) -> Option<Matrix<T>> { + if self.rows != other.rows || self.cols != other.cols { + return None; + } + + let mut result = Matrix::new(self.rows, self.cols, vec![T::zero(); self.rows * self.cols]); + + for i in 0..self.rows { + for j in 0..self.cols { + let diff = self.get(i, j) - other.get(i, j); + result.set(i, j, diff); + } + } + Some(result) + } + + pub fn scale(&self, scalar: T) -> Matrix<T> { + let mut result = Matrix::new(self.rows, self.cols, vec![T::zero(); self.rows * self.cols]); + + for i in 0..self.rows { + for j in 0..self.cols { + let scaled_value = self.get(i, j) * scalar; + result.set(i, j, scaled_value); + } + } + result + } +} + +impl<T: Float> ops::Index<(usize, usize)> for Matrix<T> { + type Output = T; + + fn index(&self, index: (usize, usize)) -> &Self::Output { + &self.data[index.0 * self.cols + index.1] + } +} + +impl<T: Float> ops::IndexMut<(usize, usize)> for Matrix<T> { + fn index_mut(&mut self, index: (usize, usize)) -> &mut Self::Output { + &mut self.data[index.0 * self.cols + index.1] + } +} + +impl<T: Float + fmt::Debug> fmt::Debug for Matrix<T> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + for i in 0..self.rows { + for j in 0..self.cols { + write!(f, "{:?} ", self.get(i, j))?; + } + writeln!(f)?; + } + Ok(()) + } +} + +impl<T: Float + fmt::Display> fmt::Display for Matrix<T> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + for i in 0..self.rows { + for j in 0..self.cols { + write!(f, "{:>8} ", self.get(i, j))?; + } + writeln!(f)?; + } + Ok(()) + } +} + +impl<T: Float> ops::Add<&Matrix<T>> for Matrix<T> { + type Output = Option<Matrix<T>>; + + fn add(self, other: &Matrix<T>) -> Self::Output { + self.add_to(other) + } +} + +impl<T: Float> ops::Sub<&Matrix<T>> for Matrix<T> { + type Output = Option<Matrix<T>>; + + fn sub(self, other: &Matrix<T>) -> Self::Output { + self.subtract(other) + } +} + +impl<T: Float> ops::Mul<T> for Matrix<T> { + type Output = Matrix<T>; + + fn mul(self, scalar: T) -> Self::Output { + self.scale(scalar) + } +} + +impl<T: Float> ops::Div<T> for Matrix<T> { + type Output = Matrix<T>; + + fn div(self, scalar: T) -> Self::Output { + self.scale(T::one() / scalar) + } +} + +impl<T: Float> ops::AddAssign<&Matrix<T>> for Matrix<T> { + fn add_assign(&mut self, other: &Matrix<T>) { + if let Some(result) = self.add_to(other) { + *self = result; + } + } +} + +impl<T: Float> ops::SubAssign<&Matrix<T>> for Matrix<T> { + fn sub_assign(&mut self, other: &Matrix<T>) { + if let Some(result) = self.subtract(other) { + *self = result; + } + } +} + +impl<T: Float> ops::MulAssign<T> for Matrix<T> { + fn mul_assign(&mut self, scalar: T) { + *self = self.scale(scalar); + } +} + +impl<T: Float> ops::DivAssign<T> for Matrix<T> { + fn div_assign(&mut self, scalar: T) { + *self = self.scale(T::one() / scalar); + } +} diff --git a/libmu/src/maths/mod.rs b/libmu/src/maths/mod.rs index a6ccc9e..3d144d7 100644 --- a/libmu/src/maths/mod.rs +++ b/libmu/src/maths/mod.rs @@ -1,5 +1,7 @@ pub mod complex; +pub mod matrix; pub mod numeric_types; pub use complex::*; +pub use matrix::*; pub use numeric_types::*; diff --git a/libmu/src/maths/numeric_types.rs b/libmu/src/maths/numeric_types.rs index f9833dc..f1355de 100644 --- a/libmu/src/maths/numeric_types.rs +++ b/libmu/src/maths/numeric_types.rs @@ -1,6 +1,5 @@ -use core::ops; - use super::Complex; +use core::ops; pub trait Numeric: Copy diff --git a/mu/src/main.rs b/mu/src/main.rs index f328e4d..e3f23b2 100644 --- a/mu/src/main.rs +++ b/mu/src/main.rs @@ -1 +1,18 @@ -fn main() {} +use libmu::Matrix; + +fn main() { + let mat1 = Matrix::<f32>::new(2, 2, vec![1.0, 2.0, 3.0, 4.0]); + let mat2 = Matrix::<f32>::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); +} |
