solver/models/
kelly_hjb.rs1use super::control::{ControlProblem, StateDerivatives};
28use super::market_making::MarketMakingControl;
29
30#[derive(Clone)]
31pub struct KellyHjb {
32 pub sigma: f64,
34 pub kappa: f64,
36 pub a: f64,
38 pub q_min: f64,
40 pub q_max: f64,
42 pub dq: f64,
44}
45
46impl KellyHjb {
47 pub fn new(sigma: f64, kappa: f64, a: f64) -> Self {
48 Self {
49 sigma,
50 kappa,
51 a,
52 q_min: -10.0,
53 q_max: 10.0,
54 dq: 1.0,
55 }
56 }
57
58 pub fn with_inventory_bounds(mut self, q_min: f64, q_max: f64) -> Self {
59 self.q_min = q_min;
60 self.q_max = q_max;
61 self
62 }
63
64 pub fn get_spreads(&self, derivs: &StateDerivatives<2>) -> (f64, f64) {
70 let base = 1.0 / self.kappa;
71
72 let dv_buy = derivs.fwd[0];
73 let dv_sell = derivs.bwd[0];
74
75 let delta_bid = base - dv_buy;
76 let delta_ask = base + dv_sell;
77
78 let min_spread = -5.0;
79 let max_spread = 10.0;
80 (
81 delta_bid.max(min_spread).min(max_spread),
82 delta_ask.max(min_spread).min(max_spread),
83 )
84 }
85
86 pub fn fill_rate_base(&self, _state: &[f64; 2]) -> f64 {
88 self.a
89 }
90
91 pub fn fill_rate_decay(&self) -> f64 {
93 self.kappa
94 }
95}
96
97impl ControlProblem<2> for KellyHjb {
98 type Control = MarketMakingControl;
99
100 fn optimize(&self, _t: f64, state: &[f64; 2], derivs: &StateDerivatives<2>) -> Self::Control {
101 let q = state[0];
102 let (delta_bid, delta_ask) = self.get_spreads(derivs);
103
104 let lambda_bid = if q >= self.q_max {
105 0.0
106 } else {
107 self.a * (-self.kappa * delta_bid).exp()
108 };
109 let lambda_ask = if q <= self.q_min {
110 0.0
111 } else {
112 self.a * (-self.kappa * delta_ask).exp()
113 };
114
115 MarketMakingControl::new(lambda_bid, lambda_ask)
116 }
117
118 fn running_reward(&self, _t: f64, _state: &[f64; 2], control: &Self::Control) -> f64 {
119 (control.bid_intensity + control.ask_intensity) / self.kappa
121 }
122
123 fn bsde_driver(
124 &self,
125 t: f64,
126 state: &[f64; 2],
127 control: &Self::Control,
128 derivs: &StateDerivatives<2>,
129 _dt: f64,
130 ) -> f64 {
131 self.driver(t, state, control, derivs)
132 }
133
134 fn generator(
135 &self,
136 _t: f64,
137 state: &[f64; 2],
138 _control: &Self::Control,
139 derivs: &StateDerivatives<2>,
140 ) -> f64 {
141 let x = state[1];
145 let sigma2 = self.sigma.powi(2);
146
147 0.5 * sigma2 * x.powi(2) * derivs.hessian[1] + sigma2 * x * derivs.grad[1] - 0.5 * sigma2
148 }
149
150 fn terminal(&self, state: &[f64; 2]) -> f64 {
151 let x = state[1].max(1e-12);
152 let q = state[0];
153 (x + q.max(-x + 1e-12)).ln()
154 }
155
156 fn discount_rate(&self, _state: &[f64; 2]) -> f64 {
157 0.0
158 }
159
160 fn constant_discount_rate(&self) -> Option<f64> {
161 Some(0.0)
162 }
163
164 fn next_step(&self, _t: f64, state: &[f64; 2], dt: f64, noise: &[f64; 2]) -> [f64; 2] {
165 let q = state[0];
166 let x = state[1].max(1e-8);
167 let sqrt_dt = dt.sqrt();
168
169 let sigma2 = self.sigma.powi(2);
172 let dx_drift = x * sigma2 * dt;
173 let dx_diff = -x * self.sigma * sqrt_dt * noise[1];
174 let x_next = (x + dx_drift + dx_diff).max(1e-8);
175
176 [q, x_next]
177 }
178
179 fn is_reduced_value(&self) -> bool {
180 true
181 }
182
183 fn is_diffusion_dimension(&self, dim: usize) -> bool {
184 dim == 1
185 }
186
187 fn gradient_step(&self, dim: usize) -> f64 {
188 if dim == 0 {
189 self.dq.abs().max(1e-8)
190 } else {
191 1.0
192 }
193 }
194}
195
196#[cfg(test)]
197mod tests {
198 use super::*;
199
200 fn zero_derivs() -> StateDerivatives<2> {
201 StateDerivatives::new([0.0; 2], [0.0; 2])
202 }
203
204 fn default_model() -> KellyHjb {
205 KellyHjb {
206 sigma: 0.3,
207 kappa: 1.5,
208 a: 140.0,
209 q_min: -10.0,
210 q_max: 10.0,
211 dq: 1.0,
212 }
213 }
214
215 #[test]
216 fn zero_gradient_gives_base_spread() {
217 let m = default_model();
218 let (bid, ask) = m.get_spreads(&zero_derivs());
219 let base = 1.0 / m.kappa;
220 assert!((bid - base).abs() < 1e-12);
221 assert!((ask - base).abs() < 1e-12);
222 }
223
224 #[test]
225 fn base_spread_matches_risk_neutral_as() {
226 let m = default_model();
227 let base = 1.0 / m.kappa;
228 let as_base = {
229 let gamma = 1e-9;
230 (1.0 + gamma / m.kappa).ln() / gamma
231 };
232 assert!((base - as_base).abs() < 1e-6);
233 }
234
235 #[test]
236 fn symmetric_at_zero_inventory() {
237 let m = default_model();
238 let ctrl = ControlProblem::optimize(&m, 0.0, &[0.0, 5.0], &zero_derivs());
239 assert!(
240 (ctrl.bid_intensity - ctrl.ask_intensity).abs() < 1e-12,
241 "bid/ask intensities should be equal at q=0"
242 );
243 }
244
245 #[test]
246 fn positive_reward_at_zero_inventory() {
247 let m = default_model();
248 let ctrl = ControlProblem::optimize(&m, 0.0, &[0.0, 5.0], &zero_derivs());
249 let reward = ControlProblem::running_reward(&m, 0.0, &[0.0, 5.0], &ctrl);
250 assert!(
251 reward > 0.0,
252 "reward should be positive at q=0: hamiltonian > source"
253 );
254 }
255
256 #[test]
257 fn terminal_gives_log_x_plus_q() {
258 let m = default_model();
259 let v = ControlProblem::terminal(&m, &[2.0, 10.0]);
260 assert!((v - 12.0_f64.ln()).abs() < 1e-10);
261 }
262
263 #[test]
264 fn terminal_handles_negative_total_clamp() {
265 let m = default_model();
266 let v = ControlProblem::terminal(&m, &[-4.0, 2.0]);
268 assert!(v.is_finite());
269 assert!(v < 0.0);
270 }
271
272 #[test]
273 fn dims_tagged_correctly() {
274 let m = default_model();
275 assert!(!ControlProblem::is_diffusion_dimension(&m, 0));
276 assert!(ControlProblem::is_diffusion_dimension(&m, 1));
277 }
278
279 #[test]
280 fn fill_rate_methods_return_expected() {
281 let m = default_model();
282 assert!((m.fill_rate_base(&[0.0, 5.0]) - 140.0).abs() < 1e-12);
283 assert!((m.fill_rate_decay() - 1.5).abs() < 1e-12);
284 }
285}