use ndarray::{ArrayBase, DataMut, Ix2};
use num_traits::{One, Zero};
use crate::Scalar;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[allow(dead_code)]
pub enum Side {
Left,
Right,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[allow(dead_code)]
pub enum Pivot {
Variable,
Top,
Bottom,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[allow(dead_code)]
pub enum Direction {
Forward,
Backward,
}
#[allow(clippy::many_single_char_names)]
#[allow(clippy::too_many_lines)]
#[allow(dead_code)]
pub fn lasr<A, S>(
side: Side,
pivot: Pivot,
direct: Direction,
c: &[A::Real],
s: &[A::Real],
a: &mut ArrayBase<S, Ix2>,
) where
A: Scalar,
S: DataMut<Elem = A>,
{
let (m, n) = a.dim();
if m == 0 || n == 0 {
return;
}
match side {
Side::Left => {
assert_eq!(c.len(), m.saturating_sub(1), "c must have length M-1");
assert_eq!(s.len(), m.saturating_sub(1), "s must have length M-1");
match pivot {
Pivot::Variable => match direct {
Direction::Forward => {
for j in 0..m - 1 {
let ctemp = c[j];
let stemp = s[j];
if ctemp != A::Real::one() || stemp != A::Real::zero() {
for i in 0..n {
let temp = a[(j + 1, i)];
a[(j + 1, i)] = ctemp.into() * temp - stemp.into() * a[(j, i)];
a[(j, i)] = stemp.into() * temp + ctemp.into() * a[(j, i)];
}
}
}
}
Direction::Backward => {
for j in (0..m - 1).rev() {
let ctemp = c[j];
let stemp = s[j];
if ctemp != A::Real::one() || stemp != A::Real::zero() {
for i in 0..n {
let temp = a[(j + 1, i)];
a[(j + 1, i)] = ctemp.into() * temp - stemp.into() * a[(j, i)];
a[(j, i)] = stemp.into() * temp + ctemp.into() * a[(j, i)];
}
}
}
}
},
Pivot::Top => match direct {
Direction::Forward => {
for j in 1..m {
let ctemp = c[j - 1];
let stemp = s[j - 1];
if ctemp != A::Real::one() || stemp != A::Real::zero() {
for i in 0..n {
let temp = a[(j, i)];
a[(j, i)] = ctemp.into() * temp - stemp.into() * a[(0, i)];
a[(0, i)] = stemp.into() * temp + ctemp.into() * a[(0, i)];
}
}
}
}
Direction::Backward => {
for j in (1..m).rev() {
let ctemp = c[j - 1];
let stemp = s[j - 1];
if ctemp != A::Real::one() || stemp != A::Real::zero() {
for i in 0..n {
let temp = a[(j, i)];
a[(j, i)] = ctemp.into() * temp - stemp.into() * a[(0, i)];
a[(0, i)] = stemp.into() * temp + ctemp.into() * a[(0, i)];
}
}
}
}
},
Pivot::Bottom => match direct {
Direction::Forward => {
for j in 0..m - 1 {
let ctemp = c[j];
let stemp = s[j];
if ctemp != A::Real::one() || stemp != A::Real::zero() {
for i in 0..n {
let temp = a[(j, i)];
a[(j, i)] = stemp.into() * a[(m - 1, i)] + ctemp.into() * temp;
a[(m - 1, i)] =
ctemp.into() * a[(m - 1, i)] - stemp.into() * temp;
}
}
}
}
Direction::Backward => {
for j in (0..m - 1).rev() {
let ctemp = c[j];
let stemp = s[j];
if ctemp != A::Real::one() || stemp != A::Real::zero() {
for i in 0..n {
let temp = a[(j, i)];
a[(j, i)] = stemp.into() * a[(m - 1, i)] + ctemp.into() * temp;
a[(m - 1, i)] =
ctemp.into() * a[(m - 1, i)] - stemp.into() * temp;
}
}
}
}
},
}
}
Side::Right => {
assert_eq!(c.len(), n.saturating_sub(1), "c must have length N-1");
assert_eq!(s.len(), n.saturating_sub(1), "s must have length N-1");
match pivot {
Pivot::Variable => match direct {
Direction::Forward => {
for j in 0..n - 1 {
let ctemp = c[j];
let stemp = s[j];
if ctemp != A::Real::one() || stemp != A::Real::zero() {
for i in 0..m {
let temp = a[(i, j + 1)];
a[(i, j + 1)] = ctemp.into() * temp - stemp.into() * a[(i, j)];
a[(i, j)] = stemp.into() * temp + ctemp.into() * a[(i, j)];
}
}
}
}
Direction::Backward => {
for j in (0..n - 1).rev() {
let ctemp = c[j];
let stemp = s[j];
if ctemp != A::Real::one() || stemp != A::Real::zero() {
for i in 0..m {
let temp = a[(i, j + 1)];
a[(i, j + 1)] = ctemp.into() * temp - stemp.into() * a[(i, j)];
a[(i, j)] = stemp.into() * temp + ctemp.into() * a[(i, j)];
}
}
}
}
},
Pivot::Top => match direct {
Direction::Forward => {
for j in 1..n {
let ctemp = c[j - 1];
let stemp = s[j - 1];
if ctemp != A::Real::one() || stemp != A::Real::zero() {
for i in 0..m {
let temp = a[(i, j)];
a[(i, j)] = ctemp.into() * temp - stemp.into() * a[(i, 0)];
a[(i, 0)] = stemp.into() * temp + ctemp.into() * a[(i, 0)];
}
}
}
}
Direction::Backward => {
for j in (1..n).rev() {
let ctemp = c[j - 1];
let stemp = s[j - 1];
if ctemp != A::Real::one() || stemp != A::Real::zero() {
for i in 0..m {
let temp = a[(i, j)];
a[(i, j)] = ctemp.into() * temp - stemp.into() * a[(i, 0)];
a[(i, 0)] = stemp.into() * temp + ctemp.into() * a[(i, 0)];
}
}
}
}
},
Pivot::Bottom => match direct {
Direction::Forward => {
for j in 0..n - 1 {
let ctemp = c[j];
let stemp = s[j];
if ctemp != A::Real::one() || stemp != A::Real::zero() {
for i in 0..m {
let temp = a[(i, j)];
a[(i, j)] = stemp.into() * a[(i, n - 1)] + ctemp.into() * temp;
a[(i, n - 1)] =
ctemp.into() * a[(i, n - 1)] - stemp.into() * temp;
}
}
}
}
Direction::Backward => {
for j in (0..n - 1).rev() {
let ctemp = c[j];
let stemp = s[j];
if ctemp != A::Real::one() || stemp != A::Real::zero() {
for i in 0..m {
let temp = a[(i, j)];
a[(i, j)] = stemp.into() * a[(i, n - 1)] + ctemp.into() * temp;
a[(i, n - 1)] =
ctemp.into() * a[(i, n - 1)] - stemp.into() * temp;
}
}
}
}
},
}
}
}
}
#[cfg(test)]
mod tests {
use approx::assert_abs_diff_eq;
use ndarray::arr2;
use num_complex::Complex64;
use super::*;
#[test]
fn left_variable_forward_real() {
let mut a = arr2(&[[1.0, 2.0], [3.0, 4.0]]);
let c = vec![0.0];
let s = vec![1.0];
lasr::<f64, _>(
Side::Left,
Pivot::Variable,
Direction::Forward,
&c,
&s,
&mut a,
);
assert_abs_diff_eq!(a[(0, 0)], 3.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(0, 1)], 4.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(1, 0)], -1.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(1, 1)], -2.0, epsilon = 1e-10);
}
#[test]
fn left_variable_forward_complex() {
let mut a = arr2(&[
[Complex64::new(1.0, 1.0), Complex64::new(2.0, 2.0)],
[Complex64::new(3.0, 3.0), Complex64::new(4.0, 4.0)],
]);
let c = vec![0.0];
let s = vec![1.0];
lasr::<Complex64, _>(
Side::Left,
Pivot::Variable,
Direction::Forward,
&c,
&s,
&mut a,
);
assert_abs_diff_eq!(a[(0, 0)].re, 3.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(0, 0)].im, 3.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(0, 1)].re, 4.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(0, 1)].im, 4.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(1, 0)].re, -1.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(1, 0)].im, -1.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(1, 1)].re, -2.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(1, 1)].im, -2.0, epsilon = 1e-10);
}
#[test]
fn right_variable_forward_real() {
let mut a = arr2(&[[1.0, 2.0], [3.0, 4.0]]);
let c = vec![0.0];
let s = vec![1.0];
lasr::<f64, _>(
Side::Right,
Pivot::Variable,
Direction::Forward,
&c,
&s,
&mut a,
);
assert_abs_diff_eq!(a[(0, 0)], 2.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(1, 0)], 4.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(0, 1)], -1.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(1, 1)], -3.0, epsilon = 1e-10);
}
#[test]
fn left_top_forward_real() {
let mut a = arr2(&[[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]);
let c = vec![0.0, 1.0]; let s = vec![1.0, 0.0];
lasr::<f64, _>(Side::Left, Pivot::Top, Direction::Forward, &c, &s, &mut a);
assert_abs_diff_eq!(a[(0, 0)], 3.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(0, 1)], 4.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(1, 0)], -1.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(1, 1)], -2.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(2, 0)], 5.0, epsilon = 1e-10); assert_abs_diff_eq!(a[(2, 1)], 6.0, epsilon = 1e-10); }
#[test]
fn left_bottom_forward_real() {
let mut a = arr2(&[[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]);
let c = vec![0.0, 1.0]; let s = vec![1.0, 0.0];
lasr::<f64, _>(
Side::Left,
Pivot::Bottom,
Direction::Forward,
&c,
&s,
&mut a,
);
assert_abs_diff_eq!(a[(0, 0)], 5.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(0, 1)], 6.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(1, 0)], 3.0, epsilon = 1e-10); assert_abs_diff_eq!(a[(1, 1)], 4.0, epsilon = 1e-10); assert_abs_diff_eq!(a[(2, 0)], -1.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(2, 1)], -2.0, epsilon = 1e-10);
}
#[test]
fn left_variable_backward_real() {
let mut a = arr2(&[[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]);
let c = vec![1.0, 0.0]; let s = vec![0.0, 1.0];
lasr::<f64, _>(
Side::Left,
Pivot::Variable,
Direction::Backward,
&c,
&s,
&mut a,
);
assert_abs_diff_eq!(a[(0, 0)], 1.0, epsilon = 1e-10); assert_abs_diff_eq!(a[(0, 1)], 2.0, epsilon = 1e-10); assert_abs_diff_eq!(a[(1, 0)], 5.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(1, 1)], 6.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(2, 0)], -3.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(2, 1)], -4.0, epsilon = 1e-10);
}
#[test]
fn empty_matrix() {
let mut a = arr2::<f64, _>(&[[]]);
let c = vec![];
let s = vec![];
lasr::<f64, _>(
Side::Left,
Pivot::Variable,
Direction::Forward,
&c,
&s,
&mut a,
);
}
#[test]
fn identity_rotation() {
let mut a = arr2(&[[1.0, 2.0], [3.0, 4.0]]);
let c = vec![1.0]; let s = vec![0.0];
lasr::<f64, _>(
Side::Left,
Pivot::Variable,
Direction::Forward,
&c,
&s,
&mut a,
);
assert_abs_diff_eq!(a[(0, 0)], 1.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(0, 1)], 2.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(1, 0)], 3.0, epsilon = 1e-10);
assert_abs_diff_eq!(a[(1, 1)], 4.0, epsilon = 1e-10);
}
#[test]
#[should_panic(expected = "c must have length M-1")]
fn wrong_c_length_left() {
let mut a = arr2(&[[1.0, 2.0], [3.0, 4.0]]);
let c = vec![0.8, 0.6]; let s = vec![0.6];
lasr::<f64, _>(
Side::Left,
Pivot::Variable,
Direction::Forward,
&c,
&s,
&mut a,
);
}
#[test]
#[should_panic(expected = "c must have length N-1")]
fn wrong_c_length_right() {
let mut a = arr2(&[[1.0, 2.0], [3.0, 4.0]]);
let c = vec![0.8, 0.6]; let s = vec![0.6];
lasr::<f64, _>(
Side::Right,
Pivot::Variable,
Direction::Forward,
&c,
&s,
&mut a,
);
}
}