use super::*;
pub fn solve<M, T>(a: &M, b: &DVector<T>, max_iter: usize, tol: T) -> Option<DVector<T>>
where
M: SpMatVecMul<T>,
T: SimdRealField + PartialOrd + Copy
{
let mut solver = BiConjugateGradient {
x: DVector::<T>::zeros(a.nrows()),
r: DVector::<T>::zeros(a.nrows()),
r_hat: DVector::<T>::zeros(a.nrows()),
p: DVector::<T>::zeros(a.nrows()),
v: DVector::<T>::zeros(a.nrows()),
iter: 0,
tol,
max_iter,
converged: false,
};
solver.init(a, b, None);
if solver.solve_iterations(a, b, max_iter) {
Some(solver.x.clone())
} else {
None
}
}
pub fn solve_with_initial_guess<M, T>(a: &M, b: &DVector<T>, x: &mut DVector<T>, max_iter: usize, tol: T) -> bool
where
M: SpMatVecMul<T>,
T: SimdRealField + PartialOrd + Copy
{
let mut solver = BiConjugateGradient {
x: x.clone(),
r: DVector::<T>::zeros(a.nrows()),
r_hat: DVector::<T>::zeros(a.nrows()),
p: DVector::<T>::zeros(a.nrows()),
v: DVector::<T>::zeros(a.nrows()),
iter: 0,
tol,
max_iter,
converged: false,
};
solver.init(a, b, Some(x));
let converged = solver.solve_iterations(a, b, max_iter);
*x = solver.x.clone();
converged
}
pub struct BiConjugateGradient<T> {
pub x: DVector<T>,
pub r: DVector<T>,
pub r_hat: DVector<T>,
pub p: DVector<T>,
pub v: DVector<T>,
pub iter: usize,
pub tol: T,
pub max_iter: usize,
pub converged: bool,
}
impl<M, T> IterativeSolver<M, DVector<T>, T> for BiConjugateGradient<T>
where
M: SpMatVecMul<T>,
T: SimdRealField + PartialOrd + Copy,
{
fn init(&mut self, a: &M, b: &DVector<T>, x0: Option<&DVector<T>>) {
let n = a.nrows();
self.x = match x0 {
Some(x0) => x0.clone(),
None => DVector::<T>::zeros(n),
};
self.r = b - &a.mul_vec(&self.x);
self.r_hat = self.r.clone();
self.p = self.r.clone();
self.v = DVector::<T>::zeros(n);
self.iter = 0;
self.converged = false;
}
fn step(&mut self, a: &M, _b: &DVector<T>) -> bool {
if self.converged { return true; }
let r_dot = self.r.dot(&self.r_hat);
if self.r.magnitude() <= self.tol {
self.converged = true;
return true;
}
self.v = a.mul_vec(&self.p);
let alpha = r_dot / self.r_hat.dot(&self.v);
self.x.axpy(alpha, &self.p, T::one());
let s = &self.r - &self.v * alpha;
if s.max() <= self.tol {
self.r = s;
self.converged = true;
return true;
}
let t = a.mul_vec(&s);
let omega = t.dot(&s) / t.dot(&t);
self.x.axpy(omega, &s, T::one());
let new_r = &s - &t * omega;
if new_r.max() <= self.tol {
self.r = new_r;
self.converged = true;
return true;
}
let new_r_dot = self.r_hat.dot(&new_r);
let beta = (new_r_dot / r_dot) * (alpha / omega);
self.p = &new_r + (&self.p - &self.v * omega) * beta;
self.r = new_r;
self.iter += 1;
false
}
fn reset(&mut self) {
self.x.fill(T::zero());
self.r.fill(T::zero());
self.r_hat.fill(T::zero());
self.p.fill(T::zero());
self.v.fill(T::zero());
self.iter = 0;
self.converged = false;
}
fn hard_reset(&mut self) {
self.x = DVector::<T>::zeros(0);
self.r = DVector::<T>::zeros(0);
self.r_hat = DVector::<T>::zeros(0);
self.p = DVector::<T>::zeros(0);
self.v = DVector::<T>::zeros(0);
self.iter = 0;
self.converged = false;
}
fn soft_reset(&mut self) {
self.r.fill(T::zero());
self.r_hat.fill(T::zero());
self.p.fill(T::zero());
self.v.fill(T::zero());
self.iter = 0;
self.converged = false;
}
fn solution(&self) -> &DVector<T> {
&self.x
}
fn iterations(&self) -> usize {
self.iter
}
}