rusolver/
sparse_direct.rs1use crate::{Matrix,MatrixView,SolverError,Tolerance};
6use crate::numerics::{checked,finite};
7use std::collections::BTreeMap;
8#[derive(Clone,Copy,Debug)]
9pub struct SparseLuOptions {pub pivot_tolerance:Tolerance,pub max_factor_nonzeros:usize}
10impl Default for SparseLuOptions{fn default()->Self{Self{pivot_tolerance:Tolerance::default(),max_factor_nonzeros:10_000_000}}}
11#[derive(Clone,Debug)]
12pub struct SparseLu {rows:Vec<BTreeMap<usize,f64>>,pivots:Vec<usize>,nonzeros:usize}
13impl SparseLu{
14 pub fn factor_csr(n:usize,offsets:&[usize],columns:&[usize],values:&[f64],options:SparseLuOptions)->Result<Self,SolverError>{
17 validate_csr(n,n,offsets,columns,values)?;options.pivot_tolerance.validate()?;
18 let mut rows=Vec::new();rows.try_reserve_exact(n).map_err(|_|SolverError::Allocation)?;
19 let mut nonzeros=0;let mut scale=0.0f64;
20 for i in 0..n{
21 let mut row=BTreeMap::new();
22 for p in offsets[i]..offsets[i+1]{let j=columns[p];let v=checked(row.get(&j).copied().unwrap_or(0.0)+values[p],"CSR duplicate sum")?;
23 if v==0.0{row.remove(&j);}else{row.insert(j,v);}}
24 nonzeros+=row.len();check_limit(nonzeros,options.max_factor_nonzeros)?;
25 for &v in row.values(){scale=scale.max(v.abs());}rows.push(row);
26 }
27 let threshold=options.pivot_tolerance.threshold(scale)?;let mut pivots=Vec::new();
28 for k in 0..n{
29 let mut p=k;let mut best=0.0f64;
30 for i in k..n{let v=rows[i].get(&k).copied().unwrap_or(0.0).abs();if v>best{best=v;p=i;}}
31 if best<=threshold{return Err(SolverError::Singular{index:k,pivot:best,threshold});}
32 rows.swap(k,p);pivots.push(p);
33 let pivot=rows[k][&k];
34 let upper:Vec<(usize,f64)>=rows[k].range(k+1..).map(|(&j,&x)|(j,x)).collect();
35 for i in k+1..n{
36 let Some(value)=rows[i].get(&k).copied()else{continue};
37 let multiplier=checked(value/pivot,"sparse LU multiplier")?;
38 if multiplier==0.0{rows[i].remove(&k);nonzeros-=1;}else{rows[i].insert(k,multiplier);}
39 for &(j,u)in &upper{
40 let before=rows[i].get(&j).copied();let value=checked((-multiplier).mul_add(u,before.unwrap_or(0.0)),"sparse LU fill")?;
41 if value==0.0{if rows[i].remove(&j).is_some(){nonzeros-=1;}}
42 else{if before.is_none(){check_limit(nonzeros+1,options.max_factor_nonzeros)?;nonzeros+=1;}rows[i].insert(j,value);}
43 }
44 }
45 }Ok(Self{rows,pivots,nonzeros})
46 }
47 pub fn order(&self)->usize{self.rows.len()}
48 pub fn factor_nonzeros(&self)->usize{self.nonzeros}
49 pub fn pivots(&self)->&[usize]{&self.pivots}
50 pub fn packed_row(&self,i:usize)->Option<&BTreeMap<usize,f64>>{self.rows.get(i)}
53 pub fn solve(&self,b:MatrixView<'_>)->Result<Matrix,SolverError>{
54 let(n,r)=(self.order(),b.columns());if b.rows()!=n{return Err(SolverError::Shape("sparse LU RHS"));}
55 let mut x=b.to_owned()?;
56 for(k,&p)in self.pivots.iter().enumerate(){for c in 0..r{x.data.swap(k*r+c,p*r+c);}}
57 for i in 0..n{for(&j,&a)in self.rows[i].range(..i){for c in 0..r{x.data[i*r+c]=checked((-a).mul_add(x.data[j*r+c],x.data[i*r+c]),"sparse forward solve")?;}}}
58 for i in(0..n).rev(){for(&j,&a)in self.rows[i].range(i+1..){for c in 0..r{x.data[i*r+c]=checked((-a).mul_add(x.data[j*r+c],x.data[i*r+c]),"sparse back solve")?;}}
59 let d=self.rows[i][&i];for c in 0..r{x.data[i*r+c]=checked(x.data[i*r+c]/d,"sparse diagonal solve")?;}}
60 Ok(x)
61 }
62 pub fn solve_transpose(&self,b:MatrixView<'_>)->Result<Matrix,SolverError>{
63 let(n,r)=(self.order(),b.columns());if b.rows()!=n{return Err(SolverError::Shape("sparse transpose RHS"));}let mut x=b.to_owned()?;
64 for i in 0..n{for c in 0..r{x.data[i*r+c]=checked(x.data[i*r+c]/self.rows[i][&i],"sparse U transpose")?;}
66 for(&j,&a)in self.rows[i].range(i+1..){for c in 0..r{x.data[j*r+c]=checked((-a).mul_add(x.data[i*r+c],x.data[j*r+c]),"sparse U transpose update")?;}}}
67 for i in(0..n).rev(){for(&j,&a)in self.rows[i].range(..i){for c in 0..r{x.data[j*r+c]=checked((-a).mul_add(x.data[i*r+c],x.data[j*r+c]),"sparse L transpose update")?;}}}
68 for k in(0..n).rev(){for c in 0..r{x.data.swap(k*r+c,self.pivots[k]*r+c);}}Ok(x)
69 }
70 #[cfg(feature="sparse")]
71 pub fn from_rusparse(a:&rusparse::CsrMatrix<'_>,options:SparseLuOptions)->Result<Self,SolverError>{
72 let base=if a.index_base()==rusparse::IndexBase::One{1}else{0};
73 let offsets:Vec<usize>=a.row_offsets().iter().map(|&x|x as usize-base).collect();
74 let columns:Vec<usize>=a.column_indices().iter().map(|&x|x as usize-base).collect();
75 let values:Vec<f64>=a.values().iter().map(|&x|f64::from(x)).collect();
76 if a.rows()!=a.columns(){return Err(SolverError::Shape("square ruSPARSE input required"));}
77 Self::factor_csr(a.rows(),&offsets,&columns,&values,options)
78 }
79}
80fn check_limit(required:usize,limit:usize)->Result<(),SolverError>{if required>limit{Err(SolverError::WorkspaceLimit{required,limit})}else{Ok(())}}
81pub(crate)fn validate_csr(rows:usize,cols:usize,offsets:&[usize],columns:&[usize],values:&[f64])->Result<(),SolverError>{
82 if cols==0||rows.checked_add(1)!=Some(offsets.len())||offsets.first()!=Some(&0)
83 ||offsets.last()!=Some(&values.len())||columns.len()!=values.len(){return Err(SolverError::Shape("CSR dimensions/offsets"));}
84 if offsets.windows(2).any(|w|w[0]>w[1])||columns.iter().any(|&j|j>=cols){return Err(SolverError::Shape("CSR invalid index"));}
85 finite(values)
86}