From 0e1ed5e2f733c42e900868cf6a67974b9f23d60e Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Tue, 14 Apr 2026 23:49:34 +0200 Subject: [PATCH] =?UTF-8?q?feat:=20weight=20layout=2026=E2=86=9242=20+=20Q?= =?UTF-8?q?-attn/selectivity/liquid-tau/VSN/GLU=20buffers?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit NUM_WEIGHT_TENSORS expanded for 8 VSN bottleneck + 8 GLU gate tensors. Q-attention (624 params), selectivity (257 params) in separate buffers. Liquid tau: 4 pinned device-mapped floats for per-branch modulation. VSN/GLU scratch buffers allocated for forward/backward. 6 CUDA kernel functions loaded from experience_kernels cubin. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 181 +++++++++++++++++- 1 file changed, 175 insertions(+), 6 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index ce8554243..6c7af7f4d 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -19,7 +19,7 @@ //! `train_step()` call: //! //! - **`graph_forward`**: zero → cuBLAS forward → C51 loss → C51 grad → cuBLAS backward -//! - **`graph_adam`**: grad_norm → Adam → unflatten (26 d2d copies) +//! - **`graph_adam`**: grad_norm → Adam → unflatten (42 d2d copies) //! //! Between the two graph replays, external code can inject auxiliary gradients //! (IQN trunk, attention, ensemble) into `grad_buf` via SAXPY. The Adam graph @@ -43,8 +43,9 @@ //! ## Parameter layout //! //! The flat parameter/gradient/moment buffers use the same layout as the CUDA -//! `GOFF_*` defines: 26 weight tensors concatenated in order (w_s1, b_s1, w_s2, -//! b_s2, ..., w_b3out, b_b3out, w_bn, b_bn). See `compute_param_sizes()`. +//! `GOFF_*` defines: 42 weight tensors concatenated in order (w_s1, b_s1, w_s2, +//! b_s2, ..., w_b3out, b_b3out, w_bn, b_bn, w_vsn1_0..w_vsn2_3, w_gate_0..b_gate_3). +//! See `compute_param_sizes()`. use std::sync::Arc; @@ -349,14 +350,17 @@ impl Default for CausalInterventionConfig { } /// Number of weight tensors in the flat parameter buffer. -/// 20 original (DQN network) + 4 branch 3 (magnitude) + 2 bottleneck (w_bn, b_bn) = 26. -pub(crate) const NUM_WEIGHT_TENSORS: usize = 26; +/// 20 original (DQN network) + 4 branch 3 (magnitude) + 2 bottleneck (w_bn, b_bn) +/// + 8 VSN bottleneck (R=16) + 8 GLU gate = 42. +pub(crate) const NUM_WEIGHT_TENSORS: usize = 42; /// Compute the size (element count) of each weight tensor. /// /// Tensors 0-19: standard DQN network (shared trunk + value + branches 0-2). /// Tensors 20-23: branch 3 (urgency) weights. /// Tensors 24-25: temporal causal bottleneck (0 elements when bottleneck_dim=0). +/// Tensors 26-33: Variable Selection Network bottleneck (R=16). +/// Tensors 34-41: GLU gate weights and biases. /// /// When bottleneck is active, w_s1 input dimension changes from state_dim to /// (bottleneck_dim + portfolio_dim) where portfolio_dim = state_dim - market_dim. @@ -399,6 +403,24 @@ pub(crate) fn compute_param_sizes(cfg: &GpuDqnTrainConfig) -> [usize; NUM_WEIGHT // #31 Temporal Causal Bottleneck (0 when disabled) bn_dim * market_dim, // [24] w_bn [bottleneck_dim, market_dim] bn_dim, // [25] b_bn [bottleneck_dim] + // ── Variable Selection bottleneck (R=16) ── + cfg.shared_h2 * 16, // [26] w_vsn1_0 [SH2, R] + 16 * cfg.shared_h2, // [27] w_vsn2_0 [R, SH2] + cfg.shared_h2 * 16, // [28] w_vsn1_1 + 16 * cfg.shared_h2, // [29] w_vsn2_1 + cfg.shared_h2 * 16, // [30] w_vsn1_2 + 16 * cfg.shared_h2, // [31] w_vsn2_2 + cfg.shared_h2 * 16, // [32] w_vsn1_3 + 16 * cfg.shared_h2, // [33] w_vsn2_3 + // ── GLU gate weights ── + cfg.adv_h * cfg.shared_h2, // [34] w_gate_0 [AH, SH2] + cfg.adv_h, // [35] b_gate_0 + cfg.adv_h * (cfg.shared_h2 + 3), // [36] w_gate_1 [AH, SH2+3] + cfg.adv_h, // [37] b_gate_1 + cfg.adv_h * cfg.shared_h2, // [38] w_gate_2 + cfg.adv_h, // [39] b_gate_2 + cfg.adv_h * cfg.shared_h2, // [40] w_gate_3 + cfg.adv_h, // [41] b_gate_3 ] } @@ -593,6 +615,37 @@ pub struct GpuDqnTrainer { /// Kernel: strided_accumulate — extracts d_h_s2 from d_mag_concat in backward. strided_accumulate_kernel: CudaFunction, + // ── Cross-Branch Q Attention ── + q_attn_params: CudaSlice, // [624] + q_attn_adam_m: CudaSlice, // [624] + q_attn_adam_v: CudaSlice, // [624] + q_attn_adam_step: i32, + q_coord_buf: CudaSlice, // [B, 12] + q_attn_kernel: CudaFunction, + + // ── Selectivity gate ── + sel_params: CudaSlice, // [SH2+1] + sel_adam_m: CudaSlice, // [SH2+1] + sel_adam_v: CudaSlice, // [SH2+1] + sel_grad: CudaSlice, // [SH2+1] + sel_adam_step: i32, + sel_out_buf: CudaSlice, // [B] + sel_fwd_kernel: CudaFunction, + sel_bwd_kernel: CudaFunction, + + // ── Liquid tau ── + per_branch_q_gap_ema: [f32; 4], + liquid_mod_pinned: *mut f32, // 4 floats pinned device-mapped + liquid_mod_dev_ptr: u64, + + // ── VSN + GLU scratch ── + vsn_masked_buf: CudaSlice, // [B, SH2] + vsn_kernel: CudaFunction, + glu_gate_pre_buf: [CudaSlice; 4], // 4 × [B, AH] + glu_value_buf: [CudaSlice; 4], // 4 × [B, AH] + glu_combine_kernel: CudaFunction, + glu_backward_kernel: CudaFunction, + save_current_lp: CudaSlice, // [B, NUM_BRANCHES(4), NUM_ATOMS] save_projected: CudaSlice, // [B, NUM_BRANCHES(4), NUM_ATOMS] @@ -1007,6 +1060,9 @@ impl Drop for GpuDqnTrainer { if !self.spread_velocity_pinned.is_null() { let _ = unsafe { cudarc::driver::result::free_host(self.spread_velocity_pinned.cast()) }; } + if !self.liquid_mod_pinned.is_null() { + let _ = unsafe { cudarc::driver::result::free_host(self.liquid_mod_pinned.cast()) }; + } if !self.t_pinned.is_null() { let _ = unsafe { cudarc::driver::result::free_host(self.t_pinned.cast()) }; } @@ -2578,6 +2634,80 @@ impl GpuDqnTrainer { .map_err(|e| MLError::ModelError(format!("strided_accumulate load: {e}")))?; info!("GpuDqnTrainer: mag_concat_qdir + strided_accumulate kernels loaded"); + // ── Cross-Branch Q Attention buffers ────────────────────────── + let q_attn_params = alloc_f32(&stream, 624, "q_attn_params")?; + let q_attn_adam_m = stream.alloc_zeros::(624) + .map_err(|e| MLError::ModelError(format!("alloc q_attn_adam_m: {e}")))?; + let q_attn_adam_v = stream.alloc_zeros::(624) + .map_err(|e| MLError::ModelError(format!("alloc q_attn_adam_v: {e}")))?; + let q_coord_buf = alloc_f32(&stream, b * 12, "q_coord_buf")?; + + // ── Selectivity gate buffers ────────────────────────────────── + let sel_dim = config.shared_h2 + 1; + let sel_params = alloc_f32(&stream, sel_dim, "sel_params")?; + let sel_adam_m = stream.alloc_zeros::(sel_dim) + .map_err(|e| MLError::ModelError(format!("alloc sel_adam_m: {e}")))?; + let sel_adam_v = stream.alloc_zeros::(sel_dim) + .map_err(|e| MLError::ModelError(format!("alloc sel_adam_v: {e}")))?; + let sel_grad = stream.alloc_zeros::(sel_dim) + .map_err(|e| MLError::ModelError(format!("alloc sel_grad: {e}")))?; + let sel_out_buf = alloc_f32(&stream, b, "sel_out_buf")?; + + // ── Liquid tau: 4 pinned device-mapped floats ───────────────── + let liquid_mod_pinned: *mut f32 = unsafe { + let flags = cudarc::driver::sys::CU_MEMHOSTALLOC_DEVICEMAP; + cudarc::driver::result::malloc_host(4 * std::mem::size_of::(), flags) + .map_err(|e| MLError::ModelError(format!("pinned liquid_mod alloc: {e}")))? + as *mut f32 + }; + unsafe { + *liquid_mod_pinned = 1.0; + *liquid_mod_pinned.add(1) = 1.0; + *liquid_mod_pinned.add(2) = 1.0; + *liquid_mod_pinned.add(3) = 1.0; + } + let liquid_mod_dev_ptr = unsafe { + let mut dev_ptr: u64 = 0; + cudarc::driver::sys::cuMemHostGetDevicePointer_v2( + &mut dev_ptr as *mut u64, + liquid_mod_pinned.cast(), + 0, + ); + dev_ptr + }; + + // ── VSN + GLU scratch buffers ───────────────────────────────── + let vsn_masked_buf = alloc_f32(&stream, b * config.shared_h2, "vsn_masked_buf")?; + let glu_gate_pre_buf = [ + alloc_f32(&stream, b * config.adv_h, "glu_gate0")?, + alloc_f32(&stream, b * config.adv_h, "glu_gate1")?, + alloc_f32(&stream, b * config.adv_h, "glu_gate2")?, + alloc_f32(&stream, b * config.adv_h, "glu_gate3")?, + ]; + let glu_value_buf = [ + alloc_f32(&stream, b * config.adv_h, "glu_val0")?, + alloc_f32(&stream, b * config.adv_h, "glu_val1")?, + alloc_f32(&stream, b * config.adv_h, "glu_val2")?, + alloc_f32(&stream, b * config.adv_h, "glu_val3")?, + ]; + + // ── Load Q-attn/selectivity/VSN/GLU kernels from experience_kernels cubin ── + let cpbi_module = stream.context().load_cubin(EXPECTED_Q_CUBIN.to_vec()) + .map_err(|e| MLError::ModelError(format!("cpbi cubin: {e}")))?; + let q_attn_kernel = cpbi_module.load_function("q_cross_branch_attn") + .map_err(|e| MLError::ModelError(format!("q_cross_branch_attn load: {e}")))?; + let sel_fwd_kernel = cpbi_module.load_function("selectivity_gate_fwd") + .map_err(|e| MLError::ModelError(format!("selectivity_gate_fwd load: {e}")))?; + let sel_bwd_kernel = cpbi_module.load_function("selectivity_gate_bwd") + .map_err(|e| MLError::ModelError(format!("selectivity_gate_bwd load: {e}")))?; + let vsn_kernel = cpbi_module.load_function("vsn_bottleneck_fwd") + .map_err(|e| MLError::ModelError(format!("vsn_bottleneck_fwd load: {e}")))?; + let glu_combine_kernel = cpbi_module.load_function("glu_gate_combine") + .map_err(|e| MLError::ModelError(format!("glu_gate_combine load: {e}")))?; + let glu_backward_kernel = cpbi_module.load_function("glu_gate_backward") + .map_err(|e| MLError::ModelError(format!("glu_gate_backward load: {e}")))?; + info!("GpuDqnTrainer: Q-attn + selectivity + VSN + GLU kernels loaded"); + // ── Compile CQL penalty kernel (if enabled) ────────────────────── let cql_logit_grad_kernel = if config.cql_alpha > 0.0 { match compile_cql_logit_grad_kernel(&stream) { @@ -3048,6 +3178,29 @@ impl GpuDqnTrainer { d_mag_concat_buf, mag_concat_kernel, strided_accumulate_kernel, + q_attn_params, + q_attn_adam_m, + q_attn_adam_v, + q_attn_adam_step: 0, + q_coord_buf, + q_attn_kernel, + sel_params, + sel_adam_m, + sel_adam_v, + sel_grad, + sel_adam_step: 0, + sel_out_buf, + sel_fwd_kernel, + sel_bwd_kernel, + per_branch_q_gap_ema: [0.0; 4], + liquid_mod_pinned, + liquid_mod_dev_ptr, + vsn_masked_buf, + vsn_kernel, + glu_gate_pre_buf, + glu_value_buf, + glu_combine_kernel, + glu_backward_kernel, save_current_lp, save_projected, per_sample_loss_buf, @@ -4466,7 +4619,7 @@ impl GpuDqnTrainer { /// ensemble) into `grad_buf` between the two graph replays. /// /// Graph A (`graph_forward`): zero → cuBLAS forward → C51 loss → C51 grad → cuBLAS backward - /// Graph B (`graph_adam`): grad_norm → Adam → unflatten (26 d2d copies) + /// Graph B (`graph_adam`): grad_norm → Adam → unflatten (42 d2d copies) fn capture_training_graphs( &mut self, online_d: &DuelingWeightSet, @@ -6073,6 +6226,22 @@ impl GpuDqnTrainer { (0, 0), // [23] b_b3out (cfg.bottleneck_dim, cfg.market_dim), // [24] w_bn (0, 0), // [25] b_bn + (cfg.shared_h2, 16), // [26] w_vsn1_0 (Xavier) + (0, 0), // [27] w_vsn2_0 (ZERO — uniform softmax) + (cfg.shared_h2, 16), // [28] w_vsn1_1 + (0, 0), // [29] w_vsn2_1 + (cfg.shared_h2, 16), // [30] w_vsn1_2 + (0, 0), // [31] w_vsn2_2 + (cfg.shared_h2, 16), // [32] w_vsn1_3 + (0, 0), // [33] w_vsn2_3 + (0, 0), // [34] w_gate_0 (ZERO — sigmoid=0.5) + (0, 0), // [35] b_gate_0 + (0, 0), // [36] w_gate_1 + (0, 0), // [37] b_gate_1 + (0, 0), // [38] w_gate_2 + (0, 0), // [39] b_gate_2 + (0, 0), // [40] w_gate_3 + (0, 0), // [41] b_gate_3 ]; // Build flat host buffer: Xavier init for weights, zeros for biases + padding.