use std::cell::RefCell;
use crate::math::comparison::close_enough;
use crate::types::{Real, Size};
use super::gaussianorthogonalpolynomial::GaussianOrthogonalPolynomial;
pub(super) fn memoized(
cache: &RefCell<Vec<Real>>,
n: Size,
compute: impl FnOnce() -> Real,
) -> Real {
{
let mut c = cache.borrow_mut();
if c.len() <= n {
c.resize(n + 1, Real::NAN);
}
if !c[n].is_nan() {
return c[n];
}
}
let value = compute();
cache.borrow_mut()[n] = value;
value
}
pub trait MomentBasedPolynomial {
fn moment(&self, i: Size) -> Real;
fn w(&self, x: Real) -> Real;
}
pub struct MomentBasedGaussianPolynomial<P: MomentBasedPolynomial> {
poly: P,
b: RefCell<Vec<Real>>,
c: RefCell<Vec<Real>>,
z: RefCell<Vec<Vec<Real>>>,
}
impl<P: MomentBasedPolynomial> MomentBasedGaussianPolynomial<P> {
pub fn new(poly: P) -> Self {
MomentBasedGaussianPolynomial {
poly,
b: RefCell::new(Vec::new()),
c: RefCell::new(Vec::new()),
z: RefCell::new(vec![Vec::new()]),
}
}
fn ensure_z(&self, k: Size, i: Size) {
let mut z = self.z.borrow_mut();
let cols = z[0].len().max(i + 1);
for row in z.iter_mut() {
if row.len() < cols {
row.resize(cols, Real::NAN);
}
}
if z.len() <= k {
z.resize(k + 1, vec![Real::NAN; cols]);
}
}
fn z(&self, k: Size, i: Size) -> Real {
self.ensure_z(k, i);
let cached = self.z.borrow()[k][i];
if !cached.is_nan() {
return cached;
}
let value = if k == 0 {
self.poly.moment(i)
} else {
let mut value = self.z(k - 1, i + 1) - self.alpha_(k - 1) * self.z(k - 1, i);
if k >= 2 {
value -= self.beta_(k - 1) * self.z(k - 2, i);
}
value
};
self.z.borrow_mut()[k][i] = value;
value
}
fn alpha_(&self, u: Size) -> Real {
memoized(&self.b, u, || {
if u == 0 {
self.poly.moment(1)
} else {
-self.z(u - 1, u) / self.z(u - 1, u - 1) + self.z(u, u + 1) / self.z(u, u)
}
})
}
fn beta_(&self, u: Size) -> Real {
if u == 0 {
return 1.0;
}
memoized(&self.c, u, || self.z(u, u) / self.z(u - 1, u - 1))
}
}
impl<P: MomentBasedPolynomial> GaussianOrthogonalPolynomial for MomentBasedGaussianPolynomial<P> {
fn mu_0(&self) -> Real {
let m0 = self.poly.moment(0);
assert!(close_enough(m0, 1.0), "zero moment must be one");
m0
}
fn alpha(&self, i: Size) -> Real {
self.alpha_(i)
}
fn beta(&self, i: Size) -> Real {
self.beta_(i)
}
fn w(&self, x: Real) -> Real {
self.poly.w(x)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::math::integrals::gaussianorthogonalpolynomial::GaussLaguerrePolynomial;
struct MomentBasedGaussLaguerrePolynomial;
impl MomentBasedPolynomial for MomentBasedGaussLaguerrePolynomial {
fn moment(&self, i: Size) -> Real {
if i == 0 {
1.0
} else {
i as Real * self.moment(i - 1)
}
}
fn w(&self, x: Real) -> Real {
(-x).exp()
}
}
#[test]
fn moment_based_polynomial_reproduces_laguerre_recurrence() {
let g = GaussLaguerrePolynomial::new(0.0).expect("0 > -1");
let k = MomentBasedGaussianPolynomial::new(MomentBasedGaussLaguerrePolynomial);
let tol = 1e-12;
for i in 0..10 {
let diff_alpha = (k.alpha(i) - g.alpha(i)).abs();
assert!(
diff_alpha <= tol,
"failed to reproduce alpha for Laguerre quadrature: \
calculated {}, expected {}, diff {}",
k.alpha(i),
g.alpha(i),
diff_alpha
);
if i > 0 {
let diff_beta = (k.beta(i) - g.beta(i)).abs();
assert!(
diff_beta <= tol,
"failed to reproduce beta for Laguerre quadrature: \
calculated {}, expected {}, diff {}",
k.beta(i),
g.beta(i),
diff_beta
);
}
}
}
#[test]
fn mu_0_is_the_zeroth_moment() {
let k = MomentBasedGaussianPolynomial::new(MomentBasedGaussLaguerrePolynomial);
assert_eq!(k.mu_0(), 1.0);
}
}