solver/models/
avellaneda_drift.rs1use super::traits::{ControlOutput, Gradients, Model};
2
3#[derive(Clone)]
8pub struct AvellanedaDrift {
9 pub gamma: f64,
10 pub sigma: f64,
11 pub kappa: f64,
12 pub a: f64,
13 pub mu: f64,
15}
16
17impl AvellanedaDrift {
18 pub fn get_spreads(&self, grads: &Gradients<2>) -> (f64, f64) {
20 let base_spread = (1.0 / self.gamma) * (1.0 + self.gamma / self.kappa).ln();
21
22 let dv_dq_buy = grads.fwd[0];
23 let dv_dq_sell = grads.bwd[0];
24
25 let mut delta_bid = base_spread - dv_dq_buy;
26 let mut delta_ask = base_spread + dv_dq_sell;
27
28 let min_spread = -5.0;
29 delta_bid = delta_bid.max(min_spread);
30 delta_ask = delta_ask.max(min_spread);
31
32 (delta_bid, delta_ask)
33 }
34}
35
36impl Model<2> for AvellanedaDrift {
37 type Process = ();
38
39 fn process(&self) {}
40
41 fn optimize(&self, state: &[f64; 2], grads: &Gradients<2>) -> ControlOutput<2> {
42 let q = state[0];
43
44 let (d_bid, d_ask) = self.get_spreads(grads);
45
46 let lambda_bid = self.a * (-self.kappa * d_bid).exp();
47 let lambda_ask = self.a * (-self.kappa * d_ask).exp();
48
49 let hamiltonian_val = (lambda_bid + lambda_ask) / (self.gamma + self.kappa);
50
51 let jump_correction = lambda_bid * grads.fwd[0] - lambda_ask * grads.bwd[0];
52
53 let risk_penalty = -0.5 * self.gamma * self.sigma.powi(2) * q.powi(2);
54
55 let drift_effect = q * self.mu;
57
58 ControlOutput {
59 lambda_plus: [lambda_bid, 0.0],
60 lambda_minus: [lambda_ask, 0.0],
61 flow: hamiltonian_val + risk_penalty + drift_effect - jump_correction,
62 }
63 }
64
65 fn terminal(&self, _state: &[f64; 2]) -> f64 {
66 0.0
67 }
68
69 fn constant_discount_rate(&self) -> Option<f64> {
70 Some(0.0)
71 }
72
73 fn next_step(&self, current_state: &[f64; 2], dt: f64, noise: &[f64; 2]) -> [f64; 2] {
74 let mut next = *current_state;
75 let sqrt_dt = dt.sqrt();
76 next[1] += self.sigma * sqrt_dt * noise[1];
79 next
80 }
81
82 fn is_diffusion_dimension(&self, dim: usize) -> bool {
83 dim == 1 }
85
86 fn is_integer_dimension(&self, dim: usize) -> bool {
87 dim == 0 }
89
90 fn fill_rate_base(&self, _state: &[f64; 2]) -> f64 {
91 self.a
92 }
93
94 fn fill_rate_decay(&self) -> f64 {
95 self.kappa
96 }
97}