Skip to main content

solver/numeric/finite_difference/elliptic/
solver.rs

1//! Stationary solver for elliptic finite-difference problems.
2//!
3//! The solver assembles the sparse matrix for
4//!
5//! ```text
6//! -T u + reaction * u = f
7//! ```
8//!
9//! where `T` is the transport operator built from
10//! [`crate::numeric::finite_difference::discretization::Transport`]. The
11//! negative sign follows the standard elliptic convention: for a diffusion
12//! coefficient `D`, `T u = D u''` and the problem reads `-D u'' + rho u = f`.
13//!
14//! Boundary conditions are enforced with ghost points. Dirichlet conditions
15//! replace the boundary row with the identity; Neumann and Robin conditions
16//! fold the missing neighbour coefficient into the existing interior
17//! neighbour and add a right-hand-side correction.
18
19use super::problem::{EllipticControlProblem, EllipticProblem};
20use crate::core::grid::Grid;
21use crate::linalg::csr::CsrMatrix;
22use crate::models::control::ControlProblem;
23use crate::numeric::finite_difference::discretization::{
24    BoundaryCondition, BoundaryConditions, Transport,
25};
26use crate::numeric::finite_difference::operator::Operator;
27
28/// Configuration for the stationary solver.
29#[derive(Clone, Copy, Debug)]
30pub struct StationarySolver {
31    /// Successive over-relaxation parameter.
32    pub sor_omega: f64,
33    /// SOR convergence tolerance.
34    pub tol: f64,
35    /// Maximum SOR iterations.
36    pub max_iter: usize,
37    /// Maximum policy-iteration sweeps.
38    pub max_policy_iter: usize,
39    /// Policy-iteration convergence tolerance on the value-function change.
40    pub policy_tol: f64,
41}
42
43impl Default for StationarySolver {
44    fn default() -> Self {
45        Self {
46            sor_omega: 1.2,
47            tol: 1e-8,
48            max_iter: 20_000,
49            max_policy_iter: 100,
50            policy_tol: 1e-8,
51        }
52    }
53}
54
55impl StationarySolver {
56    /// Creates a new stationary solver with default settings.
57    ///
58    /// # Examples
59    ///
60    /// ```
61    /// use solver::numeric::finite_difference::elliptic::StationarySolver;
62    /// let solver = StationarySolver::new();
63    /// assert_eq!(solver.sor_omega, 1.2);
64    /// ```
65    pub fn new() -> Self {
66        Self::default()
67    }
68
69    /// Selects the SOR relaxation parameter.
70    ///
71    /// # Examples
72    ///
73    /// ```
74    /// use solver::numeric::finite_difference::elliptic::StationarySolver;
75    /// let solver = StationarySolver::new().with_sor_omega(1.5);
76    /// assert_eq!(solver.sor_omega, 1.5);
77    /// ```
78    pub fn with_sor_omega(mut self, omega: f64) -> Self {
79        self.sor_omega = omega;
80        self
81    }
82
83    /// Selects the SOR tolerance.
84    ///
85    /// # Examples
86    ///
87    /// ```
88    /// use solver::numeric::finite_difference::elliptic::StationarySolver;
89    /// let solver = StationarySolver::new().with_tol(1e-6);
90    /// assert_eq!(solver.tol, 1e-6);
91    /// ```
92    pub fn with_tol(mut self, tol: f64) -> Self {
93        self.tol = tol;
94        self
95    }
96
97    /// Selects the maximum number of policy-iteration sweeps.
98    ///
99    /// # Examples
100    ///
101    /// ```
102    /// use solver::numeric::finite_difference::elliptic::StationarySolver;
103    /// let solver = StationarySolver::new().with_max_policy_iter(50);
104    /// assert_eq!(solver.max_policy_iter, 50);
105    /// ```
106    pub fn with_max_policy_iter(mut self, max_policy_iter: usize) -> Self {
107        self.max_policy_iter = max_policy_iter;
108        self
109    }
110
111    /// Selects the policy-iteration convergence tolerance.
112    ///
113    /// # Examples
114    ///
115    /// ```
116    /// use solver::numeric::finite_difference::elliptic::StationarySolver;
117    /// let solver = StationarySolver::new().with_policy_tol(1e-6);
118    /// assert_eq!(solver.policy_tol, 1e-6);
119    /// ```
120    pub fn with_policy_tol(mut self, policy_tol: f64) -> Self {
121        self.policy_tol = policy_tol;
122        self
123    }
124
125    /// Solves `-T u + reaction * u = f` on `grid`.
126    ///
127    /// # Examples
128    ///
129    /// ```
130    /// use solver::core::grid::Grid;
131    /// use solver::models::control::StateDerivatives;
132    /// use solver::numeric::finite_difference::discretization::{
133    ///     BoundaryCondition, BoundaryConditions, DimensionKind, Transport,
134    /// };
135    /// use solver::numeric::finite_difference::elliptic::{
136    ///     EllipticProblem, StationarySolver,
137    /// };
138    ///
139    /// struct Poisson { dx: f64 }
140    /// impl EllipticProblem<1> for Poisson {
141    ///     fn dimension_kind(&self, _dim: usize) -> DimensionKind {
142    ///         DimensionKind::Diffusion
143    ///     }
144    ///     fn transport(&self, _s: &[f64; 1], _d: &StateDerivatives<1>) -> Transport<1> {
145    ///         let c = 1.0 / (self.dx * self.dx);
146    ///         Transport::new([c], [c], 0.0)
147    ///     }
148    ///     fn rhs(&self, _s: &[f64; 1]) -> f64 { 0.0 }
149    ///     fn boundary_conditions(&self) -> BoundaryConditions<1> {
150    ///         BoundaryConditions::new(
151    ///             [BoundaryCondition::Dirichlet(0.0)],
152    ///             [BoundaryCondition::Dirichlet(1.0)],
153    ///         )
154    ///     }
155    /// }
156    ///
157    /// let grid = Grid::<1>::new([11], [0.0], [1.0]);
158    /// let u = StationarySolver::new().solve(&grid, &Poisson { dx: grid.dx[0] });
159    /// assert!((u[0] - 0.0).abs() < 1e-12);
160    /// assert!((u[10] - 1.0).abs() < 1e-12);
161    /// ```
162    pub fn solve<const N: usize, P: EllipticProblem<N>>(
163        &self,
164        grid: &Grid<N>,
165        problem: &P,
166    ) -> Vec<f64> {
167        let size = grid.total_size();
168        let mut mat = Operator::initialize_matrix(grid);
169        let derivs = Operator::compute_state_derivatives(grid, &vec![0.0; size]);
170
171        let transports: Vec<_> = (0..size)
172            .map(|i| problem.transport(&grid_state(grid, i), &derivs[i]))
173            .collect();
174        let reactions: Vec<_> = (0..size)
175            .map(|i| problem.reaction(&grid_state(grid, i)))
176            .collect();
177        let contribs = boundary_contribs(grid, &transports, &problem.boundary_conditions());
178
179        assemble_matrix(&mut mat, grid, &reactions, &contribs);
180
181        let mut b = vec![0.0; size];
182        assemble_rhs(&mut b, grid, problem, &contribs);
183
184        if N == 1 {
185            return mat.solve_tridiagonal_system(&b);
186        }
187
188        let mut x = vec![0.0; size];
189        mat.solve_sor(&b, &mut x, self.tol, self.max_iter, self.sor_omega);
190        x
191    }
192
193    /// Solves a stationary stochastic optimal control problem by policy
194    /// iteration.
195    ///
196    /// The problem is `0 = sup_u { f + L^u V - r V }`. At each sweep the solver
197    /// optimizes the control at the current derivative bundle, assembles
198    /// `-T V + r V = source`, solves the linear system, and repeats until the
199    /// value function stops changing.
200    ///
201    /// # Examples
202    ///
203    /// ```
204    /// use solver::core::grid::Grid;
205    /// use solver::models::control::{ControlProblem, StateDerivatives};
206    /// use solver::numeric::finite_difference::discretization::{
207    ///     BoundaryConditions, DimensionKind, Transport,
208    /// };
209    /// use solver::numeric::finite_difference::elliptic::{
210    ///     EllipticControlProblem, StationarySolver,
211    /// };
212    ///
213    /// struct Constant;
214    /// impl ControlProblem<1> for Constant {
215    ///     type Control = f64;
216    ///     fn optimize(&self, _t: f64, _s: &[f64; 1], _d: &StateDerivatives<1>) -> f64 { 0.0 }
217    ///     fn running_reward(&self, _t: f64, _s: &[f64; 1], _c: &f64) -> f64 { 0.0 }
218    ///     fn generator(&self, _t: f64, _s: &[f64; 1], _c: &f64, _d: &StateDerivatives<1>) -> f64 { 0.0 }
219    ///     fn terminal(&self, _s: &[f64; 1]) -> f64 { 0.0 }
220    ///     fn discount_rate(&self, _s: &[f64; 1]) -> f64 { 1.0 }
221    ///     fn next_step(&self, _t: f64, s: &[f64; 1], _dt: f64, _n: &[f64; 1]) -> [f64; 1] { *s }
222    /// }
223    /// impl EllipticControlProblem<1> for Constant {
224    ///     fn dimension_kind(&self, _dim: usize) -> DimensionKind { DimensionKind::Diffusion }
225    ///     fn transport(&self, _s: &[f64; 1], _c: &f64, _d: &StateDerivatives<1>) -> Transport<1> {
226    ///         Transport::new([0.0], [0.0], 0.0)
227    ///     }
228    /// }
229    ///
230    /// let grid = Grid::<1>::new([11], [0.0], [1.0]);
231    /// let v = StationarySolver::new().solve_control(&grid, &Constant);
232    /// assert!(v.iter().all(|x| x.abs() < 1e-8));
233    /// ```
234    pub fn solve_control<const N: usize, P: EllipticControlProblem<N> + Sync>(
235        &self,
236        grid: &Grid<N>,
237        problem: &P,
238    ) -> Vec<f64>
239    where
240        P::Control: Send + Sync,
241    {
242        let size = grid.total_size();
243        let mut v = vec![0.0; size];
244
245        for _ in 0..self.max_policy_iter {
246            let derivs = Operator::compute_state_derivatives(grid, &v);
247
248            let mut transports = Vec::with_capacity(size);
249            let mut reactions = Vec::with_capacity(size);
250            let mut sources = Vec::with_capacity(size);
251
252            for (i, deriv) in derivs.iter().enumerate() {
253                let state = grid_state(grid, i);
254                let control = ControlProblem::optimize(problem, 0.0, &state, deriv);
255                let transport = EllipticControlProblem::transport(problem, &state, &control, deriv);
256                transports.push(transport);
257                reactions.push(ControlProblem::discount_rate(problem, &state));
258                sources.push(transport.source);
259            }
260
261            let contribs = boundary_contribs(grid, &transports, &problem.boundary_conditions());
262            let mut mat = Operator::initialize_matrix(grid);
263            assemble_matrix(&mut mat, grid, &reactions, &contribs);
264
265            let mut b = sources.clone();
266            apply_boundary_source(&mut b, &contribs);
267
268            let next = if N == 1 {
269                mat.solve_tridiagonal_system(&b)
270            } else {
271                let mut x = vec![0.0; size];
272                mat.solve_sor(&b, &mut x, self.tol, self.max_iter, self.sor_omega);
273                x
274            };
275
276            let delta = next
277                .iter()
278                .zip(v.iter())
279                .map(|(a, b)| (a - b).abs())
280                .fold(0.0, f64::max);
281
282            v = next;
283            if delta < self.policy_tol {
284                break;
285            }
286        }
287
288        v
289    }
290}
291
292/// Per-node result of applying boundary conditions to the transport stencil.
293///
294/// `plus_coeff` and `minus_coeff` are the coefficients to use in the operator
295/// after any ghost-point folding. `dirichlet` pins the node value when set.
296/// `diag_extra` and `rhs_extra` collect the ghost-point corrections.
297#[derive(Clone, Copy, Debug)]
298struct BoundaryContrib<const N: usize> {
299    dirichlet: Option<f64>,
300    plus_coeff: [f64; N],
301    minus_coeff: [f64; N],
302    diag_extra: f64,
303    rhs_extra: f64,
304}
305
306/// Computes boundary contributions for every node.
307///
308/// For an interior coordinate the effective coefficients equal the transport
309/// coefficients. For a lower Neumann or Robin boundary, the missing backward
310/// coefficient is folded into the forward neighbour. For an upper Neumann or
311/// Robin boundary, the forward coefficient is folded into the backward
312/// neighbour. Dirichlet pins the node and wins over any other condition.
313fn boundary_contribs<const N: usize>(
314    grid: &Grid<N>,
315    transports: &[Transport<N>],
316    boundaries: &BoundaryConditions<N>,
317) -> Vec<BoundaryContrib<N>> {
318    let tensor = &grid.tensor_info;
319    let size = grid.total_size();
320
321    (0..size)
322        .map(|i| {
323            let coords = tensor.get_coords(i);
324            let t = transports[i];
325            let mut plus_coeff = [0.0; N];
326            let mut minus_coeff = [0.0; N];
327            let mut diag_extra = 0.0;
328            let mut rhs_extra = 0.0;
329            let mut dirichlet = None;
330
331            for (dim, coord) in coords.iter().enumerate() {
332                let shape = tensor.shape[dim];
333                let upper = *coord == shape - 1;
334                let lower = *coord == 0;
335                let h = grid.dx[dim];
336
337                if !upper {
338                    plus_coeff[dim] = t.plus[dim];
339                }
340                if !lower {
341                    minus_coeff[dim] = t.minus[dim];
342                }
343
344                let cond = if upper {
345                    Some(boundaries.upper[dim])
346                } else if lower {
347                    Some(boundaries.lower[dim])
348                } else {
349                    None
350                };
351
352                match cond {
353                    Some(BoundaryCondition::Dirichlet(value)) => {
354                        if dirichlet.is_none() {
355                            dirichlet = Some(value);
356                        }
357                    }
358                    Some(BoundaryCondition::Neumann(value)) => {
359                        if lower && shape > 1 {
360                            plus_coeff[dim] += t.minus[dim];
361                            rhs_extra -= 2.0 * h * value * t.minus[dim];
362                        } else if upper && shape > 1 {
363                            minus_coeff[dim] += t.plus[dim];
364                            rhs_extra += 2.0 * h * value * t.plus[dim];
365                        }
366                    }
367                    Some(BoundaryCondition::Robin { value, alpha, beta }) => {
368                        if beta == 0.0 {
369                            // Degenerate Robin; leave the natural stencil.
370                            continue;
371                        }
372                        if lower && shape > 1 {
373                            plus_coeff[dim] += t.minus[dim];
374                            diag_extra -= 2.0 * h * alpha / beta * t.minus[dim];
375                            rhs_extra -= 2.0 * h * value / beta * t.minus[dim];
376                        } else if upper && shape > 1 {
377                            minus_coeff[dim] += t.plus[dim];
378                            diag_extra += 2.0 * h * alpha / beta * t.plus[dim];
379                            rhs_extra += 2.0 * h * value / beta * t.plus[dim];
380                        }
381                    }
382                    None => {}
383                }
384            }
385
386            BoundaryContrib {
387                dirichlet,
388                plus_coeff,
389                minus_coeff,
390                diag_extra,
391                rhs_extra,
392            }
393        })
394        .collect()
395}
396
397/// Assembles `-T + reaction * I` using the precomputed boundary contributions.
398fn assemble_matrix<const N: usize>(
399    mat: &mut CsrMatrix,
400    grid: &Grid<N>,
401    reactions: &[f64],
402    contribs: &[BoundaryContrib<N>],
403) {
404    let tensor = &grid.tensor_info;
405    let size = grid.total_size();
406    let row_ptr = &mat.row_ptr;
407    let values_ptr = mat.values.as_mut_ptr() as usize;
408
409    for i in 0..size {
410        let coords = tensor.get_coords(i);
411        let c = &contribs[i];
412
413        let start = row_ptr[i];
414        let end = row_ptr[i + 1];
415        // SAFETY: rows are disjoint by construction.
416        let row_values = unsafe {
417            std::slice::from_raw_parts_mut((values_ptr as *mut f64).add(start), end - start)
418        };
419
420        let mut diag = reactions[i] + c.diag_extra;
421        let mut idx = 0;
422
423        for (dim, &coord) in coords.iter().enumerate() {
424            let upper = coord == tensor.shape[dim] - 1;
425            let lower = coord == 0;
426
427            if !upper {
428                row_values[idx] = -c.plus_coeff[dim];
429                diag += c.plus_coeff[dim];
430                idx += 1;
431            }
432            if !lower {
433                row_values[idx] = -c.minus_coeff[dim];
434                diag += c.minus_coeff[dim];
435                idx += 1;
436            }
437        }
438
439        if c.dirichlet.is_some() {
440            // Replace the Dirichlet row by the identity. `idx` points at the
441            // diagonal slot because `Operator::initialize_matrix` appends the
442            // diagonal after all forward and backward off-diagonal entries.
443            for v in row_values.iter_mut() {
444                *v = 0.0;
445            }
446            row_values[idx] = 1.0;
447        } else {
448            row_values[idx] = diag;
449        }
450    }
451}
452
453/// Builds the right-hand side from the forcing, ghost-point corrections, and
454/// Dirichlet values.
455fn assemble_rhs<const N: usize, P: EllipticProblem<N>>(
456    b: &mut [f64],
457    grid: &Grid<N>,
458    problem: &P,
459    contribs: &[BoundaryContrib<N>],
460) {
461    for (i, item) in b.iter_mut().enumerate() {
462        let state = grid_state(grid, i);
463        let c = &contribs[i];
464        *item = match c.dirichlet {
465            Some(value) => value,
466            None => problem.rhs(&state) + c.rhs_extra,
467        };
468    }
469}
470
471/// Applies boundary contributions to an already-computed source vector.
472///
473/// Used by the control path, where the interior source comes from
474/// [`Transport::source`] rather than [`EllipticProblem::rhs`].
475fn apply_boundary_source<const N: usize>(b: &mut [f64], contribs: &[BoundaryContrib<N>]) {
476    for (item, c) in b.iter_mut().zip(contribs) {
477        *item = match c.dirichlet {
478            Some(value) => value,
479            None => *item + c.rhs_extra,
480        };
481    }
482}
483
484fn grid_state<const N: usize>(grid: &Grid<N>, index: usize) -> [f64; N] {
485    let coords = grid.tensor_info.get_coords(index);
486    let mut state = [0.0; N];
487    for (d, val) in state.iter_mut().enumerate() {
488        *val = grid.min[d] + (coords[d] as f64) * grid.dx[d];
489    }
490    state
491}
492
493#[cfg(test)]
494mod tests {
495    use super::*;
496    use crate::models::control::StateDerivatives;
497    use crate::numeric::finite_difference::discretization::DimensionKind;
498
499    /// One-dimensional diffusion problem `-u'' = rhs` on [0, 1].
500    struct Diffusion1D {
501        dx: f64,
502        rhs_value: f64,
503        lower: BoundaryCondition,
504        upper: BoundaryCondition,
505    }
506
507    impl EllipticProblem<1> for Diffusion1D {
508        fn dimension_kind(&self, _dim: usize) -> DimensionKind {
509            DimensionKind::Diffusion
510        }
511
512        fn transport(&self, _state: &[f64; 1], _derivs: &StateDerivatives<1>) -> Transport<1> {
513            let c = 1.0 / (self.dx * self.dx);
514            Transport::new([c], [c], 0.0)
515        }
516
517        fn rhs(&self, _state: &[f64; 1]) -> f64 {
518            self.rhs_value
519        }
520
521        fn boundary_conditions(&self) -> BoundaryConditions<1> {
522            BoundaryConditions::new([self.lower], [self.upper])
523        }
524    }
525
526    #[test]
527    fn neumann_zero_boundary_matches_constant_solution() {
528        // `-u'' = 0`, u'(0) = 0, u(1) = 1.
529        // Exact: u(x) = 1.
530        let grid = Grid::<1>::new([101], [0.0], [1.0]);
531        let problem = Diffusion1D {
532            dx: grid.dx[0],
533            rhs_value: 0.0,
534            lower: BoundaryCondition::Neumann(0.0),
535            upper: BoundaryCondition::Dirichlet(1.0),
536        };
537        let u = StationarySolver::new().solve(&grid, &problem);
538        for (i, value) in u.iter().enumerate() {
539            let x = grid.min[0] + (i as f64) * grid.dx[0];
540            assert!((value - 1.0).abs() < 1e-4, "x={x} numerical={value}");
541        }
542    }
543
544    #[test]
545    fn nonzero_neumann_matches_exact_linear_solution() {
546        // `-u'' = 1`, u'(0) = 1, u(1) = 0.
547        // Exact: u(x) = x - x^2/2 - 0.5.
548        let grid = Grid::<1>::new([101], [0.0], [1.0]);
549        let problem = Diffusion1D {
550            dx: grid.dx[0],
551            rhs_value: 1.0,
552            lower: BoundaryCondition::Neumann(1.0),
553            upper: BoundaryCondition::Dirichlet(0.0),
554        };
555        let u = StationarySolver::new().solve(&grid, &problem);
556
557        for (i, value) in u.iter().enumerate() {
558            let x = grid.min[0] + (i as f64) * grid.dx[0];
559            let exact = x - 0.5 * x * x - 0.5;
560            assert!(
561                (value - exact).abs() < 1e-4,
562                "x={x} numerical={value} exact={exact}"
563            );
564        }
565    }
566
567    #[test]
568    fn robin_boundary_matches_exact_linear_solution() {
569        // `-u'' = 0`, u(0) - u'(0) = 0, u(1) = 1.
570        // Exact: u(x) = (x + 1) / 2.
571        let grid = Grid::<1>::new([101], [0.0], [1.0]);
572        let problem = Diffusion1D {
573            dx: grid.dx[0],
574            rhs_value: 0.0,
575            lower: BoundaryCondition::Robin {
576                value: 0.0,
577                alpha: 1.0,
578                beta: -1.0,
579            },
580            upper: BoundaryCondition::Dirichlet(1.0),
581        };
582        let u = StationarySolver::new().solve(&grid, &problem);
583
584        for (i, value) in u.iter().enumerate() {
585            let x = grid.min[0] + (i as f64) * grid.dx[0];
586            let exact = (x + 1.0) / 2.0;
587            assert!(
588                (value - exact).abs() < 1e-4,
589                "x={x} numerical={value} exact={exact}"
590            );
591        }
592    }
593}