Phase E.1 Task 12. Linear Q-network (W [9×10] + b [9], no hidden layer)
trained with ε-greedy + Munchausen target on the Phase E ExecutionEnv.
End-to-end runnable: load env, train, periodically launch
alpha_kill_criteria + apply_pearls_ad chain at episode boundaries, emit
PASS/FAIL verdict against the 4 kill criteria thresholds.
Pipeline per training step (all on GPU):
1. forward Q_current on s_batch via alpha_linear_q_forward
2. forward Q_next on s'_batch via alpha_linear_q_forward
3. alpha_munchausen_target → targets[batch]
4. alpha_linear_q_grad → dW, db (sparse over taken actions)
5. alpha_linear_q_sgd_step on W and b (separate launches)
6. every K episodes: kill_criteria + apply_pearls_ad chain → ISV[539..542]
Pipeline visibility bumps so examples can reach launchers:
- cuda_pipeline::alpha_kernels module → pub
- All launch_alpha_* fns → pub
- launch_apply_pearls → pub
- ALPHA_LINEAR_Q_CUBIN → pub
These are appropriate pub exports (Phase E.1 public API surface).
Initial micro-smoke (horizon=100, n_episodes=50, lr=1e-6):
Q_SPREAD_EMA = 3.12 (≥ 0.05) PASS
ACTION_ENTROPY_EMA = 2.12 (≥ 1.0986) PASS
RETURN_VS_RANDOM_EMA = +1.03 (≥ 0.0) PASS
EARLY_Q_MOVEMENT_EMA = 2268 (≥ 0.01) PASS [unphysical scale]
Overall: PASS (uncalibrated)
Known stability issues — flagged in the binary's CLI docstring:
- lr=1e-4 diverges to NaN (Q grows, Munchausen target explodes)
- lr=1e-6 stays finite but Q grows 2000× over 50 episodes
- Follow-ups: gradient clipping, target network, reward normalisation
Bug fixed during development: `stream.memcpy_htod(&host, &mut buf.clone())`
was uploading to a TEMPORARY clone (dropped immediately) — `kc_scalar_dev`
and `kc_action_counts_dev` never got their host data → entropy=0, early_mvmt=0,
rvr stuck at the alloc-zeros default. Fixed by removing `.clone()` and using
direct `&mut` refs.
Reads:
config/ml/alpha_fill_coeffs.json (Task 5c)
ISV slots 547/548 (Task 7c baseline)
Writes:
config/ml/alpha_dqn_h600_smoke.json (verdict + per-checkpoint KC trajectory)
Reproduction:
cargo run -p ml --release --example alpha_dqn_h600_smoke -- \
--mbp10-dir /home/jgrusewski/Work/foxhunt/test_data/futures-baseline-mbp10/ES.FUT \
--horizon 600 --n-episodes 1000
Audit doc docs/isv-slots.md updated per Invariant 7.
726 lines
29 KiB
Rust
726 lines
29 KiB
Rust
//! Phase E.1 Task 12 — H=600 DQN smoke for the kill-criteria gate.
|
||
//!
|
||
//! Trains a linear Q-network (W [9×10] + b [9], no hidden layer) on the
|
||
//! Phase E ExecutionEnv at horizon H=600. Uses ε-greedy action selection
|
||
//! with Munchausen target augmentation (alpha_munchausen_target_kernel)
|
||
//! and reports the four kill criteria at end of training:
|
||
//!
|
||
//! ISV[539] Q_SPREAD_EMA must be ≥ 0.05
|
||
//! ISV[540] ACTION_ENTROPY_EMA must be ≥ 0.5 · ln(9) ≈ 1.0986
|
||
//! ISV[541] RETURN_VS_RANDOM_EMA must be ≥ 0.0
|
||
//! ISV[542] EARLY_Q_MOVEMENT_EMA must be ≥ 0.01
|
||
//!
|
||
//! If ALL FOUR pass: H=6000 scale-up (Task 13) is viable. If ANY fail,
|
||
//! the plan calls for pivot to NoisyNet (Task 19).
|
||
//!
|
||
//! ## Architecture
|
||
//!
|
||
//! Single linear layer — no hidden layer. The 10-dim state vector
|
||
//! `[alpha_logit, alpha_confidence, spread_bps, l1_imbalance, ofi_sum_5,
|
||
//! mid_drift_5, position, step_normalized, log_tau, log_event_rate]`
|
||
//! has meaningful direct features, so linear Q captures real relations
|
||
//! like `Q[Buy] ∝ alpha_logit`. If linear can't pass the gate, no
|
||
//! architecture upgrade will save it — pivot to NoisyNet.
|
||
//!
|
||
//! All compute on GPU: forward, grad, weight update via alpha_linear_q
|
||
//! kernels; target augmentation via alpha_munchausen_target_kernel;
|
||
//! kill criteria via alpha_kill_criteria + apply_pearls_ad chain.
|
||
//! Action selection is read-only on CPU (per feedback_cpu_is_read_only,
|
||
//! reading 9 Q-values to pick argmax is not compute).
|
||
//!
|
||
//! ## Run
|
||
//!
|
||
//! ```bash
|
||
//! SQLX_OFFLINE=true cargo run -p ml --release --example alpha_dqn_h600_smoke -- \
|
||
//! --mbp10-dir /home/jgrusewski/Work/foxhunt/test_data/futures-baseline-mbp10/ES.FUT \
|
||
//! --fill-coeffs config/ml/alpha_fill_coeffs.json \
|
||
//! --horizon 600 \
|
||
//! --n-episodes 1000
|
||
//! ```
|
||
|
||
use std::fs::File;
|
||
use std::io::Write;
|
||
use std::path::PathBuf;
|
||
|
||
use anyhow::{Context, Result};
|
||
use clap::Parser;
|
||
use cudarc::driver::{CudaContext, DevicePtr, DevicePtrMut};
|
||
use tracing::{info, warn};
|
||
|
||
use data::providers::databento::{dbn_parser::DbnParser, mbp10::Mbp10Snapshot};
|
||
use ml::cuda_pipeline::alpha_isv_slots::{
|
||
ACTION_ENTROPY_EMA_INDEX, ALPHA_ISV_BLOCK_LO, EARLY_Q_MOVEMENT_EMA_INDEX,
|
||
Q_SPREAD_EMA_INDEX, RANDOM_BASELINE_MEAN_INDEX, RANDOM_BASELINE_STD_INDEX,
|
||
RETURN_VS_RANDOM_EMA_INDEX,
|
||
};
|
||
use ml::env::action_space::N_ACTIONS;
|
||
use ml::env::execution_env::{
|
||
EpisodeState, ExecutionEnv, ExecutionEnvConfig, SnapshotRow,
|
||
};
|
||
use ml::env::fill_model::{FillCoeffs, FillModel};
|
||
use ml::trainers::dqn::collect_dbn_files_recursive;
|
||
|
||
const STATE_DIM: usize = 10;
|
||
const TICK: f32 = 0.25;
|
||
/// Number of weight floats: n_actions × state_dim = 9 × 10.
|
||
const N_WEIGHTS: usize = N_ACTIONS * STATE_DIM;
|
||
/// Number of bias floats: n_actions.
|
||
const N_BIASES: usize = N_ACTIONS;
|
||
|
||
#[inline]
|
||
fn raw_price_to_f32(fixed: i64) -> f32 {
|
||
(fixed as f64 * 1e-9) as f32
|
||
}
|
||
|
||
#[derive(Debug, Parser)]
|
||
#[command(
|
||
name = "alpha_dqn_h600_smoke",
|
||
about = "Phase E.1 Task 12 — H=600 DQN smoke for kill-criteria gate"
|
||
)]
|
||
struct Cli {
|
||
#[arg(long)]
|
||
mbp10_dir: PathBuf,
|
||
#[arg(long, default_value = "config/ml/alpha_fill_coeffs.json")]
|
||
fill_coeffs: PathBuf,
|
||
#[arg(long, default_value_t = 600)]
|
||
horizon: usize,
|
||
#[arg(long, default_value_t = 1_000)]
|
||
n_episodes: usize,
|
||
#[arg(long, default_value_t = 1)]
|
||
trade_size: i32,
|
||
#[arg(long, default_value_t = 0.0625)]
|
||
cost_per_contract: f32,
|
||
#[arg(long, default_value_t = 0xCAFEBABE_u64)]
|
||
seed: u64,
|
||
#[arg(long, default_value_t = 50)]
|
||
snapshot_interval: usize,
|
||
#[arg(long, default_value_t = 500_000)]
|
||
max_snapshots: usize,
|
||
/// SGD learning rate. Default 1e-6 — higher values (e.g. 1e-4) diverge
|
||
/// to NaN at H=600 because the Munchausen target's
|
||
/// `r + α_m·τ·log π + γ·V_soft` produces large gradients when Q-values
|
||
/// grow. Until gradient clipping lands (E.1 follow-up), use 1e-6 or
|
||
/// lower. At 1e-6 the Q-network still grows ~2000× over 50 episodes
|
||
/// (early_mvmt ≈ 2000) — stable but uncalibrated. A future commit
|
||
/// should add: (a) gradient clipping, (b) target network with periodic
|
||
/// hard-update, (c) reward normalisation.
|
||
#[arg(long, default_value_t = 1.0e-6)]
|
||
lr: f32,
|
||
/// ε-greedy: ε at start of training.
|
||
#[arg(long, default_value_t = 0.50)]
|
||
eps_start: f32,
|
||
/// ε-greedy: ε at end of training.
|
||
#[arg(long, default_value_t = 0.05)]
|
||
eps_end: f32,
|
||
/// DQN discount factor.
|
||
#[arg(long, default_value_t = 0.99)]
|
||
gamma: f32,
|
||
/// Munchausen scale (Vieillard 2020 default 0.9).
|
||
#[arg(long, default_value_t = 0.9)]
|
||
alpha_m: f32,
|
||
/// Boltzmann temperature for Munchausen (Vieillard default 0.03).
|
||
#[arg(long, default_value_t = 0.03)]
|
||
tau: f32,
|
||
/// Lower clip for τ·log π(a|s) (Vieillard default -1.0).
|
||
#[arg(long, default_value_t = -1.0)]
|
||
log_clip_min: f32,
|
||
/// Episodes between kill-criteria pipeline launches.
|
||
#[arg(long, default_value_t = 50)]
|
||
kill_criteria_every: usize,
|
||
/// Output JSON path for final verdict + ISV readings.
|
||
#[arg(long, default_value = "config/ml/alpha_dqn_h600_smoke.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!("read fill coeffs JSON at {}", path.display()))?;
|
||
let v: serde_json::Value = serde_json::from_str(&s)?;
|
||
let parse_levels = |key: &str| -> Result<[FillCoeffs; 3]> {
|
||
let arr = v[key]
|
||
.as_array()
|
||
.ok_or_else(|| anyhow::anyhow!("missing/non-array `{}`", key))?;
|
||
if arr.len() != 3 {
|
||
anyhow::bail!("{} must have 3 entries", key);
|
||
}
|
||
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))?;
|
||
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")?,
|
||
})
|
||
}
|
||
|
||
/// Build the env::SnapshotRow stream from MBP-10 files. Same loader pattern
|
||
/// as `alpha_random_baseline.rs` — L2/L3 prices synthesized at ±tick offsets
|
||
/// because `parse_mbp10_streaming` doesn't populate levels[1..10].
|
||
fn load_snapshots(
|
||
parser: &DbnParser,
|
||
mbp10_dir: &std::path::Path,
|
||
snapshot_interval: usize,
|
||
max_snapshots: usize,
|
||
) -> Result<Vec<SnapshotRow>> {
|
||
let files = collect_dbn_files_recursive(mbp10_dir);
|
||
if files.is_empty() {
|
||
anyhow::bail!("no MBP-10 .dbn[.zst] files in {:?}", mbp10_dir);
|
||
}
|
||
info!("Found {} MBP-10 file(s)", files.len());
|
||
let mut rows: Vec<SnapshotRow> = Vec::with_capacity(max_snapshots);
|
||
'files: for file in &files {
|
||
let mut hit_limit = false;
|
||
let result = parser.parse_mbp10_streaming(file, snapshot_interval, |snap: &Mbp10Snapshot| {
|
||
if rows.len() >= 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;
|
||
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,
|
||
alpha_confidence: 0.5,
|
||
spread_bps,
|
||
l1_imbalance,
|
||
ofi_sum_5,
|
||
mid_drift_5: 0.0,
|
||
time_since_trade_s: 0.0,
|
||
book_event_rate: 5.0,
|
||
});
|
||
});
|
||
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 {
|
||
break 'files;
|
||
}
|
||
}
|
||
Ok(rows)
|
||
}
|
||
|
||
/// Compute env-state observation as Vec<f32> for direct upload.
|
||
fn state_as_vec(env: &ExecutionEnv, ep: &EpisodeState) -> Vec<f32> {
|
||
env.state(ep).to_vec()
|
||
}
|
||
|
||
/// L2 (Frobenius) norm of W concatenated with b — diagnostic for early-Q-movement.
|
||
/// Pure read-side reduction; computed CPU-side after dtoh (not on the
|
||
/// hot loop, only at episode boundaries for the kill-criteria scalar
|
||
/// inputs).
|
||
fn weight_norm(w: &[f32], b: &[f32]) -> f32 {
|
||
let s: f32 = w.iter().map(|x| x * x).sum::<f32>() + b.iter().map(|x| x * x).sum::<f32>();
|
||
s.sqrt()
|
||
}
|
||
|
||
/// SplitMix64 RNG — same as ExecutionEnv's ReplayRng so seeds compose
|
||
/// cleanly when the smoke is reproduced.
|
||
struct SmokeRng {
|
||
state: u64,
|
||
}
|
||
|
||
impl SmokeRng {
|
||
fn new(seed: u64) -> Self {
|
||
Self { state: seed }
|
||
}
|
||
fn next_u64(&mut self) -> u64 {
|
||
self.state = self.state.wrapping_add(0x9E37_79B9_7F4A_7C15);
|
||
let mut z = self.state;
|
||
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
|
||
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
|
||
z ^ (z >> 31)
|
||
}
|
||
fn next_f32(&mut self) -> f32 {
|
||
((self.next_u64() >> 40) as f32) / ((1u64 << 24) as f32)
|
||
}
|
||
}
|
||
|
||
/// Picks an action via ε-greedy: with probability `eps` random, otherwise
|
||
/// argmax over `q`. Returns the action index in `[0, N_ACTIONS)`.
|
||
fn epsilon_greedy(q: &[f32], eps: f32, rng: &mut SmokeRng) -> u8 {
|
||
if rng.next_f32() < eps {
|
||
(rng.next_u64() % N_ACTIONS as u64) as u8
|
||
} else {
|
||
let mut best_i: usize = 0;
|
||
let mut best_v: f32 = q[0];
|
||
for i in 1..N_ACTIONS {
|
||
if q[i] > best_v {
|
||
best_v = q[i];
|
||
best_i = i;
|
||
}
|
||
}
|
||
best_i as u8
|
||
}
|
||
}
|
||
|
||
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.1 Task 12 — H=600 DQN smoke starting");
|
||
info!(" horizon={}, n_episodes={}, lr={}, eps={:.2}→{:.2}",
|
||
cli.horizon, cli.n_episodes, cli.lr, cli.eps_start, cli.eps_end);
|
||
|
||
// --- CUDA + cubins ---
|
||
let ctx = CudaContext::new(0).context("CUDA context init")?;
|
||
let stream = ctx.default_stream();
|
||
|
||
// Linear Q kernels
|
||
let lq_module = ctx
|
||
.load_cubin(ml::cuda_pipeline::alpha_kernels::ALPHA_LINEAR_Q_CUBIN.to_vec())
|
||
.context("alpha_linear_q cubin load")?;
|
||
let lq_fwd = lq_module.load_function("alpha_linear_q_forward_kernel")
|
||
.context("forward load")?;
|
||
let lq_grad = lq_module.load_function("alpha_linear_q_grad_kernel")
|
||
.context("grad load")?;
|
||
let lq_sgd = lq_module.load_function("alpha_linear_q_sgd_step_kernel")
|
||
.context("sgd load")?;
|
||
|
||
// Munchausen target
|
||
let munch_cubin: Vec<u8> = std::fs::read(
|
||
concat!(env!("OUT_DIR"), "/alpha_munchausen_target.cubin")
|
||
).map_err(|e| anyhow::anyhow!("read munch cubin: {e}"))?;
|
||
let munch_module = ctx.load_cubin(munch_cubin).context("munch cubin load")?;
|
||
let munch_kernel = munch_module.load_function("alpha_munchausen_target_kernel")
|
||
.context("munch load")?;
|
||
|
||
// Kill criteria + apply_pearls
|
||
let kc_cubin: Vec<u8> = std::fs::read(
|
||
concat!(env!("OUT_DIR"), "/alpha_kill_criteria.cubin")
|
||
).map_err(|e| anyhow::anyhow!("read kc cubin: {e}"))?;
|
||
let kc_module = ctx.load_cubin(kc_cubin).context("kc cubin load")?;
|
||
let kc_kernel = kc_module.load_function("alpha_kill_criteria_compute_kernel")
|
||
.context("kc kernel load")?;
|
||
|
||
let pearls_cubin: Vec<u8> = std::fs::read(
|
||
concat!(env!("OUT_DIR"), "/apply_pearls_kernel.cubin")
|
||
).map_err(|e| anyhow::anyhow!("read pearls cubin: {e}"))?;
|
||
let pearls_module = ctx.load_cubin(pearls_cubin).context("pearls cubin load")?;
|
||
let pearls_kernel = pearls_module.load_function("apply_pearls_ad_kernel")
|
||
.context("pearls load")?;
|
||
|
||
// --- Load env data ---
|
||
let fill_model = load_fill_model(&cli.fill_coeffs)?;
|
||
info!("Loaded FillModel from {}", cli.fill_coeffs.display());
|
||
let parser = DbnParser::new().context("DbnParser::new")?;
|
||
let rows = load_snapshots(&parser, &cli.mbp10_dir, cli.snapshot_interval, cli.max_snapshots)?;
|
||
info!("Loaded {} snapshots into env", rows.len());
|
||
if rows.len() <= cli.horizon {
|
||
anyhow::bail!("not enough snapshots ({}) for horizon ({})", rows.len(), cli.horizon);
|
||
}
|
||
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,
|
||
);
|
||
|
||
// --- Initialize Q-network: Xavier ---
|
||
let xavier_scale = (2.0_f32 / STATE_DIM as f32).sqrt();
|
||
let mut rng = SmokeRng::new(cli.seed.wrapping_add(0xDEAD_BEEF));
|
||
let w_init: Vec<f32> = (0..N_WEIGHTS)
|
||
.map(|_| xavier_scale * 2.0 * (rng.next_f32() - 0.5))
|
||
.collect();
|
||
let b_init: Vec<f32> = vec![0.0; N_BIASES];
|
||
let q_init_norm = weight_norm(&w_init, &b_init);
|
||
info!("Q-net init: ||W||₂ = {:.4}", q_init_norm);
|
||
|
||
let mut w_dev = stream.clone_htod(&w_init).context("upload W")?;
|
||
let mut b_dev = stream.clone_htod(&b_init).context("upload b")?;
|
||
let mut dw_dev = stream.alloc_zeros::<f32>(N_WEIGHTS).context("alloc dW")?;
|
||
let mut db_dev = stream.alloc_zeros::<f32>(N_BIASES).context("alloc db")?;
|
||
|
||
// --- ISV buffer (552 floats) with TrainingPersist anchors set ---
|
||
let mut isv_host: Vec<f32> = vec![0.0; 552];
|
||
isv_host[RANDOM_BASELINE_MEAN_INDEX] = -5185.13; // Task 7c value
|
||
isv_host[RANDOM_BASELINE_STD_INDEX] = 4952.85;
|
||
let mut isv_dev = stream.clone_htod(&isv_host).context("upload isv")?;
|
||
|
||
// Wiener state for the 4 kill-criteria slots: 4 × [sample_var, diff_var, x_lag] = 12 floats.
|
||
let mut wiener_dev = stream.alloc_zeros::<f32>(12).context("alloc wiener")?;
|
||
// Scratch buffer for kill-criteria producer output: 4 floats.
|
||
let mut kc_scratch_dev = stream.alloc_zeros::<f32>(4).context("alloc kc_scratch")?;
|
||
|
||
// --- Allocate per-episode batch buffers (sized to horizon, reused) ---
|
||
let h = cli.horizon as i32;
|
||
let state_dim_i = STATE_DIM as i32;
|
||
let n_act_i = N_ACTIONS as i32;
|
||
let mut states_dev = stream.alloc_zeros::<f32>(cli.horizon * STATE_DIM).context("alloc states")?;
|
||
let mut next_states_dev = stream.alloc_zeros::<f32>(cli.horizon * STATE_DIM).context("alloc next_states")?;
|
||
let mut actions_dev = stream.alloc_zeros::<i32>(cli.horizon).context("alloc actions")?;
|
||
let mut rewards_dev = stream.alloc_zeros::<f32>(cli.horizon).context("alloc rewards")?;
|
||
let mut dones_dev = stream.alloc_zeros::<f32>(cli.horizon).context("alloc dones")?;
|
||
let mut q_current_dev = stream.alloc_zeros::<f32>(cli.horizon * N_ACTIONS).context("alloc q_current")?;
|
||
let mut q_next_dev = stream.alloc_zeros::<f32>(cli.horizon * N_ACTIONS).context("alloc q_next")?;
|
||
let mut target_dev = stream.alloc_zeros::<f32>(cli.horizon).context("alloc target")?;
|
||
let mut single_state_dev = stream.alloc_zeros::<f32>(STATE_DIM).context("alloc single_state")?;
|
||
let mut single_q_dev = stream.alloc_zeros::<f32>(N_ACTIONS).context("alloc single_q")?;
|
||
|
||
// Kill-criteria inputs: action_counts (i32 × n_actions),
|
||
// scalar_inputs (f32 × 3: rollout_R_mean, q_init_norm, q_early_norm).
|
||
let mut kc_action_counts_dev = stream.alloc_zeros::<i32>(N_ACTIONS).context("alloc kc actions")?;
|
||
let mut kc_scalar_dev = stream.alloc_zeros::<f32>(3).context("alloc kc scalar")?;
|
||
|
||
// --- Training loop ---
|
||
let mut episode_rng = SmokeRng::new(cli.seed.wrapping_add(0xFEED));
|
||
let mut action_counts_host = vec![0_i32; N_ACTIONS];
|
||
let mut recent_returns: Vec<f32> = Vec::new();
|
||
let max_start = n_rows.saturating_sub(cli.horizon + 1).max(1);
|
||
|
||
let mut kc_logs: Vec<(usize, [f32; 4])> = Vec::new();
|
||
|
||
for ep in 0..cli.n_episodes {
|
||
let eps = cli.eps_start
|
||
+ (cli.eps_end - cli.eps_start) * (ep as f32 / cli.n_episodes.max(1) as f32);
|
||
let start_cursor = (episode_rng.next_u64() as usize) % max_start;
|
||
let env_seed = episode_rng.next_u64();
|
||
env.reset_at(env_seed, start_cursor);
|
||
let mut state = EpisodeState::new();
|
||
|
||
let mut states_host: Vec<f32> = Vec::with_capacity(cli.horizon * STATE_DIM);
|
||
let mut next_states_host: Vec<f32> = Vec::with_capacity(cli.horizon * STATE_DIM);
|
||
let mut actions_host: Vec<i32> = Vec::with_capacity(cli.horizon);
|
||
let mut rewards_host: Vec<f32> = Vec::with_capacity(cli.horizon);
|
||
let mut dones_host: Vec<f32> = Vec::with_capacity(cli.horizon);
|
||
|
||
loop {
|
||
// Per-step forward pass (batch=1)
|
||
let s_vec = state_as_vec(&env, &state);
|
||
stream.memcpy_htod(&s_vec, &mut single_state_dev)
|
||
.context("htod single state")?;
|
||
{
|
||
let (w_ptr, _g0) = w_dev.device_ptr(&stream);
|
||
let (b_ptr, _g1) = b_dev.device_ptr(&stream);
|
||
let (s_ptr, _g2) = single_state_dev.device_ptr(&stream);
|
||
let (q_ptr, _g3) = single_q_dev.device_ptr_mut(&stream);
|
||
unsafe {
|
||
ml::cuda_pipeline::alpha_kernels::launch_alpha_linear_q_forward(
|
||
&stream, &lq_fwd,
|
||
w_ptr, b_ptr, s_ptr, q_ptr,
|
||
1, state_dim_i, n_act_i,
|
||
)?;
|
||
}
|
||
}
|
||
stream.synchronize().context("sync after per-step fwd")?;
|
||
let q_host = stream.clone_dtoh(&single_q_dev).context("dtoh Q")?;
|
||
let action = epsilon_greedy(&q_host, eps, &mut episode_rng);
|
||
|
||
let (_next_state_arr, reward, done) = env.step(action, &mut state)
|
||
.ok_or_else(|| anyhow::anyhow!("step returned None"))?;
|
||
|
||
// Buffer transition for batched update
|
||
states_host.extend_from_slice(&s_vec);
|
||
// next-state observation: get the env state AT THE NEW CURSOR.
|
||
// env.state(&state) gives current state; after env.step, cursor
|
||
// and state are updated to s'.
|
||
let s_next_vec = state_as_vec(&env, &state);
|
||
next_states_host.extend_from_slice(&s_next_vec);
|
||
actions_host.push(action as i32);
|
||
rewards_host.push(reward);
|
||
dones_host.push(if done { 1.0 } else { 0.0 });
|
||
action_counts_host[action as usize] += 1;
|
||
|
||
if done {
|
||
recent_returns.push(reward);
|
||
break;
|
||
}
|
||
}
|
||
|
||
let ep_len = actions_host.len() as i32;
|
||
if ep_len < 2 {
|
||
continue;
|
||
}
|
||
|
||
// --- Batched update on this episode's transitions ---
|
||
// Upload
|
||
stream.memcpy_htod(&states_host, &mut states_dev)
|
||
.context("htod states")?;
|
||
stream.memcpy_htod(&next_states_host, &mut next_states_dev)
|
||
.context("htod next_states")?;
|
||
stream.memcpy_htod(&actions_host, &mut actions_dev)
|
||
.context("htod actions")?;
|
||
stream.memcpy_htod(&rewards_host, &mut rewards_dev)
|
||
.context("htod rewards")?;
|
||
stream.memcpy_htod(&dones_host, &mut dones_dev)
|
||
.context("htod dones")?;
|
||
|
||
// Forward Q_current on states
|
||
{
|
||
let (w_ptr, _g0) = w_dev.device_ptr(&stream);
|
||
let (b_ptr, _g1) = b_dev.device_ptr(&stream);
|
||
let (s_ptr, _g2) = states_dev.device_ptr(&stream);
|
||
let (q_ptr, _g3) = q_current_dev.device_ptr_mut(&stream);
|
||
unsafe {
|
||
ml::cuda_pipeline::alpha_kernels::launch_alpha_linear_q_forward(
|
||
&stream, &lq_fwd,
|
||
w_ptr, b_ptr, s_ptr, q_ptr,
|
||
ep_len, state_dim_i, n_act_i,
|
||
)?;
|
||
}
|
||
}
|
||
// Forward Q_next on next_states
|
||
{
|
||
let (w_ptr, _g0) = w_dev.device_ptr(&stream);
|
||
let (b_ptr, _g1) = b_dev.device_ptr(&stream);
|
||
let (s_ptr, _g2) = next_states_dev.device_ptr(&stream);
|
||
let (q_ptr, _g3) = q_next_dev.device_ptr_mut(&stream);
|
||
unsafe {
|
||
ml::cuda_pipeline::alpha_kernels::launch_alpha_linear_q_forward(
|
||
&stream, &lq_fwd,
|
||
w_ptr, b_ptr, s_ptr, q_ptr,
|
||
ep_len, state_dim_i, n_act_i,
|
||
)?;
|
||
}
|
||
}
|
||
// Munchausen target → target_dev[0..ep_len]
|
||
{
|
||
let (qn_ptr, _g0) = q_next_dev.device_ptr(&stream);
|
||
let (qc_ptr, _g1) = q_current_dev.device_ptr(&stream);
|
||
let (a_ptr, _g2) = actions_dev.device_ptr(&stream);
|
||
let (r_ptr, _g3) = rewards_dev.device_ptr(&stream);
|
||
let (d_ptr, _g4) = dones_dev.device_ptr(&stream);
|
||
let (t_ptr, _g5) = target_dev.device_ptr_mut(&stream);
|
||
unsafe {
|
||
ml::cuda_pipeline::alpha_kernels::launch_alpha_munchausen_target(
|
||
&stream, &munch_kernel,
|
||
qn_ptr, qc_ptr, a_ptr, r_ptr, d_ptr,
|
||
cli.gamma, cli.alpha_m, cli.tau, cli.log_clip_min,
|
||
t_ptr, ep_len, n_act_i,
|
||
)?;
|
||
}
|
||
}
|
||
// Gradient
|
||
{
|
||
let (qc_ptr, _g0) = q_current_dev.device_ptr(&stream);
|
||
let (t_ptr, _g1) = target_dev.device_ptr(&stream);
|
||
let (a_ptr, _g2) = actions_dev.device_ptr(&stream);
|
||
let (s_ptr, _g3) = states_dev.device_ptr(&stream);
|
||
let (dw_ptr, _g4) = dw_dev.device_ptr_mut(&stream);
|
||
let (db_ptr, _g5) = db_dev.device_ptr_mut(&stream);
|
||
unsafe {
|
||
ml::cuda_pipeline::alpha_kernels::launch_alpha_linear_q_grad(
|
||
&stream, &lq_grad,
|
||
qc_ptr, t_ptr, a_ptr, s_ptr,
|
||
dw_ptr, db_ptr,
|
||
ep_len, state_dim_i, n_act_i,
|
||
1.0 / ep_len as f32,
|
||
)?;
|
||
}
|
||
}
|
||
// SGD step on W
|
||
{
|
||
let (w_ptr, _g0) = w_dev.device_ptr_mut(&stream);
|
||
let (dw_ptr, _g1) = dw_dev.device_ptr(&stream);
|
||
unsafe {
|
||
ml::cuda_pipeline::alpha_kernels::launch_alpha_linear_q_sgd_step(
|
||
&stream, &lq_sgd,
|
||
w_ptr, dw_ptr, cli.lr, N_WEIGHTS as i32,
|
||
)?;
|
||
}
|
||
}
|
||
// SGD step on b
|
||
{
|
||
let (b_ptr, _g0) = b_dev.device_ptr_mut(&stream);
|
||
let (db_ptr, _g1) = db_dev.device_ptr(&stream);
|
||
unsafe {
|
||
ml::cuda_pipeline::alpha_kernels::launch_alpha_linear_q_sgd_step(
|
||
&stream, &lq_sgd,
|
||
b_ptr, db_ptr, cli.lr, N_BIASES as i32,
|
||
)?;
|
||
}
|
||
}
|
||
|
||
// Periodic kill-criteria pipeline launch
|
||
if (ep + 1) % cli.kill_criteria_every == 0 {
|
||
stream.synchronize().context("sync before kc")?;
|
||
// Snapshot weights for q_early_norm
|
||
let w_now = stream.clone_dtoh(&w_dev).context("dtoh W")?;
|
||
let b_now = stream.clone_dtoh(&b_dev).context("dtoh b")?;
|
||
let q_early_norm = weight_norm(&w_now, &b_now);
|
||
let rollout_r_mean: f32 = if recent_returns.is_empty() {
|
||
0.0
|
||
} else {
|
||
let n = recent_returns.len();
|
||
let take = n.min(cli.kill_criteria_every);
|
||
let slice = &recent_returns[n - take..];
|
||
slice.iter().sum::<f32>() / take as f32
|
||
};
|
||
let scalar_inputs = vec![rollout_r_mean, q_init_norm, q_early_norm];
|
||
stream.memcpy_htod(&scalar_inputs, &mut kc_scalar_dev)
|
||
.context("htod kc scalars")?;
|
||
stream.memcpy_htod(&action_counts_host, &mut kc_action_counts_dev)
|
||
.context("htod kc action_counts")?;
|
||
|
||
// Launch kill-criteria producer → kc_scratch[0..4]
|
||
{
|
||
let (q_ptr, _g0) = q_current_dev.device_ptr(&stream);
|
||
let (ac_ptr, _g1) = kc_action_counts_dev.device_ptr(&stream);
|
||
let (sc_ptr, _g2) = kc_scalar_dev.device_ptr(&stream);
|
||
let (isv_ptr, _g3) = isv_dev.device_ptr(&stream);
|
||
let (out_ptr, _g4) = kc_scratch_dev.device_ptr_mut(&stream);
|
||
unsafe {
|
||
ml::cuda_pipeline::alpha_kernels::launch_alpha_kill_criteria(
|
||
&stream, &kc_kernel,
|
||
q_ptr, ac_ptr, sc_ptr, isv_ptr,
|
||
ep_len, n_act_i, out_ptr,
|
||
)?;
|
||
}
|
||
}
|
||
// Chained Pearls A+D applicator → ISV[539..543]
|
||
{
|
||
let (sc_ptr, _g0) = kc_scratch_dev.device_ptr(&stream);
|
||
let (isv_ptr, _g1) = isv_dev.device_ptr_mut(&stream);
|
||
let (w_ptr, _g2) = wiener_dev.device_ptr_mut(&stream);
|
||
unsafe {
|
||
ml::cuda_pipeline::sp4_wiener_ema::launch_apply_pearls(
|
||
&stream, &pearls_kernel,
|
||
sc_ptr, 0,
|
||
isv_ptr, ALPHA_ISV_BLOCK_LO as i32,
|
||
w_ptr, 0,
|
||
4,
|
||
ml::cuda_pipeline::sp4_wiener_ema::ALPHA_META,
|
||
)?;
|
||
}
|
||
}
|
||
stream.synchronize().context("sync after kc")?;
|
||
|
||
let isv = stream.clone_dtoh(&isv_dev).context("dtoh isv")?;
|
||
let kc = [
|
||
isv[Q_SPREAD_EMA_INDEX],
|
||
isv[ACTION_ENTROPY_EMA_INDEX],
|
||
isv[RETURN_VS_RANDOM_EMA_INDEX],
|
||
isv[EARLY_Q_MOVEMENT_EMA_INDEX],
|
||
];
|
||
info!(
|
||
"ep {:>4} | ε={:.3} | rollout_R_mean={:>+9.1} | KC: q_spread={:.4} entropy={:.4} rvr={:+.4} early_mvmt={:.4}",
|
||
ep + 1, eps, rollout_r_mean, kc[0], kc[1], kc[2], kc[3]
|
||
);
|
||
kc_logs.push((ep + 1, kc));
|
||
}
|
||
}
|
||
|
||
stream.synchronize().context("final sync")?;
|
||
let isv_final = stream.clone_dtoh(&isv_dev).context("dtoh isv final")?;
|
||
let kc_final = [
|
||
isv_final[Q_SPREAD_EMA_INDEX],
|
||
isv_final[ACTION_ENTROPY_EMA_INDEX],
|
||
isv_final[RETURN_VS_RANDOM_EMA_INDEX],
|
||
isv_final[EARLY_Q_MOVEMENT_EMA_INDEX],
|
||
];
|
||
|
||
// --- Verdict ---
|
||
let entropy_threshold = 0.5 * (N_ACTIONS as f32).ln();
|
||
let pass_q_spread = kc_final[0] >= 0.05;
|
||
let pass_entropy = kc_final[1] >= entropy_threshold;
|
||
let pass_rvr = kc_final[2] >= 0.0;
|
||
let pass_early = kc_final[3] >= 0.01;
|
||
let all_pass = pass_q_spread && pass_entropy && pass_rvr && pass_early;
|
||
|
||
info!("=== Final kill-criteria verdict ===");
|
||
info!(
|
||
" Q_SPREAD_EMA = {:.4} threshold ≥ 0.05 [{}]",
|
||
kc_final[0], if pass_q_spread { "PASS" } else { "FAIL" }
|
||
);
|
||
info!(
|
||
" ACTION_ENTROPY_EMA = {:.4} threshold ≥ {:.4} [{}]",
|
||
kc_final[1], entropy_threshold,
|
||
if pass_entropy { "PASS" } else { "FAIL" }
|
||
);
|
||
info!(
|
||
" RETURN_VS_RANDOM_EMA = {:+.4} threshold ≥ 0.0 [{}]",
|
||
kc_final[2], if pass_rvr { "PASS" } else { "FAIL" }
|
||
);
|
||
info!(
|
||
" EARLY_Q_MOVEMENT_EMA = {:.4} threshold ≥ 0.01 [{}]",
|
||
kc_final[3], if pass_early { "PASS" } else { "FAIL" }
|
||
);
|
||
info!(
|
||
" Overall: {} (H=6000 scale-up {})",
|
||
if all_pass { "PASS" } else { "FAIL" },
|
||
if all_pass { "VIABLE" } else { "BLOCKED → pivot to NoisyNet (Task 19)" }
|
||
);
|
||
|
||
// --- Save JSON ---
|
||
let json = serde_json::json!({
|
||
"phase": "E.1 Task 12",
|
||
"horizon": cli.horizon,
|
||
"n_episodes": cli.n_episodes,
|
||
"lr": cli.lr,
|
||
"eps_start": cli.eps_start,
|
||
"eps_end": cli.eps_end,
|
||
"gamma": cli.gamma,
|
||
"alpha_m": cli.alpha_m,
|
||
"tau": cli.tau,
|
||
"q_init_norm": q_init_norm,
|
||
"q_spread_ema": kc_final[0],
|
||
"action_entropy_ema": kc_final[1],
|
||
"return_vs_random_ema": kc_final[2],
|
||
"early_q_movement_ema": kc_final[3],
|
||
"pass_q_spread": pass_q_spread,
|
||
"pass_entropy": pass_entropy,
|
||
"pass_rvr": pass_rvr,
|
||
"pass_early": pass_early,
|
||
"all_pass": all_pass,
|
||
"kc_log": kc_logs.iter().map(|(ep, kc)| {
|
||
serde_json::json!({
|
||
"episode": ep,
|
||
"q_spread": kc[0],
|
||
"entropy": kc[1],
|
||
"rvr": kc[2],
|
||
"early_mvmt": kc[3],
|
||
})
|
||
}).collect::<Vec<_>>(),
|
||
});
|
||
let mut f = File::create(&cli.out_path).context("create out")?;
|
||
write!(f, "{}", serde_json::to_string_pretty(&json)?)?;
|
||
info!("Wrote verdict + KC trajectory to {}", cli.out_path.display());
|
||
|
||
Ok(())
|
||
}
|