Skip to main content

rusolver/
complex.rs

1// SPDX-License-Identifier: Apache-2.0
2//! Host complex FP64 LU, Hermitian Cholesky and thin Householder QR.
3//! Each component is f64 (128 storage bits per complex value); not CUDA Complex64.
4use crate::{SolverError,Tolerance};
5use std::ops::{Add,Sub,Mul,Div,Neg};
6#[derive(Clone,Copy,Debug,Default,PartialEq)]
7pub struct Complex64 {pub re:f64,pub im:f64}
8impl Complex64 {
9    pub const ZERO:Self=Self{re:0.0,im:0.0}; pub const ONE:Self=Self{re:1.0,im:0.0};
10    pub const fn new(re:f64,im:f64)->Self{Self{re,im}}
11    pub fn conj(self)->Self{Self::new(self.re,-self.im)}
12    pub fn abs(self)->f64{self.re.hypot(self.im)}
13    pub fn is_finite(self)->bool{self.re.is_finite()&&self.im.is_finite()}
14    pub fn scale(self,s:f64)->Self{Self::new(self.re*s,self.im*s)}
15}
16impl Add for Complex64 {type Output=Self;fn add(self,b:Self)->Self{Self::new(self.re+b.re,self.im+b.im)}}
17impl Sub for Complex64 {type Output=Self;fn sub(self,b:Self)->Self{Self::new(self.re-b.re,self.im-b.im)}}
18impl Neg for Complex64 {type Output=Self;fn neg(self)->Self{Self::new(-self.re,-self.im)}}
19impl Mul for Complex64 {type Output=Self;fn mul(self,b:Self)->Self{Self::new(self.re.mul_add(b.re,-self.im*b.im),self.re.mul_add(b.im,self.im*b.re))}}
20impl Div for Complex64 {
21    type Output=Self;
22    fn div(self,b:Self)->Self {
23        // Scaled denominator avoids squaring |b|. Algorithms check output finiteness.
24        let s=b.re.abs().max(b.im.abs()); let br=b.re/s;let bi=b.im/s;
25        let d=br*br+bi*bi;let ar=self.re/s;let ai=self.im/s;
26        Self::new((ar*br+ai*bi)/d,(ai*br-ar*bi)/d)
27    }
28}
29fn check(x:Complex64)->Result<Complex64,SolverError>{if x.is_finite(){Ok(x)}else{Err(SolverError::Arithmetic("complex intermediate"))}}
30#[derive(Clone,Debug,PartialEq)]
31pub struct ComplexMatrix {pub(crate) rows:usize,pub(crate) cols:usize,pub(crate) data:Vec<Complex64>}
32impl ComplexMatrix {
33    pub fn new(rows:usize,cols:usize,data:Vec<Complex64>)->Result<Self,SolverError>{
34        if rows==0||cols==0{return Err(SolverError::Shape("empty complex matrix"));}
35        let n=rows.checked_mul(cols).ok_or(SolverError::SizeOverflow)?;
36        if n!=data.len(){return Err(SolverError::Shape("complex matrix length"));}
37        for(i,x)in data.iter().enumerate(){if !x.is_finite(){return Err(SolverError::NonFinite{index:i});}}
38        Ok(Self{rows,cols,data})
39    }
40    pub fn zeros(rows:usize,cols:usize)->Result<Self,SolverError>{
41        let n=rows.checked_mul(cols).filter(|&n|n<=isize::MAX as usize/16).ok_or(SolverError::SizeOverflow)?;
42        let mut v=Vec::new();v.try_reserve_exact(n).map_err(|_|SolverError::Allocation)?;v.resize(n,Complex64::ZERO);
43        Self::new(rows,cols,v)
44    }
45    pub fn identity(n:usize)->Result<Self,SolverError>{let mut a=Self::zeros(n,n)?;for i in 0..n{a.data[i*n+i]=Complex64::ONE;}Ok(a)}
46    pub fn rows(&self)->usize{self.rows} pub fn columns(&self)->usize{self.cols}
47    pub fn values(&self)->&[Complex64]{&self.data}
48    pub fn adjoint(&self)->Result<Self,SolverError>{let mut b=Self::zeros(self.cols,self.rows)?;for i in 0..self.rows{for j in 0..self.cols{b.data[j*self.rows+i]=self.data[i*self.cols+j].conj();}}Ok(b)}
49    fn square(&self)->Result<usize,SolverError>{if self.rows!=self.cols{Err(SolverError::Shape("complex square matrix required"))}else{Ok(self.rows)}}
50    fn max_abs(&self)->f64{self.data.iter().fold(0.0f64,|s,x|s.max(x.abs()))}
51}
52#[derive(Clone,Debug)]
53pub struct ComplexLu {packed:ComplexMatrix,pivots:Vec<usize>}
54impl ComplexLu {
55    /// P A = L U; partial row pivoting, no input mutation.
56    pub fn factor(a:&ComplexMatrix,tolerance:Tolerance)->Result<Self,SolverError>{
57        let n=a.square()?;let cutoff=tolerance.threshold(a.max_abs())?;let mut packed=a.clone();let mut pivots=Vec::new();
58        for k in 0..n{
59            let mut p=k;for i in k+1..n{if packed.data[i*n+k].abs()>packed.data[p*n+k].abs(){p=i;}}
60            let pivot=packed.data[p*n+k].abs();if pivot<=cutoff{return Err(SolverError::Singular{index:k,pivot,threshold:cutoff});}
61            pivots.push(p);if p!=k{for j in 0..n{packed.data.swap(k*n+j,p*n+j);}}
62            for i in k+1..n{let r=check(packed.data[i*n+k]/packed.data[k*n+k])?;packed.data[i*n+k]=r;
63                for j in k+1..n{packed.data[i*n+j]=check(packed.data[i*n+j]-r*packed.data[k*n+j])?;}}
64        }Ok(Self{packed,pivots})
65    }
66    pub fn packed(&self)->&ComplexMatrix{&self.packed} pub fn pivots(&self)->&[usize]{&self.pivots}
67    pub fn solve(&self,b:&ComplexMatrix)->Result<ComplexMatrix,SolverError>{self.solve_impl(b,false)}
68    pub fn solve_adjoint(&self,b:&ComplexMatrix)->Result<ComplexMatrix,SolverError>{self.solve_impl(b,true)}
69    fn solve_impl(&self,b:&ComplexMatrix,adjoint:bool)->Result<ComplexMatrix,SolverError>{
70        let n=self.packed.rows;let r=b.cols;if b.rows!=n{return Err(SolverError::Shape("complex LU RHS"));}let mut x=b.clone();
71        if !adjoint{
72            for(k,&p)in self.pivots.iter().enumerate(){for c in 0..r{x.data.swap(k*r+c,p*r+c);}}
73            for i in 0..n{for c in 0..r{let mut s=x.data[i*r+c];for j in 0..i{s=check(s-self.packed.data[i*n+j]*x.data[j*r+c])?;}x.data[i*r+c]=s;}}
74            for i in(0..n).rev(){for c in 0..r{let mut s=x.data[i*r+c];for j in i+1..n{s=check(s-self.packed.data[i*n+j]*x.data[j*r+c])?;}x.data[i*r+c]=check(s/self.packed.data[i*n+i])?;}}
75        }else{
76            for i in 0..n{for c in 0..r{let mut s=x.data[i*r+c];for j in 0..i{s=check(s-self.packed.data[j*n+i].conj()*x.data[j*r+c])?;}x.data[i*r+c]=check(s/self.packed.data[i*n+i].conj())?;}}
77            for i in(0..n).rev(){for c in 0..r{let mut s=x.data[i*r+c];for j in i+1..n{s=check(s-self.packed.data[j*n+i].conj()*x.data[j*r+c])?;}x.data[i*r+c]=s;}}
78            for k in(0..n).rev(){for c in 0..r{x.data.swap(k*r+c,self.pivots[k]*r+c);}}
79        }Ok(x)
80    }
81}
82#[derive(Clone,Debug)]pub struct ComplexCholesky{lower:ComplexMatrix}
83impl ComplexCholesky{
84    /// A=L L^H. Lower triangle defines the matrix after Hermitian validation.
85    pub fn factor(a:&ComplexMatrix,tolerance:Tolerance)->Result<Self,SolverError>{
86        let n=a.square()?;tolerance.validate()?;let cutoff=tolerance.threshold(a.max_abs())?;
87        for i in 0..n{if a.data[i*n+i].im.abs()>cutoff{return Err(SolverError::NotSymmetric{row:i,column:i});}
88            for j in 0..i{let x=a.data[i*n+j];let y=a.data[j*n+i].conj();let s=x.abs().max(y.abs());
89                if s>0.0&&(x.scale(1.0/s)-y.scale(1.0/s)).abs()>(tolerance.absolute/s).max(tolerance.relative){return Err(SolverError::NotSymmetric{row:i,column:j});}}}
90        let mut l=ComplexMatrix::zeros(n,n)?;
91        for j in 0..n{
92            let mut d=a.data[j*n+j].re;for k in 0..j{let z=l.data[j*n+k];d-=z.re*z.re+z.im*z.im;}
93            if !d.is_finite(){return Err(SolverError::Arithmetic("complex Cholesky pivot"));}
94            if d<=0.0{return Err(SolverError::NotPositiveDefinite{index:j,pivot:d});}
95            l.data[j*n+j]=Complex64::new(d.sqrt(),0.0);
96            for i in j+1..n{let mut s=a.data[i*n+j];for k in 0..j{s=check(s-l.data[i*n+k]*l.data[j*n+k].conj())?;}
97                l.data[i*n+j]=check(s/l.data[j*n+j])?;}
98        }Ok(Self{lower:l})
99    }
100    pub fn lower(&self)->&ComplexMatrix{&self.lower}
101    pub fn solve(&self,b:&ComplexMatrix)->Result<ComplexMatrix,SolverError>{
102        let(n,r)=(self.lower.rows,b.cols);if b.rows!=n{return Err(SolverError::Shape("complex Cholesky RHS"));}let mut x=b.clone();
103        for i in 0..n{for c in 0..r{let mut s=x.data[i*r+c];for j in 0..i{s=check(s-self.lower.data[i*n+j]*x.data[j*r+c])?;}x.data[i*r+c]=check(s/self.lower.data[i*n+i])?;}}
104        for i in(0..n).rev(){for c in 0..r{let mut s=x.data[i*r+c];for j in i+1..n{s=check(s-self.lower.data[j*n+i].conj()*x.data[j*r+c])?;}x.data[i*r+c]=check(s/self.lower.data[i*n+i])?;}}Ok(x)
105    }
106}
107#[derive(Clone,Debug)]pub struct ComplexQr{q:ComplexMatrix,r:ComplexMatrix,rank:usize}
108impl ComplexQr{
109    /// Unpivoted thin Householder QR, m>=n. Two-dimensional host data only.
110    /// Rank is a diagonal diagnostic, not a rank-revealing decomposition.
111    pub fn factor(a:&ComplexMatrix,tolerance:Tolerance)->Result<Self,SolverError>{
112        let(m,n)=(a.rows,a.cols);if m<n{return Err(SolverError::Shape("complex QR requires m>=n"));}
113        let cutoff=tolerance.threshold(a.max_abs())?;let mut r=a.clone();
114        let mut reflectors:Vec<Vec<Complex64>>=Vec::new();
115        for k in 0..n{
116            let norm=crate::numerics::norm_iter((k..m).map(|i|r.data[i*n+k].abs()))?;
117            let mut v=vec![Complex64::ZERO;m-k];
118            if norm>0.0{
119                let x0=r.data[k*n+k];let phase=if x0.abs()==0.0{Complex64::ONE}else{x0.scale(1.0/x0.abs())};
120                for i in k..m{v[i-k]=r.data[i*n+k].scale(1.0/norm);}v[0]=v[0]+phase;
121                let vn=crate::numerics::norm_iter(v.iter().map(|z|z.abs()))?;for z in &mut v{*z=z.scale(1.0/vn);}
122                for j in k..n{
123                    let mut d=Complex64::ZERO;for i in k..m{d=check(d+v[i-k].conj()*r.data[i*n+j])?;}
124                    d=d.scale(2.0);for i in k..m{r.data[i*n+j]=check(r.data[i*n+j]-v[i-k]*d)?;}
125                }
126                for i in k+1..m{r.data[i*n+k]=Complex64::ZERO;}
127            }reflectors.push(v);
128        }
129        let mut q=ComplexMatrix::zeros(m,n)?;for j in 0..n{q.data[j*n+j]=Complex64::ONE;}
130        for k in(0..n).rev(){let v=&reflectors[k];for j in 0..n{
131            let mut d=Complex64::ZERO;for i in k..m{d=check(d+v[i-k].conj()*q.data[i*n+j])?;}
132            d=d.scale(2.0);for i in k..m{q.data[i*n+j]=check(q.data[i*n+j]-v[i-k]*d)?;}
133        }}
134        let mut thin=ComplexMatrix::zeros(n,n)?;for i in 0..n{for j in i..n{thin.data[i*n+j]=r.data[i*n+j];}}
135        let rank=(0..n).filter(|&j|thin.data[j*n+j].abs()>cutoff).count();Ok(Self{q,r:thin,rank})
136    }
137    pub fn q(&self)->&ComplexMatrix{&self.q} pub fn r(&self)->&ComplexMatrix{&self.r}pub fn rank(&self)->usize{self.rank}
138    pub fn least_squares(&self,b:&ComplexMatrix)->Result<ComplexMatrix,SolverError>{
139        let(m,n,p)=(self.q.rows,self.q.cols,b.cols);if b.rows!=m{return Err(SolverError::Shape("complex QR RHS"));}
140        if self.rank<n{return Err(SolverError::RankDeficient{rank:self.rank,columns:n});}
141        let mut x=ComplexMatrix::zeros(n,p)?;for j in 0..n{for c in 0..p{
142            let mut s=Complex64::ZERO;for i in 0..m{s=check(s+self.q.data[i*n+j].conj()*b.data[i*p+c])?;}x.data[j*p+c]=s;
143        }}for i in(0..n).rev(){for c in 0..p{let mut s=x.data[i*p+c];for j in i+1..n{s=check(s-self.r.data[i*n+j]*x.data[j*p+c])?;}x.data[i*p+c]=check(s/self.r.data[i*n+i])?;}}Ok(x)
144    }
145}