use oxiblas_core::scalar::{ComplexScalar, Field, Real, Scalar};
use oxiblas_matrix::{Mat, MatRef};
use super::complex_bidiag::ComplexBidiagFactors;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BidiagError {
EmptyMatrix,
DimensionMismatch,
InvalidParameter,
}
impl core::fmt::Display for BidiagError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::EmptyMatrix => write!(f, "Matrix is empty"),
Self::DimensionMismatch => write!(f, "Dimension mismatch"),
Self::InvalidParameter => write!(f, "Invalid parameter"),
}
}
}
impl std::error::Error for BidiagError {}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BidiagVect {
Q,
P,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Side {
Left,
Right,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Trans {
NoTrans,
Trans,
}
#[derive(Debug, Clone)]
pub struct BidiagFactors<T: Scalar> {
pub work: Mat<T>,
pub d: Vec<T>,
pub e: Vec<T>,
pub tauq: Vec<T>,
pub taup: Vec<T>,
pub m: usize,
pub n: usize,
}
impl<T: Field + Real + bytemuck::Zeroable> BidiagFactors<T> {
pub fn compute(a: MatRef<'_, T>) -> Result<Self, BidiagError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Err(BidiagError::EmptyMatrix);
}
if m >= n {
Self::compute_tall(a)
} else {
Self::compute_wide(a)
}
}
pub fn compute_auto(a: MatRef<'_, T>) -> Result<Self, BidiagError> {
const AUTO_BLOCK_THRESHOLD: usize = 64;
let m = a.nrows();
let n = a.ncols();
let min_dim = m.min(n);
if min_dim >= AUTO_BLOCK_THRESHOLD {
Self::compute_blocked(a)
} else {
Self::compute(a)
}
}
pub fn compute_blocked(a: MatRef<'_, T>) -> Result<Self, BidiagError> {
let nb = crate::workspace::optimal_block_size_bidiag(a.nrows(), a.ncols());
Self::compute_blocked_with_block_size(a, nb)
}
pub fn compute_blocked_with_block_size(
a: MatRef<'_, T>,
block_size: usize,
) -> Result<Self, BidiagError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Err(BidiagError::EmptyMatrix);
}
let min_dim = m.min(n);
if min_dim <= block_size || min_dim <= 2 {
return Self::compute(a);
}
if m >= n {
Self::compute_tall_blocked(a, block_size)
} else {
Self::compute_wide_blocked(a, block_size)
}
}
fn compute_tall_blocked(a: MatRef<'_, T>, block_size: usize) -> Result<Self, BidiagError> {
let m = a.nrows();
let n = a.ncols();
let mut work = Mat::zeros(m, n);
for i in 0..m {
for j in 0..n {
work[(i, j)] = a[(i, j)];
}
}
let mut tauq = vec![T::zero(); n];
let num_p = n.saturating_sub(1);
let mut taup = vec![T::zero(); num_p];
let mut d = vec![T::zero(); n];
let mut e = vec![T::zero(); num_p];
let nb = block_size.min(n);
let mut j = 0;
while j < n {
let jb = nb.min(n - j);
for jj in j..(j + jb) {
let (tau, beta) = householder_left(&mut work, jj, m, n);
d[jj] = beta;
tauq[jj] = tau;
apply_householder_left(&mut work, jj, m, n, tau);
if jj < n - 1 {
let (tau, beta) = householder_right(&mut work, jj, m, n);
e[jj] = beta;
taup[jj] = tau;
apply_householder_right(&mut work, jj, m, n, tau);
}
}
j += jb;
}
Ok(Self {
work,
d,
e,
tauq,
taup,
m,
n,
})
}
fn compute_wide_blocked(a: MatRef<'_, T>, block_size: usize) -> Result<Self, BidiagError> {
let m = a.nrows();
let n = a.ncols();
let mut work = Mat::zeros(m, n);
for i in 0..m {
for j in 0..n {
work[(i, j)] = a[(i, j)];
}
}
let mut tauq = vec![T::zero(); m];
let mut taup = vec![T::zero(); m];
let mut d = vec![T::zero(); m];
let num_e = if m > 0 { m - 1 } else { 0 };
let mut e = vec![T::zero(); num_e];
let nb = block_size.min(m);
let mut j = 0;
while j < m {
let jb = nb.min(m - j);
for jj in j..(j + jb) {
let (tau_p, beta_d) = householder_right_wide(&mut work, jj, m, n);
d[jj] = beta_d;
taup[jj] = tau_p;
apply_householder_right_wide(&mut work, jj, m, n, tau_p);
if jj < m - 1 {
let (tau_q, beta_e) = householder_left_wide(&mut work, jj, m, n);
e[jj] = beta_e;
tauq[jj] = tau_q;
apply_householder_left_wide(&mut work, jj, m, n, tau_q);
}
}
j += jb;
}
if m > 0 {
tauq[m - 1] = T::zero();
}
Ok(Self {
work,
d,
e,
tauq,
taup,
m,
n,
})
}
fn compute_tall(a: MatRef<'_, T>) -> Result<Self, BidiagError> {
let m = a.nrows();
let n = a.ncols();
let mut work = Mat::zeros(m, n);
for i in 0..m {
for j in 0..n {
work[(i, j)] = a[(i, j)];
}
}
let mut tauq = vec![T::zero(); n];
let num_p = n.saturating_sub(1);
let mut taup = vec![T::zero(); num_p];
let mut d = vec![T::zero(); n];
let mut e = vec![T::zero(); num_p];
for j in 0..n {
let (tau, beta) = householder_left(&mut work, j, m, n);
d[j] = beta;
tauq[j] = tau;
apply_householder_left(&mut work, j, m, n, tau);
if j < n - 1 {
let (tau, beta) = householder_right(&mut work, j, m, n);
e[j] = beta;
taup[j] = tau;
apply_householder_right(&mut work, j, m, n, tau);
}
}
Ok(Self {
work,
d,
e,
tauq,
taup,
m,
n,
})
}
fn compute_wide(a: MatRef<'_, T>) -> Result<Self, BidiagError> {
let m = a.nrows();
let n = a.ncols();
let mut work = Mat::zeros(m, n);
for i in 0..m {
for j in 0..n {
work[(i, j)] = a[(i, j)];
}
}
let mut tauq = vec![T::zero(); m];
let mut taup = vec![T::zero(); m];
let mut d = vec![T::zero(); m];
let num_e = if m > 0 { m - 1 } else { 0 };
let mut e = vec![T::zero(); num_e];
for j in 0..m {
let (tau_p, beta_d) = householder_right_wide(&mut work, j, m, n);
d[j] = beta_d;
taup[j] = tau_p;
apply_householder_right_wide(&mut work, j, m, n, tau_p);
if j < m - 1 {
let (tau_q, beta_e) = householder_left_wide(&mut work, j, m, n);
e[j] = beta_e;
tauq[j] = tau_q;
apply_householder_left_wide(&mut work, j, m, n, tau_q);
}
}
if m > 0 {
tauq[m - 1] = T::zero();
}
Ok(Self {
work,
d,
e,
tauq,
taup,
m,
n,
})
}
pub fn diagonal(&self) -> &[T] {
&self.d
}
pub fn superdiagonal(&self) -> &[T] {
&self.e
}
pub fn generate(&self, vect: BidiagVect) -> Result<Mat<T>, BidiagError> {
match vect {
BidiagVect::Q => self.generate_q(),
BidiagVect::P => self.generate_p(),
}
}
fn generate_q(&self) -> Result<Mat<T>, BidiagError> {
if self.m >= self.n {
let mut q = Mat::zeros(self.m, self.m);
for i in 0..self.m {
q[(i, i)] = T::one();
}
for j in 0..self.n {
let tau = self.tauq[j];
if tau != T::zero() {
for r in 0..self.m {
let mut w = q[(r, j)];
for i in (j + 1)..self.m {
w = w + q[(r, i)] * self.work[(i, j)];
}
let tw = tau * w;
q[(r, j)] = q[(r, j)] - tw;
for i in (j + 1)..self.m {
q[(r, i)] = q[(r, i)] - tw * self.work[(i, j)];
}
}
}
}
let mut q_thin = Mat::zeros(self.m, self.n);
for i in 0..self.m {
for j in 0..self.n {
q_thin[(i, j)] = q[(i, j)];
}
}
Ok(q_thin)
} else {
let mut q = Mat::zeros(self.m, self.m);
for i in 0..self.m {
q[(i, i)] = T::one();
}
let num_q = self.m.saturating_sub(1);
for j in 0..num_q {
let tau = self.tauq[j];
if tau != T::zero() {
let start = j + 1;
for r in 0..self.m {
let mut w = q[(r, start)];
for i in (start + 1)..self.m {
w = w + q[(r, i)] * self.work[(i, j)];
}
let tw = tau * w;
q[(r, start)] = q[(r, start)] - tw;
for i in (start + 1)..self.m {
q[(r, i)] = q[(r, i)] - tw * self.work[(i, j)];
}
}
}
}
Ok(q)
}
}
fn generate_p(&self) -> Result<Mat<T>, BidiagError> {
if self.m >= self.n {
let mut v = Mat::zeros(self.n, self.n);
for i in 0..self.n {
v[(i, i)] = T::one();
}
let num_p = self.taup.len();
for j in 0..num_p {
let tau = self.taup[j];
if tau != T::zero() {
let start = j + 1;
for r in 0..self.n {
let mut w = v[(r, start)];
for i in (start + 1)..self.n {
w = w + v[(r, i)] * self.work[(j, i)];
}
let tw = tau * w;
v[(r, start)] = v[(r, start)] - tw;
for i in (start + 1)..self.n {
v[(r, i)] = v[(r, i)] - tw * self.work[(j, i)];
}
}
}
}
Ok(v)
} else {
let mut v = Mat::zeros(self.n, self.n);
for i in 0..self.n {
v[(i, i)] = T::one();
}
for j in 0..self.m {
let tau = self.taup[j];
if tau != T::zero() {
for r in 0..self.n {
let mut w = v[(r, j)];
for i in (j + 1)..self.n {
w = w + v[(r, i)] * self.work[(j, i)];
}
let tw = tau * w;
v[(r, j)] = v[(r, j)] - tw;
for i in (j + 1)..self.n {
v[(r, i)] = v[(r, i)] - tw * self.work[(j, i)];
}
}
}
}
let mut p_thin = Mat::zeros(self.m, self.n);
for i in 0..self.m {
for j in 0..self.n {
p_thin[(i, j)] = v[(i, j)];
}
}
Ok(p_thin)
}
}
pub fn apply(
&self,
vect: BidiagVect,
side: Side,
trans: Trans,
c: MatRef<'_, T>,
) -> Result<Mat<T>, BidiagError> {
match vect {
BidiagVect::Q => self.apply_q(side, trans, c),
BidiagVect::P => self.apply_p(side, trans, c),
}
}
fn apply_q(&self, side: Side, trans: Trans, c: MatRef<'_, T>) -> Result<Mat<T>, BidiagError> {
let c_rows = c.nrows();
let c_cols = c.ncols();
match side {
Side::Left => {
let q_rows = self.m;
if c_rows != q_rows {
return Err(BidiagError::DimensionMismatch);
}
}
Side::Right => {
let q_cols = if self.m >= self.n { self.n } else { self.m };
if c_cols != q_cols && c_cols != self.m {
return Err(BidiagError::DimensionMismatch);
}
}
}
let mut result = Mat::zeros(c_rows, c_cols);
for i in 0..c_rows {
for j in 0..c_cols {
result[(i, j)] = c[(i, j)];
}
}
if self.m >= self.n {
self.apply_q_tall(&mut result, side, trans)?;
} else {
self.apply_q_wide(&mut result, side, trans)?;
}
Ok(result)
}
fn apply_q_tall(&self, c: &mut Mat<T>, side: Side, trans: Trans) -> Result<(), BidiagError> {
let c_rows = c.nrows();
let c_cols = c.ncols();
let apply_forward = matches!(trans, Trans::NoTrans);
match side {
Side::Left => {
if apply_forward {
for j in 0..self.n {
let tau = self.tauq[j];
if tau != T::zero() {
for col in 0..c_cols {
let mut w = c[(j, col)];
for i in (j + 1)..c_rows {
w = w + self.work[(i, j)] * c[(i, col)];
}
let tw = tau * w;
c[(j, col)] = c[(j, col)] - tw;
for i in (j + 1)..c_rows {
c[(i, col)] = c[(i, col)] - tw * self.work[(i, j)];
}
}
}
}
} else {
for j in (0..self.n).rev() {
let tau = self.tauq[j];
if tau != T::zero() {
for col in 0..c_cols {
let mut w = c[(j, col)];
for i in (j + 1)..c_rows {
w = w + self.work[(i, j)] * c[(i, col)];
}
let tw = tau * w;
c[(j, col)] = c[(j, col)] - tw;
for i in (j + 1)..c_rows {
c[(i, col)] = c[(i, col)] - tw * self.work[(i, j)];
}
}
}
}
}
}
Side::Right => {
let apply_forward_right = !apply_forward;
if apply_forward_right {
for j in 0..self.n {
let tau = self.tauq[j];
if tau != T::zero() {
for row in 0..c_rows {
let mut w = c[(row, j)];
for i in (j + 1)..c_cols.min(self.m) {
w = w + c[(row, i)] * self.work[(i, j)];
}
let tw = tau * w;
c[(row, j)] = c[(row, j)] - tw;
for i in (j + 1)..c_cols.min(self.m) {
c[(row, i)] = c[(row, i)] - tw * self.work[(i, j)];
}
}
}
}
} else {
for j in (0..self.n).rev() {
let tau = self.tauq[j];
if tau != T::zero() {
for row in 0..c_rows {
let mut w = c[(row, j)];
for i in (j + 1)..c_cols.min(self.m) {
w = w + c[(row, i)] * self.work[(i, j)];
}
let tw = tau * w;
c[(row, j)] = c[(row, j)] - tw;
for i in (j + 1)..c_cols.min(self.m) {
c[(row, i)] = c[(row, i)] - tw * self.work[(i, j)];
}
}
}
}
}
}
}
Ok(())
}
fn apply_q_wide(&self, c: &mut Mat<T>, side: Side, trans: Trans) -> Result<(), BidiagError> {
let c_rows = c.nrows();
let c_cols = c.ncols();
let apply_forward = matches!(trans, Trans::NoTrans);
let num_q = self.m.saturating_sub(1);
match side {
Side::Left => {
if apply_forward {
for j in 0..num_q {
let tau = self.tauq[j];
if tau != T::zero() {
let start = j + 1;
for col in 0..c_cols {
let mut w = c[(start, col)];
for i in (start + 1)..c_rows {
w = w + self.work[(i, j)] * c[(i, col)];
}
let tw = tau * w;
c[(start, col)] = c[(start, col)] - tw;
for i in (start + 1)..c_rows {
c[(i, col)] = c[(i, col)] - tw * self.work[(i, j)];
}
}
}
}
} else {
for j in (0..num_q).rev() {
let tau = self.tauq[j];
if tau != T::zero() {
let start = j + 1;
for col in 0..c_cols {
let mut w = c[(start, col)];
for i in (start + 1)..c_rows {
w = w + self.work[(i, j)] * c[(i, col)];
}
let tw = tau * w;
c[(start, col)] = c[(start, col)] - tw;
for i in (start + 1)..c_rows {
c[(i, col)] = c[(i, col)] - tw * self.work[(i, j)];
}
}
}
}
}
}
Side::Right => {
let apply_forward_right = !apply_forward;
if apply_forward_right {
for j in 0..num_q {
let tau = self.tauq[j];
if tau != T::zero() {
let start = j + 1;
for row in 0..c_rows {
let mut w = c[(row, start)];
for i in (start + 1)..c_cols {
w = w + c[(row, i)] * self.work[(i, j)];
}
let tw = tau * w;
c[(row, start)] = c[(row, start)] - tw;
for i in (start + 1)..c_cols {
c[(row, i)] = c[(row, i)] - tw * self.work[(i, j)];
}
}
}
}
} else {
for j in (0..num_q).rev() {
let tau = self.tauq[j];
if tau != T::zero() {
let start = j + 1;
for row in 0..c_rows {
let mut w = c[(row, start)];
for i in (start + 1)..c_cols {
w = w + c[(row, i)] * self.work[(i, j)];
}
let tw = tau * w;
c[(row, start)] = c[(row, start)] - tw;
for i in (start + 1)..c_cols {
c[(row, i)] = c[(row, i)] - tw * self.work[(i, j)];
}
}
}
}
}
}
}
Ok(())
}
fn apply_p(&self, side: Side, trans: Trans, c: MatRef<'_, T>) -> Result<Mat<T>, BidiagError> {
let c_rows = c.nrows();
let c_cols = c.ncols();
let mut result = Mat::zeros(c_rows, c_cols);
for i in 0..c_rows {
for j in 0..c_cols {
result[(i, j)] = c[(i, j)];
}
}
if self.m >= self.n {
self.apply_p_tall(&mut result, side, trans)?;
} else {
self.apply_p_wide(&mut result, side, trans)?;
}
Ok(result)
}
fn apply_p_tall(&self, c: &mut Mat<T>, side: Side, trans: Trans) -> Result<(), BidiagError> {
let c_rows = c.nrows();
let c_cols = c.ncols();
let num_p = self.taup.len();
let apply_forward = matches!(trans, Trans::NoTrans);
match side {
Side::Left => {
if !apply_forward {
for j in 0..num_p {
let tau = self.taup[j];
if tau != T::zero() {
let start = j + 1;
for col in 0..c_cols {
let mut w = c[(start, col)];
for i in (start + 1)..c_rows.min(self.n) {
w = w + self.work[(j, i)] * c[(i, col)];
}
let tw = tau * w;
c[(start, col)] = c[(start, col)] - tw;
for i in (start + 1)..c_rows.min(self.n) {
c[(i, col)] = c[(i, col)] - tw * self.work[(j, i)];
}
}
}
}
} else {
for j in (0..num_p).rev() {
let tau = self.taup[j];
if tau != T::zero() {
let start = j + 1;
for col in 0..c_cols {
let mut w = c[(start, col)];
for i in (start + 1)..c_rows.min(self.n) {
w = w + self.work[(j, i)] * c[(i, col)];
}
let tw = tau * w;
c[(start, col)] = c[(start, col)] - tw;
for i in (start + 1)..c_rows.min(self.n) {
c[(i, col)] = c[(i, col)] - tw * self.work[(j, i)];
}
}
}
}
}
}
Side::Right => {
if apply_forward {
for j in 0..num_p {
let tau = self.taup[j];
if tau != T::zero() {
let start = j + 1;
for row in 0..c_rows {
let mut w = c[(row, start)];
for i in (start + 1)..c_cols.min(self.n) {
w = w + c[(row, i)] * self.work[(j, i)];
}
let tw = tau * w;
c[(row, start)] = c[(row, start)] - tw;
for i in (start + 1)..c_cols.min(self.n) {
c[(row, i)] = c[(row, i)] - tw * self.work[(j, i)];
}
}
}
}
} else {
for j in (0..num_p).rev() {
let tau = self.taup[j];
if tau != T::zero() {
let start = j + 1;
for row in 0..c_rows {
let mut w = c[(row, start)];
for i in (start + 1)..c_cols.min(self.n) {
w = w + c[(row, i)] * self.work[(j, i)];
}
let tw = tau * w;
c[(row, start)] = c[(row, start)] - tw;
for i in (start + 1)..c_cols.min(self.n) {
c[(row, i)] = c[(row, i)] - tw * self.work[(j, i)];
}
}
}
}
}
}
}
Ok(())
}
fn apply_p_wide(&self, c: &mut Mat<T>, side: Side, trans: Trans) -> Result<(), BidiagError> {
let c_rows = c.nrows();
let c_cols = c.ncols();
let num_p = self.taup.len();
let apply_forward = matches!(trans, Trans::NoTrans);
match side {
Side::Left => {
if !apply_forward {
for j in 0..num_p {
let tau = self.taup[j];
if tau != T::zero() {
for col in 0..c_cols {
let mut w = c[(j, col)];
for i in (j + 1)..c_rows.min(self.n) {
w = w + self.work[(j, i)] * c[(i, col)];
}
let tw = tau * w;
c[(j, col)] = c[(j, col)] - tw;
for i in (j + 1)..c_rows.min(self.n) {
c[(i, col)] = c[(i, col)] - tw * self.work[(j, i)];
}
}
}
}
} else {
for j in (0..num_p).rev() {
let tau = self.taup[j];
if tau != T::zero() {
for col in 0..c_cols {
let mut w = c[(j, col)];
for i in (j + 1)..c_rows.min(self.n) {
w = w + self.work[(j, i)] * c[(i, col)];
}
let tw = tau * w;
c[(j, col)] = c[(j, col)] - tw;
for i in (j + 1)..c_rows.min(self.n) {
c[(i, col)] = c[(i, col)] - tw * self.work[(j, i)];
}
}
}
}
}
}
Side::Right => {
if apply_forward {
for j in 0..num_p {
let tau = self.taup[j];
if tau != T::zero() {
for row in 0..c_rows {
let mut w = c[(row, j)];
for i in (j + 1)..c_cols {
w = w + c[(row, i)] * self.work[(j, i)];
}
let tw = tau * w;
c[(row, j)] = c[(row, j)] - tw;
for i in (j + 1)..c_cols {
c[(row, i)] = c[(row, i)] - tw * self.work[(j, i)];
}
}
}
}
} else {
for j in (0..num_p).rev() {
let tau = self.taup[j];
if tau != T::zero() {
for row in 0..c_rows {
let mut w = c[(row, j)];
for i in (j + 1)..c_cols {
w = w + c[(row, i)] * self.work[(j, i)];
}
let tw = tau * w;
c[(row, j)] = c[(row, j)] - tw;
for i in (j + 1)..c_cols {
c[(row, i)] = c[(row, i)] - tw * self.work[(j, i)];
}
}
}
}
}
}
}
Ok(())
}
}
fn householder_left<T: Field + Real>(work: &mut Mat<T>, j: usize, m: usize, _n: usize) -> (T, T) {
let mut norm_sq = T::zero();
for i in j..m {
norm_sq = norm_sq + work[(i, j)] * work[(i, j)];
}
let norm = Real::sqrt(norm_sq);
if norm == T::zero() {
return (T::zero(), T::zero());
}
let x_j = work[(j, j)];
let beta = if x_j >= T::zero() { -norm } else { norm };
let tau = (beta - x_j) / beta;
let scale = T::one() / (x_j - beta);
for i in (j + 1)..m {
work[(i, j)] = work[(i, j)] * scale;
}
(tau, beta)
}
fn householder_right<T: Field + Real>(work: &mut Mat<T>, j: usize, _m: usize, n: usize) -> (T, T) {
let start_col = j + 1;
let mut norm_sq = T::zero();
for i in start_col..n {
norm_sq = norm_sq + work[(j, i)] * work[(j, i)];
}
let norm = Real::sqrt(norm_sq);
if norm == T::zero() {
return (T::zero(), T::zero());
}
let x_j = work[(j, start_col)];
let beta = if x_j >= T::zero() { -norm } else { norm };
let tau = (beta - x_j) / beta;
let scale = T::one() / (x_j - beta);
for i in (start_col + 1)..n {
work[(j, i)] = work[(j, i)] * scale;
}
(tau, beta)
}
fn apply_householder_left<T: Field + Real>(
work: &mut Mat<T>,
j: usize,
m: usize,
n: usize,
tau: T,
) {
if tau == T::zero() {
return;
}
for col in (j + 1)..n {
let mut w = work[(j, col)];
for i in (j + 1)..m {
w = w + work[(i, j)] * work[(i, col)];
}
let tw = tau * w;
work[(j, col)] = work[(j, col)] - tw;
for i in (j + 1)..m {
work[(i, col)] = work[(i, col)] - tw * work[(i, j)];
}
}
}
fn apply_householder_right<T: Field + Real>(
work: &mut Mat<T>,
j: usize,
m: usize,
n: usize,
tau: T,
) {
if tau == T::zero() {
return;
}
let start_col = j + 1;
for row in (j + 1)..m {
let mut w = work[(row, start_col)];
for i in (start_col + 1)..n {
w = w + work[(j, i)] * work[(row, i)];
}
let tw = tau * w;
work[(row, start_col)] = work[(row, start_col)] - tw;
for i in (start_col + 1)..n {
work[(row, i)] = work[(row, i)] - tw * work[(j, i)];
}
}
}
fn householder_right_wide<T: Field + Real>(
work: &mut Mat<T>,
j: usize,
_m: usize,
n: usize,
) -> (T, T) {
let mut norm_sq = T::zero();
for i in j..n {
norm_sq = norm_sq + work[(j, i)] * work[(j, i)];
}
let norm = Real::sqrt(norm_sq);
if norm == T::zero() {
return (T::zero(), T::zero());
}
let x_j = work[(j, j)];
let beta = if x_j >= T::zero() { -norm } else { norm };
let tau = (beta - x_j) / beta;
let scale = T::one() / (x_j - beta);
for i in (j + 1)..n {
work[(j, i)] = work[(j, i)] * scale;
}
(tau, beta)
}
fn apply_householder_right_wide<T: Field + Real>(
work: &mut Mat<T>,
j: usize,
m: usize,
n: usize,
tau: T,
) {
if tau == T::zero() {
return;
}
for row in (j + 1)..m {
let mut w = work[(row, j)];
for i in (j + 1)..n {
w = w + work[(j, i)] * work[(row, i)];
}
let tw = tau * w;
work[(row, j)] = work[(row, j)] - tw;
for i in (j + 1)..n {
work[(row, i)] = work[(row, i)] - tw * work[(j, i)];
}
}
}
fn householder_left_wide<T: Field + Real>(
work: &mut Mat<T>,
j: usize,
m: usize,
_n: usize,
) -> (T, T) {
let start = j + 1;
let mut norm_sq = T::zero();
for i in start..m {
norm_sq = norm_sq + work[(i, j)] * work[(i, j)];
}
let norm = Real::sqrt(norm_sq);
if norm == T::zero() {
return (T::zero(), T::zero());
}
let x_j = work[(start, j)];
let beta = if x_j >= T::zero() { -norm } else { norm };
let tau = (beta - x_j) / beta;
let scale = T::one() / (x_j - beta);
for i in (start + 1)..m {
work[(i, j)] = work[(i, j)] * scale;
}
(tau, beta)
}
fn apply_householder_left_wide<T: Field + Real>(
work: &mut Mat<T>,
j: usize,
m: usize,
n: usize,
tau: T,
) {
if tau == T::zero() {
return;
}
let start = j + 1;
for col in (j + 1)..n {
let mut w = work[(start, col)];
for i in (start + 1)..m {
w = w + work[(i, j)] * work[(i, col)];
}
let tw = tau * w;
work[(start, col)] = work[(start, col)] - tw;
for i in (start + 1)..m {
work[(i, col)] = work[(i, col)] - tw * work[(i, j)];
}
}
}
pub fn gebrd<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
) -> Result<BidiagFactors<T>, BidiagError> {
BidiagFactors::compute(a)
}
pub fn ormbr<T: Field + Real + bytemuck::Zeroable>(
factors: &BidiagFactors<T>,
vect: BidiagVect,
side: Side,
trans: Trans,
c: MatRef<'_, T>,
) -> Result<Mat<T>, BidiagError> {
factors.apply(vect, side, trans, c)
}
pub fn orgbr<T: Field + Real + bytemuck::Zeroable>(
factors: &BidiagFactors<T>,
vect: BidiagVect,
) -> Result<Mat<T>, BidiagError> {
factors.generate(vect)
}
pub fn unmbr<T: Field + ComplexScalar + bytemuck::Zeroable>(
factors: &ComplexBidiagFactors<T>,
vect: BidiagVect,
side: Side,
trans: Trans,
c: MatRef<'_, T>,
) -> Result<Mat<T>, BidiagError>
where
T::Real: Real,
{
factors.apply(vect, side, trans, c)
}
pub fn ungbr<T: Field + ComplexScalar + bytemuck::Zeroable>(
factors: &ComplexBidiagFactors<T>,
vect: BidiagVect,
) -> Result<Mat<T>, BidiagError>
where
T::Real: Real,
{
factors.generate(vect)
}
#[cfg(test)]
mod tests {
use super::*;
fn approx_eq(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
fn matrix_approx_eq(a: &Mat<f64>, b: &Mat<f64>, tol: f64) -> bool {
if a.nrows() != b.nrows() || a.ncols() != b.ncols() {
return false;
}
for i in 0..a.nrows() {
for j in 0..a.ncols() {
if !approx_eq(a[(i, j)], b[(i, j)], tol) {
return false;
}
}
}
true
}
#[test]
fn test_gebrd_tall() {
let a = Mat::from_rows(&[
&[1.0f64, 2.0, 3.0],
&[4.0, 5.0, 6.0],
&[7.0, 8.0, 9.0],
&[10.0, 11.0, 12.0],
]);
let factors = gebrd(a.as_ref()).unwrap();
assert_eq!(factors.d.len(), 3);
assert_eq!(factors.e.len(), 2);
assert_eq!(factors.tauq.len(), 3);
assert_eq!(factors.taup.len(), 2);
}
#[test]
fn test_gebrd_square() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 9.0]]);
let factors = gebrd(a.as_ref()).unwrap();
assert_eq!(factors.d.len(), 3);
assert_eq!(factors.e.len(), 2);
}
#[test]
fn test_gebrd_wide() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0, 4.0], &[5.0, 6.0, 7.0, 8.0]]);
let factors = gebrd(a.as_ref()).unwrap();
assert_eq!(factors.d.len(), 2);
assert_eq!(factors.e.len(), 1);
}
#[test]
fn test_orgbr_q_tall() {
let a = Mat::from_rows(&[
&[1.0f64, 2.0, 3.0],
&[4.0, 5.0, 6.0],
&[7.0, 8.0, 9.0],
&[10.0, 11.0, 12.0],
]);
let factors = gebrd(a.as_ref()).unwrap();
let q = orgbr(&factors, BidiagVect::Q).unwrap();
assert_eq!(q.nrows(), 4);
assert_eq!(q.ncols(), 3);
for i in 0..3 {
for j in 0..3 {
let mut sum = 0.0;
for k in 0..4 {
sum += q[(k, i)] * q[(k, j)];
}
let expected = if i == j { 1.0 } else { 0.0 };
assert!(
approx_eq(sum, expected, 1e-10),
"Q^T*Q[{},{}] = {}, expected {}",
i,
j,
sum,
expected
);
}
}
}
#[test]
fn test_orgbr_p_tall() {
let a = Mat::from_rows(&[
&[1.0f64, 2.0, 3.0],
&[4.0, 5.0, 6.0],
&[7.0, 8.0, 9.0],
&[10.0, 11.0, 12.0],
]);
let factors = gebrd(a.as_ref()).unwrap();
let p = orgbr(&factors, BidiagVect::P).unwrap();
assert_eq!(p.nrows(), 3);
assert_eq!(p.ncols(), 3);
for i in 0..3 {
for j in 0..3 {
let mut sum = 0.0;
for k in 0..3 {
sum += p[(k, i)] * p[(k, j)];
}
let expected = if i == j { 1.0 } else { 0.0 };
assert!(
approx_eq(sum, expected, 1e-10),
"P^T*P[{},{}] = {}, expected {}",
i,
j,
sum,
expected
);
}
}
}
#[test]
fn test_bidiag_reconstruction() {
let a = Mat::from_rows(&[
&[1.0f64, 2.0, 3.0],
&[4.0, 5.0, 6.0],
&[7.0, 8.0, 9.0],
&[10.0, 11.0, 12.0],
]);
let factors = gebrd(a.as_ref()).unwrap();
let q = orgbr(&factors, BidiagVect::Q).unwrap();
let p = orgbr(&factors, BidiagVect::P).unwrap();
let n = factors.d.len();
let mut b = Mat::zeros(4, 3);
for i in 0..n {
b[(i, i)] = factors.d[i];
}
for i in 0..factors.e.len() {
b[(i, i + 1)] = factors.e[i];
}
let mut bp = Mat::zeros(4, 3);
for i in 0..4 {
for j in 0..3 {
let mut sum = 0.0;
for k in 0..3 {
sum += b[(i, k)] * p[(j, k)]; }
bp[(i, j)] = sum;
}
}
let mut reconstructed = Mat::zeros(4, 3);
for i in 0..4 {
for j in 0..3 {
let mut sum = 0.0;
for k in 0..3 {
sum += q[(i, k)] * bp[(k, j)];
}
reconstructed[(i, j)] = sum;
}
}
for i in 0..4 {
for j in 0..3 {
assert!(
approx_eq(reconstructed[(i, j)], a[(i, j)], 1e-10),
"reconstructed[{},{}] = {}, a = {}",
i,
j,
reconstructed[(i, j)],
a[(i, j)]
);
}
}
}
#[test]
fn test_ormbr_q_left_notrans() {
let a = Mat::from_rows(&[
&[1.0f64, 2.0, 3.0],
&[4.0, 5.0, 6.0],
&[7.0, 8.0, 9.0],
&[10.0, 11.0, 12.0],
]);
let factors = gebrd(a.as_ref()).unwrap();
let _q = orgbr(&factors, BidiagVect::Q).unwrap();
let c = Mat::from_rows(&[&[1.0f64, 0.0], &[0.0, 1.0], &[1.0, 1.0], &[0.0, 0.0]]);
let result = ormbr(
&factors,
BidiagVect::Q,
Side::Left,
Trans::NoTrans,
c.as_ref(),
)
.unwrap();
assert_eq!(result.nrows(), 4);
assert_eq!(result.ncols(), 2);
}
#[test]
fn test_ormbr_p_right_notrans() {
let a = Mat::from_rows(&[
&[1.0f64, 2.0, 3.0],
&[4.0, 5.0, 6.0],
&[7.0, 8.0, 9.0],
&[10.0, 11.0, 12.0],
]);
let factors = gebrd(a.as_ref()).unwrap();
let c = Mat::from_rows(&[&[1.0f64, 0.0, 1.0], &[0.0, 1.0, 0.0]]);
let result = ormbr(
&factors,
BidiagVect::P,
Side::Right,
Trans::NoTrans,
c.as_ref(),
)
.unwrap();
assert_eq!(result.nrows(), 2);
assert_eq!(result.ncols(), 3);
}
#[test]
fn test_ormbr_roundtrip() {
let a = Mat::from_rows(&[
&[1.0f64, 2.0, 3.0],
&[4.0, 5.0, 6.0],
&[7.0, 8.0, 9.0],
&[10.0, 11.0, 12.0],
]);
let factors = gebrd(a.as_ref()).unwrap();
let c = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0], &[5.0, 6.0], &[7.0, 8.0]]);
let qt_c = ormbr(
&factors,
BidiagVect::Q,
Side::Left,
Trans::Trans,
c.as_ref(),
)
.unwrap();
let q_qt_c = ormbr(
&factors,
BidiagVect::Q,
Side::Left,
Trans::NoTrans,
qt_c.as_ref(),
)
.unwrap();
assert!(matrix_approx_eq(&q_qt_c, &c, 1e-10));
}
#[test]
fn test_gebrd_1x1() {
let a = Mat::from_rows(&[&[5.0f64]]);
let factors = gebrd(a.as_ref()).unwrap();
assert_eq!(factors.d.len(), 1);
assert_eq!(factors.e.len(), 0);
assert!(approx_eq(factors.d[0].abs(), 5.0, 1e-10));
}
#[test]
fn test_gebrd_2x2() {
let a = Mat::from_rows(&[&[3.0f64, 4.0], &[0.0, 5.0]]);
let factors = gebrd(a.as_ref()).unwrap();
assert_eq!(factors.d.len(), 2);
assert_eq!(factors.e.len(), 1);
}
#[test]
fn test_wide_matrix_reconstruction() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0, 4.0], &[5.0, 6.0, 7.0, 8.0]]);
let factors = gebrd(a.as_ref()).unwrap();
let q = orgbr(&factors, BidiagVect::Q).unwrap();
let p = orgbr(&factors, BidiagVect::P).unwrap();
assert_eq!(q.nrows(), 2);
assert_eq!(q.ncols(), 2);
assert_eq!(p.nrows(), 2);
assert_eq!(p.ncols(), 4);
let mut b = Mat::zeros(2, 4);
for i in 0..factors.d.len() {
b[(i, i)] = factors.d[i];
}
for i in 0..factors.e.len() {
b[(i, i + 1)] = factors.e[i];
}
}
#[test]
fn test_ormbr_consistency_with_explicit() {
let a = Mat::from_rows(&[
&[1.0f64, 2.0, 3.0],
&[4.0, 5.0, 6.0],
&[7.0, 8.0, 9.0],
&[10.0, 11.0, 12.0],
]);
let factors = gebrd(a.as_ref()).unwrap();
let _q = orgbr(&factors, BidiagVect::Q).unwrap();
let c = Mat::from_rows(&[&[1.0f64], &[2.0], &[3.0], &[4.0]]);
let result_ormbr = ormbr(
&factors,
BidiagVect::Q,
Side::Left,
Trans::NoTrans,
c.as_ref(),
)
.unwrap();
assert_eq!(result_ormbr.nrows(), 4);
assert_eq!(result_ormbr.ncols(), 1);
}
#[test]
fn test_f32_bidiag() {
let a = Mat::from_rows(&[&[1.0f32, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 9.0]]);
let factors = gebrd(a.as_ref()).unwrap();
assert_eq!(factors.d.len(), 3);
assert_eq!(factors.e.len(), 2);
}
#[test]
fn test_gebrd_blocked_tall() {
let a = Mat::from_rows(&[
&[1.0f64, 2.0, 3.0, 4.0],
&[5.0, 6.0, 7.0, 8.0],
&[9.0, 10.0, 11.0, 12.0],
&[13.0, 14.0, 15.0, 16.0],
&[17.0, 18.0, 19.0, 20.0],
]);
let factors = BidiagFactors::compute_blocked(a.as_ref()).unwrap();
assert_eq!(factors.d.len(), 4);
assert_eq!(factors.e.len(), 3);
assert_eq!(factors.tauq.len(), 4);
assert_eq!(factors.taup.len(), 3);
}
#[test]
fn test_gebrd_blocked_wide() {
let a = Mat::from_rows(&[
&[1.0f64, 2.0, 3.0, 4.0, 5.0],
&[6.0, 7.0, 8.0, 9.0, 10.0],
&[11.0, 12.0, 13.0, 14.0, 15.0],
]);
let factors = BidiagFactors::compute_blocked(a.as_ref()).unwrap();
assert_eq!(factors.d.len(), 3);
assert_eq!(factors.e.len(), 2);
}
#[test]
fn test_gebrd_blocked_reconstruction() {
let a = Mat::from_rows(&[
&[1.0f64, 2.0, 3.0, 4.0],
&[5.0, 6.0, 7.0, 8.0],
&[9.0, 10.0, 11.0, 12.0],
&[13.0, 14.0, 15.0, 16.0],
&[17.0, 18.0, 19.0, 20.0],
]);
let factors = BidiagFactors::compute_blocked_with_block_size(a.as_ref(), 2).unwrap();
let q = orgbr(&factors, BidiagVect::Q).unwrap();
let p = orgbr(&factors, BidiagVect::P).unwrap();
assert_eq!(q.nrows(), 5);
assert_eq!(q.ncols(), 4);
assert_eq!(p.nrows(), 4);
assert_eq!(p.ncols(), 4);
let n = factors.d.len();
let mut b = Mat::zeros(5, 4);
for i in 0..n {
b[(i, i)] = factors.d[i];
}
for i in 0..factors.e.len() {
b[(i, i + 1)] = factors.e[i];
}
let mut bp = Mat::zeros(5, 4);
for i in 0..5 {
for j in 0..4 {
let mut sum = 0.0;
for k in 0..4 {
sum += b[(i, k)] * p[(j, k)]; }
bp[(i, j)] = sum;
}
}
let mut reconstructed = Mat::zeros(5, 4);
for i in 0..5 {
for j in 0..4 {
let mut sum = 0.0;
for k in 0..4 {
sum += q[(i, k)] * bp[(k, j)];
}
reconstructed[(i, j)] = sum;
}
}
for i in 0..5 {
for j in 0..4 {
assert!(
approx_eq(reconstructed[(i, j)], a[(i, j)], 1e-9),
"blocked reconstructed[{},{}] = {}, a = {}",
i,
j,
reconstructed[(i, j)],
a[(i, j)]
);
}
}
}
#[test]
fn test_gebrd_blocked_vs_unblocked() {
let n = 50;
let mut a: Mat<f64> = Mat::zeros(n, n);
for i in 0..n {
for j in 0..n {
a[(i, j)] = ((i * 17 + j * 31) % 100) as f64 / 100.0 + 0.1;
}
}
let factors_unblocked = BidiagFactors::compute(a.as_ref()).unwrap();
let factors_blocked =
BidiagFactors::compute_blocked_with_block_size(a.as_ref(), 8).unwrap();
for i in 0..n {
assert!(
approx_eq(
factors_unblocked.d[i].abs(),
factors_blocked.d[i].abs(),
1e-9
),
"d[{}]: unblocked = {}, blocked = {}",
i,
factors_unblocked.d[i],
factors_blocked.d[i]
);
}
for i in 0..(n - 1) {
assert!(
approx_eq(
factors_unblocked.e[i].abs(),
factors_blocked.e[i].abs(),
1e-9
),
"e[{}]: unblocked = {}, blocked = {}",
i,
factors_unblocked.e[i],
factors_blocked.e[i]
);
}
}
#[test]
fn test_gebrd_blocked_small() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 9.0]]);
let factors = BidiagFactors::compute_blocked(a.as_ref()).unwrap();
assert_eq!(factors.d.len(), 3);
assert_eq!(factors.e.len(), 2);
}
}