solver/numeric/finite_difference/
operator.rs1use crate::core::grid::Grid;
2use crate::linalg::csr::CsrMatrix;
3use crate::models::traits::{ControlOutput, Gradients, Model};
4use rayon::prelude::*;
5
6pub struct Operator;
7
8impl Operator {
9 pub fn compute_gradients<const N: usize>(grid: &Grid<N>, v: &[f64]) -> Vec<Gradients<N>> {
10 let size = v.len();
11 let tensor = &grid.tensor_info;
12
13 (0..size)
14 .into_par_iter()
15 .map(|i| {
16 let coords = tensor.get_coords(i);
17 let mut fwd = [0.0; N];
18 let mut bwd = [0.0; N];
19
20 for (dim, &stride) in tensor.strides.iter().enumerate() {
21 fwd[dim] = if grid.dx[dim] == 0.0 {
22 0.0
23 } else if coords[dim] < tensor.shape[dim] - 1 {
24 (v[i + stride] - v[i]) / grid.dx[dim]
25 } else if coords[dim] > 0 {
26 (v[i] - v[i - stride]) / grid.dx[dim]
27 } else {
28 0.0
29 };
30
31 bwd[dim] = if grid.dx[dim] == 0.0 {
32 0.0
33 } else if coords[dim] > 0 {
34 (v[i] - v[i - stride]) / grid.dx[dim]
35 } else if coords[dim] < tensor.shape[dim] - 1 {
36 (v[i + stride] - v[i]) / grid.dx[dim]
37 } else {
38 0.0
39 };
40 }
41
42 Gradients { fwd, bwd }
43 })
44 .collect()
45 }
46
47 pub fn initialize_matrix<const N: usize>(grid: &Grid<N>) -> CsrMatrix {
48 let size = grid.total_size();
49 let tensor = &grid.tensor_info;
50
51 let est_entries = size * (2 * N + 1);
52 let mut mat = CsrMatrix::new(size, est_entries);
53
54 for i in 0..size {
55 let coords = tensor.get_coords(i);
56
57 for (dim, &stride) in tensor.strides.iter().enumerate() {
58 if coords[dim] < tensor.shape[dim] - 1 {
59 mat.add_entry(i + stride, 0.0);
60 }
61
62 if coords[dim] > 0 {
63 mat.add_entry(i - stride, 0.0);
64 }
65 }
66
67 mat.add_entry(i, 1.0);
68 mat.finish_row();
69 }
70 mat
71 }
72
73 pub fn update_matrix<const N: usize, M: Model<N> + Sync>(
78 mat: &mut CsrMatrix,
79 dt: f64,
80 grid: &Grid<N>,
81 controls: &[ControlOutput<N>],
82 model: &M,
83 ) {
84 let size = grid.total_size();
85 let tensor = &grid.tensor_info;
86 let row_ptr = &mat.row_ptr;
87
88 let values_ptr = mat.values.as_mut_ptr() as usize;
89
90 let constant_r = model.constant_discount_rate();
91
92 (0..size).into_par_iter().for_each(|i| {
93 let coords = tensor.get_coords(i);
94
95 let r = if let Some(cr) = constant_r {
96 cr
97 } else {
98 let mut real_state = [0.0; N];
99 for d in 0..N {
100 real_state[d] = grid.min[d] + (coords[d] as f64) * grid.dx[d];
101 }
102 model.discount_rate(&real_state)
103 };
104
105 let ctrl = &controls[i];
106
107 let start_idx = row_ptr[i];
108 let end_idx = row_ptr[i + 1];
109
110 let row_values = unsafe {
112 std::slice::from_raw_parts_mut(
113 (values_ptr as *mut f64).add(start_idx),
114 end_idx - start_idx,
115 )
116 };
117
118 let mut val_idx = 0;
119 let mut diag_val = 1.0;
120
121 for (dim, coord) in coords.iter().enumerate() {
122 if *coord < tensor.shape[dim] - 1 {
123 let lambda = ctrl.lambda_plus[dim];
124 let coeff = dt * lambda;
125 row_values[val_idx] = -coeff;
126 diag_val += coeff;
127 val_idx += 1;
128 }
129
130 if *coord > 0 {
131 let lambda = ctrl.lambda_minus[dim];
132 let coeff = dt * lambda;
133 row_values[val_idx] = -coeff;
134 diag_val += coeff;
135 val_idx += 1;
136 }
137 }
138
139 diag_val += dt * r;
140 row_values[val_idx] = diag_val;
141 });
142 }
143}