use crate::error::SolveError;
const GAUSS_X: [f64; 5] = [
-0.906_179_845_938_664,
-0.538_469_310_105_683_1,
0.0,
0.538_469_310_105_683_1,
0.906_179_845_938_664,
];
const GAUSS_W: [f64; 5] = [
0.236_926_885_056_189_1,
0.478_628_670_499_366_5,
0.568_888_888_888_888_9,
0.478_628_670_499_366_5,
0.236_926_885_056_189_1,
];
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum Bc {
Dirichlet(f64),
Neumann(f64),
Robin { alpha: f64, g: f64 },
}
#[derive(Debug, Clone, PartialEq)]
pub struct Fem1dSolution {
pub a: f64,
pub b: f64,
pub degree: usize,
pub values: Vec<f64>,
}
impl Fem1dSolution {
pub fn new(a: f64, b: f64, degree: usize, values: Vec<f64>) -> Result<Self, SolveError> {
if !(degree == 1 || degree == 2) {
return Err(SolveError::InvalidArgument("degree must be 1 or 2"));
}
if !(a.is_finite() && b.is_finite()) || b <= a {
return Err(SolveError::InvalidArgument("need a finite interval with a < b"));
}
if values.len() < degree + 1 || !(values.len() - 1).is_multiple_of(degree) {
return Err(SolveError::InvalidArgument("value count does not match the degree"));
}
Ok(Self { a, b, degree, values })
}
pub fn elements(&self) -> usize {
(self.values.len() - 1) / self.degree
}
pub fn h(&self) -> f64 {
(self.b - self.a) / self.elements() as f64
}
pub fn nodes(&self) -> Vec<f64> {
let step = (self.b - self.a) / (self.values.len() - 1) as f64;
(0..self.values.len()).map(|i| self.a + i as f64 * step).collect()
}
fn locate(&self, x: f64) -> (usize, f64) {
let ne = self.elements();
let h = self.h();
let raw = ((x - self.a) / h).floor();
let e = if raw < 0.0 {
0
} else if raw >= ne as f64 {
ne - 1
} else {
raw as usize
};
let left = self.a + e as f64 * h;
(e, 2.0 * (x - left) / h - 1.0)
}
pub fn eval(&self, x: f64) -> f64 {
let (e, xi) = self.locate(x);
let base = e * self.degree;
shape(self.degree, xi)
.iter()
.enumerate()
.map(|(k, n)| n * self.values[base + k])
.sum()
}
pub fn eval_derivative(&self, x: f64) -> f64 {
let (e, xi) = self.locate(x);
let base = e * self.degree;
let scale = 2.0 / self.h();
shape_derivative(self.degree, xi)
.iter()
.enumerate()
.map(|(k, d)| d * scale * self.values[base + k])
.sum()
}
}
fn shape(degree: usize, xi: f64) -> Vec<f64> {
if degree == 1 {
vec![0.5 * (1.0 - xi), 0.5 * (1.0 + xi)]
} else {
vec![0.5 * xi * (xi - 1.0), 1.0 - xi * xi, 0.5 * xi * (xi + 1.0)]
}
}
fn shape_derivative(degree: usize, xi: f64) -> Vec<f64> {
if degree == 1 {
vec![-0.5, 0.5]
} else {
vec![xi - 0.5, -2.0 * xi, xi + 0.5]
}
}
struct Banded {
n: usize,
half: usize,
data: Vec<f64>,
}
impl Banded {
fn new(n: usize, half: usize) -> Self {
Self { n, half, data: vec![0.0; n * (half + 1)] }
}
fn get(&self, i: usize, j: usize) -> f64 {
let (lo, hi) = if i <= j { (i, j) } else { (j, i) };
if hi - lo > self.half {
0.0
} else {
self.data[lo * (self.half + 1) + (hi - lo)]
}
}
fn add(&mut self, i: usize, j: usize, v: f64) {
let (lo, hi) = if i <= j { (i, j) } else { (j, i) };
self.data[lo * (self.half + 1) + (hi - lo)] += v;
}
fn set(&mut self, i: usize, j: usize, v: f64) {
let (lo, hi) = if i <= j { (i, j) } else { (j, i) };
self.data[lo * (self.half + 1) + (hi - lo)] = v;
}
fn row_sums(&self) -> Vec<f64> {
(0..self.n)
.map(|i| {
let lo = i.saturating_sub(self.half);
let hi = (i + self.half).min(self.n - 1);
(lo..=hi).map(|j| self.get(i, j)).sum()
})
.collect()
}
fn diagonal_scale(&self) -> f64 {
(0..self.n).map(|i| self.get(i, i).abs()).fold(0.0, f64::max)
}
fn ldl_solve(&self, rhs: &[f64]) -> Result<Vec<f64>, SolveError> {
let n = self.n;
let m = self.half;
let scale = self.diagonal_scale().max(f64::MIN_POSITIVE);
let mut l = vec![0.0; n * m];
let mut d = vec![0.0; n];
let at = |l: &[f64], i: usize, k: usize| -> f64 {
if i == k {
1.0
} else if i > k && i - k <= m {
l[i * m + (i - k - 1)]
} else {
0.0
}
};
for j in 0..n {
let mut dj = self.get(j, j);
for k in j.saturating_sub(m)..j {
let ljk = at(&l, j, k);
dj -= ljk * ljk * d[k];
}
if dj.abs() <= 1e-13 * scale {
return Err(SolveError::Singular);
}
d[j] = dj;
for i in (j + 1)..(j + m + 1).min(n) {
let mut s = self.get(i, j);
for k in i.saturating_sub(m)..j {
s -= at(&l, i, k) * at(&l, j, k) * d[k];
}
l[i * m + (i - j - 1)] = s / dj;
}
}
let mut y = rhs.to_vec();
for i in 0..n {
for k in i.saturating_sub(m)..i {
y[i] -= at(&l, i, k) * y[k];
}
}
for i in 0..n {
y[i] /= d[i];
}
for i in (0..n).rev() {
for k in (i + 1)..(i + m + 1).min(n) {
y[i] -= at(&l, k, i) * y[k];
}
}
Ok(y)
}
}
fn solve_degree(
p: &dyn Fn(f64) -> f64,
q: &dyn Fn(f64) -> f64,
f: &dyn Fn(f64) -> f64,
a: f64,
b: f64,
bc: (Bc, Bc),
n: usize,
degree: usize,
) -> Result<Vec<f64>, SolveError> {
if n == 0 {
return Err(SolveError::InvalidArgument("need at least one element"));
}
if !(a.is_finite() && b.is_finite()) || b <= a {
return Err(SolveError::InvalidArgument("need a finite interval with a < b"));
}
let h = (b - a) / n as f64;
let nodes = degree * n + 1;
let mut mat = Banded::new(nodes, degree);
let mut rhs = vec![0.0; nodes];
for e in 0..n {
let left = a + e as f64 * h;
let base = e * degree;
for (&xi, &w) in GAUSS_X.iter().zip(GAUSS_W.iter()) {
let x = left + 0.5 * (xi + 1.0) * h;
let pv = p(x);
let qv = q(x);
let fv = f(x);
if !(pv.is_finite() && qv.is_finite() && fv.is_finite()) {
return Err(SolveError::InvalidArgument("coefficients must be finite"));
}
if pv <= 0.0 {
return Err(SolveError::InvalidArgument("p must be positive"));
}
let sh = shape(degree, xi);
let dsh = shape_derivative(degree, xi);
let stiff = 2.0 * w * pv / h;
let mass = 0.5 * w * qv * h;
let load = 0.5 * w * fv * h;
for j in 0..=degree {
rhs[base + j] += load * sh[j];
for k in j..=degree {
mat.add(base + j, base + k, stiff * dsh[j] * dsh[k] + mass * sh[j] * sh[k]);
}
}
}
}
for (end, cond) in [(0usize, bc.0), (nodes - 1, bc.1)] {
match cond {
Bc::Dirichlet(_) => {}
Bc::Neumann(g) => {
if !g.is_finite() {
return Err(SolveError::InvalidArgument("boundary data must be finite"));
}
rhs[end] += g;
}
Bc::Robin { alpha, g } => {
if !(alpha.is_finite() && g.is_finite()) {
return Err(SolveError::InvalidArgument("boundary data must be finite"));
}
mat.add(end, end, alpha);
rhs[end] += g;
}
}
}
let has_dirichlet = matches!(bc.0, Bc::Dirichlet(_)) || matches!(bc.1, Bc::Dirichlet(_));
if !has_dirichlet {
let scale = mat.diagonal_scale().max(f64::MIN_POSITIVE);
if mat.row_sums().iter().all(|s| s.abs() <= 1e-12 * scale) {
return Err(SolveError::Singular);
}
}
for (end, cond) in [(0usize, bc.0), (nodes - 1, bc.1)] {
if let Bc::Dirichlet(g) = cond {
if !g.is_finite() {
return Err(SolveError::InvalidArgument("boundary data must be finite"));
}
let lo = end.saturating_sub(degree);
let hi = (end + degree).min(nodes - 1);
for j in lo..=hi {
if j != end {
rhs[j] -= mat.get(j, end) * g;
mat.set(j, end, 0.0);
}
}
mat.set(end, end, 1.0);
rhs[end] = g;
}
}
mat.ldl_solve(&rhs)
}
pub fn fem_1d_poisson(
f: &dyn Fn(f64) -> f64,
a: f64,
b: f64,
bc: (Bc, Bc),
n: usize,
) -> Result<Vec<f64>, SolveError> {
solve_degree(&|_| 1.0, &|_| 0.0, f, a, b, bc, n, 1)
}
pub fn fem_1d_general(
p: &dyn Fn(f64) -> f64,
q: &dyn Fn(f64) -> f64,
f: &dyn Fn(f64) -> f64,
a: f64,
b: f64,
bc: (Bc, Bc),
n: usize,
) -> Result<Vec<f64>, SolveError> {
solve_degree(p, q, f, a, b, bc, n, 1)
}
pub fn fem_1d_quadratic(
p: &dyn Fn(f64) -> f64,
q: &dyn Fn(f64) -> f64,
f: &dyn Fn(f64) -> f64,
a: f64,
b: f64,
bc: (Bc, Bc),
n: usize,
) -> Result<Vec<f64>, SolveError> {
solve_degree(p, q, f, a, b, bc, n, 2)
}
fn integrate_by_element(u_h: &Fem1dSolution, g: &dyn Fn(f64) -> f64) -> f64 {
let h = u_h.h();
let mut total = 0.0;
for e in 0..u_h.elements() {
let left = u_h.a + e as f64 * h;
for (&xi, &w) in GAUSS_X.iter().zip(GAUSS_W.iter()) {
let x = left + 0.5 * (xi + 1.0) * h;
total += 0.5 * w * h * g(x);
}
}
total
}
pub fn fem_1d_error_l2(u_h: &Fem1dSolution, u_exact: &dyn Fn(f64) -> f64) -> f64 {
integrate_by_element(u_h, &|x| {
let e = u_exact(x) - u_h.eval(x);
e * e
})
.max(0.0)
.sqrt()
}
pub fn fem_1d_error_h1_seminorm(u_h: &Fem1dSolution, du_exact: &dyn Fn(f64) -> f64) -> f64 {
integrate_by_element(u_h, &|x| {
let e = du_exact(x) - u_h.eval_derivative(x);
e * e
})
.max(0.0)
.sqrt()
}
pub fn fem_1d_error_h1(
u_h: &Fem1dSolution,
u_exact: &dyn Fn(f64) -> f64,
du_exact: &dyn Fn(f64) -> f64,
) -> f64 {
let l2 = fem_1d_error_l2(u_h, u_exact);
let semi = fem_1d_error_h1_seminorm(u_h, du_exact);
l2.hypot(semi)
}
pub fn convergence_rate(errors: &[f64], hs: &[f64]) -> Result<f64, SolveError> {
if errors.len() != hs.len() {
return Err(SolveError::DimensionMismatch { expected: errors.len(), got: hs.len() });
}
if errors.len() < 2 {
return Err(SolveError::InvalidArgument("need at least two refinements"));
}
if errors.iter().chain(hs.iter()).any(|v| !v.is_finite() || *v <= 0.0) {
return Err(SolveError::InvalidArgument("errors and spacings must be positive"));
}
let n = errors.len() as f64;
let lx: Vec<f64> = hs.iter().map(|h| h.ln()).collect();
let ly: Vec<f64> = errors.iter().map(|e| e.ln()).collect();
let mx = lx.iter().sum::<f64>() / n;
let my = ly.iter().sum::<f64>() / n;
let sxx: f64 = lx.iter().map(|x| (x - mx) * (x - mx)).sum();
let sxy: f64 = lx.iter().zip(ly.iter()).map(|(x, y)| (x - mx) * (y - my)).sum();
if sxx <= 0.0 {
return Err(SolveError::InvalidArgument("need at least two distinct spacings"));
}
Ok(sxy / sxx)
}
#[cfg(test)]
mod tests {
use super::*;
const PI: f64 = std::f64::consts::PI;
fn wrap(a: f64, b: f64, degree: usize, v: Vec<f64>) -> Fem1dSolution {
Fem1dSolution::new(a, b, degree, v).unwrap()
}
#[test]
fn linear_elements_are_nodally_exact_for_poisson() {
let u = |x: f64| 2.0 * x - x * x - x * x * x;
for n in [3, 7, 40] {
let v = fem_1d_poisson(
&|x: f64| 2.0 + 6.0 * x,
0.0,
1.0,
(Bc::Dirichlet(0.0), Bc::Dirichlet(0.0)),
n,
)
.unwrap();
for (i, got) in v.iter().enumerate() {
let x = i as f64 / n as f64;
assert!(
(got - u(x)).abs() < 1e-14,
"node {i} of {n} was off by {}",
got - u(x)
);
}
}
}
#[test]
fn nodal_exactness_degrades_only_by_the_load_quadrature() {
let mut errors = Vec::new();
let hs = [1.0 / 3.0, 1.0 / 4.0, 1.0 / 5.0];
for n in [3usize, 4, 5] {
let v = fem_1d_poisson(
&|x: f64| PI * PI * (PI * x).sin(),
0.0,
1.0,
(Bc::Dirichlet(0.0), Bc::Dirichlet(0.0)),
n,
)
.unwrap();
let worst = v
.iter()
.enumerate()
.map(|(i, got)| {
let x = i as f64 / n as f64;
(got - (PI * x).sin()).abs()
})
.fold(0.0, f64::max);
assert!(worst < 1e-10, "{n} elements were off by {worst}");
errors.push(worst);
}
let rate = convergence_rate(&errors, &hs).unwrap();
assert!(rate > 8.0, "nodal error fell off only as h^{rate}, not as the quadrature does");
}
#[test]
fn the_patch_test_passes_at_both_degrees() {
let linear = fem_1d_poisson(&|_| 0.0, 0.0, 2.0, (Bc::Dirichlet(2.0), Bc::Dirichlet(8.0)), 5)
.unwrap();
for (i, got) in linear.iter().enumerate() {
let x = 2.0 * i as f64 / 5.0;
assert!((got - (2.0 + 3.0 * x)).abs() < 1e-12);
}
let quad = fem_1d_quadratic(
&|_| 1.0,
&|_| 0.0,
&|_| -2.0,
0.0,
1.0,
(Bc::Dirichlet(0.0), Bc::Dirichlet(1.0)),
4,
)
.unwrap();
for (i, got) in quad.iter().enumerate() {
let x = i as f64 / 8.0;
assert!((got - x * x).abs() < 1e-12, "midside {i} was off by {}", got - x * x);
}
}
#[test]
fn a_flux_condition_is_imposed_with_the_outward_normal() {
let right = fem_1d_poisson(&|_| 0.0, 0.0, 1.0, (Bc::Dirichlet(0.0), Bc::Neumann(3.0)), 6)
.unwrap();
assert!((right[6] - 3.0).abs() < 1e-12, "got {}", right[6]);
let left = fem_1d_poisson(&|_| 0.0, 0.0, 1.0, (Bc::Neumann(3.0), Bc::Dirichlet(0.0)), 6)
.unwrap();
assert!((left[0] - 3.0).abs() < 1e-12, "got {}", left[0]);
}
#[test]
fn a_pure_flux_problem_is_reported_as_singular() {
let e = fem_1d_poisson(&|_| 1.0, 0.0, 1.0, (Bc::Neumann(0.5), Bc::Neumann(-0.5)), 8);
assert_eq!(e, Err(SolveError::Singular));
assert!(fem_1d_general(
&|_| 1.0,
&|_| 1.0,
&|_| 1.0,
0.0,
1.0,
(Bc::Neumann(0.5), Bc::Neumann(-0.5)),
8
)
.is_ok());
assert!(fem_1d_poisson(
&|_| 1.0,
0.0,
1.0,
(Bc::Neumann(0.0), Bc::Robin { alpha: 2.0, g: 1.0 }),
8
)
.is_ok());
}
#[test]
fn a_robin_end_reproduces_its_own_algebra() {
for alpha in [0.25, 1.0, 40.0] {
let g = 2.0;
let v = fem_1d_poisson(
&|_| 0.0,
0.0,
1.0,
(Bc::Dirichlet(0.0), Bc::Robin { alpha, g }),
5,
)
.unwrap();
let expect = g / (1.0 + alpha);
assert!((v[5] - expect).abs() < 1e-12, "alpha {alpha}: got {} want {expect}", v[5]);
}
}
#[test]
fn a_variable_coefficient_converges_at_second_order() {
let p = |x: f64| 1.0 + x;
let f = |x: f64| -PI * (PI * x).cos() + (1.0 + x) * PI * PI * (PI * x).sin();
let u = |x: f64| (PI * x).sin();
let mut errors = Vec::new();
let mut hs = Vec::new();
for n in [10, 20, 40, 80] {
let v = fem_1d_general(
&p,
&|_| 0.0,
&f,
0.0,
1.0,
(Bc::Dirichlet(0.0), Bc::Dirichlet(0.0)),
n,
)
.unwrap();
errors.push(fem_1d_error_l2(&wrap(0.0, 1.0, 1, v), &u));
hs.push(1.0 / n as f64);
}
let rate = convergence_rate(&errors, &hs).unwrap();
assert!((rate - 2.0).abs() < 0.05, "L2 rate was {rate}");
}
#[test]
fn quadratic_elements_are_exact_at_vertices_but_not_at_midsides() {
let n = 8;
let v = fem_1d_quadratic(
&|_| 1.0,
&|_| 0.0,
&|x: f64| PI * PI * (PI * x).sin(),
0.0,
1.0,
(Bc::Dirichlet(0.0), Bc::Dirichlet(0.0)),
n,
)
.unwrap();
let err = |i: usize| {
let x = i as f64 / (2 * n) as f64;
(v[i] - (PI * x).sin()).abs()
};
let vertex = (0..=2 * n).step_by(2).map(err).fold(0.0, f64::max);
let midside = (1..2 * n).step_by(2).map(err).fold(0.0, f64::max);
assert!(vertex < 1e-13, "vertices were off by {vertex}");
assert!(midside > 1e3 * vertex, "midsides were as exact as the vertices");
}
#[test]
fn quadratic_elements_converge_one_order_faster() {
let u = |x: f64| (PI * x).sin();
let du = |x: f64| PI * (PI * x).cos();
let f = |x: f64| PI * PI * (PI * x).sin();
let (mut l2, mut h1, mut hs) = (Vec::new(), Vec::new(), Vec::new());
for n in [4, 8, 16, 32] {
let v = fem_1d_quadratic(
&|_| 1.0,
&|_| 0.0,
&f,
0.0,
1.0,
(Bc::Dirichlet(0.0), Bc::Dirichlet(0.0)),
n,
)
.unwrap();
let s = wrap(0.0, 1.0, 2, v);
l2.push(fem_1d_error_l2(&s, &u));
h1.push(fem_1d_error_h1_seminorm(&s, &du));
hs.push(1.0 / n as f64);
}
let rl2 = convergence_rate(&l2, &hs).unwrap();
let rh1 = convergence_rate(&h1, &hs).unwrap();
assert!((rl2 - 3.0).abs() < 0.05, "P2 L2 rate was {rl2}");
assert!((rh1 - 2.0).abs() < 0.05, "P2 H1 rate was {rh1}");
}
#[test]
fn the_convergence_rate_recovers_an_exact_power_law() {
let hs: Vec<f64> = (1..6).map(|k| 0.5f64.powi(k)).collect();
for k in [1.0, 2.0, 3.5] {
let errors: Vec<f64> = hs.iter().map(|h| 7.0 * h.powf(k)).collect();
let got = convergence_rate(&errors, &hs).unwrap();
assert!((got - k).abs() < 1e-10, "wanted {k}, got {got}");
}
assert!(convergence_rate(&[1.0], &[1.0]).is_err());
assert!(convergence_rate(&[1.0, 2.0], &[1.0]).is_err());
assert!(convergence_rate(&[1.0, 0.0], &[1.0, 0.5]).is_err());
assert!(convergence_rate(&[1.0, 2.0], &[0.5, 0.5]).is_err());
}
#[test]
fn the_solution_wrapper_interpolates_and_differentiates() {
let s = wrap(0.0, 1.0, 1, vec![0.0, 1.0, 4.0]);
assert!((s.eval(0.25) - 0.5).abs() < 1e-14);
assert!((s.eval_derivative(0.75) - 6.0).abs() < 1e-14);
assert_eq!(s.elements(), 2);
assert!((s.h() - 0.5).abs() < 1e-15);
assert_eq!(s.nodes().len(), 3);
let q = wrap(0.0, 1.0, 2, vec![0.0, 0.25, 1.0]);
assert!((q.eval(0.3) - 0.09).abs() < 1e-14);
assert!((q.eval_derivative(0.3) - 0.6).abs() < 1e-14);
assert!(Fem1dSolution::new(0.0, 1.0, 3, vec![0.0; 4]).is_err());
assert!(Fem1dSolution::new(1.0, 0.0, 1, vec![0.0; 4]).is_err());
assert!(Fem1dSolution::new(0.0, 1.0, 2, vec![0.0; 4]).is_err());
}
#[test]
fn bad_arguments_are_refused() {
let ok = (Bc::Dirichlet(0.0), Bc::Dirichlet(0.0));
assert!(fem_1d_poisson(&|_| 1.0, 0.0, 1.0, ok, 0).is_err());
assert!(fem_1d_poisson(&|_| 1.0, 1.0, 1.0, ok, 4).is_err());
assert!(fem_1d_poisson(&|_| f64::NAN, 0.0, 1.0, ok, 4).is_err());
assert!(fem_1d_general(&|_| -1.0, &|_| 0.0, &|_| 1.0, 0.0, 1.0, ok, 4).is_err());
assert!(fem_1d_poisson(&|_| 1.0, 0.0, 1.0, (Bc::Dirichlet(f64::NAN), ok.1), 4).is_err());
assert!(
fem_1d_poisson(&|_| 1.0, 0.0, 1.0, (Bc::Neumann(f64::INFINITY), ok.1), 4).is_err()
);
}
#[test]
fn a_reaction_term_is_assembled_with_the_right_sign() {
let n = 60;
let v = fem_1d_general(
&|_| 1.0,
&|_| 1.0,
&|_| 0.0,
0.0,
1.0,
(Bc::Dirichlet(1.0), Bc::Dirichlet(std::f64::consts::E)),
n,
)
.unwrap();
let s = wrap(0.0, 1.0, 1, v);
let err = fem_1d_error_l2(&s, &|x| x.exp());
assert!(err < 1e-4, "L2 error was {err}");
assert!((s.eval(0.5) - 0.5f64.exp()).abs() < 1e-4);
}
}