solver/numeric/bsde/
solution.rs1use super::config::BasisFunctionType;
2use crate::models::control::{ControlProblem, StateDerivatives};
3use crate::models::market_making::MarketMakingControl;
4use crate::numeric::basis::{Basis, PolynomialBasis, ScaledBasis};
5
6pub struct BsdeSolution<const N: usize, C> {
11 pub coefficients: Vec<Vec<f64>>,
14 pub basis_type: BasisFunctionType,
16 pub dt: f64,
18 pub steps: usize,
20 pub x_init: [f64; N],
22 pub feature_scale: Vec<f64>,
24 pub clamp_range: Option<(f64, f64)>,
26 pub gradient_steps: [f64; N],
28 pub v0: f64,
30 pub control_t0: C,
32}
33
34impl<const N: usize, C> BsdeSolution<N, C> {
35 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 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 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 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 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
112pub(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 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 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 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 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}