use crate::errors::QlResult;
use crate::math::array::Array;
use crate::methods::finitedifferences::meshers::FdmMesher;
use crate::shared::Shared;
use crate::types::{Real, Size};
use crate::{ensure, require};
use super::fdmlinearop::FdmLinearOp;
use super::fdmlinearopiterator::FdmLinearOpIterator;
use super::fdmlinearoplayout::FdmLinearOpLayout;
#[derive(Clone)]
pub struct TripleBandLinearOp {
direction: Size,
i0: Vec<Size>,
i2: Vec<Size>,
reverse_index: Vec<Size>,
lower: Vec<Real>,
diag: Vec<Real>,
upper: Vec<Real>,
mesher: Shared<dyn FdmMesher>,
}
impl TripleBandLinearOp {
pub fn new(direction: Size, mesher: Shared<dyn FdmMesher>) -> Self {
let layout = Shared::clone(mesher.layout());
let size = layout.size();
let mut new_dim = layout.dim().to_vec();
new_dim.swap(0, direction);
let mut new_spacing = FdmLinearOpLayout::new(new_dim).spacing().to_vec();
new_spacing.swap(0, direction);
let mut i0 = vec![0; size];
let mut i2 = vec![0; size];
let mut reverse_index = vec![0; size];
let mut position = layout.begin();
while position.index() < size {
let i = position.index();
i0[i] = layout.neighbourhood(&position, direction, -1);
i2[i] = layout.neighbourhood(&position, direction, 1);
let transposed: Size = position
.coordinates()
.iter()
.zip(&new_spacing)
.map(|(coordinate, stride)| coordinate * stride)
.sum();
reverse_index[transposed] = i;
position.advance();
}
TripleBandLinearOp {
direction,
i0,
i2,
reverse_index,
lower: vec![0.0; size],
diag: vec![0.0; size],
upper: vec![0.0; size],
mesher,
}
}
pub fn with_bands(
direction: Size,
mesher: Shared<dyn FdmMesher>,
mut bands: impl FnMut(&dyn FdmMesher, &FdmLinearOpIterator) -> (Real, Real, Real),
) -> Self {
let mut operator = Self::new(direction, Shared::clone(&mesher));
let layout = Shared::clone(mesher.layout());
let mut position = layout.begin();
while position.index() < layout.size() {
let i = position.index();
let (lower, diag, upper) = bands(&*mesher, &position);
operator.lower[i] = lower;
operator.diag[i] = diag;
operator.upper[i] = upper;
position.advance();
}
operator
}
pub fn direction(&self) -> Size {
self.direction
}
pub fn mesher(&self) -> &Shared<dyn FdmMesher> {
&self.mesher
}
#[allow(clippy::needless_range_loop)]
pub fn solve_splitting(&self, r: &Array, a: Real, b: Real) -> QlResult<Array> {
let size = self.size();
require!(r.size() == size, "inconsistent size of rhs");
let mut result = Array::with_size(size);
let mut tmp = vec![0.0; size];
let mut previous = self.reverse_index[0];
let mut bet = 1.0 / (a * self.diag[previous] + b);
require!(bet != 0.0, "division by zero");
result[previous] = r[previous] * bet;
for j in 1..size {
let current = self.reverse_index[j];
tmp[j] = a * self.upper[previous] * bet;
bet = b + a * (self.diag[current] - tmp[j] * self.lower[current]);
ensure!(bet != 0.0, "division by zero");
bet = 1.0 / bet;
result[current] = (r[current] - a * self.lower[current] * result[previous]) * bet;
previous = current;
}
for j in (1..size.saturating_sub(1)).rev() {
result[self.reverse_index[j]] -= tmp[j + 1] * result[self.reverse_index[j + 1]];
}
if size > 1 {
result[self.reverse_index[0]] -= tmp[1] * result[self.reverse_index[1]];
}
Ok(result)
}
pub fn axpyb(&mut self, a: &Array, x: &TripleBandLinearOp, y: &TripleBandLinearOp, b: &Array) {
let size = self.size();
assert!(broadcasts_over(a, size), "inconsistent size of a");
assert!(broadcasts_over(b, size), "inconsistent size of b");
let a_stride = if a.size() > 1 { 1 } else { 0 };
let b_stride = if b.size() > 1 { 1 } else { 0 };
for i in 0..size {
let (lower, mut diag, upper) = if a.is_empty() {
(y.lower[i], y.diag[i], y.upper[i])
} else {
let s = a[i * a_stride];
(
y.lower[i] + s * x.lower[i],
y.diag[i] + s * x.diag[i],
y.upper[i] + s * x.upper[i],
)
};
if !b.is_empty() {
diag += b[i * b_stride];
}
self.lower[i] = lower;
self.diag[i] = diag;
self.upper[i] = upper;
}
}
pub fn mult(&self, u: &Array) -> TripleBandLinearOp {
assert_eq!(u.size(), self.size(), "inconsistent size of u");
self.map_bands(|i, lower, diag, upper| {
let s = u[i];
(lower * s, diag * s, upper * s)
})
}
pub fn mult_r(&self, u: &Array) -> TripleBandLinearOp {
let size = self.size();
assert_eq!(u.size(), size, "inconsistent size of rhs");
self.map_bands(|i, lower, diag, upper| {
let previous = if i > 0 { u[i - 1] } else { 1.0 };
let next = if i + 1 < size { u[i + 1] } else { 1.0 };
(lower * previous, diag * u[i], upper * next)
})
}
pub fn add_op(&self, m: &TripleBandLinearOp) -> TripleBandLinearOp {
self.map_bands(|i, lower, diag, upper| {
(lower + m.lower[i], diag + m.diag[i], upper + m.upper[i])
})
}
pub fn add_diagonal(&self, u: &Array) -> TripleBandLinearOp {
assert_eq!(u.size(), self.size(), "inconsistent size of u");
self.map_bands(|i, lower, diag, upper| (lower, diag + u[i], upper))
}
fn map_bands(
&self,
bands: impl Fn(Size, Real, Real, Real) -> (Real, Real, Real),
) -> TripleBandLinearOp {
let mut result = self.clone();
for i in 0..self.size() {
let (lower, diag, upper) = bands(i, self.lower[i], self.diag[i], self.upper[i]);
result.lower[i] = lower;
result.diag[i] = diag;
result.upper[i] = upper;
}
result
}
fn size(&self) -> Size {
self.diag.len()
}
}
fn broadcasts_over(values: &Array, size: Size) -> bool {
values.is_empty() || values.size() == 1 || values.size() == size
}
impl FdmLinearOp for TripleBandLinearOp {
fn apply(&self, r: &Array) -> Array {
assert_eq!(r.size(), self.size(), "inconsistent length of r");
let mut result = Array::with_size(r.size());
for i in 0..self.size() {
result[i] =
r[self.i0[i]] * self.lower[i] + r[i] * self.diag[i] + r[self.i2[i]] * self.upper[i];
}
result
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::methods::finitedifferences::meshers::UniformGridMesher;
use crate::methods::finitedifferences::operators::{first_derivative_op, second_derivative_op};
use crate::shared::shared;
const DIM: [Size; 2] = [3, 4];
const DIRECTION: Size = 1;
const TOL: Real = 1e-12;
fn mesher(dim: &[Size]) -> Shared<dyn FdmMesher> {
let layout = shared(FdmLinearOpLayout::new(dim.to_vec()));
let boundaries = vec![(0.0, 1.0); dim.len()];
shared(UniformGridMesher::new(layout, &boundaries).unwrap())
}
fn band_values(index: Size, coordinate: Size, extent: Size) -> (Real, Real, Real) {
let i = index as Real;
let lower = if coordinate == 0 { 0.0 } else { -1.0 - 0.1 * i };
let upper = if coordinate == extent - 1 {
0.0
} else {
1.0 + 0.2 * i
};
(lower, 8.0 + 0.3 * i, upper)
}
fn banded(dim: &[Size], direction: Size) -> TripleBandLinearOp {
let mesher = mesher(dim);
let extent = dim[direction];
TripleBandLinearOp::with_bands(direction, mesher, move |_, position| {
band_values(position.index(), position.coordinates()[direction], extent)
})
}
fn alternate(operator: &TripleBandLinearOp) -> TripleBandLinearOp {
TripleBandLinearOp::with_bands(
operator.direction(),
Shared::clone(operator.mesher()),
|_, position| {
let i = position.index() as Real;
(0.5 - 0.05 * i, 2.0 * i, 1.0 + 0.15 * i)
},
)
}
fn assert_close(actual: &Array, expected: &Array) {
assert_eq!(actual.size(), expected.size());
for i in 0..actual.size() {
assert!(
(actual[i] - expected[i]).abs() <= TOL,
"element {i}: {} != {}",
actual[i],
expected[i]
);
}
}
#[test]
fn reverse_index_walks_the_direction_fastest() {
let operator = banded(&DIM, DIRECTION);
let size = DIM.iter().product::<Size>();
let extent = DIM[DIRECTION];
let mut seen = vec![false; size];
for &index in &operator.reverse_index {
assert!(!seen[index], "reverse_index repeats {index}");
seen[index] = true;
}
for j in 0..size {
if j % extent != extent - 1 {
assert_eq!(
operator.i2[operator.reverse_index[j]],
operator.reverse_index[j + 1],
"chain broken between {j} and {}",
j + 1
);
}
}
}
#[test]
fn reverse_index_is_the_identity_along_the_first_direction() {
let operator = banded(&DIM, 0);
let expected: Vec<Size> = (0..DIM.iter().product::<Size>()).collect();
assert_eq!(operator.reverse_index, expected);
}
#[test]
fn apply_matches_a_hand_computed_product() {
let mesher = mesher(&[4]);
let operator =
TripleBandLinearOp::with_bands(0, mesher, |_, position| match position.index() {
0 => (0.0, 10.0, 2.0),
1 => (1.0, 20.0, 2.0),
2 => (1.0, 30.0, 2.0),
_ => (1.0, 40.0, 0.0),
});
let r = Array::from([1.0, 2.0, 3.0, 4.0]);
assert_eq!(operator.apply(&r), Array::from([14.0, 47.0, 100.0, 163.0]));
}
#[test]
fn apply_reaches_the_neighbours_the_layout_names() {
let operator = banded(&DIM, DIRECTION);
let layout = Shared::clone(operator.mesher().layout());
let r = Array::incremental(layout.size(), 1.0, 3.0);
let mut expected = Array::with_size(layout.size());
for position in layout.iter() {
let i = position.index();
let (lower, diag, upper) =
band_values(i, position.coordinates()[DIRECTION], DIM[DIRECTION]);
expected[i] = r[layout.neighbourhood(&position, DIRECTION, -1)] * lower
+ r[i] * diag
+ r[layout.neighbourhood(&position, DIRECTION, 1)] * upper;
}
assert_close(&operator.apply(&r), &expected);
}
#[test]
fn solve_splitting_inverts_apply() {
for direction in 0..DIM.len() {
let operator = banded(&DIM, direction);
let x = Array::incremental(operator.size(), 2.0, -0.5);
let recovered = operator
.solve_splitting(&operator.apply(&x), 1.0, 0.0)
.unwrap();
assert_close(&recovered, &x);
}
}
#[test]
fn solve_splitting_solves_the_shifted_system() {
let operator = banded(&DIM, DIRECTION);
let x = Array::incremental(operator.size(), -1.0, 0.7);
let (a, b) = (0.25, 1.5);
let r = &(&operator.apply(&x) * a) + &(&x * b);
let recovered = operator.solve_splitting(&r, a, b).unwrap();
assert_close(&recovered, &x);
}
#[test]
fn solve_splitting_rejects_a_mismatched_rhs() {
let operator = banded(&DIM, DIRECTION);
let err = operator
.solve_splitting(&Array::with_size(operator.size() - 1), 1.0, 0.0)
.unwrap_err();
assert_eq!(err.message(), "inconsistent size of rhs");
}
#[test]
fn axpyb_broadcasts_a_and_b_over_the_grid() {
let operator = banded(&DIM, DIRECTION);
let size = operator.size();
let r = Array::incremental(size, 1.0, 2.0);
let applied = operator.apply(&r);
let scalar = 3.0;
let vector = Array::incremental(size, 0.5, 0.25);
let shift = Array::incremental(size, -2.0, 0.125);
let cases: [(Array, Array); 4] = [
(Array::new(), Array::new()),
(Array::new(), Array::filled(1, 0.75)),
(Array::filled(1, scalar), Array::new()),
(vector.clone(), shift.clone()),
];
for (a, b) in cases {
let mut target = TripleBandLinearOp::new(DIRECTION, Shared::clone(operator.mesher()));
target.axpyb(&a, &operator, &operator, &b);
let mut expected = Array::with_size(size);
for i in 0..size {
let scale = if a.is_empty() {
0.0
} else if a.size() == 1 {
a[0]
} else {
a[i]
};
let bias = if b.is_empty() {
0.0
} else if b.size() == 1 {
b[0]
} else {
b[i]
};
expected[i] = applied[i] * (1.0 + scale) + bias * r[i];
}
assert_close(&target.apply(&r), &expected);
}
}
#[test]
fn axpyb_combines_two_distinct_operators() {
let x = banded(&DIM, DIRECTION);
let y = alternate(&x);
let size = x.size();
let r = Array::incremental(size, 3.0, -1.0);
let a = Array::incremental(size, 1.0, 0.5);
let mut target = TripleBandLinearOp::new(DIRECTION, Shared::clone(x.mesher()));
target.axpyb(&a, &x, &y, &Array::new());
let (applied_x, applied_y) = (x.apply(&r), y.apply(&r));
let expected: Array = (0..size)
.map(|i| applied_y[i] + a[i] * applied_x[i])
.collect();
assert_close(&target.apply(&r), &expected);
}
#[test]
#[should_panic(expected = "inconsistent size of a")]
fn axpyb_rejects_an_unbroadcastable_a() {
let operator = banded(&DIM, DIRECTION);
let mut target = TripleBandLinearOp::new(DIRECTION, Shared::clone(operator.mesher()));
target.axpyb(&Array::with_size(2), &operator, &operator, &Array::new());
}
#[test]
#[should_panic(expected = "inconsistent length of r")]
fn apply_rejects_a_mismatched_argument() {
let operator = banded(&DIM, DIRECTION);
operator.apply(&Array::with_size(operator.size() + 1));
}
#[test]
fn mult_scales_each_row() {
let operator = banded(&DIM, DIRECTION);
let size = operator.size();
let r = Array::incremental(size, 1.0, 1.5);
let u = Array::incremental(size, 0.5, 0.25);
let applied = operator.apply(&r);
let expected: Array = (0..size).map(|i| u[i] * applied[i]).collect();
assert_close(&operator.mult(&u).apply(&r), &expected);
}
#[test]
fn mult_r_scales_each_column() {
let operator = banded(&[4], 0);
let r = Array::from([1.0, 2.0, 3.0, 4.0]);
let u = Array::from([0.5, -1.5, 2.0, 3.5]);
let scaled: Array = (0..r.size()).map(|i| u[i] * r[i]).collect();
assert_close(&operator.mult_r(&u).apply(&r), &operator.apply(&scaled));
}
#[test]
fn mult_r_substitutes_one_at_the_flat_array_ends() {
let operator = TripleBandLinearOp::with_bands(0, mesher(&[3]), |_, _| (1.0, 1.0, 1.0));
let scaled = operator.mult_r(&Array::with_size(3));
assert_eq!(
scaled.apply(&Array::from([1.0, 2.0, 3.0])),
Array::from([2.0, 0.0, 2.0])
);
}
#[test]
fn add_op_sums_the_operators() {
let operator = banded(&DIM, DIRECTION);
let other = alternate(&operator);
let r = Array::incremental(operator.size(), 2.0, -0.75);
let expected = &operator.apply(&r) + &other.apply(&r);
assert_close(&operator.add_op(&other).apply(&r), &expected);
}
#[test]
fn add_diagonal_adds_to_the_diagonal() {
let operator = banded(&DIM, DIRECTION);
let size = operator.size();
let r = Array::incremental(size, 1.0, 0.5);
let u = Array::incremental(size, -1.0, 0.3);
let applied = operator.apply(&r);
let expected: Array = (0..size).map(|i| applied[i] + u[i] * r[i]).collect();
assert_close(&operator.add_diagonal(&u).apply(&r), &expected);
}
#[test]
#[should_panic(expected = "inconsistent size of rhs")]
fn mult_r_rejects_a_mismatched_argument() {
let operator = banded(&DIM, DIRECTION);
operator.mult_r(&Array::with_size(operator.size() - 1));
}
#[test]
fn triple_band_map_solve_matches_quantlib() {
let layout = shared(FdmLinearOpLayout::new(vec![100, 400]));
let mesher: Shared<dyn FdmMesher> =
shared(UniformGridMesher::new(Shared::clone(&layout), &[(0.0, 1.0); 2]).unwrap());
let one = Array::filled(1, 1.0);
let u: Array = (0..layout.size())
.map(|i| (0.1 * i as Real).sin() + (0.35 * i as Real).cos())
.collect();
let assert_recovers = |recovered: &Array| {
for i in 0..u.size() {
assert!(
(u[i] - recovered[i]).abs() <= 1e-6,
"solve and apply are not consistent at {i}: {} != {}",
recovered[i],
u[i]
);
}
};
let mut dy = first_derivative_op(1, Shared::clone(&mesher));
let dy_before = dy.clone();
dy.axpyb(&Array::filled(1, 2.0), &dy_before, &dy_before, &one);
let copy_of_dy = dy.clone();
assert_recovers(&dy.solve_splitting(©_of_dy.apply(&u), 1.0, 0.0).unwrap());
let mut dx = first_derivative_op(0, Shared::clone(&mesher));
let dx_before = dx.clone();
dx.axpyb(&Array::new(), &dx_before, &dx_before, &one);
let copy_of_dx = dx.clone();
assert_recovers(&dx.solve_splitting(©_of_dx.apply(&u), 1.0, 0.0).unwrap());
let mut dxx = second_derivative_op(0, Shared::clone(&mesher));
let dxx_before = dxx.clone();
dxx.axpyb(&Array::filled(1, 0.5), &dxx_before, &dx, &one);
let copy_of_dxx = dxx.clone();
assert_recovers(
&dxx.solve_splitting(©_of_dxx.apply(&u), 1.0, 0.0)
.unwrap(),
);
let _ = copy_of_dxx.add_op(&second_derivative_op(1, Shared::clone(&mesher)));
let copy_of_dxx = dxx.clone();
assert_recovers(
&dxx.solve_splitting(©_of_dxx.apply(&u), 1.0, 0.0)
.unwrap(),
);
}
}