1use super::control::{ControlProblem, StateDerivatives};
2use super::market_making::MarketMakingControl;
3use super::normal_cdf;
4use crate::numeric::finite_difference::discretization::{DimensionKind, Transport};
5use crate::numeric::finite_difference::pde::PdeProblem;
6
7pub use super::TerminalCondition;
9
10#[derive(Clone, Copy, Debug)]
23pub struct Heston {
24 pub gamma: f64,
26 pub kappa_m: f64, pub a: f64, pub v_kappa: f64, pub v_theta: f64, pub v_xi: f64, pub rho: f64,
38 pub price_scale: f64,
40 pub dv: f64, pub dq: f64,
44 pub q_min: f64,
46 pub q_max: f64,
48 pub terminal_condition: TerminalCondition,
50}
51
52impl Heston {
53 pub fn new(gamma: f64, kappa_m: f64, a: f64) -> Self {
54 Self {
55 gamma,
56 kappa_m,
57 a,
58 v_kappa: 2.0,
59 v_theta: 0.25,
60 v_xi: 0.3,
61 rho: 0.0,
62 price_scale: 1.0,
63 dv: 1.0,
64 dq: 1.0,
65 q_min: f64::NEG_INFINITY,
66 q_max: f64::INFINITY,
67 terminal_condition: TerminalCondition::Zero,
68 }
69 }
70
71 pub fn with_terminal_condition(mut self, terminal_condition: TerminalCondition) -> Self {
72 self.terminal_condition = terminal_condition;
73 self
74 }
75
76 pub fn with_inventory_bounds(mut self, q_min: f64, q_max: f64) -> Self {
77 self.q_min = q_min;
78 self.q_max = q_max;
79 self
80 }
81
82 pub fn with_variance_params(mut self, v_kappa: f64, v_theta: f64, v_xi: f64) -> Self {
83 self.v_kappa = v_kappa;
84 self.v_theta = v_theta;
85 self.v_xi = v_xi;
86 self
87 }
88
89 pub fn with_rho(mut self, rho: f64) -> Self {
90 self.rho = rho;
91 self
92 }
93
94 pub fn with_price_scale(mut self, price_scale: f64) -> Self {
95 self.price_scale = price_scale.abs().max(1e-8);
96 self
97 }
98
99 pub fn with_dv(mut self, dv: f64) -> Self {
100 self.dv = dv.abs().max(1e-8);
101 self
102 }
103
104 pub fn with_dq(mut self, dq: f64) -> Self {
105 self.dq = dq.abs().max(1e-8);
106 self
107 }
108
109 pub fn with_grid_steps(mut self, dx: &[f64; 2]) -> Self {
112 self.dq = dx[0].abs().max(1e-8);
113 self.dv = dx[1].abs().max(1e-8);
114 self
115 }
116
117 pub fn get_spreads(&self, derivs: &StateDerivatives<2>) -> (f64, f64) {
118 let base_spread = (1.0 / self.gamma) * (1.0 + self.gamma / self.kappa_m).ln();
125
126 let dv_dq_buy = derivs.fwd[0];
127 let dv_dq_sell = derivs.bwd[0];
128
129 let mut delta_bid = base_spread - dv_dq_buy;
130 let mut delta_ask = base_spread + dv_dq_sell;
131
132 let min_spread = -5.0;
133 let max_spread = 10.0;
134 delta_bid = delta_bid.max(min_spread).min(max_spread);
135 delta_ask = delta_ask.max(min_spread).min(max_spread);
136
137 (delta_bid, delta_ask)
138 }
139
140 pub fn fill_rate_base(&self, _state: &[f64; 2]) -> f64 {
142 self.a
143 }
144
145 pub fn fill_rate_decay(&self) -> f64 {
147 self.kappa_m
148 }
149}
150
151impl ControlProblem<2> for Heston {
152 type Control = MarketMakingControl;
153
154 fn optimize(&self, _t: f64, state: &[f64; 2], derivs: &StateDerivatives<2>) -> Self::Control {
155 let q = state[0];
156 let base_spread = (1.0 / self.gamma) * (1.0 + self.gamma / self.kappa_m).ln();
157 let delta_bid = (base_spread - derivs.fwd[0]).clamp(-5.0, 10.0);
158 let delta_ask = (base_spread + derivs.bwd[0]).clamp(-5.0, 10.0);
159
160 let lambda_bid = if q >= self.q_max {
161 0.0
162 } else {
163 self.a * (-self.kappa_m * delta_bid).exp()
164 };
165 let lambda_ask = if q <= self.q_min {
166 0.0
167 } else {
168 self.a * (-self.kappa_m * delta_ask).exp()
169 };
170
171 MarketMakingControl::new(lambda_bid, lambda_ask)
172 }
173
174 fn running_reward(&self, _t: f64, _state: &[f64; 2], control: &Self::Control) -> f64 {
175 (control.bid_intensity + control.ask_intensity) / (self.gamma + self.kappa_m)
176 }
177
178 fn bsde_driver(
179 &self,
180 _t: f64,
181 state: &[f64; 2],
182 control: &Self::Control,
183 _derivs: &StateDerivatives<2>,
184 dt: f64,
185 ) -> f64 {
186 let q = state[0];
195 let v = state[1];
196 let local = -0.5 * self.gamma * self.price_scale.powi(2) * v * q.powi(2);
197 let rate_cap = 1.0 / dt.max(1e-12);
198 let reward = (control.bid_intensity.min(rate_cap) + control.ask_intensity.min(rate_cap))
199 / (self.gamma + self.kappa_m);
200 reward + local
201 }
202
203 fn generator(
204 &self,
205 _t: f64,
206 state: &[f64; 2],
207 _control: &Self::Control,
208 derivs: &StateDerivatives<2>,
209 ) -> f64 {
210 let q = state[0];
214 let v = state[1];
215 let mu_v = self.v_kappa * (self.v_theta - v)
216 - self.gamma * self.price_scale * self.rho * self.v_xi * v * q;
217 let sigma2_v = self.v_xi.powi(2) * v;
218
219 -0.5 * self.gamma * self.price_scale.powi(2) * v * q.powi(2)
220 + mu_v * derivs.grad[1]
221 + 0.5 * sigma2_v * derivs.hessian[1]
222 }
223
224 fn terminal(&self, state: &[f64; 2]) -> f64 {
225 match self.terminal_condition {
226 TerminalCondition::Zero => 0.0,
227 TerminalCondition::LiquidationCost => {
228 let q = state[0];
229 let base_spread = (1.0 / self.gamma) * (1.0 + self.gamma / self.kappa_m).ln();
230 -q.abs() * base_spread
231 }
232 }
233 }
234
235 fn discount_rate(&self, _state: &[f64; 2]) -> f64 {
236 0.0
237 }
238
239 fn constant_discount_rate(&self) -> Option<f64> {
240 Some(0.0)
241 }
242
243 fn next_step(&self, _t: f64, state: &[f64; 2], dt: f64, noise: &[f64; 2]) -> [f64; 2] {
244 let mut next = *state;
245
246 let u = normal_cdf(noise[0]);
247 let base_spread = (1.0 / self.gamma) * (1.0 + self.gamma / self.kappa_m).ln();
248 let lambda_bid = self.a * (-self.kappa_m * base_spread).exp();
249 let lambda_ask = self.a * (-self.kappa_m * base_spread).exp();
250 let p_bid = (lambda_bid * dt).clamp(0.0, 1.0);
251 let p_ask = (lambda_ask * dt).clamp(0.0, 1.0);
252 if u < p_bid {
253 next[0] += 1.0;
254 } else if u > 1.0 - p_ask {
255 next[0] -= 1.0;
256 }
257
258 let v = state[1];
259 let v_pos = v.max(1e-5);
260 let mu_v = self.v_kappa * (self.v_theta - v_pos)
261 - self.gamma * self.price_scale * self.rho * self.v_xi * v_pos * next[0];
262 let drift = mu_v * dt;
263 let diffusion = self.v_xi * v_pos.sqrt() * dt.sqrt() * noise[1];
264 let mut next_v = v_pos + drift + diffusion;
265 if next_v < 0.0 {
266 next_v = -next_v;
267 }
268 next[1] = next_v;
269 next
270 }
271
272 fn next_step_controlled(
273 &self,
274 _t: f64,
275 state: &[f64; 2],
276 control: &Self::Control,
277 dt: f64,
278 noise: &[f64; 2],
279 ) -> [f64; 2] {
280 let mut next = *state;
281
282 let u = normal_cdf(noise[0]);
285 let p_bid = (control.bid_intensity * dt).clamp(0.0, 1.0);
286 let p_ask = (control.ask_intensity * dt).clamp(0.0, 1.0);
287 if u < p_bid {
288 next[0] += 1.0;
289 } else if u > 1.0 - p_ask {
290 next[0] -= 1.0;
291 }
292
293 let v = state[1];
294 let v_pos = v.max(1e-5);
295 let mu_v = self.v_kappa * (self.v_theta - v_pos)
296 - self.gamma * self.price_scale * self.rho * self.v_xi * v_pos * next[0];
297 let drift = mu_v * dt;
298 let diffusion = self.v_xi * v_pos.sqrt() * dt.sqrt() * noise[1];
299 let mut next_v = v_pos + drift + diffusion;
300 if next_v < 0.0 {
301 next_v = -next_v;
302 }
303 next[1] = next_v;
304 next
305 }
306
307 fn is_diffusion_dimension(&self, dim: usize) -> bool {
308 dim == 1
309 }
310
311 fn gradient_step(&self, dim: usize) -> f64 {
312 if dim == 1 {
313 self.dv.abs().max(1e-8)
314 } else {
315 1.0
316 }
317 }
318}
319
320impl PdeProblem<2> for Heston {
321 fn dimension_kind(&self, dim: usize) -> DimensionKind {
322 match dim {
323 0 => DimensionKind::DiscreteJump,
324 _ => DimensionKind::Diffusion,
325 }
326 }
327
328 fn transport(
329 &self,
330 _t: f64,
331 state: &[f64; 2],
332 control: &Self::Control,
333 derivs: &StateDerivatives<2>,
334 ) -> Transport<2> {
335 let q = state[0];
336 let v = state[1].max(1e-5);
337 let mu_v = self.v_kappa * (self.v_theta - v)
338 - self.gamma * self.price_scale * self.rho * self.v_xi * v * q;
339 let sigma2_v = self.v_xi.powi(2) * v;
340 let dv = self.dv.abs().max(1e-12);
341
342 let diff_term = sigma2_v / (2.0 * dv * dv);
343 let drift_abs = mu_v.abs() / dv;
344 let (plus_v, minus_v) = if mu_v >= 0.0 {
345 (diff_term + drift_abs, diff_term)
346 } else {
347 (diff_term, diff_term + drift_abs)
348 };
349
350 let hamiltonian =
351 (control.bid_intensity + control.ask_intensity) / (self.gamma + self.kappa_m);
352 let jump_transport = self.dq
353 * (control.bid_intensity * derivs.fwd[0] - control.ask_intensity * derivs.bwd[0]);
354 let local = -0.5 * self.gamma * self.price_scale.powi(2) * v * q.powi(2);
355
356 Transport::new(
357 [control.bid_intensity, plus_v],
358 [control.ask_intensity, minus_v],
359 hamiltonian + local - jump_transport,
360 )
361 }
362}
363
364#[cfg(test)]
365mod tests {
366 use super::*;
367
368 fn zero_derivs() -> StateDerivatives<2> {
369 StateDerivatives::new([0.0; 2], [0.0; 2])
370 }
371
372 fn default_model() -> Heston {
373 Heston {
374 gamma: 0.5,
375 kappa_m: 1.5,
376 a: 10.0,
377 v_kappa: 2.0,
378 v_theta: 0.25,
379 v_xi: 0.3,
380 rho: 0.0,
381 price_scale: 1.0,
382 dv: 0.02,
383 dq: 1.0,
384 q_min: f64::NEG_INFINITY,
385 q_max: f64::INFINITY,
386 terminal_condition: TerminalCondition::Zero,
387 }
388 }
389
390 #[test]
391 fn zero_gradient_gives_base_spread() {
392 let m = default_model();
393 let derivs = zero_derivs();
394 let (bid, ask) = m.get_spreads(&derivs);
395 let base = (1.0 / m.gamma) * (1.0 + m.gamma / m.kappa_m).ln();
396 assert!((bid - base).abs() < 1e-12);
397 assert!((ask - base).abs() < 1e-12);
398 }
399
400 #[test]
401 fn bsde_driver_excludes_variance_transport() {
402 let m = default_model();
403 let state = [2.0, 0.25];
404 let control = ControlProblem::optimize(&m, 0.0, &state, &zero_derivs());
405
406 let mut derivs = zero_derivs();
408 derivs.grad[1] = 1.0;
409 derivs.hessian[1] = -2.0;
410
411 let dt = 0.01;
412 let q = state[0];
413 let v = state[1];
414 let local = -0.5 * m.gamma * m.price_scale.powi(2) * v * q.powi(2);
415 let rate_cap = 1.0 / dt;
416 let reward = (control.bid_intensity.min(rate_cap) + control.ask_intensity.min(rate_cap))
417 / (m.gamma + m.kappa_m);
418 let expected = reward + local;
419
420 assert_eq!(
421 ControlProblem::bsde_driver(&m, 0.0, &state, &control, &derivs, dt),
422 expected
423 );
424 assert!(
425 ControlProblem::bsde_driver(&m, 0.0, &state, &control, &derivs, dt)
426 != ControlProblem::driver(&m, 0.0, &state, &control, &derivs),
427 "full driver adds variance transport, bsde driver must not"
428 );
429 }
430
431 #[test]
432 fn next_step_controlled_uses_optimal_intensities() {
433 let m = default_model();
434 let state = [1.0, 0.25];
435
436 let ctrl = MarketMakingControl::new(0.0, 0.0);
437 let next = ControlProblem::next_step_controlled(&m, 0.0, &state, &ctrl, 0.01, &[0.0, 0.0]);
438 assert_eq!(next[0], 1.0);
439
440 let ctrl = MarketMakingControl::new(1000.0, 0.0);
441 let next = ControlProblem::next_step_controlled(&m, 0.0, &state, &ctrl, 0.01, &[0.0, 0.0]);
442 assert_eq!(next[0], 2.0);
443 }
444
445 #[test]
446 fn base_spread_independent_of_variance_at_zero_gradient() {
447 let m = default_model();
448 let derivs = zero_derivs();
449 let (bid_low, ask_low) = m.get_spreads(&derivs);
450 let (bid_high, ask_high) = m.get_spreads(&derivs);
451 assert!((bid_low - bid_high).abs() < 1e-12);
452 assert!((ask_low - ask_high).abs() < 1e-12);
453 }
454
455 #[test]
456 fn risk_penalty_scales_with_variance() {
457 let m = default_model();
458 let derivs = zero_derivs();
459 let gen_low = ControlProblem::generator(
460 &m,
461 0.0,
462 &[2.0, 0.05],
463 &MarketMakingControl::new(0.0, 0.0),
464 &derivs,
465 );
466 let gen_high = ControlProblem::generator(
467 &m,
468 0.0,
469 &[2.0, 0.50],
470 &MarketMakingControl::new(0.0, 0.0),
471 &derivs,
472 );
473 assert!(
474 gen_high < gen_low,
475 "Higher variance should decrease the generator: low={}, high={}",
476 gen_low,
477 gen_high
478 );
479 }
480
481 #[test]
482 fn variance_transport_uses_upwinding() {
483 let m = default_model();
484 let mut derivs = zero_derivs();
485 derivs.grad[1] = 1.0;
486 let control = MarketMakingControl::new(0.0, 0.0);
487 let g_low = ControlProblem::generator(&m, 0.0, &[0.0, 0.05], &control, &derivs);
488 let g_high = ControlProblem::generator(&m, 0.0, &[0.0, 0.80], &control, &derivs);
489 assert!(
490 g_low > g_high,
491 "Positive variance drift at low v should raise the generator more than at high v"
492 );
493 }
494
495 #[test]
496 fn rho_affects_variance_drift() {
497 let mut derivs = zero_derivs();
498 derivs.grad[1] = 1.0;
499 let state = [2.0, 0.25];
500 let control = MarketMakingControl::new(0.0, 0.0);
501
502 let g0 = ControlProblem::generator(
503 &Heston {
504 rho: 0.0,
505 ..default_model()
506 },
507 0.0,
508 &state,
509 &control,
510 &derivs,
511 );
512 let g_pos = ControlProblem::generator(
513 &Heston {
514 rho: 0.5,
515 ..default_model()
516 },
517 0.0,
518 &state,
519 &control,
520 &derivs,
521 );
522 let g_neg = ControlProblem::generator(
523 &Heston {
524 rho: -0.5,
525 ..default_model()
526 },
527 0.0,
528 &state,
529 &control,
530 &derivs,
531 );
532
533 assert!(
534 g_neg > g0 && g0 > g_pos,
535 "Positive rho lowers the effective variance drift; negative rho raises it"
536 );
537 }
538
539 #[test]
540 fn rho_has_no_effect_at_zero_inventory() {
541 let mut derivs = zero_derivs();
542 derivs.grad[1] = 1.0;
543 let state = [0.0, 0.25];
544 let control = MarketMakingControl::new(0.0, 0.0);
545
546 let g0 = ControlProblem::generator(
547 &Heston {
548 rho: 0.0,
549 ..default_model()
550 },
551 0.0,
552 &state,
553 &control,
554 &derivs,
555 );
556 let g_pos = ControlProblem::generator(
557 &Heston {
558 rho: 0.8,
559 ..default_model()
560 },
561 0.0,
562 &state,
563 &control,
564 &derivs,
565 );
566
567 assert!((g0 - g_pos).abs() < 1e-12);
568 }
569
570 #[test]
571 fn cir_next_step_stays_positive() {
572 let m = default_model();
573 let state = [0.0, 0.001];
574 let next = ControlProblem::next_step(&m, 0.0, &state, 0.01, &[0.0, -5.0]);
575 assert!(
576 next[1] > 0.0,
577 "CIR reflection should keep variance positive: {}",
578 next[1]
579 );
580 }
581
582 #[test]
583 fn next_step_mean_reverts_variance() {
584 let m = default_model();
585 let state = [0.0, 0.05];
586 let mut sum = 0.0;
587 for i in 0..1000 {
588 let noise = [0.0, (i as f64 * 0.01).sin()];
589 let next = ControlProblem::next_step(&m, 0.0, &state, 0.01, &noise);
590 sum += next[1] - state[1];
591 }
592 assert!(
593 sum / 1000.0 > 0.0,
594 "Variance should drift upward when below theta"
595 );
596 }
597
598 #[test]
599 fn builder_methods_work() {
600 let m = Heston::new(0.5, 1.5, 10.0)
601 .with_variance_params(3.0, 0.1, 0.5)
602 .with_rho(-0.3)
603 .with_dv(0.05)
604 .with_dq(2.0);
605 assert_eq!(m.v_kappa, 3.0);
606 assert_eq!(m.rho, -0.3);
607 assert_eq!(m.dv, 0.05);
608 }
609
610 #[test]
611 fn with_grid_steps_sets_both() {
612 let m = Heston::new(0.5, 1.5, 10.0).with_grid_steps(&[2.0, 0.05]);
613 assert_eq!(m.dq, 2.0);
614 assert_eq!(m.dv, 0.05);
615 }
616
617 #[test]
618 fn risk_penalty_scales_with_price_scale_squared() {
619 let derivs = zero_derivs();
620 let state = [3.0, 0.25];
621 let control = MarketMakingControl::new(0.0, 0.0);
622
623 let m1 = Heston {
624 rho: 0.0,
625 price_scale: 1.0,
626 ..default_model()
627 };
628 let m10 = Heston {
629 rho: 0.0,
630 price_scale: 10.0,
631 ..default_model()
632 };
633
634 let g1 = ControlProblem::generator(&m1, 0.0, &state, &control, &derivs);
635 let g10 = ControlProblem::generator(&m10, 0.0, &state, &control, &derivs);
636
637 let penalty1 = -0.5 * m1.gamma * m1.price_scale.powi(2) * state[1] * state[0].powi(2);
638 let penalty10 = -0.5 * m10.gamma * m10.price_scale.powi(2) * state[1] * state[0].powi(2);
639 let expected_delta = penalty10 - penalty1;
640 let observed_delta = g10 - g1;
641
642 assert!(
643 (observed_delta - expected_delta).abs() < 1e-10,
644 "Risk penalty scaling mismatch: observed_delta={}, expected_delta={}",
645 observed_delta,
646 expected_delta
647 );
648 }
649
650 #[test]
651 fn cross_term_scales_linearly_with_price_scale() {
652 let mut derivs = zero_derivs();
653 derivs.grad[1] = 1.0;
654 let state = [2.0, 0.25];
655 let control = MarketMakingControl::new(0.0, 0.0);
656 let base = Heston {
657 rho: -0.7,
658 v_xi: 0.8,
659 dv: 0.05,
660 price_scale: 1.0,
661 ..default_model()
662 };
663 let scaled = Heston {
664 price_scale: 5.0,
665 ..base
666 };
667
668 let g1 = ControlProblem::generator(&base, 0.0, &state, &control, &derivs);
669 let g5 = ControlProblem::generator(&scaled, 0.0, &state, &control, &derivs);
670
671 let diff1 =
672 g1 - (-0.5 * base.gamma * base.price_scale.powi(2) * state[1] * state[0].powi(2));
673 let diff5 =
674 g5 - (-0.5 * scaled.gamma * scaled.price_scale.powi(2) * state[1] * state[0].powi(2));
675
676 let ratio = diff5 / diff1;
677 assert!(
678 (ratio - 5.0).abs() < 1e-10,
679 "Cross-term scaling should be linear in price_scale: ratio={}",
680 ratio
681 );
682 }
683}