use crate::{EquationHandler, EssentialBcs1d, Grid1d, NaturalBcs1d, StrError};
use russell_lab::Vector;
use russell_sparse::{CooMatrix, Genie, LinSolver, Sym};
const CUR: usize = 0; const LEF: usize = 1; const RIG: usize = 2; const INI_X: usize = 0;
pub struct Fdm1d<'a> {
grid: Grid1d,
ebcs: EssentialBcs1d<'a>,
nbcs: NaturalBcs1d<'a>,
equations: EquationHandler,
molecule: Vec<f64>,
dx: f64,
genie: Genie,
symmetric: bool,
}
impl<'a> Fdm1d<'a> {
pub fn new(grid: Grid1d, ebcs: EssentialBcs1d<'a>, nbcs: NaturalBcs1d<'a>, kx: f64) -> Result<Self, StrError> {
let dx = match grid.get_dx() {
Some(dx) => dx,
None => return Err("grid must have uniform spacing"),
};
let neq = grid.nx();
let mut equations = EquationHandler::new(neq);
equations.recompute(&ebcs.get_nodes(&grid));
let dx2 = dx * dx;
let alpha = 2.0 * kx / dx2;
let beta = -kx / dx2;
Ok(Fdm1d {
grid,
ebcs,
nbcs,
equations,
molecule: vec![alpha, beta, beta],
dx,
genie: Genie::Umfpack,
symmetric: true,
})
}
pub fn set_solver_options(&mut self, genie: Genie, symmetric: bool) {
self.genie = genie;
self.symmetric = symmetric;
}
pub fn solve_sps<F>(&self, alpha: f64, source: F) -> Result<Vector, StrError>
where
F: Fn(f64) -> f64,
{
self.ebcs.validate(&self.nbcs)?;
let sym_kk_bar = self.genie.get_sym(self.symmetric);
let (kk_bar, kk_check) = self.get_matrices_sps(alpha, 0, sym_kk_bar);
let (mut a_bar, a_check, mut f_bar) = self.get_vectors_sps(source);
let kk_check = kk_check.unwrap();
kk_check.mat_vec_mul_update(&mut f_bar, -1.0, &a_check)?;
let mut solver = LinSolver::new(self.genie)?;
solver.actual.factorize(&kk_bar, None)?;
solver.actual.solve(&mut a_bar, &f_bar, false)?;
Ok(self.get_joined_vector_sps(&a_bar, &a_check))
}
pub fn solve_lmm<F>(&self, alpha: f64, source: F) -> Result<Vector, StrError>
where
F: Fn(f64) -> f64,
{
self.ebcs.validate(&self.nbcs)?;
let sym_mm = self.genie.get_sym(self.symmetric);
let neq = self.equations.neq();
let extra_nnz = neq; let (mm, _) = self.get_matrices_lmm(alpha, extra_nnz, false, sym_mm);
let (mut aa, ff) = self.get_vectors_lmm(source);
let mut solver = LinSolver::new(self.genie)?;
solver.actual.factorize(&mm, None)?;
solver.actual.solve(&mut aa, &ff, false)?;
let neq = self.equations.neq();
Ok(Vector::from(&&aa.as_data()[..neq]))
}
pub fn get_dims_sps(&self) -> (usize, usize) {
let nu = self.equations.nu();
let np = self.equations.np();
(nu, np)
}
pub fn get_dims_lmm(&self) -> (usize, usize, usize) {
let neq = self.equations.neq();
let nlag = self.equations.np();
let ndim = neq + nlag;
(neq, nlag, ndim)
}
pub fn get_grid(&self) -> &Grid1d {
&self.grid
}
pub fn get_equations(&self) -> &EquationHandler {
&self.equations
}
pub fn get_matrices_sps(&self, alpha: f64, extra_nnz: usize, sym_kk_bar: Sym) -> (CooMatrix, Option<CooMatrix>) {
let nu = self.equations.nu();
let np = self.equations.np();
let band = if sym_kk_bar.triangular() { 2 } else { 3 };
let nnz_kk_bar = band * nu + extra_nnz;
let mut kk_bar = CooMatrix::new(nu, nu, nnz_kk_bar, sym_kk_bar).unwrap();
let mut kk_check = if np == 0 {
CooMatrix::new(1, 1, 1, Sym::No).unwrap()
} else {
let nnz_kk_check = 2 * np; CooMatrix::new(nu, np, nnz_kk_check, Sym::No).unwrap()
};
self.equations.unknown().iter().for_each(|&m| {
let iu = self.equations.iu(m);
self.loop_over_bandwidth(m, |b, n| {
let mut val = self.molecule[b];
if m == n {
val += alpha; }
if !self.ebcs.periodic_along_x && (m == 0 || m == self.grid.nx() - 1) {
val /= 2.0;
}
if self.equations.is_prescribed(n) {
let jp = self.equations.ip(n);
kk_check.put(iu, jp, val).unwrap();
} else {
let skip = (sym_kk_bar == Sym::YesLower && m < n) || (sym_kk_bar == Sym::YesUpper && m > n);
if !skip {
let ju = self.equations.iu(n);
kk_bar.put(iu, ju, val).unwrap();
}
}
});
});
if np == 0 {
(kk_bar, None)
} else {
(kk_bar, Some(kk_check))
}
}
pub fn get_matrices_lmm(
&self,
alpha: f64,
extra_nnz: usize,
get_constraints_mat: bool,
sym_mm: Sym,
) -> (CooMatrix, Option<CooMatrix>) {
let (neq, nlag, ndim) = self.get_dims_lmm();
let band = if sym_mm.triangular() { 2 } else { 3 };
let nnz = band * neq + 2 * nlag + extra_nnz; let mut mm = CooMatrix::new(ndim, ndim, nnz, sym_mm).unwrap();
for m in 0..neq {
self.loop_over_bandwidth(m, |b, n| {
if (sym_mm == Sym::YesLower && m < n) || (sym_mm == Sym::YesUpper && m > n) {
return;
}
let mut val = self.molecule[b];
if m == n {
val += alpha;
}
if !self.ebcs.periodic_along_x && (m == 0 || m == self.grid.nx() - 1) {
val /= 2.0;
}
mm.put(m, n, val).unwrap();
});
}
self.equations.prescribed().iter().for_each(|&m| {
let ip = self.equations.ip(m);
match sym_mm {
Sym::YesLower => {
mm.put(neq + ip, m, 1.0).unwrap(); }
Sym::YesUpper => {
mm.put(m, neq + ip, 1.0).unwrap(); }
Sym::YesFull | Sym::No => {
mm.put(neq + ip, m, 1.0).unwrap(); mm.put(m, neq + ip, 1.0).unwrap(); }
}
});
if get_constraints_mat && nlag > 0 {
let mut cc = CooMatrix::new(nlag, neq, nlag, Sym::No).unwrap();
self.equations.prescribed().iter().for_each(|&m| {
let ip = self.equations.ip(m);
cc.put(ip, m, 1.0).unwrap(); });
(mm, Some(cc))
} else {
(mm, None)
}
}
pub fn get_vectors_sps<F>(&self, source: F) -> (Vector, Vector, Vector)
where
F: Fn(f64) -> f64,
{
let nu = self.equations.nu();
let np = self.equations.np();
let a_bar = Vector::new(nu);
let mut a_check = Vector::new(np);
let mut f_bar = Vector::new(nu);
self.equations.unknown().iter().for_each(|&m| {
let iu = self.equations.iu(m);
let x = self.grid.coord(m);
let mut den = 1.0;
if !self.ebcs.periodic_along_x {
if self.grid.is_xmin(m) {
let wn = self.nbcs.functions[0](x);
f_bar[iu] += -wn / self.dx;
den *= 2.0;
} else if self.grid.is_xmax(m) {
let wn = self.nbcs.functions[1](x);
f_bar[iu] += -wn / self.dx;
den *= 2.0;
}
}
f_bar[iu] += source(x) / den;
});
for index in 0..2 {
if self.ebcs.sides[index] {
let m = if index == 0 { 0 } else { self.grid.nx() - 1 };
let ip = self.equations.ip(m);
let x = self.grid.coord(m);
let val = self.ebcs.functions[index](x);
a_check[ip] = val;
}
}
(a_bar, a_check, f_bar)
}
pub fn get_joined_vector_sps(&self, a_bar: &Vector, a_check: &Vector) -> Vector {
let neq = self.equations.neq();
let mut a = Vector::new(neq);
self.equations.unknown().iter().for_each(|&m| {
let iu = self.equations.iu(m);
a[m] = a_bar[iu];
});
self.equations.prescribed().iter().for_each(|&m| {
let ip = self.equations.ip(m);
a[m] = a_check[ip];
});
a
}
pub fn get_vectors_lmm<F>(&self, source: F) -> (Vector, Vector)
where
F: Fn(f64) -> f64,
{
let (neq, _, ndim) = self.get_dims_lmm();
let aa = Vector::new(ndim);
let mut ff = Vector::new(ndim);
self.grid.for_each_coord(|m, x| {
let mut den = 1.0;
if !self.ebcs.periodic_along_x {
if self.grid.is_xmin(m) {
let wn = self.nbcs.functions[0](x);
ff[m] += -wn / self.dx;
den *= 2.0;
} else if self.grid.is_xmax(m) {
let wn = self.nbcs.functions[1](x);
ff[m] += -wn / self.dx;
den *= 2.0;
}
}
ff[m] += source(x) / den;
});
for index in 0..2 {
if self.ebcs.sides[index] {
let m = if index == 0 { 0 } else { self.grid.nx() - 1 };
let ip = self.equations.ip(m);
let x = self.grid.coord(m);
let val = self.ebcs.functions[index](x);
ff[neq + ip] = val;
}
}
(aa, ff)
}
pub fn loop_over_molecule<F>(&self, m: usize, mut callback: F)
where
F: FnMut(usize, f64),
{
self.loop_over_bandwidth(m, |b, n| {
callback(n, self.molecule[b]);
});
}
pub fn for_each_coord<F>(&self, mut callback: F)
where
F: FnMut(usize, f64),
{
self.grid.for_each_coord(|m, x| {
callback(m, x);
});
}
fn loop_over_bandwidth<F>(&self, m: usize, mut callback: F)
where
F: FnMut(usize, usize),
{
let fin_x = self.grid.nx() - 1;
let mut nn = [0, 0, 0];
nn[CUR] = m;
if self.ebcs.periodic_along_x {
nn[LEF] = if m != INI_X { m - 1 } else { m + fin_x };
nn[RIG] = if m != fin_x { m + 1 } else { m - fin_x };
} else {
nn[LEF] = if m != INI_X { m - 1 } else { m + 1 };
nn[RIG] = if m != fin_x { m + 1 } else { m - 1 };
}
for b in 0..3 {
callback(b, nn[b]);
}
}
}
#[cfg(test)]
mod tests {
use super::Fdm1d;
use crate::{EssentialBcs1d, Grid1d, NaturalBcs1d, Side};
use russell_lab::Matrix;
use russell_sparse::Sym;
const LEF: f64 = 1.0;
const RIG: f64 = 2.0;
fn assert_symmetric(mat: &Matrix) {
let (nrow, ncol) = mat.dims();
assert_eq!(nrow, ncol);
for i in 0..nrow {
for j in (i + 1)..ncol {
assert_eq!(mat.get(i, j), mat.get(j, i));
}
}
}
#[test]
fn new_captures_errors() {
let grid = Grid1d::new(&[0.0, 0.1, 0.4]).unwrap();
let ebcs = EssentialBcs1d::new();
let nbcs = NaturalBcs1d::new();
let fdm = Fdm1d::new(grid, ebcs, nbcs, 1.0);
assert_eq!(fdm.err(), Some("grid must have uniform spacing"));
}
#[test]
fn get_matrices_work() {
let grid = Grid1d::new_uniform(0.0, 3.0, 4).unwrap();
let mut ebcs = EssentialBcs1d::new();
let mut nbcs = NaturalBcs1d::new();
const LEF: f64 = 1.0;
let lef = |_| LEF;
assert_eq!(lef(0.0), LEF);
ebcs.set(Side::Xmin, lef); nbcs.set(Side::Xmax, |_| 0.0);
let fdm = Fdm1d::new(grid, ebcs, nbcs, 100.0).unwrap();
assert_eq!(fdm.get_dims_sps(), (3, 1));
assert_eq!(fdm.get_dims_lmm(), (4, 1, 5));
fdm.loop_over_molecule(0, |n, val_mn| {
if n == 0 {
assert_eq!(val_mn, 200.0);
} else if n == 1 {
assert_eq!(val_mn, -100.0);
} else {
assert_eq!(val_mn, -100.0);
}
});
for sym_kk_bar in [Sym::No, Sym::YesLower, Sym::YesUpper, Sym::YesFull] {
let (kk_bar, kk_check) = fdm.get_matrices_sps(0.0, 0, sym_kk_bar);
let kk_check = kk_check.unwrap();
let kk_bar_dense = kk_bar.as_dense();
assert_symmetric(&kk_bar_dense);
assert_eq!(
format!("{}", kk_bar_dense),
"┌ ┐\n\
│ 200 -100 0 │\n\
│ -100 200 -100 │\n\
│ 0 -100 100 │\n\
└ ┘"
);
assert_eq!(
format!("{}", kk_check.as_dense()),
"┌ ┐\n\
│ -100 │\n\
│ 0 │\n\
│ 0 │\n\
└ ┘"
);
}
for sym_mm in [Sym::No, Sym::YesLower, Sym::YesUpper, Sym::YesFull] {
let (mm, cc) = fdm.get_matrices_lmm(0.0, 0, true, sym_mm);
let cc = cc.unwrap();
let mm_dense = mm.as_dense();
assert_symmetric(&mm_dense);
assert_eq!(
format!("{}", cc.as_dense()),
"┌ ┐\n\
│ 1 0 0 0 │\n\
└ ┘"
);
assert_eq!(
format!("{}", mm_dense),
"┌ ┐\n\
│ 100 -100 0 0 1 │\n\
│ -100 200 -100 0 0 │\n\
│ 0 -100 200 -100 0 │\n\
│ 0 0 -100 100 0 │\n\
│ 1 0 0 0 0 │\n\
└ ┘"
);
}
}
#[test]
fn get_matrices_homogeneous_bcs_work() {
let grid = Grid1d::new_uniform(0.0, 3.0, 4).unwrap();
let mut ebcs = EssentialBcs1d::new();
let nbcs = NaturalBcs1d::new();
ebcs.set_homogeneous();
let fdm = Fdm1d::new(grid, ebcs, nbcs, 1.0).unwrap();
let (kk, cc_mat) = fdm.get_matrices_sps(0.0, 0, Sym::No);
let (aa, ee_mat) = fdm.get_matrices_lmm(0.0, 0, true, Sym::No);
let cc = cc_mat.unwrap();
let ee = ee_mat.unwrap();
assert_eq!(fdm.get_dims_sps(), (2, 2));
assert_eq!(fdm.get_dims_lmm(), (4, 2, 6));
assert_eq!(
format!("{}", kk.as_dense()),
"┌ ┐\n\
│ 2 -1 │\n\
│ -1 2 │\n\
└ ┘"
);
assert_eq!(
format!("{}", cc.as_dense()),
"┌ ┐\n\
│ -1 0 │\n\
│ 0 -1 │\n\
└ ┘"
);
assert_eq!(
format!("{}", ee.as_dense()),
"┌ ┐\n\
│ 1 0 0 0 │\n\
│ 0 0 0 1 │\n\
└ ┘"
);
assert_eq!(
format!("{}", aa.as_dense()),
"┌ ┐\n\
│ 1 -1 0 0 1 0 │\n\
│ -1 2 -1 0 0 0 │\n\
│ 0 -1 2 -1 0 0 │\n\
│ 0 0 -1 1 0 1 │\n\
│ 1 0 0 0 0 0 │\n\
│ 0 0 0 1 0 0 │\n\
└ ┘"
);
}
#[test]
fn get_matrices_periodic_bcs_work() {
let grid = Grid1d::new_uniform(0.0, 3.0, 4).unwrap();
let mut ebcs = EssentialBcs1d::new();
let nbcs = NaturalBcs1d::new();
ebcs.set_periodic(true);
let fdm = Fdm1d::new(grid, ebcs, nbcs, 1.0).unwrap();
let (kk, cc_mat) = fdm.get_matrices_sps(0.0, 0, Sym::No);
let (aa, ee_mat) = fdm.get_matrices_lmm(0.0, 0, true, Sym::No);
assert!(cc_mat.is_none());
assert!(ee_mat.is_none());
assert_eq!(fdm.get_dims_sps(), (4, 0));
assert_eq!(fdm.get_dims_lmm(), (4, 0, 4));
assert_eq!(
format!("{}", kk.as_dense()),
"┌ ┐\n\
│ 2 -1 0 -1 │\n\
│ -1 2 -1 0 │\n\
│ 0 -1 2 -1 │\n\
│ -1 0 -1 2 │\n\
└ ┘"
);
assert_eq!(
format!("{}", aa.as_dense()),
"┌ ┐\n\
│ 2 -1 0 -1 │\n\
│ -1 2 -1 0 │\n\
│ 0 -1 2 -1 │\n\
│ -1 0 -1 2 │\n\
└ ┘"
);
}
#[test]
fn get_vectors_works() {
let grid = Grid1d::new_uniform(0.0, 1.0, 5).unwrap();
let mut ebcs = EssentialBcs1d::new();
let nbcs = NaturalBcs1d::new();
ebcs.set(Side::Xmin, |_| LEF);
ebcs.set(Side::Xmax, |_| RIG);
let fdm = Fdm1d::new(grid, ebcs, nbcs, 1.0).unwrap();
let (a_bar, a_check, f_bar) = fdm.get_vectors_sps(|_| 100.0);
assert_eq!(a_bar.dim(), 3); assert_eq!(a_check.dim(), 2); assert_eq!(f_bar.dim(), 3); assert_eq!(a_bar.as_data(), &[0.0, 0.0, 0.0]);
assert_eq!(a_check.as_data(), &[LEF, RIG]);
assert_eq!(f_bar.as_data(), &[100.0, 100.0, 100.0]);
let a = fdm.get_joined_vector_sps(&a_bar, &a_check);
assert_eq!(a.dim(), 5); assert_eq!(a.as_data(), &[LEF, 0.0, 0.0, 0.0, RIG]);
let (aa, ff) = fdm.get_vectors_lmm(|_| 100.0);
assert_eq!(aa.dim(), 5 + 2); assert_eq!(ff.dim(), 5 + 2); assert_eq!(aa.as_data(), &[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0]);
assert_eq!(ff.as_data(), &[100.0 / 2.0, 100.0, 100.0, 100.0, 100.0 / 2.0, LEF, RIG]);
}
#[test]
fn get_grid_and_get_equations_work() {
let grid = Grid1d::new_uniform(0.0, 1.0, 3).unwrap();
let mut ebcs = EssentialBcs1d::new();
let nbcs = NaturalBcs1d::new();
ebcs.set_homogeneous();
let fdm = Fdm1d::new(grid, ebcs, nbcs, 1.0).unwrap();
assert_eq!(fdm.get_grid().nx(), 3);
assert_eq!(fdm.get_equations().neq(), 3);
}
}