feat(alpha): phase_e_random_baseline example + reset_at extension
Phase E.0 Task 7b. Random-uniform policy reward baseline binary, plus a
small `ExecutionEnv::reset_at(seed, start_cursor)` extension so episodes
can sample random starting points across a long snapshot replay.
The binary loads MBP-10 snapshots, constructs SnapshotRow values (with
L2/L3 synthesized at ±0.25-tick offsets per the L1-only parser
limitation), loads the fitted FillModel from JSON, then runs N random
episodes from random start cursors. Reports mean / std / quintile
percentiles + kill threshold (mean + 2σ) for E.1 to exceed.
Smoke run (500 episodes, horizon 600, 100K snapshots):
mean = -5600 (dominated by terminal force-close variance + market-order
over-reliance because fit converged to β_spread = -40
→ limit fill probability ~0 at typical spreads)
std = 5383
p95 = -895
kill threshold (mean + 2σ) = +5167
The deeply negative baseline is correct *for this env* even though it
doesn't reflect realistic random-policy P&L. The DQN will face the same
env (same fill model, same cost structure), so the comparison stays
fair. Fitter regularisation (to prevent β_spread runaway) is a Phase E.1
follow-up.
Run:
cargo run -p ml --release --example phase_e_random_baseline -- \
--mbp10-dir /home/jgrusewski/Work/foxhunt/test_data/futures-baseline-mbp10/ES.FUT \
--fill-coeffs config/ml/phase_e_fill_coeffs.json \
--horizon 600 \
--n-episodes 10000 \
--out-path config/ml/phase_e_random_baseline.json
env.reset_at also called by reset() (1-line refactor); no behavior change.
This commit is contained in:
349
crates/ml/examples/phase_e_random_baseline.rs
Normal file
349
crates/ml/examples/phase_e_random_baseline.rs
Normal file
@@ -0,0 +1,349 @@
|
||||
//! Phase E.0 Task 7 — random-uniform policy reward baseline.
|
||||
//!
|
||||
//! Loads MBP-10 snapshots into the Phase E `ExecutionEnv` with a `FillModel`
|
||||
//! deserialized from JSON (the artifact produced by Task 5b), runs N random-
|
||||
//! uniform episodes from randomly-sampled starting cursors, and reports
|
||||
//! mean / std / percentile distribution of terminal reward. This baseline
|
||||
//! defines the kill-criterion threshold the E.1 DQN smoke must exceed
|
||||
//! (`mean + 2σ` — written to ISV slots 547/548 in the trainer).
|
||||
//!
|
||||
//! ## Run
|
||||
//!
|
||||
//! ```bash
|
||||
//! SQLX_OFFLINE=true cargo run -p ml --release --example phase_e_random_baseline -- \
|
||||
//! --mbp10-dir /home/jgrusewski/Work/foxhunt/test_data/futures-baseline-mbp10/ES.FUT \
|
||||
//! --fill-coeffs config/ml/phase_e_fill_coeffs.json \
|
||||
//! --horizon 600 \
|
||||
//! --n-episodes 10000 \
|
||||
//! --out-path config/ml/phase_e_random_baseline.json
|
||||
//! ```
|
||||
//!
|
||||
//! ## L1-only data limitation (inherited)
|
||||
//!
|
||||
//! Same as the fit binary: `parse_mbp10_streaming` populates only L1. L2/L3
|
||||
//! posting prices in the synthesized `SnapshotRow` are L1 ± 0.25-tick
|
||||
//! offsets (ES futures tick size). Fill probabilities at L2/L3 still derive
|
||||
//! from the JSON-supplied coefficients (which themselves are L1 with
|
||||
//! attenuated β_0).
|
||||
|
||||
use std::fs::File;
|
||||
use std::io::Write;
|
||||
use std::path::PathBuf;
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use clap::Parser;
|
||||
use tracing::{info, warn};
|
||||
|
||||
use data::providers::databento::{
|
||||
dbn_parser::DbnParser,
|
||||
mbp10::Mbp10Snapshot,
|
||||
};
|
||||
use ml::env::action_space::N_ACTIONS;
|
||||
use ml::env::execution_env::{
|
||||
EpisodeState, ExecutionEnv, ExecutionEnvConfig, ReplayRng, SnapshotRow,
|
||||
};
|
||||
use ml::env::fill_model::{FillCoeffs, FillModel};
|
||||
use ml::trainers::dqn::collect_dbn_files_recursive;
|
||||
|
||||
/// Convert raw `Mbp10Msg.price` (i64 fixed-point at 1e9 scale per the
|
||||
/// databento DBN format) to f32 in price units. See
|
||||
/// `phase_e_fit_fill_model.rs` for the rationale (mismatch with
|
||||
/// `BidAskPair::price_to_f64`'s 1e12 convention).
|
||||
#[inline]
|
||||
fn raw_price_to_f32(fixed: i64) -> f32 {
|
||||
(fixed as f64 * 1e-9) as f32
|
||||
}
|
||||
|
||||
/// ES futures minimum price increment.
|
||||
const TICK: f32 = 0.25;
|
||||
|
||||
#[derive(Debug, Parser)]
|
||||
#[command(
|
||||
name = "phase_e_random_baseline",
|
||||
about = "Phase E.0 Task 7 — random-uniform policy reward baseline"
|
||||
)]
|
||||
struct Cli {
|
||||
/// Directory containing MBP-10 `.dbn[.zst]` files.
|
||||
#[arg(long)]
|
||||
mbp10_dir: PathBuf,
|
||||
/// Path to fitted FillModel JSON (from Task 5b).
|
||||
#[arg(long, default_value = "config/ml/phase_e_fill_coeffs.json")]
|
||||
fill_coeffs: PathBuf,
|
||||
/// Episode horizon in snapshots.
|
||||
#[arg(long, default_value_t = 600)]
|
||||
horizon: usize,
|
||||
/// Number of random episodes to simulate.
|
||||
#[arg(long, default_value_t = 10_000)]
|
||||
n_episodes: usize,
|
||||
/// Fixed trade size per fill.
|
||||
#[arg(long, default_value_t = 1)]
|
||||
trade_size: i32,
|
||||
/// Per-contract round-turn cost in price units.
|
||||
#[arg(long, default_value_t = 0.0625)]
|
||||
cost_per_contract: f32,
|
||||
/// Master RNG seed.
|
||||
#[arg(long, default_value_t = 0xCAFEBABE_u64)]
|
||||
seed: u64,
|
||||
/// MBP-10 events between snapshot emissions.
|
||||
#[arg(long, default_value_t = 50)]
|
||||
snapshot_interval: usize,
|
||||
/// Cap on snapshots loaded (memory budget).
|
||||
#[arg(long, default_value_t = 500_000)]
|
||||
max_snapshots: usize,
|
||||
/// Output JSON path.
|
||||
#[arg(long, default_value = "config/ml/phase_e_random_baseline.json")]
|
||||
out_path: PathBuf,
|
||||
}
|
||||
|
||||
fn load_fill_model(path: &std::path::Path) -> Result<FillModel> {
|
||||
let s = std::fs::read_to_string(path)
|
||||
.with_context(|| format!("reading fill coeffs JSON at {}", path.display()))?;
|
||||
let v: serde_json::Value = serde_json::from_str(&s)
|
||||
.with_context(|| format!("parsing fill coeffs JSON at {}", path.display()))?;
|
||||
let parse_levels = |key: &str| -> Result<[FillCoeffs; 3]> {
|
||||
let arr = v[key]
|
||||
.as_array()
|
||||
.ok_or_else(|| anyhow::anyhow!("missing/non-array field `{}`", key))?;
|
||||
if arr.len() != 3 {
|
||||
anyhow::bail!("`{}` must have 3 entries, got {}", key, arr.len());
|
||||
}
|
||||
let mut out: [FillCoeffs; 3] = [FillCoeffs { beta: [0.0; 5] }; 3];
|
||||
for (i, lvl) in arr.iter().enumerate() {
|
||||
let vv = lvl
|
||||
.as_array()
|
||||
.ok_or_else(|| anyhow::anyhow!("`{}`[{}] not array", key, i))?;
|
||||
if vv.len() != 5 {
|
||||
anyhow::bail!("`{}`[{}] must have 5 floats, got {}", key, i, vv.len());
|
||||
}
|
||||
for k in 0..5 {
|
||||
out[i].beta[k] = vv[k]
|
||||
.as_f64()
|
||||
.ok_or_else(|| anyhow::anyhow!("`{}`[{}][{}] not number", key, i, k))?
|
||||
as f32;
|
||||
}
|
||||
}
|
||||
Ok(out)
|
||||
};
|
||||
Ok(FillModel {
|
||||
bid_coeffs: parse_levels("bid_coeffs")?,
|
||||
ask_coeffs: parse_levels("ask_coeffs")?,
|
||||
})
|
||||
}
|
||||
|
||||
fn main() -> Result<()> {
|
||||
tracing_subscriber::fmt()
|
||||
.with_env_filter(
|
||||
tracing_subscriber::EnvFilter::try_from_default_env()
|
||||
.unwrap_or_else(|_| tracing_subscriber::EnvFilter::new("info")),
|
||||
)
|
||||
.init();
|
||||
let cli = Cli::parse();
|
||||
info!("Phase E.0 Task 7 — random-uniform baseline starting");
|
||||
info!(" mbp10_dir = {}", cli.mbp10_dir.display());
|
||||
info!(" fill_coeffs = {}", cli.fill_coeffs.display());
|
||||
info!(" horizon = {}", cli.horizon);
|
||||
info!(" n_episodes = {}", cli.n_episodes);
|
||||
info!(" trade_size = {}", cli.trade_size);
|
||||
info!(" cost = {:.4}", cli.cost_per_contract);
|
||||
info!(" seed = {:#x}", cli.seed);
|
||||
|
||||
// --- 1. Load FillModel ---
|
||||
let fill_model = load_fill_model(&cli.fill_coeffs)?;
|
||||
info!(
|
||||
"Loaded FillModel from {} (bid β_0={:+.3}, ask β_0={:+.3})",
|
||||
cli.fill_coeffs.display(),
|
||||
fill_model.bid_coeffs[0].beta[0],
|
||||
fill_model.ask_coeffs[0].beta[0],
|
||||
);
|
||||
|
||||
// --- 2. Stream MBP-10 → env::SnapshotRow ---
|
||||
let parser = DbnParser::new().context("DbnParser::new")?;
|
||||
let files = collect_dbn_files_recursive(&cli.mbp10_dir);
|
||||
if files.is_empty() {
|
||||
anyhow::bail!("no MBP-10 .dbn[.zst] files in {:?}", cli.mbp10_dir);
|
||||
}
|
||||
info!("Found {} MBP-10 file(s)", files.len());
|
||||
|
||||
let mut rows: Vec<SnapshotRow> = Vec::with_capacity(cli.max_snapshots);
|
||||
'files: for file in &files {
|
||||
let mut hit_limit = false;
|
||||
let result =
|
||||
parser.parse_mbp10_streaming(file, cli.snapshot_interval, |snap: &Mbp10Snapshot| {
|
||||
if rows.len() >= cli.max_snapshots {
|
||||
hit_limit = true;
|
||||
return;
|
||||
}
|
||||
if snap.levels.is_empty() {
|
||||
return;
|
||||
}
|
||||
let l1 = &snap.levels[0];
|
||||
let bid_l1 = raw_price_to_f32(l1.bid_px);
|
||||
let ask_l1 = raw_price_to_f32(l1.ask_px);
|
||||
if bid_l1 <= 0.0 || ask_l1 <= 0.0 || bid_l1 >= ask_l1 {
|
||||
return;
|
||||
}
|
||||
let mid = 0.5 * (bid_l1 + ask_l1);
|
||||
let spread_bps = 10_000.0 * (ask_l1 - bid_l1) / mid;
|
||||
let bid_sz = l1.bid_sz as f32;
|
||||
let ask_sz = l1.ask_sz as f32;
|
||||
let l1_imbalance = if bid_sz + ask_sz > 0.0 {
|
||||
bid_sz / (bid_sz + ask_sz)
|
||||
} else {
|
||||
0.5
|
||||
};
|
||||
let ofi_sum_5 = bid_sz - ask_sz;
|
||||
// L2/L3 posting prices synthesised at ±tick offsets (parser
|
||||
// doesn't populate levels[1..10]; document limitation).
|
||||
let bid_l = [bid_l1, bid_l1 - TICK, bid_l1 - 2.0 * TICK];
|
||||
let ask_l = [ask_l1, ask_l1 + TICK, ask_l1 + 2.0 * TICK];
|
||||
|
||||
rows.push(SnapshotRow {
|
||||
mid_price: mid,
|
||||
bid_l,
|
||||
ask_l,
|
||||
alpha_logit: 0.0, // Random policy ignores state — placeholder.
|
||||
alpha_confidence: 0.5, // placeholder.
|
||||
spread_bps,
|
||||
l1_imbalance,
|
||||
ofi_sum_5,
|
||||
mid_drift_5: 0.0, // placeholder (no rolling history maintained).
|
||||
time_since_trade_s: 0.0, // no trade tape loaded in baseline.
|
||||
book_event_rate: 5.0, // plausible default, doesn't influence reward.
|
||||
});
|
||||
});
|
||||
match result {
|
||||
Ok(c) => info!(
|
||||
" {} → {} streamed snapshots",
|
||||
file.file_name().unwrap_or_default().to_string_lossy(),
|
||||
c
|
||||
),
|
||||
Err(e) => warn!("parser error on {}: {}", file.display(), e),
|
||||
}
|
||||
if hit_limit {
|
||||
info!("max_snapshots cap reached; stopping iteration");
|
||||
break 'files;
|
||||
}
|
||||
}
|
||||
info!("Loaded {} snapshots into env", rows.len());
|
||||
if rows.len() <= cli.horizon {
|
||||
anyhow::bail!(
|
||||
"insufficient snapshots ({}) for horizon ({}); reduce horizon or raise --max-snapshots",
|
||||
rows.len(),
|
||||
cli.horizon
|
||||
);
|
||||
}
|
||||
|
||||
// --- 3. Construct env, run random episodes ---
|
||||
let n_rows = rows.len();
|
||||
let mut env = ExecutionEnv::new(
|
||||
ExecutionEnvConfig {
|
||||
horizon_snapshots: cli.horizon,
|
||||
trade_size_contracts: cli.trade_size,
|
||||
cost_per_contract: cli.cost_per_contract,
|
||||
},
|
||||
fill_model,
|
||||
rows,
|
||||
cli.seed,
|
||||
);
|
||||
|
||||
let mut episode_rng = ReplayRng::new(cli.seed.wrapping_add(0xDEAD_BEEF));
|
||||
let mut rewards: Vec<f32> = Vec::with_capacity(cli.n_episodes);
|
||||
let max_start = n_rows.saturating_sub(cli.horizon + 1).max(1);
|
||||
let mut n_fills_total: u64 = 0;
|
||||
|
||||
for ep in 0..cli.n_episodes {
|
||||
let start_cursor = (episode_rng.next_u64() as usize) % max_start;
|
||||
let env_seed = episode_rng.next_u64();
|
||||
let mut action_rng = ReplayRng::new(episode_rng.next_u64());
|
||||
env.reset_at(env_seed, start_cursor);
|
||||
let mut state = EpisodeState::new();
|
||||
let mut terminal_reward = 0.0_f32;
|
||||
loop {
|
||||
let action = (action_rng.next_u64() % N_ACTIONS as u64) as u8;
|
||||
match env.step(action, &mut state) {
|
||||
Some((_, reward, done)) => {
|
||||
if done {
|
||||
terminal_reward = reward;
|
||||
break;
|
||||
}
|
||||
}
|
||||
None => break,
|
||||
}
|
||||
}
|
||||
rewards.push(terminal_reward);
|
||||
n_fills_total += state.n_fills as u64;
|
||||
if (ep + 1) % 1000 == 0 {
|
||||
info!(
|
||||
" episode {}/{} (latest_reward={:+.4}, fills={})",
|
||||
ep + 1,
|
||||
cli.n_episodes,
|
||||
terminal_reward,
|
||||
state.n_fills
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// --- 4. Statistics ---
|
||||
let n = rewards.len() as f64;
|
||||
let mean: f64 = rewards.iter().map(|r| *r as f64).sum::<f64>() / n;
|
||||
let var: f64 = rewards
|
||||
.iter()
|
||||
.map(|r| (*r as f64 - mean).powi(2))
|
||||
.sum::<f64>()
|
||||
/ n;
|
||||
let std = var.sqrt();
|
||||
let mut sorted = rewards.clone();
|
||||
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
|
||||
let pick = |q: f64| -> f32 {
|
||||
let idx = (q * (n - 1.0)).round() as usize;
|
||||
sorted[idx.min(sorted.len() - 1)]
|
||||
};
|
||||
let p05 = pick(0.05);
|
||||
let p25 = pick(0.25);
|
||||
let p50 = pick(0.50);
|
||||
let p75 = pick(0.75);
|
||||
let p95 = pick(0.95);
|
||||
let kill_threshold = mean + 2.0 * std;
|
||||
let avg_fills_per_ep = n_fills_total as f64 / n;
|
||||
|
||||
info!(
|
||||
"=== Random-uniform baseline ({} episodes, horizon {}) ===",
|
||||
cli.n_episodes, cli.horizon
|
||||
);
|
||||
info!(" mean = {:+.6}", mean);
|
||||
info!(" std = {:.6}", std);
|
||||
info!(" p05 = {:+.6}", p05);
|
||||
info!(" p25 = {:+.6}", p25);
|
||||
info!(" p50 (median) = {:+.6}", p50);
|
||||
info!(" p75 = {:+.6}", p75);
|
||||
info!(" p95 = {:+.6}", p95);
|
||||
info!(" kill (mean + 2σ) = {:+.6}", kill_threshold);
|
||||
info!(" avg fills/ep = {:.2}", avg_fills_per_ep);
|
||||
|
||||
// --- 5. Save JSON ---
|
||||
let json = serde_json::json!({
|
||||
"phase": "E.0 Task 7",
|
||||
"n_episodes": cli.n_episodes,
|
||||
"horizon": cli.horizon,
|
||||
"trade_size_contracts": cli.trade_size,
|
||||
"cost_per_contract": cli.cost_per_contract,
|
||||
"seed": cli.seed,
|
||||
"n_snapshots_loaded": n_rows,
|
||||
"snapshot_interval": cli.snapshot_interval,
|
||||
"fill_coeffs_path": cli.fill_coeffs.to_string_lossy(),
|
||||
"mean_reward": mean,
|
||||
"std_reward": std,
|
||||
"p05_reward": p05,
|
||||
"p25_reward": p25,
|
||||
"p50_reward": p50,
|
||||
"p75_reward": p75,
|
||||
"p95_reward": p95,
|
||||
"kill_threshold_mean_plus_2sigma": kill_threshold,
|
||||
"avg_fills_per_episode": avg_fills_per_ep,
|
||||
});
|
||||
let mut f = File::create(&cli.out_path).context("create out file")?;
|
||||
write!(f, "{}", serde_json::to_string_pretty(&json)?)?;
|
||||
info!("Wrote results to {}", cli.out_path.display());
|
||||
|
||||
Ok(())
|
||||
}
|
||||
9
crates/ml/src/env/execution_env.rs
vendored
9
crates/ml/src/env/execution_env.rs
vendored
@@ -139,7 +139,14 @@ impl ExecutionEnv {
|
||||
|
||||
/// Reset episode to step 0 with a fresh RNG seed.
|
||||
pub fn reset(&mut self, seed: u64) -> [f32; STATE_DIM] {
|
||||
self.cursor = 0;
|
||||
self.reset_at(seed, 0)
|
||||
}
|
||||
|
||||
/// Reset episode at the given starting cursor with a fresh RNG seed.
|
||||
/// Used by the random-baseline binary to sample episodes from random
|
||||
/// starting points in a long snapshot replay.
|
||||
pub fn reset_at(&mut self, seed: u64, start_cursor: usize) -> [f32; STATE_DIM] {
|
||||
self.cursor = start_cursor.min(self.snapshots.len().saturating_sub(1));
|
||||
self.rng = ReplayRng::new(seed);
|
||||
self.state(&EpisodeState::new())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user