1use 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 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 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 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 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}