sp7(trainer): load loss-balance controller kernel + launcher fn

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) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-03 01:33:45 +02:00
parent 8e4d5cb45b
commit ab6fbcb668
2 changed files with 185 additions and 4 deletions

View File

@@ -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::<f32>() 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,

View File

@@ -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.
- T9T10: smoke + 50-epoch verification.