use super::{Float, Matrix}; use core::{fmt, 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; 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; fn from_matrix(matrix: &Matrix) -> Self; } pub trait VectorMatrix { fn to_matrix(&self) -> Matrix; } #[derive(Clone)] pub struct VectorImpl(Vec); 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()) } } impl VectorMatrix for RowVector { fn to_matrix(&self) -> Matrix { Matrix::new(1, self.size(), self.0.clone()) } } impl VectorMatrix for ColumnVector { fn to_matrix(&self) -> Matrix { Matrix::new(self.size(), 1, self.0.clone()) } } impl Vector for VectorImpl { fn from_matrix(matrix: &Matrix) -> Self { Self::new(matrix.data.clone()) } 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 fmt::Debug for RowVector { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!(f, "RowVector({:?})", self.0) } } impl fmt::Debug for ColumnVector { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!(f, "ColumnVector({:?})", self.0) } } impl fmt::Display for RowVector { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!( f, "[{}]", self.0 .iter() .map(|x| x.to_string()) .collect::>() .join(", ") ) } } impl fmt::Display for ColumnVector { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!(f, "{}", self.to_matrix()) } }