Skip to main content

solver/numeric/finite_difference/
operator.rs

1use 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    /// Updates the values of the existing CSR Matrix.
74    ///
75    /// Uses unsafe pointer arithmetic to allow parallel updates to the values vector.
76    /// SAFETY: row_ptr guarantees disjoint memory regions for each row.
77    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            // SAFETY: row_ptr guarantees [start_idx, end_idx) ranges are disjoint for different i.
111            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}