@@ -0,0 +1,725 @@
//! 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 ) / ( ( 1 u64 < < 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 ( ( ) )
}