Skip to main content

solver/models/
control.rs

1//! Generic stochastic optimal control abstraction.
2//!
3//! This module defines a control-problem contract that is independent of any
4//! market-making interpretation. A model supplies the running reward, the
5//! infinitesimal generator, the terminal condition, and the control optimizer;
6//! numerical solvers consume only this trait.
7
8/// Derivatives of the value function at a state.
9///
10/// * `grad[i]` is the centered first derivative `dV/dx_i`.
11/// * `hessian[i]` is the diagonal second derivative `d^2V/dx_i^2`.
12/// * `hessian_full[i][j]` is the symmetric mixed second derivative
13///   `d^2V/dx_i dx_j`; its diagonal `hessian_full[i][i]` equals `hessian[i]`.
14/// * `fwd[i]` is the forward difference `(V(x + h_i e_i) - V(x)) / h_i`.
15/// * `bwd[i]` is the backward difference `(V(x) - V(x - h_i e_i)) / h_i`.
16///
17/// Continuous controls consume `grad`/`hessian` (and `hessian_full` when
18/// their noise is correlated across coordinates); jump controls (e.g. market
19/// making on a discrete inventory) consume `fwd`/`bwd`. The diagonal
20/// `hessian` and the full matrix are kept consistent: constructors that take
21/// only a diagonal fill `hessian_full` with zeros off the diagonal.
22#[derive(Clone, Copy, Debug)]
23pub struct StateDerivatives<const N: usize> {
24    pub grad: [f64; N],
25    pub hessian: [f64; N],
26    pub hessian_full: [[f64; N]; N],
27    pub fwd: [f64; N],
28    pub bwd: [f64; N],
29}
30
31impl<const N: usize> StateDerivatives<N> {
32    /// Construct a derivative bundle from first and diagonal second
33    /// derivatives, with zero directional differences and zero off-diagonal
34    /// Hessian.
35    ///
36    /// # Examples
37    ///
38    /// ```
39    /// use solver::models::control::StateDerivatives;
40    /// let d = StateDerivatives::new([0.1], [-0.02]);
41    /// assert_eq!(d.grad, [0.1]);
42    /// assert_eq!(d.hessian, [-0.02]);
43    /// assert_eq!(d.fwd, [0.0]);
44    /// assert_eq!(d.hessian_full, [[-0.02]]);
45    /// ```
46    pub fn new(grad: [f64; N], hessian: [f64; N]) -> Self {
47        Self::with_directional(grad, hessian, [0.0; N], [0.0; N])
48    }
49
50    /// Construct a derivative bundle from all four components, with zero
51    /// off-diagonal Hessian.
52    pub fn with_directional(
53        grad: [f64; N],
54        hessian: [f64; N],
55        fwd: [f64; N],
56        bwd: [f64; N],
57    ) -> Self {
58        let mut hessian_full = [[0.0; N]; N];
59        for i in 0..N {
60            hessian_full[i][i] = hessian[i];
61        }
62        Self {
63            grad,
64            hessian,
65            hessian_full,
66            fwd,
67            bwd,
68        }
69    }
70
71    /// Construct a derivative bundle with an explicit full symmetric Hessian.
72    ///
73    /// `hessian_full[i][j]` must equal `hessian_full[j][i]`; the diagonal is
74    /// taken from `hessian_full[i][i]` and the `hessian` diagonal is kept in
75    /// sync.
76    ///
77    /// # Examples
78    ///
79    /// ```
80    /// use solver::models::control::StateDerivatives;
81    /// let d = StateDerivatives::<2>::with_full_hessian(
82    ///     [1.0, 2.0],
83    ///     [[0.5, 0.25], [0.25, -0.5]],
84    ///     [3.0, 4.0],
85    ///     [5.0, 6.0],
86    /// );
87    /// assert_eq!(d.hessian, [0.5, -0.5]);
88    /// assert_eq!(d.hessian_full[0][1], 0.25);
89    /// assert_eq!(d.hessian_full[1][0], 0.25);
90    /// ```
91    pub fn with_full_hessian(
92        grad: [f64; N],
93        hessian_full: [[f64; N]; N],
94        fwd: [f64; N],
95        bwd: [f64; N],
96    ) -> Self {
97        let mut hessian = [0.0; N];
98        for i in 0..N {
99            hessian[i] = hessian_full[i][i];
100        }
101        Self {
102            grad,
103            hessian,
104            hessian_full,
105            fwd,
106            bwd,
107        }
108    }
109}
110
111/// A finite-horizon stochastic optimal control problem.
112///
113/// The generic formulation is
114///
115/// ```text
116/// state      x in R^n
117/// control    u in U
118/// dynamics   dx = b(t,x,u) dt + sigma(t,x) dW  (plus optional jumps)
119/// objective  J = E[ integral_t^T f(s,x,u) ds + g(x_T) ]
120/// value      V(t,x) = sup_u J
121/// ```
122///
123/// The associated HJB equation is
124///
125/// ```text
126/// 0 = d_t V + sup_u { f(t,x,u) + L^u V }
127/// ```
128///
129/// where `L^u V` is the infinitesimal generator.
130///
131/// # Contract
132///
133/// The running reward `f` and the generator `L^u V` are kept separate because
134/// different solvers consume them differently:
135///
136/// * [`ControlProblem::running_reward`] returns the running reward
137///   `f(t,x,u)`. This is the driver used by the regression-based BSDE solver:
138///   the generator is already accounted for by simulating the forward SDE
139///   under the control.
140/// * [`ControlProblem::generator`] returns the infinitesimal generator
141///   `L^u V = b(t,x,u) dot grad V + 0.5 tr(sigma sigma' hess V)` evaluated at
142///   a fixed control and derivative bundle.
143/// * [`ControlProblem::driver`] is the sum `f + L^u V`, the full HJB driver
144///   consumed by finite-difference (grid) solvers. Its default implementation
145///   is `running_reward + generator` and should not be overridden.
146/// * [`ControlProblem::optimize`] returns the control that maximizes the
147///   driver for the given state and derivatives.
148///
149/// Discounting is handled separately by [`ControlProblem::discount_rate`] and
150/// must not be folded into any of these hooks.
151///
152/// # Examples
153///
154/// ```
155/// use solver::models::control::{ControlProblem, StateDerivatives};
156///
157/// /// Constant-control problem with zero reward and zero terminal value.
158/// struct Zero;
159/// impl ControlProblem<1> for Zero {
160///     type Control = f64;
161///     fn optimize(&self, _t: f64, _s: &[f64; 1], _d: &StateDerivatives<1>) -> f64 { 0.0 }
162///     fn running_reward(&self, _t: f64, _s: &[f64; 1], _c: &f64) -> f64 { 0.0 }
163///     fn generator(&self, _t: f64, _s: &[f64; 1], _c: &f64, _d: &StateDerivatives<1>) -> f64 { 0.0 }
164///     fn terminal(&self, _s: &[f64; 1]) -> f64 { 0.0 }
165///     fn next_step(&self, _t: f64, s: &[f64; 1], _dt: f64, _n: &[f64; 1]) -> [f64; 1] { *s }
166/// }
167/// ```
168pub trait ControlProblem<const N: usize> {
169    /// The action type. For Merton this is a scalar portfolio fraction; for
170    /// market making it is a pair of bid/ask intensities.
171    type Control;
172
173    /// Returns the control that maximizes the driver at `(t, state)`.
174    fn optimize(&self, t: f64, state: &[f64; N], derivs: &StateDerivatives<N>) -> Self::Control;
175
176    /// Returns the running reward `f(t,x,u)`.
177    fn running_reward(&self, t: f64, state: &[f64; N], control: &Self::Control) -> f64;
178
179    /// Returns the infinitesimal generator `L^u V` for the given control and
180    /// derivative bundle.
181    fn generator(
182        &self,
183        t: f64,
184        state: &[f64; N],
185        control: &Self::Control,
186        derivs: &StateDerivatives<N>,
187    ) -> f64;
188
189    /// Returns the full HJB driver `f(t,x,u) + L^u V`.
190    ///
191    /// The default implementation is `running_reward + generator`. Grid
192    /// solvers consume this; the regression-based BSDE solver consumes
193    /// [`ControlProblem::bsde_driver`] instead.
194    fn driver(
195        &self,
196        t: f64,
197        state: &[f64; N],
198        control: &Self::Control,
199        derivs: &StateDerivatives<N>,
200    ) -> f64 {
201        self.running_reward(t, state, control) + self.generator(t, state, control, derivs)
202    }
203
204    /// Returns the backward driver consumed by the BSDE regression solver.
205    ///
206    /// For a full (non-reduced) control problem this is the running reward
207    /// `f` alone; the infinitesimal generator is already accounted for by
208    /// simulating the forward SDE under the control. Reduced value problems
209    /// (for example the CARA market-making `theta` ansatz) override this to
210    /// include the local source terms that are not part of the forward
211    /// transport.
212    ///
213    /// `dt` is the forward time step. It lets the driver bound the running
214    /// reward at `1 / dt` per fill side, which is the largest rate the forward
215    /// Euler step can represent before its Bernoulli fill probability
216    /// `lambda * dt` saturates at `1`. Without this bound the backward reward
217    /// can count fills the forward step cannot produce, creating a
218    /// forward/backward inconsistency that destabilizes the regression at
219    /// large base intensities.
220    fn bsde_driver(
221        &self,
222        t: f64,
223        state: &[f64; N],
224        control: &Self::Control,
225        _derivs: &StateDerivatives<N>,
226        _dt: f64,
227    ) -> f64 {
228        self.running_reward(t, state, control)
229    }
230
231    /// Terminal value `g(x)` at the horizon.
232    fn terminal(&self, state: &[f64; N]) -> f64;
233
234    /// Optional pointwise constraint on the value, e.g. the early-exercise
235    /// obstacle `V >= payoff` for an American option.
236    ///
237    /// Defaults to the identity. Finite-difference solvers apply this after
238    /// each backward step; the BSDE path does not yet enforce it.
239    fn apply_constraint(&self, _state: &[f64; N], value: f64) -> f64 {
240        value
241    }
242
243    /// Discount rate `r(t,x)`. Defaults to zero.
244    fn discount_rate(&self, _state: &[f64; N]) -> f64 {
245        0.0
246    }
247
248    /// Constant discount rate hint. Defaults to `None` (state dependent).
249    fn constant_discount_rate(&self) -> Option<f64> {
250        None
251    }
252
253    /// Advances the state one step under the optimal control at forward time
254    /// `t`.
255    fn next_step(&self, t: f64, state: &[f64; N], dt: f64, noise: &[f64; N]) -> [f64; N];
256
257    /// Advances the state one step under an explicitly supplied control.
258    ///
259    /// The default ignores `control` and delegates to
260    /// [`ControlProblem::next_step`]. Problems whose forward dynamics depend
261    /// on the control (for example market-making models whose fill intensities
262    /// determine the inventory jump, or price-impact models) override this to
263    /// use `control` instead of a frozen proxy, so the coupled (Picard) BSDE
264    /// forward pass simulates the state under the current optimal control.
265    fn next_step_controlled(
266        &self,
267        t: f64,
268        state: &[f64; N],
269        _control: &Self::Control,
270        dt: f64,
271        noise: &[f64; N],
272    ) -> [f64; N] {
273        self.next_step(t, state, dt, noise)
274    }
275
276    /// Whether the value this problem solves is a reduced value.
277    ///
278    /// A reduced value problem models a value function that has been collapsed
279    /// onto a subset of the state dimensions (for example the CARA market-making
280    /// ansatz reduces the full `(S, q, X)` value to the inventory-only function
281    /// `theta(t, q)`). For such problems the forward pass must not simulate the
282    /// jumps of the collapsed dimensions: the jumps are already absorbed into the
283    /// local source carried by [`ControlProblem::bsde_driver`], so simulating them
284    /// forward would regress the jump-convolved continuation instead of the reduced
285    /// value at the same collapsed state.
286    ///
287    /// Defaults to `false` (a full, non-reduced value problem).
288    fn is_reduced_value(&self) -> bool {
289        false
290    }
291
292    /// Whether a dimension is driven by Brownian diffusion.
293    fn is_diffusion_dimension(&self, _dim: usize) -> bool {
294        false
295    }
296
297    /// Physical finite-difference step for a dimension, used by mesh-free
298    /// gradient stencils.
299    fn gradient_step(&self, _dim: usize) -> f64 {
300        1.0
301    }
302}
303
304#[cfg(test)]
305mod tests {
306    use super::*;
307
308    struct FullValue;
309
310    impl ControlProblem<1> for FullValue {
311        type Control = f64;
312        fn optimize(&self, _t: f64, _s: &[f64; 1], _d: &StateDerivatives<1>) -> f64 {
313            0.0
314        }
315        fn running_reward(&self, _t: f64, _s: &[f64; 1], _c: &f64) -> f64 {
316            0.0
317        }
318        fn generator(&self, _t: f64, _s: &[f64; 1], _c: &f64, _d: &StateDerivatives<1>) -> f64 {
319            0.0
320        }
321        fn terminal(&self, _s: &[f64; 1]) -> f64 {
322            0.0
323        }
324        fn next_step(&self, _t: f64, s: &[f64; 1], _dt: f64, _n: &[f64; 1]) -> [f64; 1] {
325            *s
326        }
327    }
328
329    struct ReducedValue;
330
331    impl ControlProblem<1> for ReducedValue {
332        type Control = f64;
333        fn optimize(&self, _t: f64, _s: &[f64; 1], _d: &StateDerivatives<1>) -> f64 {
334            0.0
335        }
336        fn running_reward(&self, _t: f64, _s: &[f64; 1], _c: &f64) -> f64 {
337            0.0
338        }
339        fn generator(&self, _t: f64, _s: &[f64; 1], _c: &f64, _d: &StateDerivatives<1>) -> f64 {
340            0.0
341        }
342        fn terminal(&self, _s: &[f64; 1]) -> f64 {
343            0.0
344        }
345        fn next_step(&self, _t: f64, s: &[f64; 1], _dt: f64, _n: &[f64; 1]) -> [f64; 1] {
346            *s
347        }
348        fn is_reduced_value(&self) -> bool {
349            true
350        }
351    }
352
353    #[test]
354    fn full_value_defaults_to_not_reduced() {
355        assert!(!ControlProblem::is_reduced_value(&FullValue));
356    }
357
358    #[test]
359    fn reduced_value_overrides_to_true() {
360        assert!(ControlProblem::is_reduced_value(&ReducedValue));
361    }
362
363    #[test]
364    fn with_full_hessian_syncs_diagonal_and_is_symmetric() {
365        let d = StateDerivatives::<2>::with_full_hessian(
366            [1.0, 2.0],
367            [[0.5, 0.25], [0.25, -0.5]],
368            [3.0, 4.0],
369            [5.0, 6.0],
370        );
371        assert_eq!(d.hessian, [0.5, -0.5]);
372        assert_eq!(d.hessian_full[0][1], 0.25);
373        assert_eq!(d.hessian_full[1][0], 0.25);
374        assert_eq!(d.grad, [1.0, 2.0]);
375        assert_eq!(d.fwd, [3.0, 4.0]);
376        assert_eq!(d.bwd, [5.0, 6.0]);
377    }
378
379    #[test]
380    fn default_constructors_have_zero_off_diagonal() {
381        let d = StateDerivatives::new([0.1, 0.2], [0.3, 0.4]);
382        assert_eq!(d.hessian_full[0][1], 0.0);
383        assert_eq!(d.hessian_full[1][0], 0.0);
384        assert_eq!(d.hessian_full[0][0], 0.3);
385        assert_eq!(d.hessian_full[1][1], 0.4);
386    }
387}