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:
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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}")))?;
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user