feat(isv): regime-conditioned Mamba2 decay — forget old patterns during transitions

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) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-17 00:42:14 +02:00
parent 1acb23cb24
commit 2aae3ceec7
3 changed files with 29 additions and 6 deletions

View File

@@ -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;
}
}
}

View File

@@ -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}")))?;
}

View File

@@ -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;
}