1use crate::data_source::SimulatedDataSource;
10use crate::engine::Engine;
11use crate::matcher::StochasticMatcher;
12use crate::strategies::external::ExternalStrategy;
13use crate::types::{Order, OrderRequest, Side};
14use market_model::process::heston::HestonProcess;
15use ndarray::Array1;
16
17#[derive(Clone, Debug)]
19pub struct VecEnvConfig {
20 pub mu: f64,
22 pub kappa: f64,
23 pub theta: f64,
24 pub sigma_v: f64,
25 pub rho: f64,
26 pub initial_price: f64,
27 pub initial_variance: f64,
28
29 pub spread: f64,
31 pub dt: f64,
32 pub max_steps: usize,
33 pub impact_factor: f64,
34
35 pub k: f64,
37 pub a: f64,
38 pub hawkes_alpha: f64,
39 pub hawkes_beta: f64,
40
41 pub initial_cash: f64,
43 pub transaction_cost: f64,
44
45 pub reward_type: RewardType,
47 pub inventory_penalty: f64,
48
49 pub matcher_dt: Option<f64>,
51
52 pub seed: u64,
54}
55
56impl Default for VecEnvConfig {
57 fn default() -> Self {
58 Self {
59 mu: 0.0,
60 kappa: 2.0,
61 theta: 0.04,
62 sigma_v: 0.3,
63 rho: -0.7,
64 initial_price: 100.0,
65 initial_variance: 0.04,
66 spread: 0.05,
67 dt: 1e-4,
68 max_steps: 10_000,
69 impact_factor: 0.0,
70 k: 1.5,
71 a: 10.0,
72 hawkes_alpha: 0.3,
73 hawkes_beta: 1.0,
74 initial_cash: 10_000.0,
75 transaction_cost: 0.0,
76 reward_type: RewardType::PnL,
77 inventory_penalty: 0.01,
78 matcher_dt: None,
79 seed: 42,
80 }
81 }
82}
83
84#[derive(Clone, Debug, PartialEq)]
85pub enum RewardType {
86 PnL,
89 DiffSharpe,
91}
92
93struct DiffSharpeState {
95 a_prev: f64, b_prev: f64, eta: f64, }
99
100impl DiffSharpeState {
101 fn new(eta: f64) -> Self {
102 Self {
103 a_prev: 0.0,
104 b_prev: 0.0,
105 eta,
106 }
107 }
108
109 fn update(&mut self, ret: f64) -> f64 {
111 let a_new = self.a_prev + self.eta * (ret - self.a_prev);
112 let b_new = self.b_prev + self.eta * (ret * ret - self.b_prev);
113 let denom = b_new - a_new * a_new;
114 let dsr = if denom > 1e-12 {
115 let denom_sqrt = denom.sqrt();
116 (b_new * (ret - self.a_prev) - 0.5 * a_new * (ret * ret - self.b_prev))
117 / (denom * denom_sqrt)
118 } else {
119 0.0
120 };
121 self.a_prev = a_new;
122 self.b_prev = b_new;
123 dsr
124 }
125}
126
127struct EnvSlot {
129 engine: Engine<StochasticMatcher, SimulatedDataSource<HestonProcess>>,
130 prev_equity: f64,
131 done: bool,
132 steps_taken: usize,
133 dsr: DiffSharpeState,
134}
135
136pub struct VecEnv {
144 config: VecEnvConfig,
145 slots: Vec<EnvSlot>,
146 states: Vec<f32>,
148 rewards: Vec<f32>,
149 dones: Vec<f32>,
150 equities: Vec<f32>,
152 positions: Vec<f32>,
153 step_fills: Vec<f32>,
154}
155
156impl VecEnv {
157 pub fn new(config: VecEnvConfig, num_envs: usize) -> Self {
158 let mut slots = Vec::with_capacity(num_envs);
159 for i in 0..num_envs {
160 slots.push(Self::make_slot(&config, i));
161 }
162
163 let mut states = vec![0.0f32; num_envs * 6];
165 for (i, slot) in slots.iter_mut().enumerate() {
166 slot.engine.step();
167 Self::write_state(slot, &config, &mut states[i * 6..(i + 1) * 6]);
168 }
169
170 let mut equities = vec![0.0f32; num_envs];
171 for val in equities.iter_mut() {
172 *val = config.initial_cash as f32;
173 }
174
175 Self {
176 config,
177 slots,
178 states,
179 rewards: vec![0.0f32; num_envs],
180 dones: vec![0.0f32; num_envs],
181 equities,
182 positions: vec![0.0f32; num_envs],
183 step_fills: vec![0.0f32; num_envs],
184 }
185 }
186
187 pub fn num_envs(&self) -> usize {
189 self.slots.len()
190 }
191
192 pub fn reset(&mut self) -> &[f32] {
194 for (i, slot) in self.slots.iter_mut().enumerate() {
195 *slot = Self::make_slot(&self.config, i);
196 slot.engine.step();
197 Self::write_state(slot, &self.config, &mut self.states[i * 6..(i + 1) * 6]);
198 self.equities[i] = self.config.initial_cash as f32;
199 self.positions[i] = 0.0;
200 self.step_fills[i] = 0.0;
201 }
202 &self.states
203 }
204
205 pub fn step(&mut self, actions: &[f32]) -> (&[f32], &[f32], &[f32]) {
215 let n = self.slots.len();
216 debug_assert_eq!(actions.len(), n * 2);
217
218 use rayon::prelude::*;
219
220 self.slots
221 .par_iter_mut()
222 .zip(self.rewards.par_iter_mut())
223 .zip(self.dones.par_iter_mut())
224 .zip(self.equities.par_iter_mut())
225 .zip(self.positions.par_iter_mut())
226 .zip(self.step_fills.par_iter_mut())
227 .zip(self.states.par_chunks_mut(6))
228 .enumerate()
229 .for_each(
230 |(i, ((((((slot, reward), done), equity), position), step_fill), state_row))| {
231 let bid_dist = actions[i * 2].max(0.0) as f64;
232 let ask_dist = actions[i * 2 + 1].max(0.0) as f64;
233
234 let mid = Self::get_mid(slot);
236
237 let bid_price = mid - bid_dist;
239 let ask_price = mid + ask_dist;
240
241 let requests = [
242 OrderRequest::CancelAll,
243 OrderRequest::New(Order::new(
244 (slot.steps_taken as u64) * 2 + 1,
245 Side::Buy,
246 bid_price,
247 1.0,
248 )),
249 OrderRequest::New(Order::new(
250 (slot.steps_taken as u64) * 2 + 2,
251 Side::Sell,
252 ask_price,
253 1.0,
254 )),
255 ];
256
257 if let Some(strat) = slot.engine.strategy_mut::<ExternalStrategy>() {
259 strat.set_pending_requests(requests.to_vec());
260 }
261
262 let running = slot.engine.step();
263 slot.steps_taken += 1;
264
265 if !running {
266 slot.done = true;
267 }
268
269 let eq = Self::compute_equity(slot);
271 let pos = Self::get_position(slot);
272
273 let r = match self.config.reward_type {
274 RewardType::PnL => {
275 let pnl = eq - slot.prev_equity;
276 pnl - self.config.inventory_penalty * pos * pos
277 }
278 RewardType::DiffSharpe => {
279 let ret = if slot.prev_equity.abs() > 1e-12 {
280 (eq - slot.prev_equity) / slot.prev_equity
281 } else {
282 0.0
283 };
284 let dsr = slot.dsr.update(ret);
285 dsr - self.config.inventory_penalty * pos * pos
286 }
287 };
288
289 slot.prev_equity = eq;
290 *reward = r as f32;
291 *done = if slot.done { 1.0 } else { 0.0 };
292 *equity = eq as f32;
293 *position = pos as f32;
294 *step_fill = slot.engine.last_step_fills().len() as f32;
295
296 if slot.done {
298 *slot = Self::make_slot(&self.config, i);
299 slot.engine.step();
300 }
301
302 Self::write_state(slot, &self.config, state_row);
303 },
304 );
305
306 (&self.states, &self.rewards, &self.dones)
307 }
308
309 pub fn get_states(&self) -> &[f32] {
311 &self.states
312 }
313
314 pub fn get_equities(&self) -> &[f32] {
316 &self.equities
317 }
318
319 pub fn get_positions(&self) -> &[f32] {
321 &self.positions
322 }
323
324 pub fn get_step_fills(&self) -> &[f32] {
326 &self.step_fills
327 }
328
329 fn make_slot(config: &VecEnvConfig, slot_index: usize) -> EnvSlot {
332 use rand::SeedableRng;
333 use rand::rngs::StdRng;
334
335 let seed = config.seed.wrapping_add(slot_index as u64);
336 let ds_rng = StdRng::seed_from_u64(seed);
337 let matcher_rng = StdRng::seed_from_u64(seed.wrapping_add(1));
338
339 let process = HestonProcess::new(
340 config.mu,
341 config.kappa,
342 config.theta,
343 config.sigma_v,
344 config.rho,
345 config.initial_price,
346 config.initial_variance,
347 );
348 let initial_state = Array1::from_vec(vec![config.initial_price, config.initial_variance]);
349 let ds = SimulatedDataSource::new(
350 process,
351 initial_state,
352 config.dt,
353 config.spread,
354 config.max_steps,
355 config.impact_factor,
356 )
357 .with_rng(ds_rng);
358 let m_dt = config.matcher_dt.unwrap_or(config.dt);
359 let base_matcher = StochasticMatcher::new(m_dt, config.k, config.a).with_rng(matcher_rng);
360 let matcher = if config.hawkes_alpha > 0.0 {
361 base_matcher.with_bilateral_hawkes(config.hawkes_alpha, config.hawkes_beta)
362 } else {
363 base_matcher
364 };
365 let strategy = ExternalStrategy::new();
366 let engine = Engine::new(
367 matcher,
368 strategy,
369 ds,
370 config.initial_cash,
371 config.transaction_cost,
372 );
373
374 EnvSlot {
375 engine,
376 prev_equity: config.initial_cash,
377 done: false,
378 steps_taken: 0,
379 dsr: DiffSharpeState::new(0.01),
380 }
381 }
382
383 fn get_mid(slot: &EnvSlot) -> f64 {
384 slot.engine
385 .get_last_state()
386 .map(|s| (s.best_bid + s.best_ask) / 2.0)
387 .unwrap_or(100.0)
388 }
389
390 fn compute_equity(slot: &EnvSlot) -> f64 {
391 let p = slot.engine.get_portfolio();
392 let mid = Self::get_mid(slot);
393 p.cash + p.position * mid
394 }
395
396 fn get_position(slot: &EnvSlot) -> f64 {
397 slot.engine.get_portfolio().position
398 }
399
400 fn write_state(slot: &EnvSlot, config: &VecEnvConfig, row: &mut [f32]) {
404 debug_assert_eq!(row.len(), 6);
405 let portfolio = slot.engine.get_portfolio();
406
407 if let Some(state) = slot.engine.get_last_state() {
408 let mid = (state.best_bid + state.best_ask) / 2.0;
409 row[0] = mid as f32;
410 row[1] = portfolio.position as f32;
411
412 let variance = state
414 .true_volatility
415 .map(|v| v * v) .unwrap_or(config.initial_variance);
417 row[2] = variance as f32;
418
419 let (buy_int, sell_int) = if let Some(ref params) = state.parameters {
421 (
422 params
423 .get("hawkes_buy_intensity")
424 .copied()
425 .unwrap_or(config.a),
426 params
427 .get("hawkes_sell_intensity")
428 .copied()
429 .unwrap_or(config.a),
430 )
431 } else {
432 (config.a, config.a)
433 };
434 row[3] = buy_int as f32;
435 row[4] = sell_int as f32;
436
437 let remaining = (config.max_steps - slot.steps_taken) as f64 * config.dt;
439 row[5] = remaining as f32;
440 } else {
441 row[0] = config.initial_price as f32;
443 row[1] = 0.0;
444 row[2] = config.initial_variance as f32;
445 row[3] = config.a as f32;
446 row[4] = config.a as f32;
447 row[5] = (config.max_steps as f64 * config.dt) as f32;
448 }
449 }
450}
451
452#[cfg(test)]
453mod tests {
454 use super::*;
455
456 fn test_config() -> VecEnvConfig {
457 VecEnvConfig {
458 max_steps: 100,
459 hawkes_alpha: 0.3,
460 hawkes_beta: 1.0,
461 ..Default::default()
462 }
463 }
464
465 #[test]
466 fn test_vec_env_creation() {
467 let env = VecEnv::new(test_config(), 4);
468 assert_eq!(env.num_envs(), 4);
469 assert_eq!(env.get_states().len(), 24); }
471
472 #[test]
473 fn test_initial_states_valid() {
474 let config = test_config();
475 let env = VecEnv::new(config.clone(), 2);
476 let s = env.get_states();
477
478 assert!(
480 (s[0] as f64 - config.initial_price).abs() < 1.0,
481 "mid_price {} far from {}",
482 s[0],
483 config.initial_price
484 );
485 assert_eq!(s[1], 0.0);
487 assert!((s[2] as f64 - config.initial_variance).abs() < 0.1);
489 assert!(s[5] > 0.0);
491 }
492
493 #[test]
494 fn test_step_returns_correct_shapes() {
495 let mut env = VecEnv::new(test_config(), 4);
496 let actions = vec![0.01f32; 8]; let (states, rewards, dones) = env.step(&actions);
498 assert_eq!(states.len(), 24);
499 assert_eq!(rewards.len(), 4);
500 assert_eq!(dones.len(), 4);
501 }
502
503 #[test]
504 fn test_episode_completion_and_reset() {
505 let mut config = test_config();
506 config.max_steps = 10; let mut env = VecEnv::new(config, 2);
508 let actions = vec![0.01f32; 4];
509
510 let mut any_done = false;
511 for _ in 0..15 {
512 let (_, _, dones) = env.step(&actions);
513 if dones.iter().any(|&d| d > 0.5) {
514 any_done = true;
515 }
516 }
517 assert!(any_done, "expected at least one episode to finish");
518 }
519
520 #[test]
521 fn test_reset_all() {
522 let mut env = VecEnv::new(test_config(), 3);
523 let actions = vec![0.01f32; 6];
524 for _ in 0..5 {
526 env.step(&actions);
527 }
528 let states = env.reset();
530 assert_eq!(states.len(), 18);
531 assert_eq!(states[1], 0.0);
533 assert_eq!(states[7], 0.0);
534 assert_eq!(states[13], 0.0);
535 }
536
537 #[test]
538 fn test_rewards_finite() {
539 let mut env = VecEnv::new(test_config(), 4);
540 let actions = vec![0.02f32; 8];
541 for _ in 0..20 {
542 let (_, rewards, _) = env.step(&actions);
543 for &r in rewards.iter() {
544 assert!(r.is_finite(), "reward must be finite, got {}", r);
545 }
546 }
547 }
548
549 #[test]
550 fn test_diff_sharpe_reward() {
551 let mut config = test_config();
552 config.reward_type = RewardType::DiffSharpe;
553 let mut env = VecEnv::new(config, 2);
554 let actions = vec![0.01f32; 4];
555 for _ in 0..20 {
556 let (_, rewards, _) = env.step(&actions);
557 for &r in rewards.iter() {
558 assert!(r.is_finite(), "DSR reward must be finite");
559 }
560 }
561 }
562
563 #[test]
564 fn test_negative_action_clamped() {
565 let mut env = VecEnv::new(test_config(), 1);
566 let actions = vec![-0.5f32, -0.5];
568 let (states, _, _) = env.step(&actions);
569 assert!(states[0].is_finite());
571 }
572
573 #[test]
574 fn test_hawkes_intensities_in_state() {
575 let config = test_config();
576 let mut env = VecEnv::new(config.clone(), 1);
577 let actions = vec![0.01f32, 0.01];
578 for _ in 0..10 {
580 env.step(&actions);
581 }
582 let s = env.get_states();
583 assert!(s[3] > 0.0, "hawkes_buy should be positive: {}", s[3]);
585 assert!(s[4] > 0.0, "hawkes_sell should be positive: {}", s[4]);
586 }
587}