Skip to main content

solver/numeric/finite_difference/
operator.rs

1use crate::core::grid::Grid;
2use crate::linalg::csr::CsrMatrix;
3use crate::models::control::StateDerivatives;
4use crate::numeric::finite_difference::discretization::Transport;
5use rayon::prelude::*;
6
7/// Sparse finite-difference operator assembly.
8pub struct Operator;
9
10impl Operator {
11    /// Computes first and diagonal second derivatives of `v` on `grid`.
12    ///
13    /// Uses central second differences on interior nodes and one-sided first
14    /// differences on boundaries. This is the derivative bundle consumed by
15    /// [`crate::models::control::ControlProblem::driver`].
16    pub fn compute_state_derivatives<const N: usize>(
17        grid: &Grid<N>,
18        v: &[f64],
19    ) -> Vec<StateDerivatives<N>> {
20        let size = v.len();
21        let tensor = &grid.tensor_info;
22
23        (0..size)
24            .into_par_iter()
25            .map(|i| {
26                let coords = tensor.get_coords(i);
27                let mut grad = [0.0; N];
28                let mut hessian = [0.0; N];
29                let mut hessian_full = [[0.0; N]; N];
30                let mut fwd = [0.0; N];
31                let mut bwd = [0.0; N];
32
33                for (dim, &stride) in tensor.strides.iter().enumerate() {
34                    let dx = grid.dx[dim];
35                    if dx == 0.0 {
36                        grad[dim] = 0.0;
37                        hessian[dim] = 0.0;
38                        fwd[dim] = 0.0;
39                        bwd[dim] = 0.0;
40                        hessian_full[dim][dim] = 0.0;
41                        continue;
42                    }
43
44                    let has_plus = coords[dim] < tensor.shape[dim] - 1;
45                    let has_minus = coords[dim] > 0;
46
47                    match (has_plus, has_minus) {
48                        (true, true) => {
49                            let v_plus = v[i + stride];
50                            let v_minus = v[i - stride];
51                            let v_curr = v[i];
52                            grad[dim] = (v_plus - v_minus) / (2.0 * dx);
53                            hessian[dim] = (v_plus - 2.0 * v_curr + v_minus) / (dx * dx);
54                            fwd[dim] = (v_plus - v_curr) / dx;
55                            bwd[dim] = (v_curr - v_minus) / dx;
56                        }
57                        (true, false) => {
58                            let v_curr = v[i];
59                            grad[dim] = (v[i + stride] - v_curr) / dx;
60                            hessian[dim] = 0.0;
61                            fwd[dim] = (v[i + stride] - v_curr) / dx;
62                            bwd[dim] = 0.0;
63                        }
64                        (false, true) => {
65                            let v_curr = v[i];
66                            grad[dim] = (v_curr - v[i - stride]) / dx;
67                            hessian[dim] = 0.0;
68                            fwd[dim] = 0.0;
69                            bwd[dim] = (v_curr - v[i - stride]) / dx;
70                        }
71                        (false, false) => {
72                            grad[dim] = 0.0;
73                            hessian[dim] = 0.0;
74                            fwd[dim] = 0.0;
75                            bwd[dim] = 0.0;
76                        }
77                    }
78                    hessian_full[dim][dim] = hessian[dim];
79                }
80
81                // Mixed second derivatives via the centered four-corner
82                // stencil, only where both dimensions have both neighbours.
83                for dim_i in 0..N {
84                    for dim_j in (dim_i + 1)..N {
85                        let dxi = grid.dx[dim_i];
86                        let dxj = grid.dx[dim_j];
87                        if dxi == 0.0 || dxj == 0.0 {
88                            continue;
89                        }
90                        let has_plus_i = coords[dim_i] < tensor.shape[dim_i] - 1;
91                        let has_minus_i = coords[dim_i] > 0;
92                        let has_plus_j = coords[dim_j] < tensor.shape[dim_j] - 1;
93                        let has_minus_j = coords[dim_j] > 0;
94                        if !(has_plus_i && has_minus_i && has_plus_j && has_minus_j) {
95                            continue;
96                        }
97                        let si = tensor.strides[dim_i];
98                        let sj = tensor.strides[dim_j];
99                        let mixed = (v[i + si + sj] - v[i + si - sj] - v[i - si + sj]
100                            + v[i - si - sj])
101                            / (4.0 * dxi * dxj);
102                        hessian_full[dim_i][dim_j] = mixed;
103                        hessian_full[dim_j][dim_i] = mixed;
104                    }
105                }
106
107                StateDerivatives::with_full_hessian(grad, hessian_full, fwd, bwd)
108            })
109            .collect()
110    }
111
112    /// Creates a CSR matrix with the sparsity pattern of the FD transport
113    /// operator on `grid`.
114    ///
115    /// Each row has up to one forward and one backward entry per dimension,
116    /// plus the diagonal. Entries are initialized to zero except the diagonal,
117    /// which starts at one.
118    pub fn initialize_matrix<const N: usize>(grid: &Grid<N>) -> CsrMatrix {
119        let size = grid.total_size();
120        let tensor = &grid.tensor_info;
121
122        let est_entries = size * (2 * N + 1);
123        let mut mat = CsrMatrix::new(size, est_entries);
124
125        for i in 0..size {
126            let coords = tensor.get_coords(i);
127
128            for (dim, &stride) in tensor.strides.iter().enumerate() {
129                if coords[dim] < tensor.shape[dim] - 1 {
130                    mat.add_entry(i + stride, 0.0);
131                }
132
133                if coords[dim] > 0 {
134                    mat.add_entry(i - stride, 0.0);
135                }
136            }
137
138            mat.add_entry(i, 1.0);
139            mat.finish_row();
140        }
141        mat
142    }
143
144    /// Updates an existing CSR matrix to `I - dt * transport`.
145    ///
146    /// The forward and backward entries are populated with `-dt * plus[dim]`
147    /// and `-dt * minus[dim]`, and the diagonal is `1 + dt * (sum of transport
148    /// rates) + dt * discount_rate`.
149    pub fn update_matrix<const N: usize>(
150        mat: &mut CsrMatrix,
151        dt: f64,
152        grid: &Grid<N>,
153        transports: &[Transport<N>],
154        discount_rates: &[f64],
155    ) {
156        let size = grid.total_size();
157        let tensor = &grid.tensor_info;
158        let row_ptr = &mat.row_ptr;
159        let values_ptr = mat.values.as_mut_ptr() as usize;
160
161        (0..size).into_par_iter().for_each(|i| {
162            let coords = tensor.get_coords(i);
163            let transport = transports[i];
164            let r = discount_rates[i];
165
166            let start_idx = row_ptr[i];
167            let end_idx = row_ptr[i + 1];
168
169            // SAFETY: row_ptr guarantees [start_idx, end_idx) ranges are
170            // disjoint for different `i`.
171            let row_values = unsafe {
172                std::slice::from_raw_parts_mut(
173                    (values_ptr as *mut f64).add(start_idx),
174                    end_idx - start_idx,
175                )
176            };
177
178            let mut val_idx = 0;
179            let mut diag_val = 1.0 + dt * r;
180
181            for (dim, coord) in coords.iter().enumerate() {
182                if *coord < tensor.shape[dim] - 1 {
183                    let coeff = dt * transport.plus[dim];
184                    row_values[val_idx] = -coeff;
185                    diag_val += coeff;
186                    val_idx += 1;
187                }
188
189                if *coord > 0 {
190                    let coeff = dt * transport.minus[dim];
191                    row_values[val_idx] = -coeff;
192                    diag_val += coeff;
193                    val_idx += 1;
194                }
195            }
196
197            row_values[val_idx] = diag_val;
198        });
199    }
200}
201
202#[cfg(test)]
203mod tests {
204    use super::*;
205
206    #[test]
207    fn mixed_second_derivative_matches_known_quadratic() {
208        // V(x0, x1) = x0 * x1 has exact mixed derivative d^2V/dx0 dx1 = 1 and
209        // zero diagonal second derivatives. The centered four-corner stencil is
210        // exact for this bilinear function, so this isolates the mixed-derivative
211        // assembly rather than the finite-difference truncation error.
212        let shape = [5, 5];
213        let min = [-1.0, -1.0];
214        let max = [1.0, 1.0];
215        let grid = Grid::<2>::new(shape, min, max);
216
217        let mut v = vec![0.0; grid.total_size()];
218        for (i, val) in v.iter_mut().enumerate() {
219            let coords = grid.tensor_info.get_coords(i);
220            let x0 = min[0] + coords[0] as f64 * grid.dx[0];
221            let x1 = min[1] + coords[1] as f64 * grid.dx[1];
222            *val = x0 * x1;
223        }
224
225        let derivs = Operator::compute_state_derivatives(&grid, &v);
226
227        // Interior node (2, 2) -> state (0.0, 0.0).
228        let center = grid.tensor_info.get_idx(&[2, 2]);
229        let d = derivs[center];
230        assert!(
231            (d.hessian_full[0][1] - 1.0).abs() < 1e-12,
232            "mixed derivative {}, expected 1.0",
233            d.hessian_full[0][1]
234        );
235        assert!(
236            (d.hessian_full[1][0] - 1.0).abs() < 1e-12,
237            "mixed derivative must be symmetric"
238        );
239        assert!((d.hessian[0]).abs() < 1e-12);
240        assert!((d.hessian[1]).abs() < 1e-12);
241    }
242}