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:
jgrusewski
2026-05-15 13:36:43 +02:00
parent 12151ccf6a
commit 9cfc6d8502
2 changed files with 357 additions and 1 deletions

View 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(())
}

View File

@@ -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())
}