diff --git a/crates/ml-alpha/build.rs b/crates/ml-alpha/build.rs index 5127cea4b..5c47f8bae 100644 --- a/crates/ml-alpha/build.rs +++ b/crates/ml-alpha/build.rs @@ -99,6 +99,7 @@ const KERNELS: &[&str] = &[ "rl_noisy_linear_backward", // NoisyNet: factored noisy linear backward — grad_mu_w/sigma_w/mu_b/sigma_b per-batch scratch for reduce_axis0 "rl_sample_tau", // CUDA graph prereq: device-side xorshift32 tau ~ U(0,1) for IQN; replaces host ChaCha8 + mapped-pinned upload "rl_sample_noise", // CUDA graph prereq: device-side factored noise f(rand) for NoisyLinear; replaces host ChaCha8 + mapped-pinned upload + "rl_write_u64", // CUDA graph prereq: single-thread u64 scalar write for device-resident ts_ns (graph-captured kernels read via pointer) "rl_increment_step", // device-resident step counter bump (ISV[548] += 1.0); graph-safe prereq — removes scalar current_step from all downstream kernel args "rl_fused_controllers", // fused kernel: 10 RL ISV controllers in one launch (gamma, tau, ppo_clip, entropy_coef, rollout_steps, per_alpha, reward_scale, ppo_ratio_clamp, gate_threshold, q_distill_lambda) — saves 9 kernel launch overheads (~40-80μs/step) ]; diff --git a/crates/ml-alpha/cuda/rl_multires_features_update.cu b/crates/ml-alpha/cuda/rl_multires_features_update.cu index 5e5889b50..02c10e194 100644 --- a/crates/ml-alpha/cuda/rl_multires_features_update.cu +++ b/crates/ml-alpha/cuda/rl_multires_features_update.cu @@ -40,7 +40,7 @@ extern "C" __global__ void rl_multires_features_update( const float* __restrict__ bid_sz, // [BOOK_LEVELS] const float* __restrict__ ask_sz, // [BOOK_LEVELS] const float* __restrict__ isv, - unsigned long long current_ts_ns, + const unsigned long long* __restrict__ current_ts_ns_ptr, int b_size ) { const int b = blockIdx.x * blockDim.x + threadIdx.x; @@ -48,6 +48,7 @@ extern "C" __global__ void rl_multires_features_update( const float mid = 0.5f * (bid_px[0] + ask_px[0]); const float old_mid = prev_mid[b]; + const unsigned long long current_ts_ns = current_ts_ns_ptr[0]; const unsigned long long old_ts = prev_ts_ns[b]; // dt in seconds (protect against zero/backward timestamps). diff --git a/crates/ml-alpha/cuda/rl_write_u64.cu b/crates/ml-alpha/cuda/rl_write_u64.cu new file mode 100644 index 000000000..e80b23a6c --- /dev/null +++ b/crates/ml-alpha/cuda/rl_write_u64.cu @@ -0,0 +1,13 @@ +// rl_write_u64.cu — Single-thread scalar write for device-resident u64. +// Launched OUTSIDE graph capture to stage changing scalar values into +// device buffers that graph-captured kernels read via pointer. +// Grid=(1,1,1), Block=(1,1,1). + +#include + +extern "C" __global__ void rl_write_u64( + uint64_t* __restrict__ dst, + uint64_t val +) { + if (threadIdx.x == 0) dst[0] = val; +} diff --git a/crates/ml-alpha/src/trainer/integrated.rs b/crates/ml-alpha/src/trainer/integrated.rs index bfb6c15e2..06e0525c3 100644 --- a/crates/ml-alpha/src/trainer/integrated.rs +++ b/crates/ml-alpha/src/trainer/integrated.rs @@ -239,6 +239,8 @@ const RL_INCREMENT_STEP_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/rl_increment_step.cubin")); const RL_SAMPLE_TAU_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/rl_sample_tau.cubin")); +const RL_WRITE_U64_CUBIN: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/rl_write_u64.cubin")); // Element-wise in-place add: dst[i] += src[i]. Used by noisy // exploration to fold NoisyLinear output into ensemble_q_d before @@ -563,6 +565,10 @@ pub struct IntegratedTrainer { // Device-side tau ~ U(0,1) for IQN (graph-safe, no host RNG). _rl_sample_tau_module: Arc, rl_sample_tau_fn: CudaFunction, + // Device-resident ts_ns for graph-captured multires kernel. + _rl_write_u64_module: Arc, + rl_write_u64_fn: CudaFunction, + ts_ns_d: CudaSlice, pub outcome_ema_d: CudaSlice, pub trade_context_d: CudaSlice, pub multires_output_d: CudaSlice, @@ -861,6 +867,81 @@ pub struct IntegratedTrainer { /// to -1; F.5 will populate before each `step_with_lobsim` from /// the loader's forward-snapshot lookahead. pub frd_labels_d: CudaSlice, + + // ── Persistent replay-step scratch buffers (CUDA Graph stable) ──── + // Pre-allocated at init so device pointers are stable across calls, + // enabling CUDA Graph capture of the RL replay training step. + // Zeroed via memset_zeros at the start of each step_synthetic / + // dqn_replay_step call. + + // Loss accumulators (size 1, atomicAdd targets). + pub ss_q_loss_d: CudaSlice, + pub ss_pi_loss_d: CudaSlice, + pub ss_pi_loss_entropy_d: CudaSlice, + pub ss_v_loss_sum_d: CudaSlice, + + // Per-batch diagnostic scratch. + pub ss_v_loss_per_batch_d: CudaSlice, + pub ss_pi_log_prob_d: CudaSlice, + pub ss_entropy_d: CudaSlice, + + // Q-head per-batch + reduced gradient buffers. + pub ss_q_grad_logits_d: CudaSlice, + pub ss_q_grad_w_per_batch_d: CudaSlice, + pub ss_q_grad_b_per_batch_d: CudaSlice, + pub ss_q_grad_h_t_d: CudaSlice, + pub ss_q_grad_w_d: CudaSlice, + pub ss_q_grad_b_d: CudaSlice, + + // π-head per-batch + reduced gradient buffers. + pub ss_pi_grad_logits_d: CudaSlice, + pub ss_pi_grad_w_per_batch_d: CudaSlice, + pub ss_pi_grad_b_per_batch_d: CudaSlice, + pub ss_pi_grad_h_t_d: CudaSlice, + pub ss_pi_grad_w_d: CudaSlice, + pub ss_pi_grad_b_d: CudaSlice, + + // V-head per-batch + reduced gradient buffers. + pub ss_v_grad_w_per_batch_d: CudaSlice, + pub ss_v_grad_b_per_batch_d: CudaSlice, + pub ss_v_grad_h_t_d: CudaSlice, + pub ss_v_grad_w_d: CudaSlice, + pub ss_v_grad_b_d: CudaSlice, + + // FRD encoder-upstream gradient. + pub ss_frd_grad_h_t_d: CudaSlice, + + // Bellman projection scratch. + pub ss_q_target_full_d: CudaSlice, + pub ss_q_target_action_d: CudaSlice, + pub ss_target_dist_d: CudaSlice, + + // FRD backward chain scratch (step_synthetic only). + pub ss_frd_grad_logits_d: CudaSlice, + pub ss_frd_loss_per_b_h_d: CudaSlice, + pub ss_frd_grad_w2_pb_d: CudaSlice, + pub ss_frd_grad_b2_pb_d: CudaSlice, + pub ss_frd_grad_hidden_d: CudaSlice, + pub ss_frd_grad_w1_pb_d: CudaSlice, + pub ss_frd_grad_b1_pb_d: CudaSlice, + pub ss_frd_grad_w1_d: CudaSlice, + pub ss_frd_grad_b1_d: CudaSlice, + pub ss_frd_grad_w2_d: CudaSlice, + pub ss_frd_grad_b2_d: CudaSlice, + + // IQN replay-step scratch (dqn_replay_step only). + pub ss_iqn_tau_target_d: CudaSlice, + pub ss_iqn_target_q_d: CudaSlice, + pub ss_iqn_loss_pb_d: CudaSlice, + pub ss_iqn_grad_q_d: CudaSlice, + pub ss_iqn_grad_w_out_pb_d: CudaSlice, + pub ss_iqn_grad_b_out_pb_d: CudaSlice, + pub ss_iqn_grad_w_embed_pb_d: CudaSlice, + pub ss_iqn_grad_b_embed_pb_d: CudaSlice, + pub ss_iqn_grad_w_out_d: CudaSlice, + pub ss_iqn_grad_b_out_d: CudaSlice, + pub ss_iqn_grad_w_embed_d: CudaSlice, + pub ss_iqn_grad_b_embed_d: CudaSlice, } impl IntegratedTrainer { @@ -1228,6 +1309,15 @@ impl IntegratedTrainer { let rl_sample_tau_fn = rl_sample_tau_module .load_function("rl_sample_tau") .context("load rl_sample_tau")?; + let rl_write_u64_module = ctx + .load_cubin(RL_WRITE_U64_CUBIN.to_vec()) + .context("load rl_write_u64 cubin")?; + let rl_write_u64_fn = rl_write_u64_module + .load_function("rl_write_u64") + .context("load rl_write_u64")?; + let ts_ns_d = stream + .alloc_zeros::(1) + .context("alloc ts_ns_d")?; let actions_to_market_targets_module = ctx .load_cubin(ACTIONS_TO_MARKET_TARGETS_CUBIN.to_vec()) .context("load actions_to_market_targets cubin")?; @@ -1650,6 +1740,143 @@ impl IntegratedTrainer { let neg_ones = vec![-1_i32; b_size * crate::rl::common::FRD_N_HORIZONS]; write_slice_i32_d_pub(&stream, &neg_ones, &mut frd_labels_d)?; + // ── Persistent replay-step scratch buffers ───────────────────── + // Pre-allocated at init for CUDA Graph pointer stability. Zeroed + // via memset_zeros at the start of each step_synthetic / + // dqn_replay_step call. + let k_dqn_ss = N_ACTIONS * Q_N_ATOMS; + let frd_h_dim = crate::rl::common::FRD_HIDDEN_DIM; + let frd_out_dim = crate::rl::frd::FRD_OUT_DIM; + let frd_n_h = crate::rl::common::FRD_N_HORIZONS; + let k_w_out = HIDDEN_DIM * N_ACTIONS; + let k_w_embed = EMBED_DIM * HIDDEN_DIM; + + // Loss accumulators (size 1). + let ss_q_loss_d = stream.alloc_zeros::(1) + .context("alloc ss_q_loss_d")?; + let ss_pi_loss_d = stream.alloc_zeros::(1) + .context("alloc ss_pi_loss_d")?; + let ss_pi_loss_entropy_d = stream.alloc_zeros::(1) + .context("alloc ss_pi_loss_entropy_d")?; + let ss_v_loss_sum_d = stream.alloc_zeros::(1) + .context("alloc ss_v_loss_sum_d")?; + + // Per-batch diagnostic scratch. + let ss_v_loss_per_batch_d = stream.alloc_zeros::(b_size) + .context("alloc ss_v_loss_per_batch_d")?; + let ss_pi_log_prob_d = stream.alloc_zeros::(b_size) + .context("alloc ss_pi_log_prob_d")?; + let ss_entropy_d = stream.alloc_zeros::(b_size) + .context("alloc ss_entropy_d")?; + + // Q-head per-batch + reduced gradient buffers. + let ss_q_grad_logits_d = stream.alloc_zeros::(b_size * k_dqn_ss) + .context("alloc ss_q_grad_logits_d")?; + let ss_q_grad_w_per_batch_d = stream + .alloc_zeros::(b_size * k_dqn_ss * HIDDEN_DIM) + .context("alloc ss_q_grad_w_per_batch_d")?; + let ss_q_grad_b_per_batch_d = stream.alloc_zeros::(b_size * k_dqn_ss) + .context("alloc ss_q_grad_b_per_batch_d")?; + let ss_q_grad_h_t_d = stream.alloc_zeros::(b_size * HIDDEN_DIM) + .context("alloc ss_q_grad_h_t_d")?; + let ss_q_grad_w_d = stream.alloc_zeros::(k_dqn_ss * HIDDEN_DIM) + .context("alloc ss_q_grad_w_d")?; + let ss_q_grad_b_d = stream.alloc_zeros::(k_dqn_ss) + .context("alloc ss_q_grad_b_d")?; + + // π-head per-batch + reduced gradient buffers. + let ss_pi_grad_logits_d = stream.alloc_zeros::(b_size * N_ACTIONS) + .context("alloc ss_pi_grad_logits_d")?; + let ss_pi_grad_w_per_batch_d = stream + .alloc_zeros::(b_size * N_ACTIONS * HIDDEN_DIM) + .context("alloc ss_pi_grad_w_per_batch_d")?; + let ss_pi_grad_b_per_batch_d = stream.alloc_zeros::(b_size * N_ACTIONS) + .context("alloc ss_pi_grad_b_per_batch_d")?; + let ss_pi_grad_h_t_d = stream.alloc_zeros::(b_size * HIDDEN_DIM) + .context("alloc ss_pi_grad_h_t_d")?; + let ss_pi_grad_w_d = stream.alloc_zeros::(N_ACTIONS * HIDDEN_DIM) + .context("alloc ss_pi_grad_w_d")?; + let ss_pi_grad_b_d = stream.alloc_zeros::(N_ACTIONS) + .context("alloc ss_pi_grad_b_d")?; + + // V-head per-batch + reduced gradient buffers. + let ss_v_grad_w_per_batch_d = stream.alloc_zeros::(b_size * HIDDEN_DIM) + .context("alloc ss_v_grad_w_per_batch_d")?; + let ss_v_grad_b_per_batch_d = stream.alloc_zeros::(b_size) + .context("alloc ss_v_grad_b_per_batch_d")?; + let ss_v_grad_h_t_d = stream.alloc_zeros::(b_size * HIDDEN_DIM) + .context("alloc ss_v_grad_h_t_d")?; + let ss_v_grad_w_d = stream.alloc_zeros::(HIDDEN_DIM) + .context("alloc ss_v_grad_w_d")?; + let ss_v_grad_b_d = stream.alloc_zeros::(1) + .context("alloc ss_v_grad_b_d")?; + + // FRD encoder-upstream gradient. + let ss_frd_grad_h_t_d = stream.alloc_zeros::(b_size * HIDDEN_DIM) + .context("alloc ss_frd_grad_h_t_d")?; + + // Bellman projection scratch. + let ss_q_target_full_d = stream.alloc_zeros::(b_size * k_dqn_ss) + .context("alloc ss_q_target_full_d")?; + let ss_q_target_action_d = stream.alloc_zeros::(b_size * Q_N_ATOMS) + .context("alloc ss_q_target_action_d")?; + let ss_target_dist_d = stream.alloc_zeros::(b_size * Q_N_ATOMS) + .context("alloc ss_target_dist_d")?; + + // FRD backward chain scratch. + let ss_frd_grad_logits_d = stream.alloc_zeros::(b_size * frd_out_dim) + .context("alloc ss_frd_grad_logits_d")?; + let ss_frd_loss_per_b_h_d = stream.alloc_zeros::(b_size * frd_n_h) + .context("alloc ss_frd_loss_per_b_h_d")?; + let ss_frd_grad_w2_pb_d = stream + .alloc_zeros::(b_size * frd_h_dim * frd_out_dim) + .context("alloc ss_frd_grad_w2_pb_d")?; + let ss_frd_grad_b2_pb_d = stream.alloc_zeros::(b_size * frd_out_dim) + .context("alloc ss_frd_grad_b2_pb_d")?; + let ss_frd_grad_hidden_d = stream.alloc_zeros::(b_size * frd_h_dim) + .context("alloc ss_frd_grad_hidden_d")?; + let ss_frd_grad_w1_pb_d = stream + .alloc_zeros::(b_size * HIDDEN_DIM * frd_h_dim) + .context("alloc ss_frd_grad_w1_pb_d")?; + let ss_frd_grad_b1_pb_d = stream.alloc_zeros::(b_size * frd_h_dim) + .context("alloc ss_frd_grad_b1_pb_d")?; + let ss_frd_grad_w1_d = stream.alloc_zeros::(HIDDEN_DIM * frd_h_dim) + .context("alloc ss_frd_grad_w1_d")?; + let ss_frd_grad_b1_d = stream.alloc_zeros::(frd_h_dim) + .context("alloc ss_frd_grad_b1_d")?; + let ss_frd_grad_w2_d = stream.alloc_zeros::(frd_h_dim * frd_out_dim) + .context("alloc ss_frd_grad_w2_d")?; + let ss_frd_grad_b2_d = stream.alloc_zeros::(frd_out_dim) + .context("alloc ss_frd_grad_b2_d")?; + + // IQN replay-step scratch. + let ss_iqn_tau_target_d = stream.alloc_zeros::(b_size * iqn_n_tau) + .context("alloc ss_iqn_tau_target_d")?; + let ss_iqn_target_q_d = stream + .alloc_zeros::(b_size * iqn_n_tau * N_ACTIONS) + .context("alloc ss_iqn_target_q_d")?; + let ss_iqn_loss_pb_d = stream.alloc_zeros::(b_size) + .context("alloc ss_iqn_loss_pb_d")?; + let ss_iqn_grad_q_d = stream + .alloc_zeros::(b_size * iqn_n_tau * N_ACTIONS) + .context("alloc ss_iqn_grad_q_d")?; + let ss_iqn_grad_w_out_pb_d = stream.alloc_zeros::(b_size * k_w_out) + .context("alloc ss_iqn_grad_w_out_pb_d")?; + let ss_iqn_grad_b_out_pb_d = stream.alloc_zeros::(b_size * N_ACTIONS) + .context("alloc ss_iqn_grad_b_out_pb_d")?; + let ss_iqn_grad_w_embed_pb_d = stream.alloc_zeros::(b_size * k_w_embed) + .context("alloc ss_iqn_grad_w_embed_pb_d")?; + let ss_iqn_grad_b_embed_pb_d = stream.alloc_zeros::(b_size * HIDDEN_DIM) + .context("alloc ss_iqn_grad_b_embed_pb_d")?; + let ss_iqn_grad_w_out_d = stream.alloc_zeros::(k_w_out) + .context("alloc ss_iqn_grad_w_out_d")?; + let ss_iqn_grad_b_out_d = stream.alloc_zeros::(N_ACTIONS) + .context("alloc ss_iqn_grad_b_out_d")?; + let ss_iqn_grad_w_embed_d = stream.alloc_zeros::(k_w_embed) + .context("alloc ss_iqn_grad_w_embed_d")?; + let ss_iqn_grad_b_embed_d = stream.alloc_zeros::(HIDDEN_DIM) + .context("alloc ss_iqn_grad_b_embed_d")?; + Ok(Self { cfg, perception, @@ -1756,6 +1983,9 @@ impl IntegratedTrainer { rl_increment_step_fn, _rl_sample_tau_module: rl_sample_tau_module, rl_sample_tau_fn, + _rl_write_u64_module: rl_write_u64_module, + rl_write_u64_fn, + ts_ns_d, outcome_ema_d, trade_context_d, multires_output_d, @@ -1856,6 +2086,57 @@ impl IntegratedTrainer { frd_b1_adam, frd_w2_adam, frd_b2_adam, + ss_q_loss_d, + ss_pi_loss_d, + ss_pi_loss_entropy_d, + ss_v_loss_sum_d, + ss_v_loss_per_batch_d, + ss_pi_log_prob_d, + ss_entropy_d, + ss_q_grad_logits_d, + ss_q_grad_w_per_batch_d, + ss_q_grad_b_per_batch_d, + ss_q_grad_h_t_d, + ss_q_grad_w_d, + ss_q_grad_b_d, + ss_pi_grad_logits_d, + ss_pi_grad_w_per_batch_d, + ss_pi_grad_b_per_batch_d, + ss_pi_grad_h_t_d, + ss_pi_grad_w_d, + ss_pi_grad_b_d, + ss_v_grad_w_per_batch_d, + ss_v_grad_b_per_batch_d, + ss_v_grad_h_t_d, + ss_v_grad_w_d, + ss_v_grad_b_d, + ss_frd_grad_h_t_d, + ss_q_target_full_d, + ss_q_target_action_d, + ss_target_dist_d, + ss_frd_grad_logits_d, + ss_frd_loss_per_b_h_d, + ss_frd_grad_w2_pb_d, + ss_frd_grad_b2_pb_d, + ss_frd_grad_hidden_d, + ss_frd_grad_w1_pb_d, + ss_frd_grad_b1_pb_d, + ss_frd_grad_w1_d, + ss_frd_grad_b1_d, + ss_frd_grad_w2_d, + ss_frd_grad_b2_d, + ss_iqn_tau_target_d, + ss_iqn_target_q_d, + ss_iqn_loss_pb_d, + ss_iqn_grad_q_d, + ss_iqn_grad_w_out_pb_d, + ss_iqn_grad_b_out_pb_d, + ss_iqn_grad_w_embed_pb_d, + ss_iqn_grad_b_embed_pb_d, + ss_iqn_grad_w_out_d, + ss_iqn_grad_b_out_d, + ss_iqn_grad_w_embed_d, + ss_iqn_grad_b_embed_d, } .with_controllers_bootstrapped()?) } @@ -3029,53 +3310,80 @@ impl IntegratedTrainer { let h_t_borrow: &CudaSlice = self.perception.h_t_view(); debug_assert_eq!(h_t_borrow.len(), b_size * HIDDEN_DIM); - // ── Step 3: allocate per-step scratch ──────────────────────── + // ── Step 3: zero persistent per-step scratch ───────────────── + // Buffers are pre-allocated struct fields for CUDA Graph pointer + // stability. memset_zeros resets them each call. let k_dqn = N_ACTIONS * Q_N_ATOMS; // Loss accumulators (atomicAdd into single floats from kernels). - let q_loss_d = self.stream.alloc_zeros::(1)?; - let pi_loss_d = self.stream.alloc_zeros::(1)?; - let pi_loss_entropy_d = self.stream.alloc_zeros::(1)?; - // Per-batch V loss scratch (reduced below). - let mut v_loss_per_batch_d = self.stream.alloc_zeros::(b_size)?; - let mut v_loss_sum_d = self.stream.alloc_zeros::(1)?; + self.stream.memset_zeros(&mut self.ss_q_loss_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_loss_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_pi_loss_d) + .map_err(|e| anyhow::anyhow!("zero ss_pi_loss_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_pi_loss_entropy_d) + .map_err(|e| anyhow::anyhow!("zero ss_pi_loss_entropy_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_v_loss_per_batch_d) + .map_err(|e| anyhow::anyhow!("zero ss_v_loss_per_batch_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_v_loss_sum_d) + .map_err(|e| anyhow::anyhow!("zero ss_v_loss_sum_d: {e}"))?; - // Per-batch diagnostic scratch (consumed by the surrogate kernel). - let mut pi_log_prob_d = self.stream.alloc_zeros::(b_size)?; - let mut entropy_d = self.stream.alloc_zeros::(b_size)?; + // Per-batch diagnostic scratch. + self.stream.memset_zeros(&mut self.ss_pi_log_prob_d) + .map_err(|e| anyhow::anyhow!("zero ss_pi_log_prob_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_entropy_d) + .map_err(|e| anyhow::anyhow!("zero ss_entropy_d: {e}"))?; - // Per-batch grad scratch (reduced across batches via reduce_axis0). - let mut q_grad_logits_d = self.stream.alloc_zeros::(b_size * k_dqn)?; - let mut q_grad_w_per_batch_d = self - .stream - .alloc_zeros::(b_size * k_dqn * HIDDEN_DIM)?; - let mut q_grad_b_per_batch_d = self.stream.alloc_zeros::(b_size * k_dqn)?; - let mut q_grad_h_t_d = self.stream.alloc_zeros::(b_size * HIDDEN_DIM)?; - let mut q_grad_w_d = self.stream.alloc_zeros::(k_dqn * HIDDEN_DIM)?; - let mut q_grad_b_d = self.stream.alloc_zeros::(k_dqn)?; + // Q-head gradient scratch. + self.stream.memset_zeros(&mut self.ss_q_grad_logits_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_grad_logits_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_q_grad_w_per_batch_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_grad_w_per_batch_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_q_grad_b_per_batch_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_grad_b_per_batch_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_q_grad_h_t_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_grad_h_t_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_q_grad_w_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_grad_w_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_q_grad_b_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_grad_b_d: {e}"))?; - let mut pi_grad_logits_d = self.stream.alloc_zeros::(b_size * N_ACTIONS)?; - let mut pi_grad_w_per_batch_d = - self.stream.alloc_zeros::(b_size * N_ACTIONS * HIDDEN_DIM)?; - let mut pi_grad_b_per_batch_d = self.stream.alloc_zeros::(b_size * N_ACTIONS)?; - let mut pi_grad_h_t_d = self.stream.alloc_zeros::(b_size * HIDDEN_DIM)?; - let mut pi_grad_w_d = self.stream.alloc_zeros::(N_ACTIONS * HIDDEN_DIM)?; - let mut pi_grad_b_d = self.stream.alloc_zeros::(N_ACTIONS)?; + // π-head gradient scratch. + self.stream.memset_zeros(&mut self.ss_pi_grad_logits_d) + .map_err(|e| anyhow::anyhow!("zero ss_pi_grad_logits_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_pi_grad_w_per_batch_d) + .map_err(|e| anyhow::anyhow!("zero ss_pi_grad_w_per_batch_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_pi_grad_b_per_batch_d) + .map_err(|e| anyhow::anyhow!("zero ss_pi_grad_b_per_batch_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_pi_grad_h_t_d) + .map_err(|e| anyhow::anyhow!("zero ss_pi_grad_h_t_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_pi_grad_w_d) + .map_err(|e| anyhow::anyhow!("zero ss_pi_grad_w_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_pi_grad_b_d) + .map_err(|e| anyhow::anyhow!("zero ss_pi_grad_b_d: {e}"))?; - let mut v_grad_w_per_batch_d = self.stream.alloc_zeros::(b_size * HIDDEN_DIM)?; - let mut v_grad_b_per_batch_d = self.stream.alloc_zeros::(b_size)?; - let mut v_grad_h_t_d = self.stream.alloc_zeros::(b_size * HIDDEN_DIM)?; - let mut v_grad_w_d = self.stream.alloc_zeros::(HIDDEN_DIM)?; - let mut v_grad_b_d = self.stream.alloc_zeros::(1)?; - // SP20 P3 FRD layer-1 backward produces grad_h_t into this - // buffer (the encoder-upstream gradient). Folded into - // grad_h_t_combined_d at Step 10 with the lambdas.frd weight. - let mut frd_grad_h_t_d = self.stream.alloc_zeros::(b_size * HIDDEN_DIM)?; + // V-head gradient scratch. + self.stream.memset_zeros(&mut self.ss_v_grad_w_per_batch_d) + .map_err(|e| anyhow::anyhow!("zero ss_v_grad_w_per_batch_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_v_grad_b_per_batch_d) + .map_err(|e| anyhow::anyhow!("zero ss_v_grad_b_per_batch_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_v_grad_h_t_d) + .map_err(|e| anyhow::anyhow!("zero ss_v_grad_h_t_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_v_grad_w_d) + .map_err(|e| anyhow::anyhow!("zero ss_v_grad_w_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_v_grad_b_d) + .map_err(|e| anyhow::anyhow!("zero ss_v_grad_b_d: {e}"))?; - // Bellman projection scratch (item 2). - let mut q_target_full_d = self.stream.alloc_zeros::(b_size * k_dqn)?; - let mut q_target_action_d = self.stream.alloc_zeros::(b_size * Q_N_ATOMS)?; - let mut target_dist_d = self.stream.alloc_zeros::(b_size * Q_N_ATOMS)?; + // FRD encoder-upstream gradient. + self.stream.memset_zeros(&mut self.ss_frd_grad_h_t_d) + .map_err(|e| anyhow::anyhow!("zero ss_frd_grad_h_t_d: {e}"))?; + + // Bellman projection scratch. + self.stream.memset_zeros(&mut self.ss_q_target_full_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_target_full_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_q_target_action_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_target_action_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_target_dist_d) + .map_err(|e| anyhow::anyhow!("zero ss_target_dist_d: {e}"))?; // ── Step 4 (Phase R7a): inputs are trainer-owned device buffers, // populated by step_with_lobsim's GPU pipeline before this call. @@ -3130,8 +3438,6 @@ impl IntegratedTrainer { // PPO surrogate forward — policy loss + entropy bonus only. // V loss flows through the V-head kernels (item 4). { - let mut pi_loss_d_mut = pi_loss_d; - let mut pi_loss_entropy_d_mut = pi_loss_entropy_d; self.policy_head .surrogate_forward( &self.pi_logits_d, @@ -3140,14 +3446,14 @@ impl IntegratedTrainer { &self.advantages_d, &self.isv_d, b_size, - &mut pi_log_prob_d, - &mut entropy_d, - &mut pi_loss_d_mut, - &mut pi_loss_entropy_d_mut, + &mut self.ss_pi_log_prob_d, + &mut self.ss_entropy_d, + &mut self.ss_pi_loss_d, + &mut self.ss_pi_loss_entropy_d, ) .context("policy_head.surrogate_forward")?; - let l_pi_host = read_scalar_d(&self.stream, &pi_loss_d_mut)?; - let _l_pi_ent_host = read_scalar_d(&self.stream, &pi_loss_entropy_d_mut)?; + let l_pi_host = read_scalar_d(&self.stream, &self.ss_pi_loss_d)?; + let _l_pi_ent_host = read_scalar_d(&self.stream, &self.ss_pi_loss_entropy_d)?; self.last_pi_loss = l_pi_host; } @@ -3163,68 +3469,67 @@ impl IntegratedTrainer { // entire Bellman target build through PER-sampled buffers per // `feedback_always_per`. self.dqn_head - .forward_target(&self.sampled_h_tp1_d, b_size, &mut q_target_full_d) + .forward_target(&self.sampled_h_tp1_d, b_size, &mut self.ss_q_target_full_d) .context("dqn_head.forward_target(sampled_h_tp1) [R7d off-policy]")?; self.dqn_head .select_action_atoms( - &q_target_full_d, + &self.ss_q_target_full_d, &self.sampled_next_actions_d, b_size, - &mut q_target_action_d, + &mut self.ss_q_target_action_d, ) .context("dqn_head.select_action_atoms (sampled_next_actions)")?; self.dqn_head .project_bellman_target( - &q_target_action_d, + &self.ss_q_target_action_d, &self.sampled_rewards_d, &self.sampled_dones_d, &self.sampled_n_step_gammas_d, &self.isv_d, b_size, - &mut target_dist_d, + &mut self.ss_target_dist_d, ) .context("dqn_head.project_bellman_target (sampled rewards/dones)")?; // ── Step 6b: DQN backward (logits → grad_w/b/h_t) ──────────── - let mut q_loss_d_mut = q_loss_d; self.dqn_head .backward_logits( &self.q_logits_d, - &target_dist_d, + &self.ss_target_dist_d, &self.sampled_actions_d, b_size, - &mut q_loss_d_mut, + &mut self.ss_q_loss_d, &mut self.td_per_sample_d, - &mut q_grad_logits_d, + &mut self.ss_q_grad_logits_d, ) .context("dqn_head.backward_logits (sampled_actions) [R7d off-policy]")?; - let l_q_host = read_scalar_d(&self.stream, &q_loss_d_mut)?; + let l_q_host = read_scalar_d(&self.stream, &self.ss_q_loss_d)?; // Mirror for next step's LR controller plateau detection. self.last_q_loss = l_q_host / (b_size as f32); self.dqn_head .backward_to_w_b_h( &self.sampled_h_t_d, - &q_grad_logits_d, + &self.ss_q_grad_logits_d, b_size, - &mut q_grad_w_per_batch_d, - &mut q_grad_b_per_batch_d, - &mut q_grad_h_t_d, + &mut self.ss_q_grad_w_per_batch_d, + &mut self.ss_q_grad_b_per_batch_d, + &mut self.ss_q_grad_h_t_d, ) .context("dqn_head.backward_to_w_b_h(sampled_h_t) [R7d off-policy]")?; - // R7d stop-grad: q_grad_h_t_d is the gradient wrt SAMPLED h_t + // R7d stop-grad: ss_q_grad_h_t_d is the gradient wrt SAMPLED h_t // (a past-step encoder output). Accumulating it into the // shared encoder via grad_h_t_combined_d would poison the // encoder with stale-state gradient signal. Standard pattern // for off-policy + shared encoder (SAC / R2D2 / IMPALA do the - // same). The buffer is allocated above for backward_to_w_b_h + // same). The buffer is pre-allocated for backward_to_w_b_h // to write into; we just don't FOLD it into the encoder grad // combine below. Encoder learns from π + V (on-policy) + // BCE/aux (supervised via step_batched) only. - let _ = &q_grad_h_t_d; + let _ = &self.ss_q_grad_h_t_d; - self.launch_reduce_axis0(&q_grad_w_per_batch_d, b_size, k_dqn * HIDDEN_DIM, &mut q_grad_w_d)?; - self.launch_reduce_axis0(&q_grad_b_per_batch_d, b_size, k_dqn, &mut q_grad_b_d)?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_q_grad_w_per_batch_d, b_size, k_dqn * HIDDEN_DIM, &mut self.ss_q_grad_w_d)?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_q_grad_b_per_batch_d, b_size, k_dqn, &mut self.ss_q_grad_b_d)?; // ── Step 7: PPO backward (logits → grad_w/b/h_t) ───────────── self.policy_head @@ -3235,7 +3540,7 @@ impl IntegratedTrainer { &self.advantages_d, &self.isv_d, b_size, - &mut pi_grad_logits_d, + &mut self.ss_pi_grad_logits_d, ) .context("policy_head.surrogate_backward_logits")?; @@ -3260,7 +3565,7 @@ impl IntegratedTrainer { .arg(&self.pi_logits_d) .arg(&self.atom_supports_d) .arg(&self.isv_d) - .arg(&mut pi_grad_logits_d) + .arg(&mut self.ss_pi_grad_logits_d) .arg(&b_size_i); unsafe { launch @@ -3277,21 +3582,16 @@ impl IntegratedTrainer { self.policy_head .backward_to_w_b_h( h_t_borrow, - &pi_grad_logits_d, + &self.ss_pi_grad_logits_d, b_size, - &mut pi_grad_w_per_batch_d, - &mut pi_grad_b_per_batch_d, - &mut pi_grad_h_t_d, + &mut self.ss_pi_grad_w_per_batch_d, + &mut self.ss_pi_grad_b_per_batch_d, + &mut self.ss_pi_grad_h_t_d, ) .context("policy_head.backward_to_w_b_h")?; - self.launch_reduce_axis0( - &pi_grad_w_per_batch_d, - b_size, - N_ACTIONS * HIDDEN_DIM, - &mut pi_grad_w_d, - )?; - self.launch_reduce_axis0(&pi_grad_b_per_batch_d, b_size, N_ACTIONS, &mut pi_grad_b_d)?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_pi_grad_w_per_batch_d, b_size, N_ACTIONS * HIDDEN_DIM, &mut self.ss_pi_grad_w_d)?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_pi_grad_b_per_batch_d, b_size, N_ACTIONS, &mut self.ss_pi_grad_b_d)?; // ── Step 8: V backward ─────────────────────────────────────── self.value_head @@ -3300,39 +3600,39 @@ impl IntegratedTrainer { &self.v_pred_d, &self.returns_d, b_size, - &mut v_loss_per_batch_d, - &mut v_grad_w_per_batch_d, - &mut v_grad_b_per_batch_d, - &mut v_grad_h_t_d, + &mut self.ss_v_loss_per_batch_d, + &mut self.ss_v_grad_w_per_batch_d, + &mut self.ss_v_grad_b_per_batch_d, + &mut self.ss_v_grad_h_t_d, ) .context("value_head.backward")?; - self.launch_reduce_axis0(&v_grad_w_per_batch_d, b_size, HIDDEN_DIM, &mut v_grad_w_d)?; - self.launch_reduce_axis0(&v_grad_b_per_batch_d, b_size, 1, &mut v_grad_b_d)?; - self.launch_reduce_axis0(&v_loss_per_batch_d, b_size, 1, &mut v_loss_sum_d)?; - let l_v_sum_host = read_scalar_d(&self.stream, &v_loss_sum_d)?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_v_grad_w_per_batch_d, b_size, HIDDEN_DIM, &mut self.ss_v_grad_w_d)?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_v_grad_b_per_batch_d, b_size, 1, &mut self.ss_v_grad_b_d)?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_v_loss_per_batch_d, b_size, 1, &mut self.ss_v_loss_sum_d)?; + let l_v_sum_host = read_scalar_d(&self.stream, &self.ss_v_loss_sum_d)?; let l_v_host = l_v_sum_host / (b_size as f32); // Mirror for next step's LR controller plateau detection. self.last_v_loss = l_v_host; // ── Step 9: Adam updates on each head's w and b ────────────── self.dqn_w_adam - .step(&mut self.dqn_head.w_d, &q_grad_w_d) + .step(&mut self.dqn_head.w_d, &self.ss_q_grad_w_d) .context("dqn_w_adam.step")?; self.dqn_b_adam - .step(&mut self.dqn_head.b_d, &q_grad_b_d) + .step(&mut self.dqn_head.b_d, &self.ss_q_grad_b_d) .context("dqn_b_adam.step")?; self.policy_w_adam - .step(&mut self.policy_head.w_d, &pi_grad_w_d) + .step(&mut self.policy_head.w_d, &self.ss_pi_grad_w_d) .context("policy_w_adam.step")?; self.policy_b_adam - .step(&mut self.policy_head.b_d, &pi_grad_b_d) + .step(&mut self.policy_head.b_d, &self.ss_pi_grad_b_d) .context("policy_b_adam.step")?; self.value_w_adam - .step(&mut self.value_head.w_d, &v_grad_w_d) + .step(&mut self.value_head.w_d, &self.ss_v_grad_w_d) .context("value_w_adam.step")?; self.value_b_adam - .step(&mut self.value_head.b_d, &v_grad_b_d) + .step(&mut self.value_head.b_d, &self.ss_v_grad_b_d) .context("value_b_adam.step")?; // ── SP20 P3 FRD backward chain (F.4) ───────────────────────── @@ -3348,87 +3648,75 @@ impl IntegratedTrainer { let frd_n_h = crate::rl::common::FRD_N_HORIZONS; let frd_h_dim = crate::rl::common::FRD_HIDDEN_DIM; let frd_out_dim = crate::rl::frd::FRD_OUT_DIM; - let mut frd_grad_logits_d = self.stream.alloc_zeros::(b_size * frd_out_dim)?; - let mut frd_loss_per_b_h_d = self.stream.alloc_zeros::(b_size * frd_n_h)?; + // Zero FRD scratch buffers. + self.stream.memset_zeros(&mut self.ss_frd_grad_logits_d) + .map_err(|e| anyhow::anyhow!("zero ss_frd_grad_logits_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_frd_loss_per_b_h_d) + .map_err(|e| anyhow::anyhow!("zero ss_frd_loss_per_b_h_d: {e}"))?; self.frd_head.softmax_ce_grad( &self.frd_logits_d, &self.frd_labels_d, - &mut frd_grad_logits_d, - &mut frd_loss_per_b_h_d, + &mut self.ss_frd_grad_logits_d, + &mut self.ss_frd_loss_per_b_h_d, b_size, )?; - let mut frd_grad_w2_pb_d = self - .stream - .alloc_zeros::(b_size * frd_h_dim * frd_out_dim)?; - let mut frd_grad_b2_pb_d = self.stream.alloc_zeros::(b_size * frd_out_dim)?; - let mut frd_grad_hidden_d = self.stream.alloc_zeros::(b_size * frd_h_dim)?; + self.stream.memset_zeros(&mut self.ss_frd_grad_w2_pb_d) + .map_err(|e| anyhow::anyhow!("zero ss_frd_grad_w2_pb_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_frd_grad_b2_pb_d) + .map_err(|e| anyhow::anyhow!("zero ss_frd_grad_b2_pb_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_frd_grad_hidden_d) + .map_err(|e| anyhow::anyhow!("zero ss_frd_grad_hidden_d: {e}"))?; self.frd_head.layer2_bwd( &self.frd_hidden_d, - &frd_grad_logits_d, - &mut frd_grad_w2_pb_d, - &mut frd_grad_b2_pb_d, - &mut frd_grad_hidden_d, + &self.ss_frd_grad_logits_d, + &mut self.ss_frd_grad_w2_pb_d, + &mut self.ss_frd_grad_b2_pb_d, + &mut self.ss_frd_grad_hidden_d, b_size, )?; - let mut frd_grad_w1_pb_d = self - .stream - .alloc_zeros::(b_size * HIDDEN_DIM * frd_h_dim)?; - let mut frd_grad_b1_pb_d = self.stream.alloc_zeros::(b_size * frd_h_dim)?; + self.stream.memset_zeros(&mut self.ss_frd_grad_w1_pb_d) + .map_err(|e| anyhow::anyhow!("zero ss_frd_grad_w1_pb_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_frd_grad_b1_pb_d) + .map_err(|e| anyhow::anyhow!("zero ss_frd_grad_b1_pb_d: {e}"))?; self.frd_head.layer1_bwd( h_t_borrow, &self.frd_hidden_d, - &frd_grad_hidden_d, - &mut frd_grad_w1_pb_d, - &mut frd_grad_b1_pb_d, - &mut frd_grad_h_t_d, + &self.ss_frd_grad_hidden_d, + &mut self.ss_frd_grad_w1_pb_d, + &mut self.ss_frd_grad_b1_pb_d, + &mut self.ss_frd_grad_h_t_d, b_size, )?; // Reduce-axis-0: batch-summed weight grads. - let mut frd_grad_w1_d = self.stream.alloc_zeros::(HIDDEN_DIM * frd_h_dim)?; - let mut frd_grad_b1_d = self.stream.alloc_zeros::(frd_h_dim)?; - let mut frd_grad_w2_d = self.stream.alloc_zeros::(frd_h_dim * frd_out_dim)?; - let mut frd_grad_b2_d = self.stream.alloc_zeros::(frd_out_dim)?; - self.launch_reduce_axis0( - &frd_grad_w1_pb_d, - b_size, - HIDDEN_DIM * frd_h_dim, - &mut frd_grad_w1_d, - )?; - self.launch_reduce_axis0( - &frd_grad_b1_pb_d, - b_size, - frd_h_dim, - &mut frd_grad_b1_d, - )?; - self.launch_reduce_axis0( - &frd_grad_w2_pb_d, - b_size, - frd_h_dim * frd_out_dim, - &mut frd_grad_w2_d, - )?; - self.launch_reduce_axis0( - &frd_grad_b2_pb_d, - b_size, - frd_out_dim, - &mut frd_grad_b2_d, - )?; + self.stream.memset_zeros(&mut self.ss_frd_grad_w1_d) + .map_err(|e| anyhow::anyhow!("zero ss_frd_grad_w1_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_frd_grad_b1_d) + .map_err(|e| anyhow::anyhow!("zero ss_frd_grad_b1_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_frd_grad_w2_d) + .map_err(|e| anyhow::anyhow!("zero ss_frd_grad_w2_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_frd_grad_b2_d) + .map_err(|e| anyhow::anyhow!("zero ss_frd_grad_b2_d: {e}"))?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_frd_grad_w1_pb_d, b_size, HIDDEN_DIM * frd_h_dim, &mut self.ss_frd_grad_w1_d)?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_frd_grad_b1_pb_d, b_size, frd_h_dim, &mut self.ss_frd_grad_b1_d)?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_frd_grad_w2_pb_d, b_size, frd_h_dim * frd_out_dim, &mut self.ss_frd_grad_w2_d)?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_frd_grad_b2_pb_d, b_size, frd_out_dim, &mut self.ss_frd_grad_b2_d)?; // Adam steps — sentinel labels yield zero grads → no // weight movement (only momentum decay). self.frd_w1_adam - .step(&mut self.frd_head.w1_d, &frd_grad_w1_d) + .step(&mut self.frd_head.w1_d, &self.ss_frd_grad_w1_d) .context("frd_w1_adam.step")?; self.frd_b1_adam - .step(&mut self.frd_head.b1_d, &frd_grad_b1_d) + .step(&mut self.frd_head.b1_d, &self.ss_frd_grad_b1_d) .context("frd_b1_adam.step")?; self.frd_w2_adam - .step(&mut self.frd_head.w2_d, &frd_grad_w2_d) + .step(&mut self.frd_head.w2_d, &self.ss_frd_grad_w2_d) .context("frd_w2_adam.step")?; self.frd_b2_adam - .step(&mut self.frd_head.b2_d, &frd_grad_b2_d) + .step(&mut self.frd_head.b2_d, &self.ss_frd_grad_b2_d) .context("frd_b2_adam.step")?; // Per-row CE loss → mean over (b, h). Mapped-pinned read. let loss_pb_h = - read_slice_d(&self.stream, &frd_loss_per_b_h_d, b_size * frd_n_h)?; + read_slice_d(&self.stream, &self.ss_frd_loss_per_b_h_d, b_size * frd_n_h)?; loss_pb_h.iter().sum::() / ((b_size * frd_n_h) as f32) }; @@ -3453,14 +3741,14 @@ impl IntegratedTrainer { accumulate_grad_h( &self.stream, &self.grad_h_accumulate_fn, - &pi_grad_h_t_d, + &self.ss_pi_grad_h_t_d, lambdas.pi, &mut self.grad_h_t_combined_d, )?; accumulate_grad_h( &self.stream, &self.grad_h_accumulate_fn, - &v_grad_h_t_d, + &self.ss_v_grad_h_t_d, lambdas.v, &mut self.grad_h_t_combined_d, )?; @@ -3472,7 +3760,7 @@ impl IntegratedTrainer { accumulate_grad_h( &self.stream, &self.grad_h_accumulate_fn, - &frd_grad_h_t_d, + &self.ss_frd_grad_h_t_d, lambdas.frd, &mut self.grad_h_t_combined_d, )?; @@ -3513,7 +3801,7 @@ impl IntegratedTrainer { self.launch_ema_update_per_step( crate::rl::isv_slots::RL_ENTROPY_OBSERVED_EMA_INDEX, RL_LR_CONTROLLER_ALPHA, - &entropy_d, + &self.ss_entropy_d, b_size, ) .context("ema_update_per_step(entropy_observed)")?; @@ -3541,7 +3829,7 @@ impl IntegratedTrainer { let mut launch = self.stream.launch_builder(&self.rl_kl_approx_b_fn); launch .arg(&self.log_pi_old_d) - .arg(&pi_log_prob_d) + .arg(&self.ss_pi_log_prob_d) .arg(&mut self.ema_input_scratch_d) .arg(&b_size_i); unsafe { @@ -3569,7 +3857,7 @@ impl IntegratedTrainer { let mut launch = self.stream.launch_builder(&self.ppo_log_ratio_abs_max_b_fn); launch .arg(&self.log_pi_old_d) - .arg(&pi_log_prob_d) + .arg(&self.ss_pi_log_prob_d) .arg(&mut self.isv_d) .arg(&b_size_i); unsafe { @@ -3589,13 +3877,12 @@ impl IntegratedTrainer { // Per-head gradient-norm EMAs (Q, π, V) — feed the // signal-modulated LR controller (rl_lr_controller's target // formula: `lr_prev × TARGET_GRAD_NORM / observed_grad_norm`). - // grad_w_*_d buffers are still in scope here (allocated at - // top of step_synthetic, consumed by the Adam steps in Step 9 - // but not freed until function return); the l2_norm launches - // here read them after the Adam step has already used them. + // ss_*_grad_w_d buffers are persistent struct fields; the + // l2_norm launches here read them after the Adam step has + // already consumed them. let k_dqn_hidden = (N_ACTIONS * Q_N_ATOMS * HIDDEN_DIM) as usize; - self.launch_l2_norm(&q_grad_w_d.clone(), k_dqn_hidden) - .context("launch_l2_norm(q_grad_w_d)")?; + self.launch_l2_norm(&self.ss_q_grad_w_d.clone(), k_dqn_hidden) + .context("launch_l2_norm(ss_q_grad_w_d)")?; self.launch_ema_update_per_step( crate::rl::isv_slots::RL_Q_GRAD_NORM_EMA_INDEX, RL_LR_CONTROLLER_ALPHA, @@ -3605,8 +3892,8 @@ impl IntegratedTrainer { .context("ema_update_per_step(q_grad_norm)")?; let n_actions_hidden = (N_ACTIONS * HIDDEN_DIM) as usize; - self.launch_l2_norm(&pi_grad_w_d.clone(), n_actions_hidden) - .context("launch_l2_norm(pi_grad_w_d)")?; + self.launch_l2_norm(&self.ss_pi_grad_w_d.clone(), n_actions_hidden) + .context("launch_l2_norm(ss_pi_grad_w_d)")?; self.launch_ema_update_per_step( crate::rl::isv_slots::RL_PI_GRAD_NORM_EMA_INDEX, RL_LR_CONTROLLER_ALPHA, @@ -3615,8 +3902,8 @@ impl IntegratedTrainer { ) .context("ema_update_per_step(pi_grad_norm)")?; - self.launch_l2_norm(&v_grad_w_d.clone(), HIDDEN_DIM) - .context("launch_l2_norm(v_grad_w_d)")?; + self.launch_l2_norm(&self.ss_v_grad_w_d.clone(), HIDDEN_DIM) + .context("launch_l2_norm(ss_v_grad_w_d)")?; self.launch_ema_update_per_step( crate::rl::isv_slots::RL_V_GRAD_NORM_EMA_INDEX, RL_LR_CONTROLLER_ALPHA, @@ -3699,20 +3986,27 @@ impl IntegratedTrainer { let k_dqn = N_ACTIONS * Q_N_ATOMS; - // Per-iter scratch — small relative to the GPU allocator - // amortisation, and freed when the function returns. - let q_loss_d = self.stream.alloc_zeros::(1)?; - let mut q_target_full_d = self.stream.alloc_zeros::(b_size * k_dqn)?; - let mut q_target_action_d = self.stream.alloc_zeros::(b_size * Q_N_ATOMS)?; - let mut target_dist_d = self.stream.alloc_zeros::(b_size * Q_N_ATOMS)?; - let mut q_grad_logits_d = self.stream.alloc_zeros::(b_size * k_dqn)?; - let mut q_grad_w_per_batch_d = self - .stream - .alloc_zeros::(b_size * k_dqn * HIDDEN_DIM)?; - let mut q_grad_b_per_batch_d = self.stream.alloc_zeros::(b_size * k_dqn)?; - let mut q_grad_h_t_d = self.stream.alloc_zeros::(b_size * HIDDEN_DIM)?; - let mut q_grad_w_d = self.stream.alloc_zeros::(k_dqn * HIDDEN_DIM)?; - let mut q_grad_b_d = self.stream.alloc_zeros::(k_dqn)?; + // Zero persistent per-iter scratch (CUDA Graph pointer stability). + self.stream.memset_zeros(&mut self.ss_q_loss_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_loss_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_q_target_full_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_target_full_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_q_target_action_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_target_action_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_target_dist_d) + .map_err(|e| anyhow::anyhow!("zero ss_target_dist_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_q_grad_logits_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_grad_logits_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_q_grad_w_per_batch_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_grad_w_per_batch_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_q_grad_b_per_batch_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_grad_b_per_batch_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_q_grad_h_t_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_grad_h_t_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_q_grad_w_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_grad_w_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_q_grad_b_d) + .map_err(|e| anyhow::anyhow!("zero ss_q_grad_b_d: {e}"))?; let b_size_i = b_size as i32; @@ -3730,7 +4024,6 @@ impl IntegratedTrainer { // reduce_axis0 + Adam step. { let n_tau = self.iqn_head.n_tau(); - let tau_len = b_size * n_tau; // Sample online tau on device. { @@ -3773,11 +4066,33 @@ impl IntegratedTrainer { ) .context("dqn_replay_step: iqn_head.expected_q")?; + // Zero IQN per-iter scratch (CUDA Graph pointer stability). + self.stream.memset_zeros(&mut self.ss_iqn_tau_target_d) + .map_err(|e| anyhow::anyhow!("zero ss_iqn_tau_target_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_iqn_target_q_d) + .map_err(|e| anyhow::anyhow!("zero ss_iqn_target_q_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_iqn_loss_pb_d) + .map_err(|e| anyhow::anyhow!("zero ss_iqn_loss_pb_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_iqn_grad_q_d) + .map_err(|e| anyhow::anyhow!("zero ss_iqn_grad_q_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_iqn_grad_w_out_pb_d) + .map_err(|e| anyhow::anyhow!("zero ss_iqn_grad_w_out_pb_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_iqn_grad_b_out_pb_d) + .map_err(|e| anyhow::anyhow!("zero ss_iqn_grad_b_out_pb_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_iqn_grad_w_embed_pb_d) + .map_err(|e| anyhow::anyhow!("zero ss_iqn_grad_w_embed_pb_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_iqn_grad_b_embed_pb_d) + .map_err(|e| anyhow::anyhow!("zero ss_iqn_grad_b_embed_pb_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_iqn_grad_w_out_d) + .map_err(|e| anyhow::anyhow!("zero ss_iqn_grad_w_out_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_iqn_grad_b_out_d) + .map_err(|e| anyhow::anyhow!("zero ss_iqn_grad_b_out_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_iqn_grad_w_embed_d) + .map_err(|e| anyhow::anyhow!("zero ss_iqn_grad_w_embed_d: {e}"))?; + self.stream.memset_zeros(&mut self.ss_iqn_grad_b_embed_d) + .map_err(|e| anyhow::anyhow!("zero ss_iqn_grad_b_embed_d: {e}"))?; + // Sample target tau on device and forward target network on sampled_h_tp1. - let mut iqn_tau_target_d = self - .stream - .alloc_zeros::(tau_len) - .context("dqn_replay_step: alloc iqn_tau_target")?; { let b_i = b_size as i32; let n_tau_i = n_tau as i32; @@ -3789,7 +4104,7 @@ impl IntegratedTrainer { let mut tau_launch = self.stream.launch_builder(&self.rl_sample_tau_fn); tau_launch .arg(&mut self.iqn_prng_state_d) - .arg(&mut iqn_tau_target_d) + .arg(&mut self.ss_iqn_tau_target_d) .arg(&b_i) .arg(&n_tau_i); unsafe { @@ -3799,120 +4114,68 @@ impl IntegratedTrainer { } } - let mut iqn_target_q_d = self - .stream - .alloc_zeros::(b_size * n_tau * N_ACTIONS) - .context("dqn_replay_step: alloc iqn_target_q")?; self.iqn_head .forward_target( &self.sampled_h_tp1_d, - &iqn_tau_target_d, + &self.ss_iqn_tau_target_d, b_size, n_tau, - &mut iqn_target_q_d, + &mut self.ss_iqn_target_q_d, ) .context("dqn_replay_step: iqn_head.forward_target(sampled_h_tp1)")?; // Quantile Huber loss + gradient w.r.t. online Q values. - let mut iqn_loss_pb_d = self - .stream - .alloc_zeros::(b_size) - .context("dqn_replay_step: alloc iqn_loss_pb")?; - let mut iqn_grad_q_d = self - .stream - .alloc_zeros::(b_size * n_tau * N_ACTIONS) - .context("dqn_replay_step: alloc iqn_grad_q")?; self.iqn_head .compute_loss( &self.iqn_q_values_d, - &iqn_target_q_d, + &self.ss_iqn_target_q_d, &self.iqn_tau_d, &self.sampled_actions_d, b_size, n_tau, n_tau, - &mut iqn_loss_pb_d, - &mut iqn_grad_q_d, + &mut self.ss_iqn_loss_pb_d, + &mut self.ss_iqn_grad_q_d, ) .context("dqn_replay_step: iqn_head.compute_loss")?; // Backward through IQN forward: grad_q → grad_w_out/b_out/w_embed/b_embed. let k_w_out = HIDDEN_DIM * N_ACTIONS; let k_w_embed = EMBED_DIM * HIDDEN_DIM; - let mut iqn_grad_w_out_pb = self - .stream - .alloc_zeros::(b_size * k_w_out) - .context("dqn_replay_step: alloc iqn_grad_w_out_pb")?; - let mut iqn_grad_b_out_pb = self - .stream - .alloc_zeros::(b_size * N_ACTIONS) - .context("dqn_replay_step: alloc iqn_grad_b_out_pb")?; - let mut iqn_grad_w_embed_pb = self - .stream - .alloc_zeros::(b_size * k_w_embed) - .context("dqn_replay_step: alloc iqn_grad_w_embed_pb")?; - let mut iqn_grad_b_embed_pb = self - .stream - .alloc_zeros::(b_size * HIDDEN_DIM) - .context("dqn_replay_step: alloc iqn_grad_b_embed_pb")?; self.iqn_head .backward( &self.sampled_h_t_d, &self.iqn_tau_d, - &iqn_grad_q_d, + &self.ss_iqn_grad_q_d, b_size, n_tau, - &mut iqn_grad_w_out_pb, - &mut iqn_grad_b_out_pb, - &mut iqn_grad_w_embed_pb, - &mut iqn_grad_b_embed_pb, + &mut self.ss_iqn_grad_w_out_pb_d, + &mut self.ss_iqn_grad_b_out_pb_d, + &mut self.ss_iqn_grad_w_embed_pb_d, + &mut self.ss_iqn_grad_b_embed_pb_d, ) .context("dqn_replay_step: iqn_head.backward")?; // Reduce across batches. - let mut iqn_grad_w_out_d = self - .stream - .alloc_zeros::(k_w_out) - .context("dqn_replay_step: alloc iqn_grad_w_out")?; - let mut iqn_grad_b_out_d = self - .stream - .alloc_zeros::(N_ACTIONS) - .context("dqn_replay_step: alloc iqn_grad_b_out")?; - let mut iqn_grad_w_embed_d = self - .stream - .alloc_zeros::(k_w_embed) - .context("dqn_replay_step: alloc iqn_grad_w_embed")?; - let mut iqn_grad_b_embed_d = self - .stream - .alloc_zeros::(HIDDEN_DIM) - .context("dqn_replay_step: alloc iqn_grad_b_embed")?; - self.launch_reduce_axis0( - &iqn_grad_w_out_pb, b_size, k_w_out, &mut iqn_grad_w_out_d, - )?; - self.launch_reduce_axis0( - &iqn_grad_b_out_pb, b_size, N_ACTIONS, &mut iqn_grad_b_out_d, - )?; - self.launch_reduce_axis0( - &iqn_grad_w_embed_pb, b_size, k_w_embed, &mut iqn_grad_w_embed_d, - )?; - self.launch_reduce_axis0( - &iqn_grad_b_embed_pb, b_size, HIDDEN_DIM, &mut iqn_grad_b_embed_d, - )?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_iqn_grad_w_out_pb_d, b_size, k_w_out, &mut self.ss_iqn_grad_w_out_d)?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_iqn_grad_b_out_pb_d, b_size, N_ACTIONS, &mut self.ss_iqn_grad_b_out_d)?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_iqn_grad_w_embed_pb_d, b_size, k_w_embed, &mut self.ss_iqn_grad_w_embed_d)?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_iqn_grad_b_embed_pb_d, b_size, HIDDEN_DIM, &mut self.ss_iqn_grad_b_embed_d)?; // IQN Adam steps (LR from same ISV[RL_IQN_LR_INDEX] as the Q head). self.iqn_w_out_adam - .step(&mut self.iqn_head.w_out_d, &iqn_grad_w_out_d) + .step(&mut self.iqn_head.w_out_d, &self.ss_iqn_grad_w_out_d) .context("dqn_replay_step: iqn_w_out_adam.step")?; self.iqn_b_out_adam - .step(&mut self.iqn_head.b_out_d, &iqn_grad_b_out_d) + .step(&mut self.iqn_head.b_out_d, &self.ss_iqn_grad_b_out_d) .context("dqn_replay_step: iqn_b_out_adam.step")?; self.iqn_w_embed_adam - .step(&mut self.iqn_head.w_embed_d, &iqn_grad_w_embed_d) + .step(&mut self.iqn_head.w_embed_d, &self.ss_iqn_grad_w_embed_d) .context("dqn_replay_step: iqn_w_embed_adam.step")?; self.iqn_b_embed_adam - .step(&mut self.iqn_head.b_embed_d, &iqn_grad_b_embed_d) + .step(&mut self.iqn_head.b_embed_d, &self.ss_iqn_grad_b_embed_d) .context("dqn_replay_step: iqn_b_embed_adam.step")?; } @@ -3938,74 +4201,68 @@ impl IntegratedTrainer { // ── 3. Bellman target build via TARGET net at h_tp1 ───────── self.dqn_head - .forward_target(&self.sampled_h_tp1_d, b_size, &mut q_target_full_d) + .forward_target(&self.sampled_h_tp1_d, b_size, &mut self.ss_q_target_full_d) .context("dqn_replay_step: dqn_head.forward_target(sampled_h_tp1)")?; self.dqn_head .select_action_atoms( - &q_target_full_d, + &self.ss_q_target_full_d, &self.sampled_next_actions_d, b_size, - &mut q_target_action_d, + &mut self.ss_q_target_action_d, ) .context("dqn_replay_step: dqn_head.select_action_atoms")?; self.dqn_head .project_bellman_target( - &q_target_action_d, + &self.ss_q_target_action_d, &self.sampled_rewards_d, &self.sampled_dones_d, &self.sampled_n_step_gammas_d, &self.isv_d, b_size, - &mut target_dist_d, + &mut self.ss_target_dist_d, ) .context("dqn_replay_step: dqn_head.project_bellman_target")?; // ── 4. Q backward (logits → grad_w/b/h_t) ─────────────────── - let mut q_loss_d_mut = q_loss_d; self.dqn_head .backward_logits( &self.q_logits_d, - &target_dist_d, + &self.ss_target_dist_d, &self.sampled_actions_d, b_size, - &mut q_loss_d_mut, + &mut self.ss_q_loss_d, &mut self.td_per_sample_d, - &mut q_grad_logits_d, + &mut self.ss_q_grad_logits_d, ) .context("dqn_replay_step: dqn_head.backward_logits")?; - let l_q_host = read_scalar_d(&self.stream, &q_loss_d_mut)?; + let l_q_host = read_scalar_d(&self.stream, &self.ss_q_loss_d)?; let l_q = l_q_host / (b_size as f32); self.dqn_head .backward_to_w_b_h( &self.sampled_h_t_d, - &q_grad_logits_d, + &self.ss_q_grad_logits_d, b_size, - &mut q_grad_w_per_batch_d, - &mut q_grad_b_per_batch_d, - &mut q_grad_h_t_d, + &mut self.ss_q_grad_w_per_batch_d, + &mut self.ss_q_grad_b_per_batch_d, + &mut self.ss_q_grad_h_t_d, ) .context("dqn_replay_step: dqn_head.backward_to_w_b_h")?; - // R7d stop-grad: discard q_grad_h_t (sampled h_t is past-step + // R7d stop-grad: discard ss_q_grad_h_t_d (sampled h_t is past-step // encoder output; folding its gradient into the encoder would // poison live training). Same semantics as step_synthetic. - let _ = &q_grad_h_t_d; + let _ = &self.ss_q_grad_h_t_d; - self.launch_reduce_axis0( - &q_grad_w_per_batch_d, - b_size, - k_dqn * HIDDEN_DIM, - &mut q_grad_w_d, - )?; - self.launch_reduce_axis0(&q_grad_b_per_batch_d, b_size, k_dqn, &mut q_grad_b_d)?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_q_grad_w_per_batch_d, b_size, k_dqn * HIDDEN_DIM, &mut self.ss_q_grad_w_d)?; + reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_q_grad_b_per_batch_d, b_size, k_dqn, &mut self.ss_q_grad_b_d)?; // ── 5. Q Adam (uses LR set by step_synthetic; we don't re-fire // the LR controller here — that runs once per env step). self.dqn_w_adam - .step(&mut self.dqn_head.w_d, &q_grad_w_d) + .step(&mut self.dqn_head.w_d, &self.ss_q_grad_w_d) .context("dqn_replay_step: dqn_w_adam.step")?; self.dqn_b_adam - .step(&mut self.dqn_head.b_d, &q_grad_b_d) + .step(&mut self.dqn_head.b_d, &self.ss_q_grad_b_d) .context("dqn_replay_step: dqn_b_adam.step")?; Ok(l_q) @@ -4535,6 +4792,24 @@ impl IntegratedTrainer { let pos_bytes_i = lobsim.pos_bytes() as i32; let b_size_i = b_size as i32; + // Stage ts_ns to device buffer for graph-captured multires kernel. + // Launched OUTSIDE graph so the scalar write doesn't break capture. + { + let ts_ns_val: u64 = last_snap.ts_ns; + let cfg = LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (1, 1, 1), + shared_mem_bytes: 0, + }; + let mut launch = self.stream.launch_builder(&self.rl_write_u64_fn); + launch + .arg(&mut self.ts_ns_d) + .arg(&ts_ns_val); + unsafe { + launch.launch(cfg).context("rl_write_u64 (ts_ns) launch")?; + } + } + // ── Graph A2: post-snapshot / pre-fill kernel pipeline ───────── // 7 kernels from session_risk through actions_to_market_targets. // Same three-state machine as prefill — captures after warmup, @@ -4871,7 +5146,6 @@ impl IntegratedTrainer { shared_mem_bytes: 0, }; let b_size_i = b_size as i32; - let ts_ns: u64 = last_snap.ts_ns; let mut launch = self.stream.launch_builder(&self.rl_multires_features_update_fn); launch @@ -4884,7 +5158,7 @@ impl IntegratedTrainer { .arg(bid_px_d) // bid_sz uses bid_px (sizes at same depth levels) .arg(ask_px_d) // ask_sz uses ask_px .arg(&self.isv_d) - .arg(&ts_ns) + .arg(&self.ts_ns_d) .arg(&b_size_i); unsafe { launch @@ -5783,36 +6057,7 @@ impl IntegratedTrainer { Ok(()) } - /// Launch the reduce_axis0 kernel: - /// per_batch: [b_size × n_tail] → out: [n_tail] - fn launch_reduce_axis0( - &self, - per_batch: &CudaSlice, - b_size: usize, - n_tail: usize, - out: &mut CudaSlice, - ) -> Result<()> { - debug_assert_eq!(per_batch.len(), b_size * n_tail); - debug_assert_eq!(out.len(), n_tail); - let n_batch_i = b_size as i32; - let n_tail_i = n_tail as i32; - let cfg = LaunchConfig { - grid_dim: ((n_tail as u32).div_ceil(32), 1, 1), - block_dim: (32, 8, 1), - shared_mem_bytes: 0, - }; - let mut launch = self.stream.launch_builder(&self.reduce_axis0_fn); - launch - .arg(per_batch) - .arg(&n_batch_i) - .arg(&n_tail_i) - .arg(out); - unsafe { - launch.launch(cfg).context("reduce_axis0 launch")?; - } - Ok(()) - } /// Phase E.2 helper: combine a per-head grad_h_t into the encoder's /// accumulator slot via `grad_h_encoder[i] += λ × grad_h_head[i]`. @@ -5974,6 +6219,40 @@ fn accumulate_grad_h( Ok(()) } +/// Free-function version of `launch_reduce_axis0` — lives outside the impl +/// so callers can hold `&mut` references to trainer-owned output fields at +/// the same time as `&stream` / `&reduce_axis0_fn` borrows (same pattern as +/// `accumulate_grad_h`). Required for CUDA Graph stable pointer migration +/// where both input and output are `self` fields. +fn reduce_axis0_free( + stream: &Arc, + reduce_fn: &CudaFunction, + per_batch: &CudaSlice, + b_size: usize, + n_tail: usize, + out: &mut CudaSlice, +) -> Result<()> { + debug_assert_eq!(per_batch.len(), b_size * n_tail); + debug_assert_eq!(out.len(), n_tail); + let n_batch_i = b_size as i32; + let n_tail_i = n_tail as i32; + let cfg = LaunchConfig { + grid_dim: ((n_tail as u32).div_ceil(32), 1, 1), + block_dim: (32, 8, 1), + shared_mem_bytes: 0, + }; + let mut launch = stream.launch_builder(reduce_fn); + launch + .arg(per_batch) + .arg(&n_batch_i) + .arg(&n_tail_i) + .arg(out); + unsafe { + launch.launch(cfg).context("reduce_axis0 launch")?; + } + Ok(()) +} + // ── helpers ────────────────────────────────────────────────────────── // // Phase R7b: `upload_f32`, `upload_i32`, and `argmax_f32` are deleted diff --git a/docs/superpowers/plans/2026-05-25-cuda-performance-p1-p2-p3.md b/docs/superpowers/plans/2026-05-25-cuda-performance-p1-p2-p3.md new file mode 100644 index 000000000..a9e540983 --- /dev/null +++ b/docs/superpowers/plans/2026-05-25-cuda-performance-p1-p2-p3.md @@ -0,0 +1,377 @@ +# CUDA Performance P1-P3 Implementation Plan + +> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** Take training throughput from 1,840 transitions/sec to 70,000+ on L40S via batch size scaling (P1), complete CUDA Graph capture (P2), and fuse the reward pipeline (P3). + +**Architecture:** P1 is a config change (b=16→256) + PER capacity scaling. P2 wraps remaining kernel sequences with cudarc's begin_capture/end_capture/launch three-state machine. P3 fuses 6 sequential reward kernels into one. All changes are in the RL trainer and CUDA kernel layer. + +**Tech Stack:** Rust 1.85, cudarc 0.19, CUDA 12.4, pre-compiled cubins via build.rs + +--- + +## File Map + +| File | Responsibility | Tasks | +|------|---------------|-------| +| `infra/k8s/argo/alpha-rl-template.yaml` | Default n-backtests param | T1 | +| `crates/ml-alpha/src/trainer/integrated.rs` | Trainer: PER scaling, Graph B+C, borrow fixes | T2, T3, T4, T5 | +| `crates/ml-alpha/cuda/rl_fused_reward_pipeline.cu` | New: fused reward kernel | T6 | +| `crates/ml-alpha/build.rs` | Register new kernel | T6 | + +--- + +## Task 1: P1 — Batch Size Scaling (Argo default + PER auto-scale) + +**Files:** +- Modify: `infra/k8s/argo/alpha-rl-template.yaml:44-45` +- Modify: `crates/ml-alpha/src/trainer/integrated.rs` (PER capacity logic) + +- [ ] **Step 1: Update Argo template default n-backtests from 16 to 128** + +```yaml + - name: n-backtests + value: "128" +``` + +- [ ] **Step 2: Auto-scale PER capacity to 4× batch size minimum** + +In `IntegratedTrainer::new()`, after `let b_size = cfg.perception.n_batch;`, replace the fixed PER capacity with: + +```rust +let per_capacity = cfg.per_capacity.max(4 * b_size); +``` + +Find where `cfg.per_capacity` is passed to the replay buffer constructor (around line 1680 where `cfg.per_capacity` appears) and replace with `per_capacity`. + +- [ ] **Step 3: Verify local smoke at b=64** + +```bash +SQLX_OFFLINE=true FOXHUNT_TEST_DATA=test_data/futures-baseline cargo test -p ml-alpha --test integrated_trainer_smoke -- --ignored --nocapture +``` + +The existing smoke uses `IntegratedTrainerConfig::default()` which has `n_batch=16`. Verify it still passes. Then manually test b=64 is allocable on RTX 3050 Ti (4GB): + +```bash +# Quick alloc test — just construct the trainer, don't step +SQLX_OFFLINE=true cargo test -p ml-alpha --test integrated_trainer_smoke -- --ignored --nocapture +``` + +- [ ] **Step 4: Apply the Argo template to cluster** + +```bash +kubectl apply -f infra/k8s/argo/alpha-rl-template.yaml -n foxhunt +``` + +- [ ] **Step 5: Commit** + +```bash +git add infra/k8s/argo/alpha-rl-template.yaml crates/ml-alpha/src/trainer/integrated.rs +git commit -m "perf(rl): scale batch size to 128 default + auto-size PER capacity" +``` + +--- + +## Task 2: P2 — Fix Graph C Borrow Checker (reduce_axis0 pattern) + +**Files:** +- Modify: `crates/ml-alpha/src/trainer/integrated.rs` (reduce_axis0_free calls) + +The E0502 errors occur because `reduce_axis0_free` takes `&self.stream` + `&self.reduce_axis0_fn` (immutable borrow of self) AND `&mut self.ss_q_grad_w_d` (mutable borrow of self) in the same call. The fix is to extract the stream and fn references into locals before the call. + +- [ ] **Step 1: Fix all reduce_axis0_free borrow conflicts** + +The pattern for every call site is: + +```rust +// BEFORE (fails borrow check): +reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_q_grad_w_per_batch_d, b_size, k_dqn * HIDDEN_DIM, &mut self.ss_q_grad_w_d)?; + +// AFTER (compiles): +let (stream, fn_ref) = (&self.stream, &self.reduce_axis0_fn); +reduce_axis0_free(stream, fn_ref, &self.ss_q_grad_w_per_batch_d, b_size, k_dqn * HIDDEN_DIM, &mut self.ss_q_grad_w_d)?; +``` + +Alternatively, extract `stream` and `reduce_axis0_fn` into locals at the TOP of `step_synthetic`: + +```rust +let stream = &self.stream; +let reduce_fn = &self.reduce_axis0_fn; +``` + +Then use `stream` and `reduce_fn` throughout. This fixes all ~15 call sites at once. + +- [ ] **Step 2: Verify compilation** + +```bash +SQLX_OFFLINE=true cargo check -p ml-alpha +``` + +Expected: 0 errors. + +- [ ] **Step 3: Run local smoke** + +```bash +SQLX_OFFLINE=true FOXHUNT_TEST_DATA=test_data/futures-baseline cargo test -p ml-alpha --test integrated_trainer_smoke -- --ignored --nocapture +``` + +- [ ] **Step 4: Commit** + +```bash +git add crates/ml-alpha/src/trainer/integrated.rs +git commit -m "fix(rl): resolve borrow-checker conflicts in pre-allocated gradient reduce" +``` + +--- + +## Task 3: P2 — Graph B Capture (post-fill reward pipeline) + +**Files:** +- Modify: `crates/ml-alpha/src/trainer/integrated.rs` + +**Prerequisite:** Task 2 (Graph C buffers compile clean). + +Graph B covers ~20 kernels from `extract_realized_pnl_delta` through `launch_var_over_abs_mean` in `step_with_lobsim`. The `rl_write_u64` kernel and `ts_ns_d` device buffer are already wired. The `postfill_graph` field already exists but currently captures the pre-fill A2 section — rename and add a `reward_graph` field. + +- [ ] **Step 1: Add `reward_graph: Option` field** + +In the struct definition (near line 426): + +```rust + prefill_graph: Option, + postfill_graph: Option, + reward_graph: Option, + graph_warmup_done: bool, +``` + +Initialize in constructor: `reward_graph: None,` + +- [ ] **Step 2: Wrap post-fill section with three-state machine** + +After `lobsim.step_fill_from_market_targets(...)` and before the `extract_realized_pnl_delta` block, add: + +```rust +if self.reward_graph.is_some() { + self.reward_graph.as_ref().unwrap().launch() + .context("reward graph launch")?; +} else { + let capturing_reward = self.postfill_graph.is_some(); + if capturing_reward { + self.stream + .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED) + .map_err(|e| anyhow::anyhow!("reward begin_capture: {e}"))?; + } + + // ... existing post-fill kernels (extract_realized_pnl_delta through + // launch_var_over_abs_mean) stay here unchanged ... + + if capturing_reward { + let graph = self.stream + .end_capture(CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH) + .context("reward end_capture")? + .ok_or_else(|| anyhow::anyhow!("reward end_capture returned None"))?; + self.reward_graph = Some(graph); + } +} // end else (reward warmup / capture) +``` + +- [ ] **Step 3: Verify no host ops inside capture region** + +Audit: no `alloc_zeros`, no `stream.synchronize()`, no `lobsim.pos_bytes()` HOST calls inside the region. The `lobsim.pos_d()`, `lobsim.bid_px_d()` calls return stable device pointers — safe. + +The `launch_rl_fused_controllers(b_size)` and `launch_var_over_abs_mean(...)` are method calls — verify they contain only kernel launches (already confirmed: pure `launch_builder` + `launch`). + +- [ ] **Step 4: Run local smoke + compute-sanitizer** + +```bash +SQLX_OFFLINE=true FOXHUNT_TEST_DATA=test_data/futures-baseline cargo test -p ml-alpha --test integrated_trainer_smoke -- --ignored --nocapture +``` + +- [ ] **Step 5: Commit** + +```bash +git add crates/ml-alpha/src/trainer/integrated.rs +git commit -m "feat(rl): CUDA Graph B capture for post-fill reward/EMA/controller pipeline (~20 kernels)" +``` + +--- + +## Task 4: P2 — Graph C Capture (replay training step) + +**Files:** +- Modify: `crates/ml-alpha/src/trainer/integrated.rs` + +**Prerequisite:** Tasks 2+3 (buffers pre-allocated, compiles clean). + +Graph C captures `step_synthetic`'s kernel sequence. The three-state machine is the same pattern. Add `training_graph: Option`. + +- [ ] **Step 1: Add `training_graph: Option` field + warmup counter** + +```rust + training_graph: Option, + training_graph_warmup_done: bool, +``` + +Initialize: `training_graph: None, training_graph_warmup_done: false,` + +- [ ] **Step 2: Wrap step_synthetic with capture** + +At the top of `step_synthetic`, after the `let b_size = ...` + `let k_dqn = ...` lines: + +```rust +if self.training_graph.is_some() { + self.training_graph.as_ref().unwrap().launch() + .context("training graph launch")?; + // Skip the rest — graph replays all kernels. + self.stream.synchronize().context("training graph sync")?; + // Read back losses from mapped-pinned... + return Ok(step_synthetic_stats_from_device(self)?); +} + +let capturing_training = self.training_graph_warmup_done; +if !self.training_graph_warmup_done { + self.training_graph_warmup_done = true; +} +if capturing_training { + self.stream + .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED) + .map_err(|e| anyhow::anyhow!("training begin_capture: {e}"))?; +} + +// ... existing step_synthetic body ... + +if capturing_training { + let graph = self.stream + .end_capture(CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH) + .context("training end_capture")? + .ok_or_else(|| anyhow::anyhow!("training end_capture returned None"))?; + self.training_graph = Some(graph); +} +``` + +- [ ] **Step 3: Extract loss readback to separate method** + +The graph replay skips the kernel dispatches but still needs to read loss scalars back. Extract the mapped-pinned loss reads at the end of step_synthetic into a helper that both the graph-replay path and the capture path call AFTER sync. + +- [ ] **Step 4: Verify no host ops in capture region** + +Check: no `alloc_zeros`, no `synchronize`, no host-side branching. The `self.launch_*` helper methods must be pure kernel launches. + +- [ ] **Step 5: Run local smoke** + +```bash +SQLX_OFFLINE=true FOXHUNT_TEST_DATA=test_data/futures-baseline cargo test -p ml-alpha --test integrated_trainer_smoke -- --ignored --nocapture +``` + +- [ ] **Step 6: Commit** + +```bash +git add crates/ml-alpha/src/trainer/integrated.rs +git commit -m "feat(rl): CUDA Graph C capture for replay training step (~25 kernels)" +``` + +--- + +## Task 5: P1 — Benchmark b=128 on L40S + +**Files:** +- None (infrastructure test) + +- [ ] **Step 1: Push and submit b=128 run** + +```bash +git push origin ml-alpha-phase-a +argo submit -n foxhunt --from=wftmpl/alpha-rl \ + -p git-branch=ml-alpha-phase-a \ + -p n-steps=5000 \ + -p n-backtests=128 +``` + +- [ ] **Step 2: Measure throughput** + +Wait for completion, check `elapsed` in logs. Calculate steps/sec and transitions/sec. + +Expected: ~700 steps/sec × 128 = ~89,600 transitions/sec (48× improvement over baseline). + +- [ ] **Step 3: If passes, submit b=256** + +```bash +argo submit -n foxhunt --from=wftmpl/alpha-rl \ + -p git-branch=ml-alpha-phase-a \ + -p n-steps=5000 \ + -p n-backtests=256 +``` + +--- + +## Task 6: P3 — Fused Reward Pipeline Kernel + +**Files:** +- Create: `crates/ml-alpha/cuda/rl_fused_reward_pipeline.cu` +- Modify: `crates/ml-alpha/build.rs` +- Modify: `crates/ml-alpha/src/trainer/integrated.rs` + +- [ ] **Step 1: Write the fused kernel** + +```cuda +// rl_fused_reward_pipeline.cu — Fuses 6 sequential post-fill kernels: +// extract_realized_pnl_delta + rl_reward_shaping + abs_copy + +// ema_update_on_done(duration) + ema_update_on_done(abs_pnl) + +// ema_update_on_done(outcome) + rl_recent_outcome_update. +// +// Single launch replaces 7 kernels. Eliminates 6 L2 cache round-trips +// on rewards_d and 6 kernel launch overheads. +// +// Grid=(1,1,1), Block=(b_size, 1, 1). Each thread handles one batch. +// Shared memory for tree-reduce of EMAs. +``` + +The kernel: +1. Reads `pos_d[b]`, `prev_realized_pnl_d[b]` → computes reward, done +2. Applies reward shaping (entry cost, hold bonus, min-hold penalty) +3. Writes `raw_rewards_d[b]` (snapshot before scale) +4. Computes `|reward|` → `reward_abs_d[b]` +5. Updates `steps_since_done_d[b]`, emits `trade_duration_emit_d[b]` +6. Tree-reduce: 3× EMA updates to ISV (mean_trade_duration, mean_abs_pnl, outcome) +7. Writes `outcome_ema_d[b]` (per-batch) +8. Writes final `rewards_d[b]`, `dones_d[b]` + +- [ ] **Step 2: Register in build.rs** + +```rust +"rl_fused_reward_pipeline", // P3: 7→1 fused reward extraction + shaping + EMA update +``` + +- [ ] **Step 3: Wire into trainer — replace 7 kernel launches** + +Replace the section from `extract_realized_pnl_delta` through `rl_recent_outcome_update` with a single launch of the fused kernel. + +- [ ] **Step 4: Run local smoke + oracle test** + +Verify loss trajectory matches the unfused version (within fp noise). + +- [ ] **Step 5: Commit** + +```bash +git add crates/ml-alpha/cuda/rl_fused_reward_pipeline.cu crates/ml-alpha/build.rs crates/ml-alpha/src/trainer/integrated.rs +git commit -m "perf(rl): fuse 7 post-fill reward kernels into one launch (P3)" +``` + +--- + +## Kill Criteria + +- P1: b=128 completes without OOM on L40S and throughput ≥ 5× vs b=16 +- P2 Graph B: local smoke passes after capture (step 1 warmup, step 2 capture, step 3+ replay) +- P2 Graph C: same three-state machine, local smoke passes +- P3: loss trajectory at step 1000 matches within 1% of unfused version + +## Expected Final Throughput + +| Config | Transitions/sec | vs Baseline | +|--------|----------------|-------------| +| Baseline (b=16, no graphs) | 1,840 | 1× | +| P1 only (b=128) | ~89,600 | 48× | +| P1 + P2 (b=128, all graphs) | ~92,000 | 50× | +| P1 + P2 + P3 (b=128, fused) | ~95,000 | 52× | +| P1 at b=256 | ~150,000+ | 80×+ | diff --git a/docs/superpowers/specs/2026-05-25-cuda-performance-roadmap.md b/docs/superpowers/specs/2026-05-25-cuda-performance-roadmap.md new file mode 100644 index 000000000..ca4ec7d31 --- /dev/null +++ b/docs/superpowers/specs/2026-05-25-cuda-performance-roadmap.md @@ -0,0 +1,143 @@ +# CUDA Performance Roadmap + +**Goal:** Maximize training throughput (steps/sec × batch_size = transitions/sec) on L40S (48GB) and H100 (80GB). + +**Current baseline:** 115 steps/sec at b=16 on L40S = **1,840 transitions/sec**. + +--- + +## P1: Batch Size Scaling (highest impact, minimal code) + +**Problem:** At b=16, kernels complete in microseconds — GPU SMs are idle 90%+ of the time waiting for the next launch. Launch overhead (~10μs per kernel) dominates. + +**Memory budget per batch:** 238 KB (dominated by `q_grad_w_per_batch` at 115 KB for the C51 backward). + +**Available VRAM:** + +| GPU | Total | Fixed overhead | Available | Max b (theoretical) | Recommended b | +|-----|-------|----------------|-----------|---------------------|---------------| +| RTX 3050 Ti | 4 GB | ~90 MB | 3.4 GB | 15,000 | 64 (smoke) | +| L40S | 48 GB | ~90 MB | 47 GB | 208,000 | 256 | +| H100 | 80 GB | ~90 MB | 79 GB | 350,000 | 512 | + +**Expected throughput at higher batch sizes (L40S):** + +| b_size | Transitions/sec | Speedup vs b=16 | Notes | +|--------|----------------|------------------|-------| +| 16 | 1,840 | 1.0× | Current — launch-limited | +| 64 | ~25,600 | ~14× | GPU starts saturating | +| 128 | ~44,800 | ~24× | Good SM utilization | +| 256 | ~71,000 | ~39× | Near peak for these kernels | +| 512 | ~92,000 | ~50× | Diminishing returns | + +**Implementation:** +- `n_backtests` CLI param already controls b_size +- PER capacity must scale: `per_capacity = max(32768, 4 * b_size)` +- Encoder K-loop seq_len is independent of batch size +- Verify: `--n-backtests 256` on L40S, measure actual throughput + +**Blockers:** None. Ship today. + +--- + +## P2: CUDA Graph Capture (done for A+A2, remaining B+C) + +**Status:** +- Graph A (pre-snapshot, 20 kernels): ✅ captured +- Graph A2 (post-snapshot/pre-fill, 7 kernels): ✅ captured +- Graph B (post-fill reward/EMA/controllers, ~20 kernels): in progress + - Blocker resolved: `ts_ns` moved to device-resident u64 via `rl_write_u64` kernel + - 51 gradient buffers being pre-allocated for Graph C +- Graph C (replay training step, ~25 kernels × K iterations): in progress + - Blocker: 51 per-step `alloc_zeros` → persistent fields (agent working) + +**Expected gain:** At b=16, ~5% (launch overhead is small fraction). At b=256, negligible (<1%). **Graph capture is insurance for when we increase K (replay-to-env ratio).** + +--- + +## P3: Kernel Fusion — Reward Pipeline + +**Problem:** The post-fill reward pipeline launches 6+ small sequential kernels on the same data: `extract_realized_pnl_delta` → `rl_reward_shaping` → `abs_copy` → `apply_reward_scale` → 3× `ema_update_on_done`. Each reads/writes rewards_d, dones_d. + +**Fix:** Fuse into `rl_fused_reward_pipeline.cu`: +- Single kernel, one block per batch +- Reads pos_d, prev_realized_pnl_d once +- Writes rewards_d, dones_d, reward_abs_d, raw_rewards_d +- Updates 3 ISV EMA slots inline +- Saves 5 kernel launches + 5× L2 cache round-trips on rewards_d + +**Expected gain:** ~50μs/step at b=16, ~200μs at b=256 (memory bandwidth limited). + +--- + +## P4: FP16 Gradient Accumulation + +**Problem:** `reduce_axis0` and Adam steps are memory-bandwidth bound. Reading/writing f32 gradients at full precision wastes half the bandwidth. + +**Fix:** Mixed-precision gradient pipeline: +- Forward/backward compute stays f32 (numerical stability) +- Per-batch gradient scratch (`*_per_batch_d`) stored as f16 +- `reduce_axis0` reads f16, accumulates f32, writes f32 reduced gradient +- Adam reads f32 gradient, updates f32 weights + +**Expected gain:** 2× bandwidth on gradient reduce + Adam = ~30% speedup on the training step (reduce_axis0 is ~40% of step_synthetic). + +--- + +## P5: Multi-Stream Overlap + +**Problem:** The step pipeline is strictly sequential: env-step → PER push/sample (host) → replay training. The host-side PER operations block the GPU. + +**Fix:** Two CUDA streams: +- Stream 0: env-step (Graphs A → fill → B) + PER push +- Stream 1: replay training (Graph C × K) from PREVIOUS step's PER sample + +Pipeline: while stream 1 trains on step N's data, stream 0 runs step N+1's env-step. PER sample for step N+1 happens during stream 1's training. + +**Expected gain:** Hides PER latency (~200μs) + env-step overlap. ~1.3× at K=1, ~1.1× at K=4 (training dominates). + +**Prerequisite:** Graph C captured (P2). + +--- + +## P6: Persistent Kernel for K-loop + +**Problem:** The replay training step runs K times per env step. Each iteration launches Graph C + syncs. At K=4, that's 4 graph launches + 4 syncs. + +**Fix:** Single persistent kernel that: +- Stays resident on SMs +- Reads a "work counter" from device memory +- For each K: reads PER sample indices from a ring buffer, runs the full training step inline +- No host interaction until all K iterations complete + +**Expected gain:** Eliminates K-1 sync points. At K=4: ~300μs saved/step. At K=8: ~700μs. + +**Prerequisite:** Graph C working (P2), multi-stream (P5). + +--- + +## P7: Warp-Cooperative Action Selection + +**Problem:** `rl_pi_action_kernel` uses 1 thread per batch (sequential CDF walk). At b=256, that's 256 blocks × 1 thread = 256 SMs used at 0.4% occupancy each. + +**Fix:** Warp-cooperative softmax + CDF walk: +- 32 threads per batch (one warp) +- Warp-shuffle for parallel softmax reduction +- Parallel prefix-sum for CDF +- Single-warp ballot for multinomial threshold crossing + +**Expected gain:** ~10× faster action selection kernel. Negligible at b=16, ~50μs at b=256. + +--- + +## Priority Order + +1. **P1 (batch size)** — Ship immediately, 14-50× throughput increase +2. **P2 (Graph B+C)** — In progress, enables P5/P6 +3. **P3 (fused reward)** — Medium effort, good constant-factor win +4. **P4 (FP16 grads)** — 30% training step speedup +5. **P5 (multi-stream)** — 1.3× overlap, needs P2 +6. **P6 (persistent K-loop)** — Eliminates sync overhead, needs P5 +7. **P7 (warp-coop action)** — Polish, only matters at large b + +**Target:** P1 alone takes us from 1,840 to ~70,000 transitions/sec on L40S (38×). Combined with P2-P4: **~100,000 transitions/sec** — 1M steps in 10 seconds.