Skip to main content

solver/numeric/bsde/
solution.rs

1use super::config::BasisFunctionType;
2use crate::models::control::{ControlProblem, StateDerivatives};
3use crate::models::market_making::MarketMakingControl;
4use crate::numeric::basis::{Basis, PolynomialBasis, ScaledBasis};
5
6/// Complete BSDE solution storing regression coefficients for all time steps.
7///
8/// Enables point evaluation of the value function and optimal controls at any
9/// `(time_index, state)` pair without re-running the solver.
10pub struct BsdeSolution<const N: usize, C> {
11    /// Regression coefficients at each backward time step.
12    /// `coefficients[t]` contains the basis weights at time index `t`.
13    pub coefficients: Vec<Vec<f64>>,
14    /// Basis function configuration.
15    pub basis_type: BasisFunctionType,
16    /// Time step size.
17    pub dt: f64,
18    /// Number of backward time steps.
19    pub steps: usize,
20    /// Initial state used for centering.
21    pub x_init: [f64; N],
22    /// Feature scaling per dimension.
23    pub feature_scale: Vec<f64>,
24    /// Clamping range for basis inputs.
25    pub clamp_range: Option<(f64, f64)>,
26    /// Gradient step sizes per dimension.
27    pub gradient_steps: [f64; N],
28    /// Value function at t=0 and x_init.
29    pub v0: f64,
30    /// Optimal control at t=0 and x_init.
31    pub control_t0: C,
32}
33
34impl<const N: usize, C> BsdeSolution<N, C> {
35    /// Reconstruct the scaled basis used during solving.
36    fn make_basis(&self) -> ScaledBasis<N, PolynomialBasis> {
37        let mut inner = self.basis_type.to_basis::<N>();
38        if let Some((min, max)) = self.clamp_range {
39            inner.clamp_range = Some((min, max));
40        }
41        ScaledBasis::new(inner, self.feature_scale.clone()).with_center(self.x_init)
42    }
43
44    /// Evaluate the value function at a given time index and state.
45    ///
46    /// `time_index` ranges from 0 (start) to `self.steps` (terminal).
47    pub fn evaluate(&self, time_index: usize, state: &[f64; N]) -> f64 {
48        if time_index >= self.coefficients.len() || self.coefficients[time_index].is_empty() {
49            return 0.0;
50        }
51        let basis = self.make_basis();
52        basis.eval_dot(state, &self.coefficients[time_index])
53    }
54
55    /// Convert a time value to the nearest time index.
56    ///
57    /// `t` is time elapsed from 0 (start of horizon).
58    pub fn time_to_index(&self, t: f64) -> usize {
59        let idx = (t / self.dt).round() as usize;
60        idx.min(self.steps.saturating_sub(1))
61    }
62}
63
64impl<const N: usize> BsdeSolution<N, MarketMakingControl> {
65    /// Evaluate the optimal market-making control at a time index and state.
66    ///
67    /// Returns `None` when no regression coefficients are available for the
68    /// requested time index (the terminal step).
69    pub fn evaluate_control<M: ControlProblem<N, Control = MarketMakingControl>>(
70        &self,
71        time_index: usize,
72        state: &[f64; N],
73        problem: &M,
74    ) -> Option<MarketMakingControl> {
75        if time_index >= self.coefficients.len() || self.coefficients[time_index].is_empty() {
76            return None;
77        }
78        let basis = self.make_basis();
79        let coeffs = &self.coefficients[time_index];
80        let derivs = state_derivatives(state, coeffs, &basis, &self.gradient_steps);
81        Some(problem.optimize(self.dt * time_index as f64, state, &derivs))
82    }
83
84    /// Evaluate optimal spreads across all time steps for a given state.
85    ///
86    /// Returns vectors of `(bid_spread, ask_spread)` for each time step from
87    /// `0` to `steps - 1`. Time steps without coefficients yield `NaN`.
88    pub fn evaluate_spread_trajectory<M: ControlProblem<N, Control = MarketMakingControl>>(
89        &self,
90        state: &[f64; N],
91        problem: &M,
92        a: f64,
93        kappa: f64,
94    ) -> Vec<(f64, f64)> {
95        let basis = self.make_basis();
96        let mut result = Vec::with_capacity(self.steps);
97        for t in 0..self.steps {
98            if self.coefficients[t].is_empty() {
99                result.push((f64::NAN, f64::NAN));
100                continue;
101            }
102            let coeffs = &self.coefficients[t];
103            let derivs = state_derivatives(state, coeffs, &basis, &self.gradient_steps);
104            let ctrl = problem.optimize(self.dt * t as f64, state, &derivs);
105            let spreads = ctrl.to_spreads(a, kappa);
106            result.push((spreads.bid_spread, spreads.ask_spread));
107        }
108        result
109    }
110}
111
112/// Finite-difference derivatives of the regression model at a state.
113pub(crate) fn state_derivatives<const N: usize>(
114    state: &[f64; N],
115    coeffs: &[f64],
116    basis: &impl Basis<N>,
117    gradient_steps: &[f64; N],
118) -> StateDerivatives<N> {
119    let v_curr = basis.eval_dot(state, coeffs);
120    let mut grad = [0.0; N];
121    let mut hessian = [0.0; N];
122    let mut hessian_full = [[0.0; N]; N];
123    let mut fwd = [0.0; N];
124    let mut bwd = [0.0; N];
125
126    // Diagonal entries and first derivatives.
127    for i in 0..N {
128        let h = gradient_steps[i];
129        let h_inv = 1.0 / h;
130
131        let mut plus = *state;
132        plus[i] += h;
133        let v_plus = basis.eval_dot(&plus, coeffs);
134
135        let mut minus = *state;
136        minus[i] -= h;
137        let v_minus = basis.eval_dot(&minus, coeffs);
138
139        grad[i] = (v_plus - v_minus) * 0.5 * h_inv;
140        hessian[i] = (v_plus - 2.0 * v_curr + v_minus) * h_inv * h_inv;
141        fwd[i] = (v_plus - v_curr) * h_inv;
142        bwd[i] = (v_curr - v_minus) * h_inv;
143        hessian_full[i][i] = hessian[i];
144    }
145
146    // Mixed second derivatives via the centered four-corner stencil
147    // d^2V/dx_i dx_j = (V_{++} - V_{+-} - V_{-+} + V_{--}) / (4 h_i h_j).
148    for i in 0..N {
149        for j in (i + 1)..N {
150            let hi = gradient_steps[i];
151            let hj = gradient_steps[j];
152            let denom = 4.0 * hi * hj;
153
154            let mut pp = *state;
155            pp[i] += hi;
156            pp[j] += hj;
157
158            let mut pm = *state;
159            pm[i] += hi;
160            pm[j] -= hj;
161
162            let mut mp = *state;
163            mp[i] -= hi;
164            mp[j] += hj;
165
166            let mut mm = *state;
167            mm[i] -= hi;
168            mm[j] -= hj;
169
170            let mixed = (basis.eval_dot(&pp, coeffs)
171                - basis.eval_dot(&pm, coeffs)
172                - basis.eval_dot(&mp, coeffs)
173                + basis.eval_dot(&mm, coeffs))
174                / denom;
175            hessian_full[i][j] = mixed;
176            hessian_full[j][i] = mixed;
177        }
178    }
179
180    StateDerivatives::with_full_hessian(grad, hessian_full, fwd, bwd)
181}
182
183#[cfg(test)]
184mod tests {
185    use super::*;
186    use crate::numeric::basis::Basis;
187
188    /// Basis with a single feature phi(x) = x0 * x1, so the represented value
189    /// is V(x) = x0 * x1 when the weight is 1.
190    struct BilinearBasis;
191
192    impl Basis<2> for BilinearBasis {
193        fn num_features(&self) -> usize {
194            1
195        }
196
197        fn evaluate(&self, state: &[f64; 2], buffer: &mut [f64]) {
198            buffer[0] = state[0] * state[1];
199        }
200    }
201
202    #[test]
203    fn mixed_second_derivative_matches_known_quadratic() {
204        // V(x0, x1) = x0 * x1 has exact mixed derivative d^2V/dx0 dx1 = 1 and
205        // zero diagonal second derivatives. The centered four-corner stencil is
206        // exact for this bilinear function, isolating the mixed-derivative
207        // assembly rather than the finite-difference truncation error.
208        let basis = BilinearBasis;
209        let state = [0.3, -0.2];
210        let steps = [1e-3, 1e-3];
211
212        let d = state_derivatives(&state, &[1.0], &basis, &steps);
213
214        assert!(
215            (d.hessian_full[0][1] - 1.0).abs() < 1e-9,
216            "mixed derivative {}, expected 1.0",
217            d.hessian_full[0][1]
218        );
219        assert!(
220            (d.hessian_full[1][0] - 1.0).abs() < 1e-9,
221            "mixed derivative must be symmetric"
222        );
223        assert!((d.hessian[0]).abs() < 1e-9);
224        assert!((d.hessian[1]).abs() < 1e-9);
225    }
226}