aboutsummaryrefslogtreecommitdiff
path: root/libpsi-core/src
diff options
context:
space:
mode:
authorhachem <im@hachem.wtf>2024-10-12 19:01:40 +0200
committerhachem <im@hachem.wtf>2024-10-12 19:01:40 +0200
commit00c4b7c44b1364381c65dccbb190d64b3ff4c1e8 (patch)
tree5169843dcebc5807fdeb03120f2a997ba844d361 /libpsi-core/src
parentb04d0578c59299598c298f3ca83af3278b13f08c (diff)
Refactor 3: Matrices and fixed complex number matrix printing
Diffstat (limited to 'libpsi-core/src')
-rw-r--r--libpsi-core/src/maths/matrix.rs113
-rw-r--r--libpsi-core/src/maths/matrix_ops.rs62
-rw-r--r--libpsi-core/src/maths/mod.rs4
-rw-r--r--libpsi-core/src/maths/numeric.rs16
4 files changed, 111 insertions, 84 deletions
diff --git a/libpsi-core/src/maths/matrix.rs b/libpsi-core/src/maths/matrix.rs
index 636df08..c374a11 100644
--- a/libpsi-core/src/maths/matrix.rs
+++ b/libpsi-core/src/maths/matrix.rs
@@ -22,6 +22,29 @@ macro_rules! matrix {
}};
}
+macro_rules! impl_matrix_ops {
+ ($($trait:ident, $method:ident, $other:ty, $output:ty, $scale_fn:ident),* $(,)?) => {
+ $(
+ impl<T: Float> core::ops::$trait<$other> for Matrix<T> {
+ type Output = $output;
+
+ fn $method(self, other: $other) -> Self::Output {
+ self.$scale_fn(other)
+ }
+ }
+ )*
+ };
+ ($($trait:ident, $method:ident, $other:ty, $scale_fn:ident),* $(,)?) => {
+ $(
+ impl<T: Float> core::ops::$trait<$other> for Matrix<T> {
+ fn $method(&mut self, other: $other) {
+ *self = self.$scale_fn(other);
+ }
+ }
+ )*
+ };
+}
+
#[derive(Clone)]
pub struct Matrix<T: Float> {
pub data: Vec<T>,
@@ -158,6 +181,34 @@ impl<T: Float> ops::IndexMut<(usize, usize)> for Matrix<T> {
}
}
+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_matrix_ops! {
+ Add, add, &Matrix<T>, Option<Matrix<T>>, add_to,
+ Sub, sub, &Matrix<T>, Option<Matrix<T>>, subtract,
+ Mul, mul, T, Matrix<T>, scale,
+ Div, div, T, Matrix<T>, scale,
+}
+
+impl_matrix_ops! {
+ MulAssign, mul_assign, T, scale,
+ DivAssign, div_assign, T, scale,
+}
+
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 {
@@ -172,11 +223,58 @@ impl<T: Float + fmt::Debug> fmt::Debug for Matrix<T> {
impl<T: Float + fmt::Display> fmt::Display for Matrix<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
- let max_width = (0..self.rows)
- .flat_map(|i| (0..self.cols).map(move |j| self.get(i, j)))
- .map(|x| format!("{:.2}", x).split('.').next().unwrap().len())
- .max()
- .unwrap_or(0);
+ let elements: Vec<String> = self.data.iter().map(ToString::to_string).collect();
+ let is_complex = elements.iter().any(|element| element.contains("i"));
+
+ let normalized: Vec<(f64, f64)> = self
+ .data
+ .iter()
+ .map(|element| {
+ let element_string = format!("{}", element);
+
+ if is_complex {
+ let element_string = element_string.trim_end_matches('i').trim();
+ let element_split: Vec<&str> = element_string.split_whitespace().collect();
+ let real = element_split[0].parse::<f64>().unwrap();
+ let imaginary = element_split
+ .get(2)
+ .map_or(0.0, |&s| s.parse::<f64>().unwrap());
+ (real, imaginary)
+ } else {
+ (element_string.parse::<f64>().unwrap(), 0.0)
+ }
+ })
+ .collect();
+
+ let max_widths = normalized
+ .iter()
+ .fold((0, 0), |(max_0, max_1), &(real, imag)| {
+ let new_max_0 = max_0.max(format!("{:.2}", real).len());
+ let new_max_1 = if is_complex {
+ max_1.max(format!("{:.2}", imag.abs()).len())
+ } else {
+ max_1
+ };
+ (new_max_0, new_max_1)
+ });
+
+ let aligned: Vec<String> = normalized
+ .iter()
+ .map(|&(real, imag)| {
+ if is_complex {
+ format!(
+ "{:>rewidth$.2} {} {:>imwidth$.2}i",
+ real,
+ if imag > 0.0 { "+" } else { "-" },
+ imag.abs(),
+ rewidth = max_widths.0,
+ imwidth = max_widths.1,
+ )
+ } else {
+ format!("{:>width$.2}", real, width = max_widths.0)
+ }
+ })
+ .collect();
for i in 0..self.rows {
if i == 0 {
@@ -188,9 +286,9 @@ impl<T: Float + fmt::Display> fmt::Display for Matrix<T> {
}
for j in 0..self.cols {
- write!(f, "{:>width$.2}", self.get(i, j), width = max_width + 3)?;
+ write!(f, "{}", aligned[i + j * self.rows])?;
if j != self.cols - 1 {
- write!(f, " ")?;
+ write!(f, ", ")?;
}
}
@@ -206,6 +304,7 @@ impl<T: Float + fmt::Display> fmt::Display for Matrix<T> {
write!(f, "\n")?;
}
}
+
Ok(())
}
}
diff --git a/libpsi-core/src/maths/matrix_ops.rs b/libpsi-core/src/maths/matrix_ops.rs
deleted file mode 100644
index 15ab28b..0000000
--- a/libpsi-core/src/maths/matrix_ops.rs
+++ /dev/null
@@ -1,62 +0,0 @@
-use super::{Float, Matrix};
-use core::ops;
-
-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/libpsi-core/src/maths/mod.rs b/libpsi-core/src/maths/mod.rs
index 701ecd2..22fb558 100644
--- a/libpsi-core/src/maths/mod.rs
+++ b/libpsi-core/src/maths/mod.rs
@@ -1,8 +1,6 @@
pub mod complex;
-pub mod numeric;
-
pub mod matrix;
-pub mod matrix_ops;
+pub mod numeric;
pub mod vector;
pub mod vector_ops;
diff --git a/libpsi-core/src/maths/numeric.rs b/libpsi-core/src/maths/numeric.rs
index 0fd4380..f5a649e 100644
--- a/libpsi-core/src/maths/numeric.rs
+++ b/libpsi-core/src/maths/numeric.rs
@@ -19,17 +19,10 @@ macro_rules! impl_numeric {
macro_rules! impl_cnumeric {
($($t:ty),*) => {
- $(
- impl Numeric for Complex<$t> {
- fn zero() -> Self {
- Complex::new(0.0, 0.0)
- }
-
- fn one() -> Self {
- Complex::new(1.0, 0.0)
- }
- }
- )*
+ $(impl Numeric for Complex<$t> {
+ fn zero() -> Self { Complex::new(0.0, 0.0) }
+ fn one() -> Self { Complex::new(1.0, 0.0) }
+ })*
};
}
@@ -96,7 +89,6 @@ pub trait Numeric:
impl_numeric!(i32, i64, f32, f64);
impl_cnumeric!(f32, f64);
-
impl_float!(f32, libm::sqrtf, libm::atan2f);
impl_float!(f64, libm::sqrt, libm::atan2);
impl_cfloat!(f32, libm::sqrtf, libm::atan2f, libm::cosf, libm::sinf);