Skip to main content

rusolver/
sparse_direct.rs

1// SPDX-License-Identifier: Apache-2.0
2//! Sparse Gaussian elimination with partial row pivoting and exact fill-in.
3//! Uses ordered sparse rows throughout; NEVER silently densifies or drops small fill.
4//! No symbolic reuse, supernodes, parallel factorization or fill-reducing ordering.
5use 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    /// Borrow zero-based CSR. Unsorted/duplicate column entries are summed; explicit
15    /// zeros are removed. offsets must contain n+1 entries and start at zero.
16    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    /// Exposes packed factors by sparse row: strict lower L, diagonal/upper U;
51    /// unit diagonal of L is implicit. No n*n output allocation.
52    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        // Column-oriented updates using the existing sparse rows; no transpose allocation.
65        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}