solver/models/
kelly_hjb.rs1use super::traits::{ControlOutput, Gradients, Model};
28
29#[derive(Clone)]
30pub struct KellyHjb {
31 pub sigma: f64,
32 pub kappa: f64,
33 pub a: f64,
34 pub q_min: f64,
35 pub q_max: f64,
36 pub dq: f64,
37}
38
39impl KellyHjb {
40 pub fn new(sigma: f64, kappa: f64, a: f64) -> Self {
41 Self {
42 sigma,
43 kappa,
44 a,
45 q_min: -10.0,
46 q_max: 10.0,
47 dq: 1.0,
48 }
49 }
50
51 pub fn with_inventory_bounds(mut self, q_min: f64, q_max: f64) -> Self {
52 self.q_min = q_min;
53 self.q_max = q_max;
54 self
55 }
56
57 pub fn get_spreads(&self, grads: &Gradients<2>) -> (f64, f64) {
63 let base = 1.0 / self.kappa;
64
65 let dv_buy = grads.fwd[0]; let dv_sell = grads.bwd[0]; let delta_bid = base - dv_buy;
69 let delta_ask = base + dv_sell;
70
71 let min_spread = -5.0;
72 let max_spread = 10.0;
73 (
74 delta_bid.max(min_spread).min(max_spread),
75 delta_ask.max(min_spread).min(max_spread),
76 )
77 }
78}
79
80impl Model<2> for KellyHjb {
81 type Process = ();
82
83 fn process(&self) {}
84
85 fn optimize(&self, state: &[f64; 2], grads: &Gradients<2>) -> ControlOutput<2> {
86 let q = state[0];
87 let x = state[1];
88
89 let (delta_bid, delta_ask) = self.get_spreads(grads);
91
92 let lambda_bid = self.a * (-self.kappa * delta_bid).exp();
93 let lambda_ask = self.a * (-self.kappa * delta_ask).exp();
94
95 let lambda_bid = if q >= self.q_max { 0.0 } else { lambda_bid };
96 let lambda_ask = if q <= self.q_min { 0.0 } else { lambda_ask };
97
98 let hamiltonian = (lambda_bid + lambda_ask) / self.kappa;
101
102 let jump_correction = self.dq * (lambda_bid * grads.fwd[0] - lambda_ask * grads.bwd[0]);
104
105 let sigma2 = self.sigma.powi(2);
110 let diff_coeff = 0.5 * sigma2 * x.powi(2);
111 let drift_x = sigma2 * x;
112
113 let dx: f64 = 1.0; let diff_term = diff_coeff / dx.powi(2);
115 let drift_term_abs = drift_x.abs() / dx;
116
117 let mut lambda_x_plus = diff_term;
118 let mut lambda_x_minus = diff_term;
119
120 if drift_x > 0.0 {
121 lambda_x_plus += drift_term_abs;
122 } else {
123 lambda_x_minus += drift_term_abs;
124 }
125
126 let const_source = -0.5 * sigma2;
128
129 ControlOutput {
130 lambda_plus: [lambda_bid, lambda_x_plus],
131 lambda_minus: [lambda_ask, lambda_x_minus],
132 flow: hamiltonian + const_source - jump_correction,
133 }
134 }
135
136 fn terminal(&self, state: &[f64; 2]) -> f64 {
137 let x = state[1].max(1e-12);
138 let q = state[0];
139 (x + q.max(-x + 1e-12)).ln()
140 }
141
142 fn constant_discount_rate(&self) -> Option<f64> {
143 Some(0.0)
144 }
145
146 fn next_step(&self, current_state: &[f64; 2], dt: f64, noise: &[f64; 2]) -> [f64; 2] {
147 let q = current_state[0];
148 let x = current_state[1].max(1e-8);
149 let sqrt_dt = dt.sqrt();
150
151 let sigma2 = self.sigma.powi(2);
154 let dx_drift = x * sigma2 * dt;
155 let dx_diff = -x * self.sigma * sqrt_dt * noise[1];
156 let x_next = (x + dx_drift + dx_diff).max(1e-8);
157
158 [q, x_next]
159 }
160
161 fn is_diffusion_dimension(&self, dim: usize) -> bool {
162 dim == 1 }
164
165 fn is_integer_dimension(&self, dim: usize) -> bool {
166 dim == 0 }
168
169 fn fill_rate_base(&self, _state: &[f64; 2]) -> f64 {
170 self.a
171 }
172
173 fn fill_rate_decay(&self) -> f64 {
174 self.kappa
175 }
176}
177
178#[cfg(test)]
179mod tests {
180 use super::super::traits::{Gradients, Model};
181 use super::*;
182
183 fn zero_grads() -> Gradients<2> {
184 Gradients {
185 fwd: [0.0; 2],
186 bwd: [0.0; 2],
187 }
188 }
189
190 fn default_model() -> KellyHjb {
191 KellyHjb {
192 sigma: 0.3,
193 kappa: 1.5,
194 a: 140.0,
195 q_min: -10.0,
196 q_max: 10.0,
197 dq: 1.0,
198 }
199 }
200
201 #[test]
202 fn zero_gradient_gives_base_spread() {
203 let m = default_model();
204 let (bid, ask) = m.get_spreads(&zero_grads());
205 let base = 1.0 / m.kappa;
206 assert!((bid - base).abs() < 1e-12);
207 assert!((ask - base).abs() < 1e-12);
208 }
209
210 #[test]
211 fn base_spread_matches_risk_neutral_as() {
212 let m = default_model();
213 let base = 1.0 / m.kappa;
214 let as_base = {
215 let gamma = 1e-9;
216 (1.0 + gamma / m.kappa).ln() / gamma
217 };
218 assert!((base - as_base).abs() < 1e-6);
219 }
220
221 #[test]
222 fn symmetric_at_zero_inventory() {
223 let m = default_model();
224 let ctrl = m.optimize(&[0.0, 5.0], &zero_grads());
225 assert!(
226 (ctrl.lambda_plus[0] - ctrl.lambda_minus[0]).abs() < 1e-12,
227 "bid/ask intensities should be equal at q=0"
228 );
229 }
230
231 #[test]
232 fn positive_flow_at_zero_inventory() {
233 let m = default_model();
234 let ctrl = m.optimize(&[0.0, 5.0], &zero_grads());
235 assert!(
236 ctrl.flow > 0.0,
237 "flow should be positive at q=0: hamiltonian > source"
238 );
239 }
240
241 #[test]
242 fn terminal_gives_log_x_plus_q() {
243 let m = default_model();
244 let v = m.terminal(&[2.0, 10.0]);
245 assert!((v - 12.0_f64.ln()).abs() < 1e-10);
246 }
247
248 #[test]
249 fn terminal_handles_negative_total_clamp() {
250 let m = default_model();
251 let v = m.terminal(&[-4.0, 2.0]);
253 assert!(v.is_finite());
254 assert!(v < 0.0);
255 }
256
257 #[test]
258 fn dims_tagged_correctly() {
259 let m = default_model();
260 assert!(!m.is_diffusion_dimension(0));
261 assert!(m.is_diffusion_dimension(1));
262 assert!(m.is_integer_dimension(0));
263 assert!(!m.is_integer_dimension(1));
264 }
265
266 #[test]
267 fn fill_rate_methods_return_expected() {
268 let m = default_model();
269 assert!((m.fill_rate_base(&[0.0, 5.0]) - 140.0).abs() < 1e-12);
270 assert!((m.fill_rate_decay() - 1.5).abs() < 1e-12);
271 }
272}