From ab6fbcb668fec347d25912b2356b68a4041508f6 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sun, 3 May 2026 01:33:45 +0200 Subject: [PATCH] sp7(trainer): load loss-balance controller kernel + launcher fn MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit LOSS_BALANCE_CONTROLLER_CUBIN static, kernel slot in GpuDqnTrainer, 6 SCRATCH_LB_* index constants (218..242), pub(crate) launch_loss_balance_controller that runs the producer + 24 apply_pearls_ad smoothing launches. SP5_SCRATCH_TOTAL bumped 218→242 with updated docblock and allocation comment to match the 24 new slots. Component pointers into grad_decomp_result_pinned: IQN at offset 0, CQL_SX at offset 6, C51 at offset 9 (3-float layout per launch_grad_decomp docstring). Producer-only — call site in training_loop.rs lands at T7 atomically with the consumer floor change and stale-doc deletion per feedback_no_partial_refactor. Co-Authored-By: Claude Opus 4.7 (1M context) --- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 184 +++++++++++++++++- docs/dqn-wire-up-audit.md | 5 +- 2 files changed, 185 insertions(+), 4 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index dc70dee93..8254f8feb 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -446,6 +446,20 @@ static SP5_HEALTH_COMPOSITION_CUBIN: &[u8] = static SP5_TRAINING_METRICS_EMA_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/training_metrics_ema_kernel.cubin")); +/// SP7 Task 5 (2026-05-03): loss-balance controller cubin. +/// 8-thread single-block kernel (2 loss heads × 4 branches). Reads three +/// 3-float views into `grad_decomp_result_dev_ptr` (IQN at element offset 0, +/// CQL_SX at offset 6, C51 at offset 9), FLATNESS_BASE, prior budgets, and +/// prior Wiener state from ISV; writes new_budget_cql[4] + new_budget_c51[4] +/// + diff_var/sample_var for each (6 groups × 4 branches = 24 floats) to +/// `producer_step_scratch_buf[218..242)`. Cold-start sentinel-aware. +/// Followed by 24 `apply_pearls_ad_kernel` launches → ISV[BUDGET_CQL_BASE, +/// BUDGET_C51_BASE, LB_DIFF_VAR_CQL_BASE, LB_SAMPLE_VAR_CQL_BASE, +/// LB_DIFF_VAR_C51_BASE, LB_SAMPLE_VAR_C51_BASE]. +/// Producer-only — call site wires at T7 per `feedback_no_partial_refactor`. +static LOSS_BALANCE_CONTROLLER_CUBIN: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/loss_balance_controller_kernel.cubin")); + /// Plan 4 Task 3 (E.3): IQN multi-quantile diagnostic EMAs into ISV[99..103). /// 4-block (256 threads/block, shmem-reduce, no atomicAdd) kernel launched /// alongside `h_s2_rms_ema_update` from `training_loop.rs`. Reads the IQN @@ -1081,8 +1095,15 @@ pub const SP5_WIENER_TOTAL_FLOATS: usize = /// 4 floats for health_composition output (SCRATCH_HEALTH_COMP_BASE=211, score/q_gap_norm/q_var_norm/grad_norm_norm at [211..215)) /// Layer D Task D3 adds: /// 3 floats for training_metrics_ema output (SCRATCH_TRAINING_METRICS_EMA_BASE=215, training_sharpe_ema/max_dd_ema/low_dd_ratio at [215..218)) -/// Combined: 71 + 16 + 16 + 4 + 4 + 20 + 8 + 8 + 8 + 8 + 8 + 4 + 4 + 20 + 4 + 4 + 4 + 4 + 3 = 218 scratch slots [0..218). -pub const SP5_SCRATCH_TOTAL: usize = 218; +/// SP7 Task 5 adds: +/// 4 floats for LB new_budget_cql[4] (SCRATCH_LB_BUDGET_CQL=218, at [218..222)) +/// 4 floats for LB new_budget_c51[4] (SCRATCH_LB_BUDGET_C51=222, at [222..226)) +/// 4 floats for LB diff_var_cql[4] (SCRATCH_LB_DIFF_VAR_CQL=226, at [226..230)) +/// 4 floats for LB sample_var_cql[4] (SCRATCH_LB_SAMPLE_VAR_CQL=230, at [230..234)) +/// 4 floats for LB diff_var_c51[4] (SCRATCH_LB_DIFF_VAR_C51=234, at [234..238)) +/// 4 floats for LB sample_var_c51[4] (SCRATCH_LB_SAMPLE_VAR_C51=238, at [238..242)) +/// Combined: 71 + 16 + 16 + 4 + 4 + 20 + 8 + 8 + 8 + 8 + 8 + 4 + 4 + 20 + 4 + 4 + 4 + 4 + 3 + 24 = 242 scratch slots [0..242). +pub const SP5_SCRATCH_TOTAL: usize = 242; /// SP5 Layer D Task D1 (rewrite, 2026-05-02): scratch index base for /// `pnl_aggregation_update`. @@ -1138,6 +1159,21 @@ pub const SCRATCH_HEALTH_COMP_BASE: usize = 211; /// EMA at `training_loop.rs:5052`). pub const SCRATCH_TRAINING_METRICS_EMA_BASE: usize = 215; +/// SP7 Task 5 (2026-05-03): scratch index base for loss-balance controller +/// new_budget_cql[4]. Smoothed by apply_pearls_ad_kernel into +/// ISV[BUDGET_CQL_BASE..+4). +pub const SCRATCH_LB_BUDGET_CQL: usize = 218; // 218..222 +/// SP7: new_budget_c51[4] → ISV[BUDGET_C51_BASE..+4). +pub const SCRATCH_LB_BUDGET_C51: usize = SCRATCH_LB_BUDGET_CQL + 4; // 222..226 +/// SP7: diff_var_cql[4] → ISV[LB_DIFF_VAR_CQL_BASE..+4). +pub const SCRATCH_LB_DIFF_VAR_CQL: usize = SCRATCH_LB_BUDGET_CQL + 8; // 226..230 +/// SP7: sample_var_cql[4] → ISV[LB_SAMPLE_VAR_CQL_BASE..+4). +pub const SCRATCH_LB_SAMPLE_VAR_CQL: usize = SCRATCH_LB_BUDGET_CQL + 12; // 230..234 +/// SP7: diff_var_c51[4] → ISV[LB_DIFF_VAR_C51_BASE..+4). +pub const SCRATCH_LB_DIFF_VAR_C51: usize = SCRATCH_LB_BUDGET_CQL + 16; // 234..238 +/// SP7: sample_var_c51[4] → ISV[LB_SAMPLE_VAR_C51_BASE..+4). +pub const SCRATCH_LB_SAMPLE_VAR_C51: usize = SCRATCH_LB_BUDGET_CQL + 20; // 238..242 + /// SP5 Task A7: scratch index base for pearl_8_trail_update trail_dist[4] output block. /// Slots [199..203): per-direction trail-stop distance (Short=0, Hold=1, Long=2, Flat=3). /// Written by `pearl_8_trail_update`; consumed by apply_pearls_ad_kernel → @@ -4384,6 +4420,16 @@ pub struct GpuDqnTrainer { /// No atomicAdd (feedback_no_atomicadd), no CPU compute (feedback_no_cpu_compute_strict). /// Loaded from `pearl_2_budget_kernel.cubin`. pearl_2_budget_kernel: CudaFunction, + /// SP7 Task 5 (2026-05-03): loss-balance controller (CQL + C51 per-branch + /// budget adapter). 8-thread single-block kernel (2 loss heads × 4 branches). + /// Reads grad_decomp_result_dev_ptr (IQN@0, CQL_SX@6, C51@9), FLATNESS_BASE, + /// prior budgets, prior Wiener state; writes 24 floats to + /// `producer_step_scratch_buf[218..242)`. Followed by 24 apply_pearls_ad_kernel + /// launches → ISV[BUDGET_CQL_BASE, BUDGET_C51_BASE, LB_*_VAR_{CQL,C51}_BASE]. + /// Producer-only — call site in training_loop.rs lands at T7 per + /// `feedback_no_partial_refactor`. Loaded from + /// `loss_balance_controller_kernel.cubin`. + loss_balance_controller_kernel: CudaFunction, // ── SP5 Task A4: Pearl 4 per-group Adam β1/β2/ε kernels ───────────── /// SP5 Task A4 (2026-05-01): auxiliary gradient cosine similarity kernel. /// Single-block 8-thread kernel (one thread per SP4 param group). @@ -11284,6 +11330,120 @@ impl GpuDqnTrainer { Ok(()) } + /// SP7 Task 5 (2026-05-03): Launch the loss-balance controller producer + + /// apply_pearls_ad_kernel chain. Must run AFTER: + /// * `launch_sp5_pearl_2_budget` (FLATNESS_BASE populated) + /// * `grad_decomp_launch_iqn` (pinned slot at element offset 0) + /// * `grad_decomp_launch_cql_sx` (pinned slot at element offset 6) + /// * `grad_decomp_launch_c51` (pinned slot at element offset 9) + /// + /// Writes: scratch[SCRATCH_LB_*..]; `apply_pearls_ad_kernel` then smooths + /// these into ISV[BUDGET_{CQL,C51}_BASE] and ISV[LB_*_VAR_{CQL,C51}_BASE]. + pub(crate) fn launch_loss_balance_controller(&self) -> Result<(), MLError> { + use crate::cuda_pipeline::sp4_wiener_ema::launch_apply_pearls; + use crate::cuda_pipeline::sp5_isv_slots::{ + SP5_SLOT_BASE, + BUDGET_CQL_BASE, BUDGET_C51_BASE, FLATNESS_BASE, + LB_DIFF_VAR_CQL_BASE, LB_SAMPLE_VAR_CQL_BASE, + LB_DIFF_VAR_C51_BASE, LB_SAMPLE_VAR_C51_BASE, + }; + + debug_assert!(self.isv_signals_dev_ptr != 0, + "launch_loss_balance_controller: isv_signals_dev_ptr must be allocated"); + debug_assert!(self.grad_decomp_result_dev_ptr != 0, + "launch_loss_balance_controller: grad_decomp_result_dev_ptr must be allocated"); + + let isv_dev = self.isv_signals_dev_ptr; + let scratch_dev = self.producer_step_scratch_buf.dev_ptr; + let wiener_dev = self.wiener_state_buf.dev_ptr; + + // grad_decomp pinned layout: 27 floats total, 9 components × 3 floats + // each ([mag, dir, trunk]). Component element offsets per + // launch_grad_decomp documentation: iqn=0, cql=3, cql_sx=6, c51=9. + // We use cql_sx (post-budget delta) for ratio parity with what landed + // in grad_buf. + let f32_size = std::mem::size_of::() as u64; + let iqn_dev = self.grad_decomp_result_dev_ptr + 0 * f32_size; + let cql_dev = self.grad_decomp_result_dev_ptr + 6 * f32_size; + let c51_dev = self.grad_decomp_result_dev_ptr + 9 * f32_size; + + // Step 1: producer kernel. + let flatness_isv_base_i32 = FLATNESS_BASE as i32; + let budget_cql_isv_base_i32 = BUDGET_CQL_BASE as i32; + let budget_c51_isv_base_i32 = BUDGET_C51_BASE as i32; + let diff_var_cql_isv_base_i32 = LB_DIFF_VAR_CQL_BASE as i32; + let sample_var_cql_isv_base_i32 = LB_SAMPLE_VAR_CQL_BASE as i32; + let diff_var_c51_isv_base_i32 = LB_DIFF_VAR_C51_BASE as i32; + let sample_var_c51_isv_base_i32 = LB_SAMPLE_VAR_C51_BASE as i32; + let sb_budget_cql_i32 = SCRATCH_LB_BUDGET_CQL as i32; + let sb_budget_c51_i32 = SCRATCH_LB_BUDGET_C51 as i32; + let sb_diff_var_cql_i32 = SCRATCH_LB_DIFF_VAR_CQL as i32; + let sb_sample_var_cql_i32 = SCRATCH_LB_SAMPLE_VAR_CQL as i32; + let sb_diff_var_c51_i32 = SCRATCH_LB_DIFF_VAR_C51 as i32; + let sb_sample_var_c51_i32 = SCRATCH_LB_SAMPLE_VAR_C51 as i32; + + unsafe { + self.stream + .launch_builder(&self.loss_balance_controller_kernel) + .arg(&iqn_dev) + .arg(&cql_dev) + .arg(&c51_dev) + .arg(&isv_dev) + .arg(&flatness_isv_base_i32) + .arg(&budget_cql_isv_base_i32) + .arg(&budget_c51_isv_base_i32) + .arg(&diff_var_cql_isv_base_i32) + .arg(&sample_var_cql_isv_base_i32) + .arg(&diff_var_c51_isv_base_i32) + .arg(&sample_var_c51_isv_base_i32) + .arg(&scratch_dev) + .arg(&sb_budget_cql_i32) + .arg(&sb_budget_c51_i32) + .arg(&sb_diff_var_cql_i32) + .arg(&sb_sample_var_cql_i32) + .arg(&sb_diff_var_c51_i32) + .arg(&sb_sample_var_c51_i32) + .launch(LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (8, 1, 1), + shared_mem_bytes: 0, + }) + .map_err(|e| MLError::ModelError(format!("loss_balance_controller_update: {e}")))?; + } + + // Step 2: apply_pearls_ad_kernel × 24 — one per ISV output slot. + // Wiener offset: (SP4_PRODUCER_COUNT + (isv_slot - SP5_SLOT_BASE)) * 3. + let base_wiener_offset = SP4_PRODUCER_COUNT as i32 * 3; + + for (isv_base, scratch_base) in [ + (BUDGET_CQL_BASE, SCRATCH_LB_BUDGET_CQL), + (BUDGET_C51_BASE, SCRATCH_LB_BUDGET_C51), + (LB_DIFF_VAR_CQL_BASE, SCRATCH_LB_DIFF_VAR_CQL), + (LB_SAMPLE_VAR_CQL_BASE, SCRATCH_LB_SAMPLE_VAR_CQL), + (LB_DIFF_VAR_C51_BASE, SCRATCH_LB_DIFF_VAR_C51), + (LB_SAMPLE_VAR_C51_BASE, SCRATCH_LB_SAMPLE_VAR_C51), + ] { + for b in 0..4_usize { + let scratch_idx = (scratch_base + b) as i32; + let isv_idx = (isv_base + b) as i32; + let wiener_off = base_wiener_offset + (isv_idx - SP5_SLOT_BASE as i32) * 3; + unsafe { + launch_apply_pearls( + &self.stream, + &self.apply_pearls_ad_kernel, + scratch_dev, scratch_idx, + isv_dev, isv_idx, + wiener_dev, wiener_off, + 1, + crate::cuda_pipeline::sp4_wiener_ema::ALPHA_META, + )?; + } + } + } + + Ok(()) + } + /// SP5 Task A4 (2026-05-01): launch the two-kernel Pearl 4 chain + /// 24 `apply_pearls_ad_kernel` launches to smooth β1/β2/ε into ISV. /// @@ -14348,6 +14508,21 @@ impl GpuDqnTrainer { .map_err(|e| MLError::ModelError(format!("training_metrics_ema_update load: {e}")))? }; + // SP7 Task 5 (2026-05-03): load loss_balance_controller_kernel. + // 8-thread single-block kernel (2 loss heads × 4 branches). Reads three + // 3-float views into grad_decomp_result_dev_ptr (IQN@0, CQL_SX@6, C51@9), + // FLATNESS_BASE, prior budgets, and prior Wiener state from ISV; writes + // 24 floats to scratch[SCRATCH_LB_BUDGET_CQL=218..242). Followed by 24 + // apply_pearls_ad_kernel launches → ISV[BUDGET_CQL_BASE, BUDGET_C51_BASE, + // LB_*_VAR_{CQL,C51}_BASE]. Producer-only — call site lands at T7 per + // feedback_no_partial_refactor.md. + let loss_balance_controller_kernel = { + let module = stream.context().load_cubin(LOSS_BALANCE_CONTROLLER_CUBIN.to_vec()) + .map_err(|e| MLError::ModelError(format!("loss_balance_controller cubin load: {e}")))?; + module.load_function("loss_balance_controller_update") + .map_err(|e| MLError::ModelError(format!("loss_balance_controller_update load: {e}")))? + }; + // SP5 Task A4 (2026-05-01): allocate grad_prev_buf_per_group [total_params f32]. // Zeroed at construction; fold-boundary reset zeroes it again (Pearl A sentinel: // zero grad_prev → cosine_sim=0 on first step → β1/β2 start at envelope midpoints). @@ -15943,7 +16118,9 @@ impl GpuDqnTrainer { // [203..207) pearl_1_ext num_atoms (Pearl 1-ext / A8) // [207..211) pnl_aggregation (Layer D D1) // [211..215) health_composition (Layer D D2) - // Total: SP5_SCRATCH_TOTAL = 215. Audit doc records every commit's growth. + // [215..218) training_metrics_ema (Layer D D3) + // [218..242) loss_balance_ctrl (SP7 Task 5 — budget_cql/c51 + 4 Wiener vars) + // Total: SP5_SCRATCH_TOTAL = 242. Audit doc records every commit's growth. let producer_step_scratch_buf = unsafe { MappedF32Buffer::new(SP5_SCRATCH_TOTAL) } .map_err(|e| MLError::ModelError(format!("SP4/SP5 producer_step_scratch_buf alloc: {e}")))?; @@ -17526,6 +17703,7 @@ impl GpuDqnTrainer { pearl_1_atom_kernel, pearl_3_sigma_kernel, pearl_2_budget_kernel, + loss_balance_controller_kernel, grad_cosine_sim_kernel, pearl_4_adam_hparams_kernel, q_skew_kurtosis_kernel, diff --git a/docs/dqn-wire-up-audit.md b/docs/dqn-wire-up-audit.md index ace0ab47f..273d8d326 100644 --- a/docs/dqn-wire-up-audit.md +++ b/docs/dqn-wire-up-audit.md @@ -3897,7 +3897,10 @@ extended. No producer kernel yet — that arrives in the next commit. - T4 (commit ⟨pending⟩): build.rs cubin manifest entry. nvcc compiles loss_balance_controller_kernel.cu to $OUT_DIR/...cubin; consumed by gpu_dqn_trainer.rs in T5. -- T5: trainer struct + launcher fn. +- T5 (commit ⟨pending⟩): trainer struct + cubin static + 6 SCRATCH_LB_* + constants + `launch_loss_balance_controller(&self)` fn (the producer + kernel launch + 24 apply_pearls_ad chain). Producer-only — call site + in training_loop.rs lands at T7 (atomic with consumer + stale doc). - T6: Pearl 2 contract change (drop CQL/C51/ENS args). - T7: launch site + sentinel-aware bootstrap with bootstrap constants matching the kernel's cold-start basis (defined in T7) + stale doc. - T9–T10: smoke + 50-epoch verification.