use crate::{EquationHandler, EssentialBcs2d, Grid2d, NaturalBcs2d, Side, 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 BOT: usize = 3; const TOP: usize = 4; const INI_X: usize = 0;
const INI_Y: usize = 0;
pub struct Fdm2d<'a> {
grid: Grid2d,
ebcs: EssentialBcs2d<'a>,
nbcs: NaturalBcs2d<'a>,
equations: EquationHandler,
molecule: Vec<f64>,
dx: f64,
dy: f64,
genie: Genie,
symmetric: bool,
}
impl<'a> Fdm2d<'a> {
pub fn new(
grid: Grid2d,
ebcs: EssentialBcs2d<'a>,
nbcs: NaturalBcs2d<'a>,
kx: f64,
ky: f64,
) -> Result<Self, StrError> {
let (dx, dy) = match grid.get_dx_dy() {
Some((dx, dy)) => (dx, dy),
None => return Err("grid must have uniform spacing"),
};
let neq = grid.size();
let mut equations = EquationHandler::new(neq);
equations.recompute(&ebcs.get_nodes(&grid));
let dx2 = dx * dx;
let dy2 = dy * dy;
let alpha = 2.0 * (kx / dx2 + ky / dy2);
let beta = -kx / dx2;
let gamma = -ky / dy2;
Ok(Fdm2d {
grid,
ebcs,
nbcs,
equations,
molecule: vec![alpha, beta, beta, gamma, gamma],
dx,
dy,
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) -> 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).unwrap();
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) -> f64,
{
self.ebcs.validate(&self.nbcs)?;
let sym_mm = self.genie.get_sym(self.symmetric);
let (mm, _) = self.get_matrices_lmm(alpha, 0, 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) -> &Grid2d {
&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 nx = self.grid.nx();
let ny = self.grid.ny();
let nu = self.equations.nu();
let np = self.equations.np();
let band = if sym_kk_bar.triangular() { 3 } else { 5 };
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 = 4 * 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;
}
let (i, j) = self.grid.get_ij(m);
if !self.ebcs.periodic_along_x && (i == 0 || i == nx - 1) {
val /= 2.0;
}
if !self.ebcs.periodic_along_y && (j == 0 || j == ny - 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 nx = self.grid.nx();
let ny = self.grid.ny();
let (neq, nlag, ndim) = self.get_dims_lmm();
let band = if sym_mm.triangular() { 3 } else { 5 };
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;
}
let (i, j) = self.grid.get_ij(m);
if !self.ebcs.periodic_along_x && (i == 0 || i == nx - 1) {
val /= 2.0;
}
if !self.ebcs.periodic_along_y && (j == 0 || j == ny - 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) -> 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, y) = self.grid.coord(m);
let mut den = 1.0;
let cf = if self.grid.is_corner(m) { 0.5 } else { 1.0 };
if !self.ebcs.periodic_along_x {
if self.grid.is_xmin(m) {
let wn = self.nbcs.functions[0](x, y);
f_bar[iu] += -cf * wn / self.dx;
den *= 2.0;
} else if self.grid.is_xmax(m) {
let wn = self.nbcs.functions[1](x, y);
f_bar[iu] += -cf * wn / self.dx;
den *= 2.0;
}
}
if !self.ebcs.periodic_along_y {
if self.grid.is_ymin(m) {
let wn = self.nbcs.functions[2](x, y);
f_bar[iu] += -cf * wn / self.dy;
den *= 2.0;
} else if self.grid.is_ymax(m) {
let wn = self.nbcs.functions[3](x, y);
f_bar[iu] += -cf * wn / self.dy;
den *= 2.0;
}
}
f_bar[iu] += source(x, y) / den;
});
for index in 0..4 {
if self.ebcs.sides[index] {
for &m in self.grid.get_nodes_on_side(Side::from_index(index)) {
let ip = self.equations.ip(m);
let (x, y) = self.grid.coord(m);
let val = self.ebcs.functions[index](x, y);
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) -> 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, y| {
let mut den = 1.0;
let cf = if self.grid.is_corner(m) { 0.5 } else { 1.0 };
if !self.ebcs.periodic_along_x {
if self.grid.is_xmin(m) {
let wn = self.nbcs.functions[0](x, y);
ff[m] += -cf * wn / self.dx;
den *= 2.0;
}
if self.grid.is_xmax(m) {
let wn = self.nbcs.functions[1](x, y);
ff[m] += -cf * wn / self.dx;
den *= 2.0;
}
}
if !self.ebcs.periodic_along_y {
if self.grid.is_ymin(m) {
let wn = self.nbcs.functions[2](x, y);
ff[m] += -cf * wn / self.dy;
den *= 2.0;
}
if self.grid.is_ymax(m) {
let wn = self.nbcs.functions[3](x, y);
ff[m] += -cf * wn / self.dy;
den *= 2.0;
}
}
ff[m] += source(x, y) / den;
});
for index in 0..4 {
if self.ebcs.sides[index] {
for &m in self.grid.get_nodes_on_side(Side::from_index(index)) {
let ip = self.equations.ip(m);
let (x, y) = self.grid.coord(m);
let val = self.ebcs.functions[index](x, y);
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, f64),
{
self.grid.for_each_coord(|m, x, y| {
callback(m, x, y);
});
}
fn loop_over_bandwidth<F>(&self, m: usize, mut callback: F)
where
F: FnMut(usize, usize),
{
let nx = self.grid.nx();
let ny = self.grid.ny();
let fin_x = nx - 1;
let fin_y = ny - 1;
let i = m % nx;
let j = m / nx;
let mut nn = [0, 0, 0, 0, 0];
nn[CUR] = m;
if self.ebcs.periodic_along_x {
nn[LEF] = if i != INI_X { m - 1 } else { m + fin_x };
nn[RIG] = if i != fin_x { m + 1 } else { m - fin_x };
} else {
nn[LEF] = if i != INI_X { m - 1 } else { m + 1 };
nn[RIG] = if i != fin_x { m + 1 } else { m - 1 };
}
if self.ebcs.periodic_along_y {
nn[BOT] = if j != INI_Y { m - nx } else { m + fin_y * nx };
nn[TOP] = if j != fin_y { m + nx } else { m - fin_y * nx };
} else {
nn[BOT] = if j != INI_Y { m - nx } else { m + nx };
nn[TOP] = if j != fin_y { m + nx } else { m - nx };
}
for b in 0..5 {
callback(b, nn[b]);
}
}
}
#[cfg(test)]
mod tests {
use super::Fdm2d;
use crate::{EssentialBcs2d, Grid2d, NaturalBcs2d, Side};
use russell_lab::Matrix;
use russell_sparse::Sym;
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 = Grid2d::new(&[0.0, 0.1, 0.4], &[0.0, 0.2, 0.5]).unwrap();
let ebcs = EssentialBcs2d::new();
let nbcs = NaturalBcs2d::new();
let fdm = Fdm2d::new(grid, ebcs, nbcs, 1.0, 1.0);
assert_eq!(fdm.err(), Some("grid must have uniform spacing"));
}
#[test]
fn get_matrices_work() {
let grid = Grid2d::new_uniform(0.0, 3.0, 0.0, 2.0, 4, 3).unwrap();
let mut ebcs = EssentialBcs2d::new();
let mut nbcs = NaturalBcs2d::new();
const LEF: f64 = 1.0;
let lef = |_, _| LEF;
assert_eq!(lef(0.0, 0.0), LEF);
ebcs.set(Side::Xmin, lef); nbcs.set(Side::Xmax, |_, _| 0.0); nbcs.set(Side::Ymin, |_, _| 0.0); nbcs.set(Side::Ymax, |_, _| 0.0);
let fdm = Fdm2d::new(grid, ebcs, nbcs, 100.0, 300.0).unwrap();
assert_eq!(fdm.get_dims_sps(), (9, 3));
assert_eq!(fdm.get_dims_lmm(), (12, 3, 15));
fdm.loop_over_molecule(0, |n, val_mn| {
if n == 0 {
assert_eq!(val_mn, 800.0);
} else if n == 1 {
assert_eq!(val_mn, -100.0);
} else if n == 4 {
assert_eq!(val_mn, -300.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\
│ 400 -50 0 -300 0 0 0 0 0 │\n\
│ -50 400 -50 0 -300 0 0 0 0 │\n\
│ 0 -50 200 0 0 -150 0 0 0 │\n\
│ -300 0 0 800 -100 0 -300 0 0 │\n\
│ 0 -300 0 -100 800 -100 0 -300 0 │\n\
│ 0 0 -150 0 -100 400 0 0 -150 │\n\
│ 0 0 0 -300 0 0 400 -50 0 │\n\
│ 0 0 0 0 -300 0 -50 400 -50 │\n\
│ 0 0 0 0 0 -150 0 -50 200 │\n\
└ ┘"
);
assert_eq!(
format!("{}", kk_check.as_dense()),
"┌ ┐\n\
│ -50 0 0 │\n\
│ 0 0 0 │\n\
│ 0 0 0 │\n\
│ 0 -100 0 │\n\
│ 0 0 0 │\n\
│ 0 0 0 │\n\
│ 0 0 -50 │\n\
│ 0 0 0 │\n\
│ 0 0 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 0 0 0 0 0 0 0 0 │\n\
│ 0 0 0 0 1 0 0 0 0 0 0 0 │\n\
│ 0 0 0 0 0 0 0 0 1 0 0 0 │\n\
└ ┘"
);
assert_eq!(
format!("{}", mm_dense),
"┌ ┐\n\
│ 200 -50 0 0 -150 0 0 0 0 0 0 0 1 0 0 │\n\
│ -50 400 -50 0 0 -300 0 0 0 0 0 0 0 0 0 │\n\
│ 0 -50 400 -50 0 0 -300 0 0 0 0 0 0 0 0 │\n\
│ 0 0 -50 200 0 0 0 -150 0 0 0 0 0 0 0 │\n\
│ -150 0 0 0 400 -100 0 0 -150 0 0 0 0 1 0 │\n\
│ 0 -300 0 0 -100 800 -100 0 0 -300 0 0 0 0 0 │\n\
│ 0 0 -300 0 0 -100 800 -100 0 0 -300 0 0 0 0 │\n\
│ 0 0 0 -150 0 0 -100 400 0 0 0 -150 0 0 0 │\n\
│ 0 0 0 0 -150 0 0 0 200 -50 0 0 0 0 1 │\n\
│ 0 0 0 0 0 -300 0 0 -50 400 -50 0 0 0 0 │\n\
│ 0 0 0 0 0 0 -300 0 0 -50 400 -50 0 0 0 │\n\
│ 0 0 0 0 0 0 0 -150 0 0 -50 200 0 0 0 │\n\
│ 1 0 0 0 0 0 0 0 0 0 0 0 0 0 0 │\n\
│ 0 0 0 0 1 0 0 0 0 0 0 0 0 0 0 │\n\
│ 0 0 0 0 0 0 0 0 1 0 0 0 0 0 0 │\n\
└ ┘"
);
}
}
#[test]
fn get_matrices_periodic_bcs_work() {
let grid = Grid2d::new_uniform(0.0, 2.0, 0.0, 3.0, 3, 4).unwrap();
let mut ebcs = EssentialBcs2d::new();
let nbcs = NaturalBcs2d::new();
ebcs.set_periodic(true, true);
let fdm = Fdm2d::new(grid, ebcs, nbcs, 1.0, 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(), (12, 0));
assert_eq!(fdm.get_dims_lmm(), (12, 0, 12));
assert_eq!(
format!("{}", kk.as_dense()),
"┌ ┐\n\
│ 4 -1 -1 -1 0 0 0 0 0 -1 0 0 │\n\
│ -1 4 -1 0 -1 0 0 0 0 0 -1 0 │\n\
│ -1 -1 4 0 0 -1 0 0 0 0 0 -1 │\n\
│ -1 0 0 4 -1 -1 -1 0 0 0 0 0 │\n\
│ 0 -1 0 -1 4 -1 0 -1 0 0 0 0 │\n\
│ 0 0 -1 -1 -1 4 0 0 -1 0 0 0 │\n\
│ 0 0 0 -1 0 0 4 -1 -1 -1 0 0 │\n\
│ 0 0 0 0 -1 0 -1 4 -1 0 -1 0 │\n\
│ 0 0 0 0 0 -1 -1 -1 4 0 0 -1 │\n\
│ -1 0 0 0 0 0 -1 0 0 4 -1 -1 │\n\
│ 0 -1 0 0 0 0 0 -1 0 -1 4 -1 │\n\
│ 0 0 -1 0 0 0 0 0 -1 -1 -1 4 │\n\
└ ┘"
);
assert_eq!(
format!("{}", aa.as_dense()),
"┌ ┐\n\
│ 4 -1 -1 -1 0 0 0 0 0 -1 0 0 │\n\
│ -1 4 -1 0 -1 0 0 0 0 0 -1 0 │\n\
│ -1 -1 4 0 0 -1 0 0 0 0 0 -1 │\n\
│ -1 0 0 4 -1 -1 -1 0 0 0 0 0 │\n\
│ 0 -1 0 -1 4 -1 0 -1 0 0 0 0 │\n\
│ 0 0 -1 -1 -1 4 0 0 -1 0 0 0 │\n\
│ 0 0 0 -1 0 0 4 -1 -1 -1 0 0 │\n\
│ 0 0 0 0 -1 0 -1 4 -1 0 -1 0 │\n\
│ 0 0 0 0 0 -1 -1 -1 4 0 0 -1 │\n\
│ -1 0 0 0 0 0 -1 0 0 4 -1 -1 │\n\
│ 0 -1 0 0 0 0 0 -1 0 -1 4 -1 │\n\
│ 0 0 -1 0 0 0 0 0 -1 -1 -1 4 │\n\
└ ┘"
);
}
#[test]
fn get_vectors_works() {
let grid = Grid2d::new_uniform(1.0, 4.0, 1.0, 3.0, 4, 3).unwrap();
let mut ebcs = EssentialBcs2d::new();
ebcs.set(Side::Xmin, |x, y| x + y);
ebcs.set(Side::Xmax, |x, y| x + y);
ebcs.set(Side::Ymin, |x, y| x + y);
ebcs.set(Side::Ymax, |x, y| x + y);
let nbcs = NaturalBcs2d::new();
let fdm = Fdm2d::new(grid, ebcs, nbcs, 1.0, 1.0).unwrap();
let nu = 2;
let np = 10;
let neq = nu + np;
let (a_bar, a_check, f_bar) = fdm.get_vectors_sps(|_, _| 100.0);
assert_eq!(a_bar.dim(), nu);
assert_eq!(a_check.dim(), np);
assert_eq!(f_bar.dim(), nu);
assert_eq!(a_bar.as_data(), &[0.0, 0.0]);
assert_eq!(
a_check.as_data(),
&[
1.0 + 1.0, 2.0 + 1.0, 3.0 + 1.0, 4.0 + 1.0, 1.0 + 2.0, 4.0 + 2.0, 1.0 + 3.0, 2.0 + 3.0, 3.0 + 3.0, 4.0 + 3.0, ]
);
assert_eq!(f_bar.as_data(), &[100.0, 100.0]);
let a = fdm.get_joined_vector_sps(&a_bar, &a_check);
assert_eq!(a.dim(), neq);
assert_eq!(
a.as_data(),
&[
1.0 + 1.0, 2.0 + 1.0, 3.0 + 1.0, 4.0 + 1.0, 1.0 + 2.0, 0.0, 0.0, 4.0 + 2.0, 1.0 + 3.0, 2.0 + 3.0, 3.0 + 3.0, 4.0 + 3.0, ]
);
let (aa, ff) = fdm.get_vectors_lmm(|_, _| 100.0);
assert_eq!(aa.dim(), neq + np);
assert_eq!(aa.as_data(), &vec![0.0; neq + np]);
assert_eq!(
ff.as_data(),
&[
25.0, 50.0, 50.0, 25.0, 50.0, 100.0, 100.0, 50.0, 25.0, 50.0, 50.0, 25.0, 1.0 + 1.0, 2.0 + 1.0, 3.0 + 1.0, 4.0 + 1.0, 1.0 + 2.0, 4.0 + 2.0, 1.0 + 3.0, 2.0 + 3.0, 3.0 + 3.0, 4.0 + 3.0, ]
);
}
#[test]
fn get_grid_and_get_equations_work() {
let grid = Grid2d::new_uniform(0.0, 1.0, 0.0, 1.0, 3, 3).unwrap();
let mut ebcs = EssentialBcs2d::new();
let nbcs = NaturalBcs2d::new();
ebcs.set_homogeneous();
let fdm = Fdm2d::new(grid, ebcs, nbcs, 1.0, 1.0).unwrap();
assert_eq!(fdm.get_grid().nx(), 3);
assert_eq!(fdm.get_equations().neq(), 9);
}
}