solver/numeric/finite_difference/
operator.rs1use 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
7pub struct Operator;
9
10impl Operator {
11 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 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 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 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 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 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 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}