diff options
| author | hachem <im@hachem.wtf> | 2026-08-24 14:48:36 +0200 |
|---|---|---|
| committer | hachem <im@hachem.wtf> | 2026-08-24 14:48:36 +0200 |
| commit | 24d639224ca11112025289065d7538851606b56e (patch) | |
| tree | e7e26c89cd20099f03ff32727ba0af2cf4762088 /src/maths | |
| parent | 548dc42d9f454cb8c29fddb88ca41f8b1b595882 (diff) | |
[chore]: unwrap project
Diffstat (limited to 'src/maths')
| -rw-r--r-- | src/maths/complex.rs | 180 | ||||
| -rw-r--r-- | src/maths/format.rs | 141 | ||||
| -rw-r--r-- | src/maths/matrix.rs | 310 | ||||
| -rw-r--r-- | src/maths/mod.rs | 14 | ||||
| -rw-r--r-- | src/maths/numeric.rs | 101 | ||||
| -rw-r--r-- | src/maths/simd.rs | 510 | ||||
| -rw-r--r-- | src/maths/vector.rs | 258 | ||||
| -rw-r--r-- | src/maths/vector_ops.rs | 107 |
8 files changed, 1621 insertions, 0 deletions
diff --git a/src/maths/complex.rs b/src/maths/complex.rs new file mode 100644 index 0000000..31eae69 --- /dev/null +++ b/src/maths/complex.rs @@ -0,0 +1,180 @@ +use crate::Float; +use core::{fmt, ops}; + +#[macro_export] +macro_rules! complex { + ($real:expr, $imaginary:expr) => { + $crate::Complex::new($real, $imaginary) + }; +} + +macro_rules! impl_ops { + ($trait:ident, $method:ident, $op:tt) => { + impl<T: Float> ops::$trait for Complex<T> { + type Output = Complex<T>; + + fn $method(self, other: Complex<T>) -> Complex<T> { + Complex { + real: self.real $op other.real, + imaginary: self.imaginary $op other.imaginary, + } + } + } + }; + + ($trait:ident, $method:ident, $op:tt, real) => { + impl<T: Float> ops::$trait<T> for Complex<T> { + type Output = Complex<T>; + + fn $method(self, other: T) -> Complex<T> { + Complex { + real: self.real $op other, + imaginary: self.imaginary, + } + } + } + }; + + ($trait_assign:ident, $method_assign:ident, $op:tt, assign) => { + impl<T: Float> ops::$trait_assign for Complex<T> { + fn $method_assign(&mut self, other: Complex<T>) { + self.real = self.real $op other.real; + self.imaginary = self.imaginary $op other.imaginary; + } + } + }; + + ($trait_assign:ident, $method_assign:ident, $op:tt, assign_real) => { + impl<T: Float> ops::$trait_assign<T> for Complex<T> { + fn $method_assign(&mut self, other: T) { + self.real = self.real $op other; + } + } + }; +} + +#[derive(Copy, Clone, PartialOrd, PartialEq)] +pub struct Complex<T: Float> { + pub real: T, + pub imaginary: T, +} + +impl<T: Float + fmt::Debug> fmt::Debug for Complex<T> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!( + f, + "Complex {{ real: {:?}, imaginary: {:?} }}", + self.real, self.imaginary + ) + } +} + +impl<T: Float + fmt::Display> fmt::Display for Complex<T> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{} + {}i", self.real, self.imaginary) + } +} + +impl<T: Float> ops::Neg for Complex<T> { + type Output = Complex<T>; + + fn neg(self) -> Complex<T> { + Complex { + real: -self.real, + imaginary: -self.imaginary, + } + } +} + +impl<T: Float> From<T> for Complex<T> { + fn from(real: T) -> Complex<T> { + Complex { + real, + imaginary: T::zero(), + } + } +} + +impl<T: Float> Complex<T> { + pub fn new(real: T, imaginary: T) -> Complex<T> { + Complex { real, imaginary } + } + + pub fn get_conjugate(&self) -> Complex<T> { + Complex { + real: self.real, + imaginary: -self.imaginary, + } + } + + pub fn conjugate(&mut self) { + self.imaginary = -self.imaginary; + } + + pub fn phase(&self) -> T { + T::atan2(self.imaginary, self.real) + } + + pub fn norm2(&self) -> T { + self.real * self.real + self.imaginary * self.imaginary + } + + pub fn abs(&self) -> T { + T::sqrt(self.norm2()) + } +} + +impl_ops!(Add, add, +); +impl_ops!(Sub, sub, -); + +impl<T: Float> ops::Mul for Complex<T> { + type Output = Complex<T>; + + fn mul(self, other: Complex<T>) -> Complex<T> { + // (a + bi) * (c + di) = (ac - bd) + (ad + bc)i + Complex { + real: self.real * other.real - self.imaginary * other.imaginary, + imaginary: self.real * other.imaginary + self.imaginary * other.real, + } + } +} + +impl<T: Float> ops::Div for Complex<T> { + type Output = Complex<T>; + + fn div(self, other: Complex<T>) -> Complex<T> { + // (a + bi) / (c + di) = ((ac + bd) + (bc - ad)i) / (c² + d²) + let denom = other.real * other.real + other.imaginary * other.imaginary; + Complex { + real: (self.real * other.real + self.imaginary * other.imaginary) / denom, + imaginary: (self.imaginary * other.real - self.real * other.imaginary) / denom, + } + } +} + +impl_ops!(AddAssign, add_assign, +, assign); +impl_ops!(SubAssign, sub_assign, -, assign); + +impl<T: Float> ops::MulAssign for Complex<T> { + fn mul_assign(&mut self, other: Complex<T>) { + let new_real = self.real * other.real - self.imaginary * other.imaginary; + let new_imag = self.real * other.imaginary + self.imaginary * other.real; + self.real = new_real; + self.imaginary = new_imag; + } +} + +impl<T: Float> ops::DivAssign for Complex<T> { + fn div_assign(&mut self, other: Complex<T>) { + let denom = other.real * other.real + other.imaginary * other.imaginary; + let new_real = (self.real * other.real + self.imaginary * other.imaginary) / denom; + let new_imag = (self.imaginary * other.real - self.real * other.imaginary) / denom; + self.real = new_real; + self.imaginary = new_imag; + } +} + +impl_ops!(Add, add, +, real); +impl_ops!(Sub, sub, -, real); +impl_ops!(Mul, mul, *, real); +impl_ops!(Div, div, /, real); diff --git a/src/maths/format.rs b/src/maths/format.rs new file mode 100644 index 0000000..b6b0191 --- /dev/null +++ b/src/maths/format.rs @@ -0,0 +1,141 @@ +use crate::Complex; + +const EPSILON: f64 = 1e-10; +const SQRT_2: f64 = std::f64::consts::SQRT_2; +const INV_SQRT_2: f64 = 0.7071067811865475; +const INV_SQRT_8: f64 = 0.3535533905932738; +const INV_SQRT_32: f64 = 0.1767766952966369; + +fn approx_eq(a: f64, b: f64) -> bool { + (a - b).abs() < EPSILON +} + +fn format_real_symbolic(v: f64) -> Option<String> { + let abs_v = v.abs(); + let sign = if v < 0.0 { "-" } else { "" }; + + if approx_eq(abs_v, 0.0) { + return Some("0".to_string()); + } + if approx_eq(abs_v, 1.0) { + return Some(format!("{}1", sign)); + } + if approx_eq(abs_v, 0.5) { + return Some(format!("{}½", sign)); + } + if approx_eq(abs_v, 0.25) { + return Some(format!("{}¼", sign)); + } + if approx_eq(abs_v, 0.75) { + return Some(format!("{}¾", sign)); + } + if approx_eq(abs_v, 0.125) { + return Some(format!("{}⅛", sign)); + } + if approx_eq(abs_v, SQRT_2) { + return Some(format!("{}√2", sign)); + } + if approx_eq(abs_v, INV_SQRT_2) { + return Some(format!("{}¹⁄√2", sign)); + } + if approx_eq(abs_v, INV_SQRT_8) { + return Some(format!("{}¹⁄√8", sign)); + } + if approx_eq(abs_v, INV_SQRT_32) { + return Some(format!("{}¹⁄√32", sign)); + } + if approx_eq(abs_v, 2.0) { + return Some(format!("{}2", sign)); + } + if approx_eq(abs_v, 1.0 / 3.0) { + return Some(format!("{}⅓", sign)); + } + if approx_eq(abs_v, 2.0 / 3.0) { + return Some(format!("{}⅔", sign)); + } + + None +} + +pub fn format_amplitude(c: &Complex<f64>) -> String { + let re = c.real; + let im = c.imaginary; + + let re_zero = approx_eq(re.abs(), 0.0); + let im_zero = approx_eq(im.abs(), 0.0); + + if re_zero && im_zero { + return "0".to_string(); + } + + if im_zero { + if let Some(s) = format_real_symbolic(re) { + return s; + } + return format!("{:.4}", re); + } + + if re_zero { + if approx_eq(im.abs(), 1.0) { + return if im > 0.0 { + "i".to_string() + } else { + "-i".to_string() + }; + } + if let Some(s) = format_real_symbolic(im) { + return format!("{}i", s); + } + return format!("{:.4}i", im); + } + + let re_str = format_real_symbolic(re).unwrap_or_else(|| format!("{:.4}", re)); + let im_str = if approx_eq(im.abs(), 1.0) { + if im > 0.0 { + "+i".to_string() + } else { + "-i".to_string() + } + } else { + let im_sym = format_real_symbolic(im.abs()); + let sign = if im > 0.0 { "+" } else { "-" }; + match im_sym { + Some(s) => format!("{}{}i", sign, s.trim_start_matches('-')), + None => format!("{}{:.4}i", sign, im.abs()), + } + }; + + format!("{}{}", re_str, im_str) +} + +pub fn format_probability(p: f64) -> String { + if approx_eq(p, 0.0) { + return "0".to_string(); + } + if approx_eq(p, 1.0) { + return "1".to_string(); + } + if approx_eq(p, 0.5) { + return "½".to_string(); + } + if approx_eq(p, 0.25) { + return "¼".to_string(); + } + if approx_eq(p, 0.75) { + return "¾".to_string(); + } + if approx_eq(p, 0.125) { + return "⅛".to_string(); + } + if approx_eq(p, 0.0625) { + return "¹⁄₁₆".to_string(); + } + if approx_eq(p, 1.0 / 3.0) { + return "⅓".to_string(); + } + if approx_eq(p, 2.0 / 3.0) { + return "⅔".to_string(); + } + + format!("{:.4}", p) +} diff --git a/src/maths/matrix.rs b/src/maths/matrix.rs new file mode 100644 index 0000000..c9d96a4 --- /dev/null +++ b/src/maths/matrix.rs @@ -0,0 +1,310 @@ +use super::Float; +use core::{fmt, ops}; + +#[macro_export] +macro_rules! matrix { + ( $( $( $x:expr ),* );* ) => {{ + let mut data = Vec::new(); + let mut rows = 0; + let mut cols = 0; + + $( + let row_data = $( $x )*; + if cols == 0 { + cols = row_data.len(); + } + assert_eq!(cols, row_data.len(), "All rows must have the same number of columns."); + data.extend(row_data); + rows += 1; + )* + + $crate::Matrix::new(rows, cols, data) + }}; +} + +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>, + pub rows: usize, + pub 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 += 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> 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 { + 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 { + 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 = element.to_string(); + + 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 { + write!(f, "┌")?; + } else if i == self.rows - 1 { + write!(f, "└")?; + } else { + write!(f, "│")?; + } + + for j in 0..self.cols { + write!(f, "{}", aligned[i + j * self.rows])?; + if j != self.cols - 1 { + write!(f, ", ")?; + } + } + + if i == 0 { + write!(f, "┐")?; + } else if i == self.rows - 1 { + write!(f, "┘")?; + } else { + write!(f, "│")?; + } + + if i != self.rows - 1 { + writeln!(f)?; + } + } + + Ok(()) + } +} diff --git a/src/maths/mod.rs b/src/maths/mod.rs new file mode 100644 index 0000000..85f4872 --- /dev/null +++ b/src/maths/mod.rs @@ -0,0 +1,14 @@ +pub mod complex; +pub mod format; +pub mod matrix; +pub mod numeric; +pub mod simd; +pub mod vector; +pub mod vector_ops; + +pub use complex::*; +pub use format::*; +pub use matrix::*; +pub use numeric::*; +pub use simd::*; +pub use vector::*; diff --git a/src/maths/numeric.rs b/src/maths/numeric.rs new file mode 100644 index 0000000..f5a649e --- /dev/null +++ b/src/maths/numeric.rs @@ -0,0 +1,101 @@ +use crate::Complex; +use core::ops; + +macro_rules! impl_numeric { + ($($t:ty),*) => { + $( + impl Numeric for $t { + fn zero() -> Self { + 0 as $t + } + + fn one() -> Self { + 1 as $t + } + } + )* + }; +} + +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) } + })* + }; +} + +macro_rules! impl_float { + ($($t:ty, $sqrt_fn:path, $atan2_fn:path),*) => { + $( + impl Float for $t { + fn sqrt(self) -> Self { + $sqrt_fn(self) + } + + fn atan2(y: Self, x: Self) -> Self { + $atan2_fn(y, x) + } + } + )* + }; +} + +macro_rules! impl_cfloat { + ($($t:ty, $sqrt_fn:path, $atan2_fn:path, $cos_fn:path, $sin_fn:path),*) => { + $( + impl Float for Complex<$t> { + fn sqrt(self) -> Self { + let r = self.abs(); + let theta = self.phase(); + + let sqrt_r = $sqrt_fn(r); + let sqrt_theta = theta / 2.0; + + Complex::new( + sqrt_r * $cos_fn(sqrt_theta), + sqrt_r * $sin_fn(sqrt_theta), + ) + } + + fn atan2(y: Self, x: Self) -> Self { + Complex::new( + $atan2_fn(y.real, x.real), + $atan2_fn(y.imaginary, x.imaginary), + ) + } + } + )* + }; +} + +pub trait Numeric: + Copy + + PartialOrd + + ops::Add<Output = Self> + + ops::Mul<Output = Self> + + ops::Sub<Output = Self> + + ops::Div<Output = Self> + + ops::Neg<Output = Self> + + ops::AddAssign + + ops::SubAssign + + ops::MulAssign + + ops::DivAssign +{ + fn zero() -> Self; + fn one() -> Self; +} + +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); +impl_cfloat!(f64, libm::sqrt, libm::atan2, libm::cos, libm::sin); + +pub trait Integer: Numeric {} +pub trait Float: Numeric { + fn sqrt(self) -> Self; + fn atan2(y: Self, x: Self) -> Self; +} diff --git a/src/maths/simd.rs b/src/maths/simd.rs new file mode 100644 index 0000000..de0370c --- /dev/null +++ b/src/maths/simd.rs @@ -0,0 +1,510 @@ +use crate::{complex, Complex}; + +#[cfg(target_arch = "x86_64")] +use std::arch::x86_64::*; + +#[cfg(target_arch = "aarch64")] +use std::arch::aarch64::*; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum SimdCapability { + None, + #[cfg(any(target_arch = "x86_64", target_arch = "x86"))] + Avx2, + #[cfg(any(target_arch = "x86_64", target_arch = "x86"))] + Avx512, + #[cfg(target_arch = "aarch64")] + Neon, +} + +impl SimdCapability { + pub fn detect() -> Self { + #[cfg(any(target_arch = "x86_64", target_arch = "x86"))] + { + if is_x86_feature_detected!("avx512f") && is_x86_feature_detected!("avx512dq") { + return SimdCapability::Avx512; + } + if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") { + return SimdCapability::Avx2; + } + } + + #[cfg(target_arch = "aarch64")] + { + return SimdCapability::Neon; + } + + #[allow(unreachable_code)] + SimdCapability::None + } + + pub fn name(&self) -> &'static str { + match self { + SimdCapability::None => "Scalar", + #[cfg(any(target_arch = "x86_64", target_arch = "x86"))] + SimdCapability::Avx2 => "AVX2+FMA", + #[cfg(any(target_arch = "x86_64", target_arch = "x86"))] + SimdCapability::Avx512 => "AVX-512", + #[cfg(target_arch = "aarch64")] + SimdCapability::Neon => "NEON", + } + } +} + +pub fn apply_single_qubit_gate_simd( + state: &mut [Complex<f64>], + gate: &[[Complex<f64>; 2]; 2], + target: usize, + num_qubits: usize, +) { + let capability = SimdCapability::detect(); + + match capability { + #[cfg(target_arch = "x86_64")] + SimdCapability::Avx2 => unsafe { + apply_single_qubit_avx2(state, gate, target, num_qubits); + }, + #[cfg(target_arch = "x86_64")] + SimdCapability::Avx512 => unsafe { + apply_single_qubit_avx512(state, gate, target, num_qubits); + }, + #[cfg(target_arch = "aarch64")] + SimdCapability::Neon => unsafe { + apply_single_qubit_neon(state, gate, target, num_qubits); + }, + _ => { + apply_single_qubit_scalar(state, gate, target, num_qubits); + } + } +} + +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "avx2", enable = "fma")] +unsafe fn apply_single_qubit_avx2( + state: &mut [Complex<f64>], + gate: &[[Complex<f64>; 2]; 2], + target: usize, + num_qubits: usize, +) { + let target_bit = num_qubits - 1 - target; + let step = 1 << target_bit; + let dim = 1 << num_qubits; + + let g00 = gate[0][0]; + let g01 = gate[0][1]; + let g10 = gate[1][0]; + let g11 = gate[1][1]; + + let pairs: Vec<(usize, usize)> = (0..dim) + .filter(|&i| (i >> target_bit) & 1 == 0) + .map(|i| (i, i | step)) + .collect(); + + let chunks = pairs.len() / 2; + + for chunk_idx in 0..chunks { + let (i0, j0) = pairs[chunk_idx * 2]; + let (i1, j1) = pairs[chunk_idx * 2 + 1]; + + let s0_re = _mm256_set_pd( + state[j1].real, + state[i1].real, + state[j0].real, + state[i0].real, + ); + let s0_im = _mm256_set_pd( + state[j1].imaginary, + state[i1].imaginary, + state[j0].imaginary, + state[i0].imaginary, + ); + + let g_re_0 = _mm256_set_pd(g01.real, g00.real, g01.real, g00.real); + let g_im_0 = _mm256_set_pd(g01.imaginary, g00.imaginary, g01.imaginary, g00.imaginary); + let g_re_1 = _mm256_set_pd(g11.real, g10.real, g11.real, g10.real); + let g_im_1 = _mm256_set_pd(g11.imaginary, g10.imaginary, g11.imaginary, g10.imaginary); + + let prod0_re = _mm256_fmsub_pd(s0_re, g_re_0, _mm256_mul_pd(s0_im, g_im_0)); + let prod0_im = _mm256_fmadd_pd(s0_re, g_im_0, _mm256_mul_pd(s0_im, g_re_0)); + + let prod1_re = _mm256_fmsub_pd(s0_re, g_re_1, _mm256_mul_pd(s0_im, g_im_1)); + let prod1_im = _mm256_fmadd_pd(s0_re, g_im_1, _mm256_mul_pd(s0_im, g_re_1)); + + let mut res0_re = [0.0f64; 4]; + let mut res0_im = [0.0f64; 4]; + let mut res1_re = [0.0f64; 4]; + let mut res1_im = [0.0f64; 4]; + + _mm256_storeu_pd(res0_re.as_mut_ptr(), prod0_re); + _mm256_storeu_pd(res0_im.as_mut_ptr(), prod0_im); + _mm256_storeu_pd(res1_re.as_mut_ptr(), prod1_re); + _mm256_storeu_pd(res1_im.as_mut_ptr(), prod1_im); + + state[i0] = complex!(res0_re[0] + res0_re[1], res0_im[0] + res0_im[1]); + state[j0] = complex!(res1_re[0] + res1_re[1], res1_im[0] + res1_im[1]); + state[i1] = complex!(res0_re[2] + res0_re[3], res0_im[2] + res0_im[3]); + state[j1] = complex!(res1_re[2] + res1_re[3], res1_im[2] + res1_im[3]); + } + + for &(i, j) in pairs.iter().skip(chunks * 2) { + let s0 = state[i]; + let s1 = state[j]; + + let new0 = complex!( + s0.real * g00.real - s0.imaginary * g00.imaginary + s1.real * g01.real + - s1.imaginary * g01.imaginary, + s0.real * g00.imaginary + + s0.imaginary * g00.real + + s1.real * g01.imaginary + + s1.imaginary * g01.real + ); + + let new1 = complex!( + s0.real * g10.real - s0.imaginary * g10.imaginary + s1.real * g11.real + - s1.imaginary * g11.imaginary, + s0.real * g10.imaginary + + s0.imaginary * g10.real + + s1.real * g11.imaginary + + s1.imaginary * g11.real + ); + + state[i] = new0; + state[j] = new1; + } +} + +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "avx512f", enable = "avx512dq")] +unsafe fn apply_single_qubit_avx512( + state: &mut [Complex<f64>], + gate: &[[Complex<f64>; 2]; 2], + target: usize, + num_qubits: usize, +) { + let target_bit = num_qubits - 1 - target; + let step = 1 << target_bit; + let dim = 1 << num_qubits; + + let g00 = gate[0][0]; + let g01 = gate[0][1]; + let g10 = gate[1][0]; + let g11 = gate[1][1]; + + let pairs: Vec<(usize, usize)> = (0..dim) + .filter(|&i| (i >> target_bit) & 1 == 0) + .map(|i| (i, i | step)) + .collect(); + + let chunks = pairs.len() / 4; + + for chunk_idx in 0..chunks { + let base = chunk_idx * 4; + let (i0, j0) = pairs[base]; + let (i1, j1) = pairs[base + 1]; + let (i2, j2) = pairs[base + 2]; + let (i3, j3) = pairs[base + 3]; + + let s0_re = _mm512_set_pd( + state[j3].real, + state[i3].real, + state[j2].real, + state[i2].real, + state[j1].real, + state[i1].real, + state[j0].real, + state[i0].real, + ); + let s0_im = _mm512_set_pd( + state[j3].imaginary, + state[i3].imaginary, + state[j2].imaginary, + state[i2].imaginary, + state[j1].imaginary, + state[i1].imaginary, + state[j0].imaginary, + state[i0].imaginary, + ); + + let g_re_0 = _mm512_set_pd( + g01.real, g00.real, g01.real, g00.real, g01.real, g00.real, g01.real, g00.real, + ); + let g_im_0 = _mm512_set_pd( + g01.imaginary, + g00.imaginary, + g01.imaginary, + g00.imaginary, + g01.imaginary, + g00.imaginary, + g01.imaginary, + g00.imaginary, + ); + let g_re_1 = _mm512_set_pd( + g11.real, g10.real, g11.real, g10.real, g11.real, g10.real, g11.real, g10.real, + ); + let g_im_1 = _mm512_set_pd( + g11.imaginary, + g10.imaginary, + g11.imaginary, + g10.imaginary, + g11.imaginary, + g10.imaginary, + g11.imaginary, + g10.imaginary, + ); + + let prod0_re = _mm512_fmsub_pd(s0_re, g_re_0, _mm512_mul_pd(s0_im, g_im_0)); + let prod0_im = _mm512_fmadd_pd(s0_re, g_im_0, _mm512_mul_pd(s0_im, g_re_0)); + let prod1_re = _mm512_fmsub_pd(s0_re, g_re_1, _mm512_mul_pd(s0_im, g_im_1)); + let prod1_im = _mm512_fmadd_pd(s0_re, g_im_1, _mm512_mul_pd(s0_im, g_re_1)); + + let mut res0_re = [0.0f64; 8]; + let mut res0_im = [0.0f64; 8]; + let mut res1_re = [0.0f64; 8]; + let mut res1_im = [0.0f64; 8]; + + _mm512_storeu_pd(res0_re.as_mut_ptr(), prod0_re); + _mm512_storeu_pd(res0_im.as_mut_ptr(), prod0_im); + _mm512_storeu_pd(res1_re.as_mut_ptr(), prod1_re); + _mm512_storeu_pd(res1_im.as_mut_ptr(), prod1_im); + + state[i0] = complex!(res0_re[0] + res0_re[1], res0_im[0] + res0_im[1]); + state[j0] = complex!(res1_re[0] + res1_re[1], res1_im[0] + res1_im[1]); + state[i1] = complex!(res0_re[2] + res0_re[3], res0_im[2] + res0_im[3]); + state[j1] = complex!(res1_re[2] + res1_re[3], res1_im[2] + res1_im[3]); + state[i2] = complex!(res0_re[4] + res0_re[5], res0_im[4] + res0_im[5]); + state[j2] = complex!(res1_re[4] + res1_re[5], res1_im[4] + res1_im[5]); + state[i3] = complex!(res0_re[6] + res0_re[7], res0_im[6] + res0_im[7]); + state[j3] = complex!(res1_re[6] + res1_re[7], res1_im[6] + res1_im[7]); + } + + for &(i, j) in pairs.iter().skip(chunks * 4) { + let s0 = state[i]; + let s1 = state[j]; + + let new0 = complex!( + s0.real * g00.real - s0.imaginary * g00.imaginary + s1.real * g01.real + - s1.imaginary * g01.imaginary, + s0.real * g00.imaginary + + s0.imaginary * g00.real + + s1.real * g01.imaginary + + s1.imaginary * g01.real + ); + + let new1 = complex!( + s0.real * g10.real - s0.imaginary * g10.imaginary + s1.real * g11.real + - s1.imaginary * g11.imaginary, + s0.real * g10.imaginary + + s0.imaginary * g10.real + + s1.real * g11.imaginary + + s1.imaginary * g11.real + ); + + state[i] = new0; + state[j] = new1; + } +} + +#[cfg(target_arch = "aarch64")] +unsafe fn apply_single_qubit_neon( + state: &mut [Complex<f64>], + gate: &[[Complex<f64>; 2]; 2], + target: usize, + num_qubits: usize, +) { + let target_bit = num_qubits - 1 - target; + let step = 1 << target_bit; + let dim = 1 << num_qubits; + + let g00 = gate[0][0]; + let g01 = gate[0][1]; + let g10 = gate[1][0]; + let g11 = gate[1][1]; + + let pairs: Vec<(usize, usize)> = (0..dim) + .filter(|&i| (i >> target_bit) & 1 == 0) + .map(|i| (i, i | step)) + .collect(); + + let chunks = pairs.len() / 2; + + for chunk_idx in 0..chunks { + let (i0, j0) = pairs[chunk_idx * 2]; + let (i1, j1) = pairs[chunk_idx * 2 + 1]; + + let s0_0 = state[i0]; + let s1_0 = state[j0]; + let s0_1 = state[i1]; + let s1_1 = state[j1]; + + let s0_re = vld1q_f64([s0_0.real, s0_1.real].as_ptr()); + let s0_im = vld1q_f64([s0_0.imaginary, s0_1.imaginary].as_ptr()); + let s1_re = vld1q_f64([s1_0.real, s1_1.real].as_ptr()); + let s1_im = vld1q_f64([s1_0.imaginary, s1_1.imaginary].as_ptr()); + + let g00_re = vdupq_n_f64(g00.real); + let g00_im = vdupq_n_f64(g00.imaginary); + let g01_re = vdupq_n_f64(g01.real); + let g01_im = vdupq_n_f64(g01.imaginary); + let g10_re = vdupq_n_f64(g10.real); + let g10_im = vdupq_n_f64(g10.imaginary); + let g11_re = vdupq_n_f64(g11.real); + let g11_im = vdupq_n_f64(g11.imaginary); + + let new0_re = vaddq_f64( + vfmsq_f64(vmulq_f64(s0_re, g00_re), s0_im, g00_im), + vfmsq_f64(vmulq_f64(s1_re, g01_re), s1_im, g01_im), + ); + let new0_im = vaddq_f64( + vfmaq_f64(vmulq_f64(s0_re, g00_im), s0_im, g00_re), + vfmaq_f64(vmulq_f64(s1_re, g01_im), s1_im, g01_re), + ); + + let new1_re = vaddq_f64( + vfmsq_f64(vmulq_f64(s0_re, g10_re), s0_im, g10_im), + vfmsq_f64(vmulq_f64(s1_re, g11_re), s1_im, g11_im), + ); + let new1_im = vaddq_f64( + vfmaq_f64(vmulq_f64(s0_re, g10_im), s0_im, g10_re), + vfmaq_f64(vmulq_f64(s1_re, g11_im), s1_im, g11_re), + ); + + state[i0] = complex!(vgetq_lane_f64(new0_re, 0), vgetq_lane_f64(new0_im, 0)); + state[j0] = complex!(vgetq_lane_f64(new1_re, 0), vgetq_lane_f64(new1_im, 0)); + state[i1] = complex!(vgetq_lane_f64(new0_re, 1), vgetq_lane_f64(new0_im, 1)); + state[j1] = complex!(vgetq_lane_f64(new1_re, 1), vgetq_lane_f64(new1_im, 1)); + } + + for &(i, j) in pairs.iter().skip(chunks * 2) { + let s0 = state[i]; + let s1 = state[j]; + + let new0 = complex!( + s0.real * g00.real - s0.imaginary * g00.imaginary + s1.real * g01.real + - s1.imaginary * g01.imaginary, + s0.real * g00.imaginary + + s0.imaginary * g00.real + + s1.real * g01.imaginary + + s1.imaginary * g01.real + ); + + let new1 = complex!( + s0.real * g10.real - s0.imaginary * g10.imaginary + s1.real * g11.real + - s1.imaginary * g11.imaginary, + s0.real * g10.imaginary + + s0.imaginary * g10.real + + s1.real * g11.imaginary + + s1.imaginary * g11.real + ); + + state[i] = new0; + state[j] = new1; + } +} + +fn apply_single_qubit_scalar( + state: &mut [Complex<f64>], + gate: &[[Complex<f64>; 2]; 2], + target: usize, + num_qubits: usize, +) { + let target_bit = num_qubits - 1 - target; + let step = 1 << target_bit; + let dim = 1 << num_qubits; + + let g00 = gate[0][0]; + let g01 = gate[0][1]; + let g10 = gate[1][0]; + let g11 = gate[1][1]; + + for i in 0..dim { + if (i >> target_bit) & 1 == 1 { + continue; + } + + let j = i | step; + let s0 = state[i]; + let s1 = state[j]; + + let new0 = complex!( + s0.real * g00.real - s0.imaginary * g00.imaginary + s1.real * g01.real + - s1.imaginary * g01.imaginary, + s0.real * g00.imaginary + + s0.imaginary * g00.real + + s1.real * g01.imaginary + + s1.imaginary * g01.real + ); + + let new1 = complex!( + s0.real * g10.real - s0.imaginary * g10.imaginary + s1.real * g11.real + - s1.imaginary * g11.imaginary, + s0.real * g10.imaginary + + s0.imaginary * g10.real + + s1.real * g11.imaginary + + s1.imaginary * g11.real + ); + + state[i] = new0; + state[j] = new1; + } +} + +pub fn apply_single_qubit_gate_simd_parallel( + state: &mut [Complex<f64>], + gate: &[[Complex<f64>; 2]; 2], + target: usize, + num_qubits: usize, +) { + use rayon::prelude::*; + + let target_bit = num_qubits - 1 - target; + let step = 1 << target_bit; + let dim = 1 << num_qubits; + + let g00 = gate[0][0]; + let g01 = gate[0][1]; + let g10 = gate[1][0]; + let g11 = gate[1][1]; + + let pairs: Vec<(usize, usize)> = (0..dim) + .filter(|&i| (i >> target_bit) & 1 == 0) + .map(|i| (i, i | step)) + .collect(); + + let results: Vec<(usize, usize, Complex<f64>, Complex<f64>)> = pairs + .par_iter() + .map(|&(i, j)| { + let s0 = state[i]; + let s1 = state[j]; + + let new0 = complex!( + s0.real * g00.real - s0.imaginary * g00.imaginary + s1.real * g01.real + - s1.imaginary * g01.imaginary, + s0.real * g00.imaginary + + s0.imaginary * g00.real + + s1.real * g01.imaginary + + s1.imaginary * g01.real + ); + + let new1 = complex!( + s0.real * g10.real - s0.imaginary * g10.imaginary + s1.real * g11.real + - s1.imaginary * g11.imaginary, + s0.real * g10.imaginary + + s0.imaginary * g10.real + + s1.real * g11.imaginary + + s1.imaginary * g11.real + ); + + (i, j, new0, new1) + }) + .collect(); + + for (i, j, new0, new1) in results { + state[i] = new0; + state[j] = new1; + } +} + +pub fn get_simd_info() -> String { + let cap = SimdCapability::detect(); + format!("SIMD: {}", cap.name()) +} diff --git a/src/maths/vector.rs b/src/maths/vector.rs new file mode 100644 index 0000000..11f3d29 --- /dev/null +++ b/src/maths/vector.rs @@ -0,0 +1,258 @@ +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<T: Float> { + fn new(data: Vec<T>) -> 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<T>) -> Self; +} + +pub trait VectorMatrix<T: Float> { + fn to_matrix(&self) -> Matrix<T>; +} + +#[derive(Clone)] +pub struct VectorImpl<T: Float, const ROWS: usize, const COLS: usize>(Vec<T>); +pub type RowVector<T> = VectorImpl<T, 1, 0>; +pub type ColumnVector<T> = VectorImpl<T, 0, 1>; + +impl<T: Float> ColumnVector<T> { + pub fn mul_matrix(&self, matrix: &Matrix<T>) -> Option<ColumnVector<T>> { + 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 += matrix.get(i, j) * self.get(j) ; + } + result.set(i, sum); + } + + Some(result) + } + + pub fn transpose(&self) -> RowVector<T> { + RowVector::new(self.0.clone()) + } +} + +impl<T: Float> RowVector<T> { + pub fn mul_matrix(&self, matrix: &Matrix<T>) -> Option<RowVector<T>> { + 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 += self.get(i) * matrix.get(i, j) ; + } + result.set(j, sum); + } + + Some(result) + } + + pub fn transpose(&self) -> ColumnVector<T> { + ColumnVector::new(self.0.clone()) + } +} + +impl<T: Float> VectorMatrix<T> for RowVector<T> { + fn to_matrix(&self) -> Matrix<T> { + Matrix::new(1, self.size(), self.0.clone()) + } +} + +impl<T: Float> VectorMatrix<T> for ColumnVector<T> { + fn to_matrix(&self) -> Matrix<T> { + Matrix::new(self.size(), 1, self.0.clone()) + } +} + +impl<T: Float, const ROWS: usize, const COLS: usize> Vector<T> for VectorImpl<T, ROWS, COLS> { + fn from_matrix(matrix: &Matrix<T>) -> Self { + Self::new(matrix.data.clone()) + } + + fn new(data: Vec<T>) -> 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<T: Float, const ROWS: usize, const COLS: usize> VectorImpl<T, ROWS, COLS> { + pub fn add_to(&self, other: &Self) -> Option<VectorImpl<T, ROWS, COLS>> { + 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<VectorImpl<T, ROWS, COLS>> { + 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<T, ROWS, COLS> { + 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<T: Float, const ROWS: usize, const COLS: usize> ops::Index<usize> + for VectorImpl<T, ROWS, COLS> +{ + type Output = T; + + fn index(&self, index: usize) -> &Self::Output { + &self.0[index] + } +} + +impl<T: Float, const ROWS: usize, const COLS: usize> ops::IndexMut<usize> + for VectorImpl<T, ROWS, COLS> +{ + fn index_mut(&mut self, index: usize) -> &mut Self::Output { + &mut self.0[index] + } +} + +impl<T: Float + fmt::Debug> fmt::Debug for RowVector<T> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "RowVector({:?})", self.0) + } +} + +impl<T: Float + fmt::Debug> fmt::Debug for ColumnVector<T> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "ColumnVector({:?})", self.0) + } +} + +impl<T: Float + fmt::Display> fmt::Display for RowVector<T> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!( + f, + "[{}]", + self.0 + .iter() + .map(|x| x.to_string()) + .collect::<Vec<String>>() + .join(", ") + ) + } +} + +impl<T: Float + fmt::Display> fmt::Display for ColumnVector<T> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}", self.to_matrix()) + } +} diff --git a/src/maths/vector_ops.rs b/src/maths/vector_ops.rs new file mode 100644 index 0000000..1e7148c --- /dev/null +++ b/src/maths/vector_ops.rs @@ -0,0 +1,107 @@ +use super::{Float, Matrix}; +use crate::{ColumnVector, RowVector, VectorImpl}; +use core::ops; + +impl<T: Float, const ROWS: usize, const COLS: usize> ops::Add<&VectorImpl<T, ROWS, COLS>> + for VectorImpl<T, ROWS, COLS> +{ + type Output = Option<VectorImpl<T, ROWS, COLS>>; + + fn add(self, other: &VectorImpl<T, ROWS, COLS>) -> Self::Output { + self.add_to(other) + } +} + +impl<T: Float, const ROWS: usize, const COLS: usize> ops::Sub<&VectorImpl<T, ROWS, COLS>> + for VectorImpl<T, ROWS, COLS> +{ + type Output = Option<VectorImpl<T, ROWS, COLS>>; + + fn sub(self, other: &VectorImpl<T, ROWS, COLS>) -> Self::Output { + self.subtract(other) + } +} + +impl<T: Float, const ROWS: usize, const COLS: usize> ops::Mul<T> for VectorImpl<T, ROWS, COLS> { + type Output = VectorImpl<T, ROWS, COLS>; + + fn mul(self, scalar: T) -> Self::Output { + self.scale(scalar) + } +} + +impl<T: Float, const ROWS: usize, const COLS: usize> ops::Div<T> for VectorImpl<T, ROWS, COLS> { + type Output = VectorImpl<T, ROWS, COLS>; + + fn div(self, scalar: T) -> Self::Output { + self.scale(T::one() / scalar) + } +} + +impl<T: Float, const ROWS: usize, const COLS: usize> ops::AddAssign<VectorImpl<T, ROWS, COLS>> + for VectorImpl<T, ROWS, COLS> +{ + fn add_assign(&mut self, other: VectorImpl<T, ROWS, COLS>) { + if let Some(result) = self.add_to(&other) { + *self = result; + } + } +} + +impl<T: Float, const ROWS: usize, const COLS: usize> ops::SubAssign<VectorImpl<T, ROWS, COLS>> + for VectorImpl<T, ROWS, COLS> +{ + fn sub_assign(&mut self, other: VectorImpl<T, ROWS, COLS>) { + if let Some(result) = self.subtract(&other) { + *self = result; + } + } +} + +impl<T: Float, const ROWS: usize, const COLS: usize> ops::MulAssign<T> + for VectorImpl<T, ROWS, COLS> +{ + fn mul_assign(&mut self, scalar: T) { + *self = self.scale(scalar); + } +} + +impl<T: Float, const ROWS: usize, const COLS: usize> ops::DivAssign<T> + for VectorImpl<T, ROWS, COLS> +{ + fn div_assign(&mut self, scalar: T) { + *self = self.scale(T::one() / scalar); + } +} + +impl<T: Float> ops::Mul<&Matrix<T>> for RowVector<T> { + type Output = Option<RowVector<T>>; + + fn mul(self, matrix: &Matrix<T>) -> Self::Output { + self.mul_matrix(matrix) + } +} + +impl<T: Float> ops::Mul<&Matrix<T>> for ColumnVector<T> { + type Output = Option<ColumnVector<T>>; + + fn mul(self, matrix: &Matrix<T>) -> Self::Output { + self.mul_matrix(matrix) + } +} + +impl<T: Float> ops::MulAssign<&Matrix<T>> for RowVector<T> { + fn mul_assign(&mut self, matrix: &Matrix<T>) { + if let Some(result) = self.mul_matrix(matrix) { + *self = result; + } + } +} + +impl<T: Float> ops::MulAssign<&Matrix<T>> for ColumnVector<T> { + fn mul_assign(&mut self, matrix: &Matrix<T>) { + if let Some(result) = self.mul_matrix(matrix) { + *self = result; + } + } +} |
