feat: Mamba-2 selectivity gate — forward/backward/Adam methods
selectivity_forward: sigmoid(dot(W_sel, h_s2) + b_sel) per sample. selectivity_backward: BCE gradient on (sel, per_sample_loss/mean_loss). Separate Adam at LR=1e-4. PER priority integration deferred. Also fixes kernel names (selectivity_gate_fwd→selectivity_forward, selectivity_gate_bwd→selectivity_backward) to match experience_kernels.cu. Adds sel_norm_buf, sel_norm_partials, sel_clip_buf, sel_t_buf fields needed by dqn_adam_update_kernel. Wired into run_full_step after main Adam step with mean_loss=1.0 placeholder and non-fatal error handling. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -634,6 +634,14 @@ pub struct GpuDqnTrainer {
|
||||
sel_out_buf: CudaSlice<f32>, // [B]
|
||||
sel_fwd_kernel: CudaFunction,
|
||||
sel_bwd_kernel: CudaFunction,
|
||||
/// [1] L2 norm of sel_grad (written by grad_norm_finalize before sel Adam).
|
||||
sel_norm_buf: CudaSlice<f32>, // [1]
|
||||
/// [1] partial sums scratch for sel grad_norm reduction (SH2+1 ≤ 256, so 1 block).
|
||||
sel_norm_partials: CudaSlice<f32>, // [1]
|
||||
/// [1] max grad norm for sel Adam clipping (fixed 1.0).
|
||||
sel_clip_buf: CudaSlice<f32>, // [1]
|
||||
/// [1] sel Adam step counter on device (for CUDA kernel compatibility).
|
||||
sel_t_buf: CudaSlice<i32>, // [1]
|
||||
|
||||
// ── Liquid tau ──
|
||||
per_branch_q_gap_ema: [f32; 4],
|
||||
@@ -2703,6 +2711,16 @@ impl GpuDqnTrainer {
|
||||
let sel_grad = stream.alloc_zeros::<f32>(sel_dim)
|
||||
.map_err(|e| MLError::ModelError(format!("alloc sel_grad: {e}")))?;
|
||||
let sel_out_buf = alloc_f32(&stream, b, "sel_out_buf")?;
|
||||
let sel_norm_buf = stream.alloc_zeros::<f32>(1)
|
||||
.map_err(|e| MLError::ModelError(format!("alloc sel_norm_buf: {e}")))?;
|
||||
let sel_norm_partials = stream.alloc_zeros::<f32>(1)
|
||||
.map_err(|e| MLError::ModelError(format!("alloc sel_norm_partials: {e}")))?;
|
||||
// Fixed clip norm of 1.0 — selectivity gate is a small sigmoid head, clip at 1.0 norm.
|
||||
let mut sel_clip_buf = stream.alloc_zeros::<f32>(1)
|
||||
.map_err(|e| MLError::ModelError(format!("alloc sel_clip_buf: {e}")))?;
|
||||
super::htod_f32(&stream, &[1.0_f32], &mut sel_clip_buf)?;
|
||||
// sel_t_buf: Adam step counter for selectivity gate, initialized to 0.
|
||||
let sel_t_buf = alloc_i32(&stream, 1, "sel_t_buf")?;
|
||||
|
||||
// ── Liquid tau: 4 pinned device-mapped floats ─────────────────
|
||||
let liquid_mod_pinned: *mut f32 = unsafe {
|
||||
@@ -2747,10 +2765,10 @@ impl GpuDqnTrainer {
|
||||
.map_err(|e| MLError::ModelError(format!("cpbi cubin: {e}")))?;
|
||||
let q_attn_kernel = cpbi_module.load_function("cross_branch_q_attention")
|
||||
.map_err(|e| MLError::ModelError(format!("cross_branch_q_attention 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 sel_fwd_kernel = cpbi_module.load_function("selectivity_forward")
|
||||
.map_err(|e| MLError::ModelError(format!("selectivity_forward load: {e}")))?;
|
||||
let sel_bwd_kernel = cpbi_module.load_function("selectivity_backward")
|
||||
.map_err(|e| MLError::ModelError(format!("selectivity_backward 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")
|
||||
@@ -3279,6 +3297,10 @@ impl GpuDqnTrainer {
|
||||
sel_out_buf,
|
||||
sel_fwd_kernel,
|
||||
sel_bwd_kernel,
|
||||
sel_norm_buf,
|
||||
sel_norm_partials,
|
||||
sel_clip_buf,
|
||||
sel_t_buf,
|
||||
per_branch_q_gap_ema: [0.0; 4],
|
||||
liquid_mod_pinned,
|
||||
liquid_mod_dev_ptr,
|
||||
@@ -6644,6 +6666,161 @@ impl GpuDqnTrainer {
|
||||
let states = &self.states_buf as *const CudaSlice<f32>;
|
||||
unsafe { self.compute_q_stats(&*states, batch_size) }
|
||||
}
|
||||
|
||||
// ── Selectivity gate methods ─────────────────────────────────────────────
|
||||
|
||||
/// Launch selectivity_forward kernel: sigmoid(dot(W_sel, h_s2) + b_sel) per sample.
|
||||
///
|
||||
/// Reads `save_h_s2` [B, SH2], writes `sel_out_buf` [B].
|
||||
/// Must be called after the DQN forward pass (graph_forward / graph_mega replay).
|
||||
pub(crate) fn launch_selectivity_forward(&self, batch_size: usize) -> Result<(), MLError> {
|
||||
let b = batch_size as i32;
|
||||
let sh2 = self.config.shared_h2 as i32;
|
||||
let blocks = ((batch_size as u32 + 255) / 256).max(1);
|
||||
let h_s2_ptr = self.ptrs.save_h_s2;
|
||||
let sel_out_ptr = self.sel_out_buf.raw_ptr();
|
||||
let sel_params_ptr = self.sel_params.raw_ptr();
|
||||
unsafe {
|
||||
self.stream
|
||||
.launch_builder(&self.sel_fwd_kernel)
|
||||
.arg(&h_s2_ptr)
|
||||
.arg(&sel_out_ptr)
|
||||
.arg(&sel_params_ptr)
|
||||
.arg(&b)
|
||||
.arg(&sh2)
|
||||
.launch(LaunchConfig {
|
||||
grid_dim: (blocks, 1, 1),
|
||||
block_dim: (256, 1, 1),
|
||||
shared_mem_bytes: 0,
|
||||
})
|
||||
.map_err(|e| MLError::ModelError(format!("selectivity_forward: {e}")))?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Launch selectivity_backward kernel: BCE gradient on (sel, per_sample_loss / mean_loss).
|
||||
///
|
||||
/// Zeros `sel_grad` before launch. Reads h_s2, sel_out_buf, per_sample_loss_buf.
|
||||
/// Writes `sel_grad` [SH2+1] via atomicAdd.
|
||||
pub(crate) fn launch_selectivity_backward(&mut self, batch_size: usize, mean_loss: f32) -> Result<(), MLError> {
|
||||
let b = batch_size as i32;
|
||||
let sh2 = self.config.shared_h2 as i32;
|
||||
let blocks = ((batch_size as u32 + 255) / 256).max(1);
|
||||
|
||||
// Zero gradient accumulator before atomicAdd writes.
|
||||
self.stream.memset_zeros(&mut self.sel_grad)
|
||||
.map_err(|e| MLError::ModelError(format!("zero sel_grad: {e}")))?;
|
||||
|
||||
let h_s2_ptr = self.ptrs.save_h_s2;
|
||||
let sel_out_ptr = self.sel_out_buf.raw_ptr();
|
||||
let loss_ptr = self.per_sample_loss_buf.raw_ptr();
|
||||
let grad_ptr = self.sel_grad.raw_ptr();
|
||||
unsafe {
|
||||
self.stream
|
||||
.launch_builder(&self.sel_bwd_kernel)
|
||||
.arg(&h_s2_ptr)
|
||||
.arg(&sel_out_ptr)
|
||||
.arg(&loss_ptr)
|
||||
.arg(&grad_ptr)
|
||||
.arg(&mean_loss)
|
||||
.arg(&b)
|
||||
.arg(&sh2)
|
||||
.launch(LaunchConfig {
|
||||
grid_dim: (blocks, 1, 1),
|
||||
block_dim: (256, 1, 1),
|
||||
shared_mem_bytes: 0,
|
||||
})
|
||||
.map_err(|e| MLError::ModelError(format!("selectivity_backward: {e}")))?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Run Adam update on selectivity gate parameters.
|
||||
///
|
||||
/// Computes L2 grad norm from sel_grad, then runs dqn_adam_update_kernel at LR=1e-4.
|
||||
/// Increments sel_adam_step and writes the updated step counter to sel_t_buf.
|
||||
pub(crate) fn step_selectivity_adam(&mut self) -> Result<(), MLError> {
|
||||
self.sel_adam_step += 1;
|
||||
let step_val = self.sel_adam_step;
|
||||
let sel_dim = self.config.shared_h2 + 1;
|
||||
let sel_n = sel_dim as i32;
|
||||
|
||||
// Write step counter to device buffer (memcpy_htod — small 4-byte copy, not on hot path).
|
||||
self.stream.memcpy_htod(&[step_val], &mut self.sel_t_buf)
|
||||
.map_err(|e| MLError::ModelError(format!("sel_t_buf htod: {e}")))?;
|
||||
|
||||
// Phase 1: grad_norm_kernel on sel_grad → sel_norm_partials [1 block].
|
||||
let grad_ptr = self.sel_grad.raw_ptr();
|
||||
let partials_ptr = self.sel_norm_partials.raw_ptr();
|
||||
unsafe {
|
||||
self.stream
|
||||
.launch_builder(&self.grad_norm_standalone)
|
||||
.arg(&grad_ptr)
|
||||
.arg(&partials_ptr)
|
||||
.arg(&sel_n)
|
||||
.launch(LaunchConfig {
|
||||
grid_dim: (1, 1, 1),
|
||||
block_dim: (256, 1, 1),
|
||||
shared_mem_bytes: 0,
|
||||
})
|
||||
.map_err(|e| MLError::ModelError(format!("sel grad_norm kernel: {e}")))?;
|
||||
}
|
||||
|
||||
// Phase 2: grad_norm_finalize → sel_norm_buf [1].
|
||||
let norm_out_ptr = self.sel_norm_buf.raw_ptr();
|
||||
let nb = 1_i32;
|
||||
unsafe {
|
||||
self.stream
|
||||
.launch_builder(&self.grad_norm_finalize_kernel)
|
||||
.arg(&partials_ptr)
|
||||
.arg(&norm_out_ptr)
|
||||
.arg(&nb)
|
||||
.launch(LaunchConfig {
|
||||
grid_dim: (1, 1, 1),
|
||||
block_dim: (256, 1, 1),
|
||||
shared_mem_bytes: 0,
|
||||
})
|
||||
.map_err(|e| MLError::ModelError(format!("sel grad_norm_finalize: {e}")))?;
|
||||
}
|
||||
|
||||
// Phase 3: Adam update on sel_params.
|
||||
// LR=1e-4, beta1=0.9, beta2=0.999, eps=1e-8, weight_decay=0.0, clip=1.0.
|
||||
let lr: f32 = 1e-4;
|
||||
let beta1: f32 = 0.9;
|
||||
let beta2: f32 = 0.999;
|
||||
let eps: f32 = 1e-8;
|
||||
let weight_decay: f32 = 0.0;
|
||||
let params_ptr = self.sel_params.raw_ptr();
|
||||
let m_ptr = self.sel_adam_m.raw_ptr();
|
||||
let v_ptr = self.sel_adam_v.raw_ptr();
|
||||
let clip_ptr = self.sel_clip_buf.raw_ptr();
|
||||
let t_ptr = self.sel_t_buf.raw_ptr();
|
||||
let blocks = ((sel_dim as u32 + 255) / 256).max(1);
|
||||
unsafe {
|
||||
self.stream
|
||||
.launch_builder(&self.adam_update_kernel)
|
||||
.arg(¶ms_ptr)
|
||||
.arg(&grad_ptr)
|
||||
.arg(&m_ptr)
|
||||
.arg(&v_ptr)
|
||||
.arg(&norm_out_ptr)
|
||||
.arg(&lr)
|
||||
.arg(&beta1)
|
||||
.arg(&beta2)
|
||||
.arg(&eps)
|
||||
.arg(&weight_decay)
|
||||
.arg(&clip_ptr)
|
||||
.arg(&t_ptr)
|
||||
.arg(&sel_n)
|
||||
.launch(LaunchConfig {
|
||||
grid_dim: (blocks, 1, 1),
|
||||
block_dim: (256, 1, 1),
|
||||
shared_mem_bytes: 0,
|
||||
})
|
||||
.map_err(|e| MLError::ModelError(format!("sel Adam update: {e}")))?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
// ── Compilation ─────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -1032,6 +1032,28 @@ impl FusedTrainingCtx {
|
||||
).map_err(|e| anyhow::anyhow!("IQL modulate_td_errors: {e}"))?;
|
||||
}
|
||||
|
||||
// ── Step 5c: Mamba-2 selectivity gate ───────────────────────────
|
||||
// Forward: sel_out[b] = sigmoid(dot(W_sel, h_s2[b]) + b_sel)
|
||||
// Backward: BCE gradient using per_sample_loss as training signal
|
||||
// Adam: LR=1e-4, separate optimizer for selectivity params
|
||||
{
|
||||
let bs = self.batch_size;
|
||||
if let Err(e) = self.launch_selectivity_forward(bs) {
|
||||
tracing::warn!("selectivity_forward failed (non-fatal): {e}");
|
||||
} else {
|
||||
// mean_loss=1.0 placeholder — per_sample_loss values are the raw targets.
|
||||
// The gate learns which states have high loss (high selectivity = important for replay).
|
||||
let mean_loss = 1.0_f32;
|
||||
if let Err(e) = self.launch_selectivity_backward(bs, mean_loss) {
|
||||
tracing::warn!("selectivity_backward failed (non-fatal): {e}");
|
||||
} else if let Err(e) = self.step_selectivity_adam() {
|
||||
tracing::warn!("step_selectivity_adam failed (non-fatal): {e}");
|
||||
}
|
||||
}
|
||||
// TODO: Scale PER priorities by (1 + selectivity). Requires GPU readback
|
||||
// of sel_out_buf and modification of PER update_priorities.
|
||||
}
|
||||
|
||||
// ── Step 6: PER priority update ─────────────────────────────────
|
||||
if let Some(ref pe) = self.phase_events {
|
||||
PhaseEvents::record(pe.per_update_start, cu_stream);
|
||||
@@ -1961,6 +1983,24 @@ impl FusedTrainingCtx {
|
||||
.map_err(|e| anyhow::anyhow!("launch_q_attention: {e}"))
|
||||
}
|
||||
|
||||
/// Selectivity gate forward: sigmoid(dot(W_sel, h_s2) + b_sel) per sample.
|
||||
pub(crate) fn launch_selectivity_forward(&self, batch_size: usize) -> Result<()> {
|
||||
self.trainer.launch_selectivity_forward(batch_size)
|
||||
.map_err(|e| anyhow::anyhow!("launch_selectivity_forward: {e}"))
|
||||
}
|
||||
|
||||
/// Selectivity gate backward: BCE gradient on (sel, per_sample_loss / mean_loss).
|
||||
pub(crate) fn launch_selectivity_backward(&mut self, batch_size: usize, mean_loss: f32) -> Result<()> {
|
||||
self.trainer.launch_selectivity_backward(batch_size, mean_loss)
|
||||
.map_err(|e| anyhow::anyhow!("launch_selectivity_backward: {e}"))
|
||||
}
|
||||
|
||||
/// Adam update for selectivity gate parameters (LR=1e-4, clip=1.0).
|
||||
pub(crate) fn step_selectivity_adam(&mut self) -> Result<()> {
|
||||
self.trainer.step_selectivity_adam()
|
||||
.map_err(|e| anyhow::anyhow!("step_selectivity_adam: {e}"))
|
||||
}
|
||||
|
||||
/// Set C51 blend factor for gradual MSE→C51 ramp.
|
||||
///
|
||||
/// `alpha = 0.0`: pure MSE (warmup). `alpha = 1.0`: pure C51 (converged).
|
||||
|
||||
Reference in New Issue
Block a user