use crate::prelude::*;
use serde::de::DeserializeOwned;
use serde::{Deserialize, Serialize};
use std::default::Default;
use std::fmt::Debug;
#[derive(Clone, Serialize, Deserialize)]
pub struct ConjugateGradient<P, S> {
b: P,
r: P,
p: P,
p_prev: P,
#[serde(skip)]
rtr: S,
#[serde(skip)]
alpha: S,
#[serde(skip)]
beta: S,
}
impl<P, S> ConjugateGradient<P, S>
where
P: Clone + Default,
S: Default,
{
pub fn new(b: P) -> Result<Self, Error> {
Ok(ConjugateGradient {
b,
r: P::default(),
p: P::default(),
p_prev: P::default(),
rtr: S::default(),
alpha: S::default(),
beta: S::default(),
})
}
pub fn p(&self) -> P {
self.p.clone()
}
pub fn p_prev(&self) -> P {
self.p_prev.clone()
}
pub fn residual(&self) -> P {
self.r.clone()
}
}
impl<P, O, S, F> Solver<O> for ConjugateGradient<P, S>
where
O: ArgminOp<Param = P, Output = P, Float = F>,
P: Clone
+ Serialize
+ DeserializeOwned
+ ArgminDot<O::Param, S>
+ ArgminSub<O::Param, O::Param>
+ ArgminScaledAdd<O::Param, S, O::Param>
+ ArgminAdd<O::Param, O::Param>
+ ArgminConj
+ ArgminMul<O::Float, O::Param>,
S: Debug + ArgminDiv<S, S> + ArgminNorm<O::Float> + ArgminConj,
F: ArgminFloat,
{
const NAME: &'static str = "Conjugate Gradient";
fn init(
&mut self,
op: &mut OpWrapper<O>,
state: &IterState<O>,
) -> Result<Option<ArgminIterData<O>>, Error> {
let init_param = state.get_param();
let ap = op.apply(&init_param)?;
let r0 = self.b.sub(&ap).mul(&(F::from_f64(-1.0).unwrap()));
self.r = r0.clone();
self.p = r0.mul(&(F::from_f64(-1.0).unwrap()));
self.rtr = self.r.dot(&self.r.conj());
Ok(None)
}
fn next_iter(
&mut self,
op: &mut OpWrapper<O>,
state: &IterState<O>,
) -> Result<ArgminIterData<O>, Error> {
self.p_prev = self.p.clone();
let apk = op.apply(&self.p)?;
self.alpha = self.rtr.div(&self.p.dot(&apk.conj()));
let new_param = state.get_param().scaled_add(&self.alpha, &self.p);
self.r = self.r.scaled_add(&self.alpha, &apk);
let rtr_n = self.r.dot(&self.r.conj());
self.beta = rtr_n.div(&self.rtr);
self.rtr = rtr_n;
self.p = self
.r
.mul(&(F::from_f64(-1.0).unwrap()))
.scaled_add(&self.beta, &self.p);
let norm = self.r.dot(&self.r.conj());
Ok(ArgminIterData::new()
.param(new_param)
.cost(norm.norm())
.kv(make_kv!("alpha" => self.alpha; "beta" => self.beta;)))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_trait_impl;
test_trait_impl!(conjugate_gradient, ConjugateGradient<Vec<f64>, f64>);
}