Skip to main content

engine/
vec_env.rs

1// cargo test --lib vec_env -- --nocapture
2//
3// Vectorized market-making environment for reinforcement learning.
4//
5// Runs N independent Heston + bilateral-Hawkes market simulations in parallel.
6// Each environment produces a 6-dimensional state and accepts a 2-dimensional
7// action (bid/ask spread distances). Finished environments auto-reset.
8
9use 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/// Parameters for constructing a vectorized environment.
18#[derive(Clone, Debug)]
19pub struct VecEnvConfig {
20    // Heston process parameters
21    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    // Market microstructure
30    pub spread: f64,
31    pub dt: f64,
32    pub max_steps: usize,
33    pub impact_factor: f64,
34
35    // Matcher parameters
36    pub k: f64,
37    pub a: f64,
38    pub hawkes_alpha: f64,
39    pub hawkes_beta: f64,
40
41    // Agent parameters
42    pub initial_cash: f64,
43    pub transaction_cost: f64,
44
45    // Reward
46    pub reward_type: RewardType,
47    pub inventory_penalty: f64,
48
49    // Optional separate matcher time step (defaults to dt if None)
50    pub matcher_dt: Option<f64>,
51
52    // Master seed for per-slot RNG derivation.
53    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    /// Step-wise PnL change with quadratic inventory penalty:
87    ///   r_t = (equity_t - equity_{t-1}) - penalty * q_t^2
88    PnL,
89    /// Differential Sharpe Ratio (Moody & Saffell 2001).
90    DiffSharpe,
91}
92
93/// State for the Differential Sharpe Ratio tracker.
94struct DiffSharpeState {
95    a_prev: f64, // exponential moving average of returns
96    b_prev: f64, // exponential moving average of squared returns
97    eta: f64,    // decay factor
98}
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    /// Update the running statistics and return the differential Sharpe ratio.
110    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
127/// One slot in the vectorized environment.
128struct EnvSlot {
129    engine: Engine<StochasticMatcher, SimulatedDataSource<HestonProcess>>,
130    prev_equity: f64,
131    done: bool,
132    steps_taken: usize,
133    dsr: DiffSharpeState,
134}
135
136/// Vectorized market-making environment.
137///
138/// Manages N independent engine instances. On each `step()`:
139/// 1. Read actions `(N, 2)` as (bid_distance, ask_distance)
140/// 2. For each env: set orders, advance engine, compute reward
141/// 3. Return `(states, rewards, dones)` as contiguous f32 arrays
142/// 4. Auto-reset finished environments
143pub struct VecEnv {
144    config: VecEnvConfig,
145    slots: Vec<EnvSlot>,
146    // Pre-allocated output buffers: states (N*6), rewards (N), dones (N)
147    states: Vec<f32>,
148    rewards: Vec<f32>,
149    dones: Vec<f32>,
150    // Auxiliary buffers for evaluation (populated before auto-reset)
151    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        // Initial observation: step each engine once to populate state
164        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    /// Number of environments.
188    pub fn num_envs(&self) -> usize {
189        self.slots.len()
190    }
191
192    /// Reset all environments and return initial states.
193    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    /// Step all environments with the given actions.
206    ///
207    /// `actions` must have length `num_envs * 2`, laid out as
208    /// `[bid_dist_0, ask_dist_0, bid_dist_1, ask_dist_1, ...]`.
209    ///
210    /// Returns `(&states, &rewards, &dones)` where:
211    /// - states: `[f32; N*6]` flattened row-major
212    /// - rewards: `[f32; N]`
213    /// - dones: `[f32; N]` (1.0 if episode ended, 0.0 otherwise)
214    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                    // Get current mid-price from last observation
235                    let mid = Self::get_mid(slot);
236
237                    // Place symmetric limit orders
238                    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                    // Set pending requests and step
258                    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                    // Compute reward
270                    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                    // Auto-reset finished envs
297                    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    /// Current states buffer (read-only).
310    pub fn get_states(&self) -> &[f32] {
311        &self.states
312    }
313
314    /// Per-env equity after the last step (captured before auto-reset).
315    pub fn get_equities(&self) -> &[f32] {
316        &self.equities
317    }
318
319    /// Per-env inventory after the last step (captured before auto-reset).
320    pub fn get_positions(&self) -> &[f32] {
321        &self.positions
322    }
323
324    /// Per-env fill count from the last step (captured before auto-reset).
325    pub fn get_step_fills(&self) -> &[f32] {
326        &self.step_fills
327    }
328
329    // ------ internal helpers ------
330
331    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    /// Write the 6D state for one env into the provided 6-element row slice.
401    ///
402    /// State layout: [mid_price, inventory, variance, hawkes_buy, hawkes_sell, time_remaining]
403    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            // Variance from the Heston ground-truth
413            let variance = state
414                .true_volatility
415                .map(|v| v * v) // vol -> variance
416                .unwrap_or(config.initial_variance);
417            row[2] = variance as f32;
418
419            // Hawkes intensities from parameters
420            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            // Remaining time as fraction
438            let remaining = (config.max_steps - slot.steps_taken) as f64 * config.dt;
439            row[5] = remaining as f32;
440        } else {
441            // Before first step
442            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); // 4 * 6
470    }
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        // mid_price should be near initial_price
479        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        // inventory should be 0
486        assert_eq!(s[1], 0.0);
487        // variance should be near initial_variance
488        assert!((s[2] as f64 - config.initial_variance).abs() < 0.1);
489        // time remaining should be > 0
490        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]; // 4 envs * 2 actions
497        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; // short episode
507        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        // Take some steps
525        for _ in 0..5 {
526            env.step(&actions);
527        }
528        // Reset
529        let states = env.reset();
530        assert_eq!(states.len(), 18);
531        // All inventories should be 0 after reset
532        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        // Negative distances should be clamped to 0
567        let actions = vec![-0.5f32, -0.5];
568        let (states, _, _) = env.step(&actions);
569        // Should not crash; mid_price should still be valid
570        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        // Step a few times to let Hawkes evolve
579        for _ in 0..10 {
580            env.step(&actions);
581        }
582        let s = env.get_states();
583        // Hawkes buy/sell intensities should be positive
584        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}