From 2aae3ceec7babb4f0721a16e28e6c1aa33c622aa Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Fri, 17 Apr 2026 00:42:14 +0200 Subject: [PATCH] =?UTF-8?q?feat(isv):=20regime-conditioned=20Mamba2=20deca?= =?UTF-8?q?y=20=E2=80=94=20forget=20old=20patterns=20during=20transitions?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit mamba2_temporal_scan forget gate multiplied by regime_stability (ISV[11]). Stable regime (stability=1.0): full history preserved. Transition (stability=0.3): gate capped at 0.3 → rapid forgetting. Model adapts to new regime within 2-3 bars. Backward matches forward. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../src/cuda_pipeline/experience_kernels.cu | 24 +++++++++++++++---- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 2 ++ .../cuda_pipeline/mamba2_temporal_kernel.cu | 9 ++++++- 3 files changed, 29 insertions(+), 6 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/experience_kernels.cu b/crates/ml/src/cuda_pipeline/experience_kernels.cu index 44bec964a..946cb1bb5 100644 --- a/crates/ml/src/cuda_pipeline/experience_kernels.cu +++ b/crates/ml/src/cuda_pipeline/experience_kernels.cu @@ -4286,7 +4286,8 @@ extern "C" __global__ void mamba2_scan_backward( int B, int K, int sh2, - int state_d + int state_d, + const float* __restrict__ isv_signals /* [12] pinned. NULL = no regime modulation. */ ) { int i = blockIdx.x * blockDim.x + threadIdx.x; if (i >= B) return; @@ -4311,6 +4312,10 @@ extern "C" __global__ void mamba2_scan_backward( } a_raw[t][s] = a_val; float gate = 1.0f / (1.0f + expf(-a_val)); + /* Regime-conditioned decay — must match forward kernel exactly */ + if (isv_signals != NULL) { + gate *= isv_signals[11]; /* regime_stability */ + } x_states[t+1][s] = gate * x_states[t][s] + b_val; } } @@ -4343,8 +4348,16 @@ extern "C" __global__ void mamba2_scan_backward( float gate = 1.0f / (1.0f + expf(-a_val)); float sigmoid_deriv = gate * (1.0f - gate); - /* d_A_t = d_x_{t+1} * x_t * sigmoid'(a_raw_t) */ - float d_gate = d_x[s] * x_states[t][s] * sigmoid_deriv; + /* Regime-conditioned decay — must match forward kernel exactly. + * effective_gate = sigmoid(a_val) * stability. + * d(effective_gate)/d(a_val) = sigmoid'(a_val) * stability. */ + float stability = 1.0f; + if (isv_signals != NULL) { + stability = isv_signals[11]; /* regime_stability */ + } + + /* d_A_t = d_x_{t+1} * x_t * sigmoid'(a_raw_t) * stability */ + float d_gate = d_x[s] * x_states[t][s] * sigmoid_deriv * stability; /* d_W_A: d_gate is scalar per (sample, state_dim), h_t is [SH2] */ for (int j = 0; j < sh2; j++) { @@ -4356,8 +4369,9 @@ extern "C" __global__ void mamba2_scan_backward( atomicAdd(&d_w_b[(long long)j * state_d + s], d_x[s] * h_t[j]); } - /* Propagate d_x backward through gate: d_x_t = d_x_{t+1} * A_t */ - d_x[s] = d_x[s] * gate; + /* Propagate d_x backward through effective gate: + * d_x_t = d_x_{t+1} * sigmoid(a_val) * stability */ + d_x[s] = d_x[s] * gate * stability; } } } diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 8beb874a5..40ed64450 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -2653,6 +2653,7 @@ impl GpuDqnTrainer { .arg(&(MAMBA2_HISTORY_K as i32)) .arg(&(sh2 as i32)) .arg(&(MAMBA2_STATE_DIM as i32)) + .arg(&self.isv_signals_dev_ptr) // regime-conditioned decay (ISV[11]) .launch(LaunchConfig::for_num_elems(batch_size as u32)) .map_err(|e| MLError::ModelError(format!("mamba2_temporal_scan: {e}")))?; } @@ -2727,6 +2728,7 @@ impl GpuDqnTrainer { .arg(&(MAMBA2_HISTORY_K as i32)) .arg(&(sh2 as i32)) .arg(&(MAMBA2_STATE_DIM as i32)) + .arg(&self.isv_signals_dev_ptr) // regime-conditioned decay (ISV[11]) .launch(LaunchConfig::for_num_elems(batch_size as u32)) .map_err(|e| MLError::ModelError(format!("mamba2_scan_backward: {e}")))?; } diff --git a/crates/ml/src/cuda_pipeline/mamba2_temporal_kernel.cu b/crates/ml/src/cuda_pipeline/mamba2_temporal_kernel.cu index f95bebb44..080efd5bf 100644 --- a/crates/ml/src/cuda_pipeline/mamba2_temporal_kernel.cu +++ b/crates/ml/src/cuda_pipeline/mamba2_temporal_kernel.cu @@ -20,7 +20,8 @@ extern "C" __global__ void mamba2_temporal_scan( int N, int K, int sh2, - int state_d + int state_d, + const float* __restrict__ isv_signals /* [12] pinned. NULL = no regime modulation. */ ) { int i = blockIdx.x * blockDim.x + threadIdx.x; if (i >= N) return; @@ -44,6 +45,12 @@ extern "C" __global__ void mamba2_temporal_scan( b_val += w_b[(long long)j * state_d + s] * h_t[j]; } float gate = 1.0f / (1.0f + expf(-a_val)); + /* Regime-conditioned decay: during transitions (low stability), + * decay faster -> forget old-regime patterns. Stable regime -> full history. */ + if (isv_signals != NULL) { + float stability = isv_signals[11]; /* regime_stability in [0, 1] */ + gate *= stability; /* gate in [0, stability] instead of [0, 1] */ + } x[s] = gate * x[s] + b_val; }