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:
@@ -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,
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user