Mamba M2: analytic backward (gradient-checked)
Hand-written VJP of the selective-scan forward (Phase-3 spec §5/§7): MambaBlock::backward(x, d_out) -> per-parameter gradients keyed by persistence name, summed over the batch. Differentiates the scan analytically (reverse-time recurrence over the cached h trajectory) rather than through the immature rtx-tensor autograd tape. Covers every parameter: in_proj, conv1d_weight, conv1d_bias, A_log (via A=-exp(A_log) ⇒ dA_log = dA·A), x_proj, dt_proj, dt_bias, D, out_proj. Adds stable sigmoid_f32 / silu_grad_f32 helpers. New test analytic_gradients_match_finite_differences: on a small well-conditioned instance, ≥30 sampled grad elements across all 9 params match central finite differences within (5e-3 + 5e-2·|fd|). All 5 selective-scan tests green. Co-Authored-By: Claude Opus 4.8 (1M context) <[email protected]>
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
00cd527ed4
commit
199130fa7d
@@ -624,6 +624,284 @@ impl MambaBlock {
|
||||
aux_info: None,
|
||||
})
|
||||
}
|
||||
|
||||
/// Analytic backward pass: given `d_out = ∂L/∂out` (same shape as
|
||||
/// the forward output, `[b,l,d_model]`), return `∂L/∂θ` for every
|
||||
/// parameter, keyed by its persistence name (`in_proj`,
|
||||
/// `conv1d_weight`, `conv1d_bias`, `A_log`, `x_proj`, `dt_proj`,
|
||||
/// `dt_bias`, `D`, `out_proj`). Grads are summed over the batch.
|
||||
///
|
||||
/// This is the hand-written VJP of [`Self::forward`] (Phase-3 spec
|
||||
/// §5/§7) — we differentiate the selective scan analytically rather
|
||||
/// than through the autograd tape. Verified by finite differences
|
||||
/// (see `tests/real_selective_scan.rs`).
|
||||
pub fn backward(&self, x: &Tensor, d_out: &Tensor) -> Result<HashMap<String, Tensor>> {
|
||||
let dims = x.shape().dims().to_vec();
|
||||
let (b, l, d_model) = (dims[0], dims[1], dims[2]);
|
||||
let d = self.config.get_d_inner();
|
||||
let n = self.config.d_state;
|
||||
let dt_rank = self.config.get_dt_rank();
|
||||
let kc = self.config.d_conv;
|
||||
let dbc = dt_rank + 2 * n;
|
||||
|
||||
let xv = x.to_vec()?;
|
||||
let dov = d_out.to_vec()?;
|
||||
let in_proj = self.in_proj.to_vec()?;
|
||||
let conv_w = self.conv1d_weight.to_vec()?;
|
||||
let conv_b = match &self.conv1d_bias {
|
||||
Some(t) => Some(t.to_vec()?),
|
||||
None => None,
|
||||
};
|
||||
let a_log = self.A_log.to_vec()?;
|
||||
let x_proj = self.x_proj.to_vec()?;
|
||||
let dt_proj = self.dt_proj.to_vec()?;
|
||||
let dt_bias = self.dt_bias.to_vec()?;
|
||||
let d_skip = self.d_skip.to_vec()?;
|
||||
let out_proj = self.out_proj.to_vec()?;
|
||||
|
||||
// A = -exp(A_log) is batch-independent.
|
||||
let mut a_mat = vec![0.0f32; d * n];
|
||||
for i in 0..d * n {
|
||||
a_mat[i] = -a_log[i].exp();
|
||||
}
|
||||
|
||||
// Gradient accumulators (param layouts).
|
||||
let mut g_in = vec![0.0f32; d_model * 2 * d];
|
||||
let mut g_cw = vec![0.0f32; d * kc];
|
||||
let mut g_cb = vec![0.0f32; d];
|
||||
let mut g_a = vec![0.0f32; d * n]; // ∂L/∂A, converted to ∂L/∂A_log at the end
|
||||
let mut g_xp = vec![0.0f32; d * dbc];
|
||||
let mut g_dtp = vec![0.0f32; dt_rank * d];
|
||||
let mut g_dtb = vec![0.0f32; d];
|
||||
let mut g_dsk = vec![0.0f32; d];
|
||||
let mut g_op = vec![0.0f32; d * d_model];
|
||||
|
||||
for bi in 0..b {
|
||||
// ---- forward recompute, caching every intermediate ----
|
||||
let mut x_in = vec![0.0f32; l * d];
|
||||
let mut z = vec![0.0f32; l * d];
|
||||
for li in 0..l {
|
||||
let xr = &xv[(bi * l + li) * d_model..(bi * l + li) * d_model + d_model];
|
||||
for j in 0..d {
|
||||
let (mut sx, mut sz) = (0.0f32, 0.0f32);
|
||||
for m in 0..d_model {
|
||||
sx += xr[m] * in_proj[m * (2 * d) + j];
|
||||
sz += xr[m] * in_proj[m * (2 * d) + d + j];
|
||||
}
|
||||
x_in[li * d + j] = sx;
|
||||
z[li * d + j] = sz;
|
||||
}
|
||||
}
|
||||
let mut conv_pre = vec![0.0f32; l * d];
|
||||
let mut u = vec![0.0f32; l * d];
|
||||
for li in 0..l {
|
||||
for j in 0..d {
|
||||
let mut acc = conv_b.as_ref().map_or(0.0f32, |cb| cb[j]);
|
||||
for kk in 0..kc {
|
||||
let src = li as isize - (kc as isize - 1) + kk as isize;
|
||||
if src >= 0 {
|
||||
acc += x_in[(src as usize) * d + j] * conv_w[j * kc + kk];
|
||||
}
|
||||
}
|
||||
conv_pre[li * d + j] = acc;
|
||||
u[li * d + j] = silu_f32(acc);
|
||||
}
|
||||
}
|
||||
let mut xdbl = vec![0.0f32; l * dbc];
|
||||
let mut bmat = vec![0.0f32; l * n];
|
||||
let mut cmat = vec![0.0f32; l * n];
|
||||
let mut dt_pre = vec![0.0f32; l * d];
|
||||
let mut delta = vec![0.0f32; l * d];
|
||||
for li in 0..l {
|
||||
for q in 0..dbc {
|
||||
let mut s = 0.0f32;
|
||||
for j in 0..d {
|
||||
s += u[li * d + j] * x_proj[j * dbc + q];
|
||||
}
|
||||
xdbl[li * dbc + q] = s;
|
||||
}
|
||||
for nn in 0..n {
|
||||
bmat[li * n + nn] = xdbl[li * dbc + dt_rank + nn];
|
||||
cmat[li * n + nn] = xdbl[li * dbc + dt_rank + n + nn];
|
||||
}
|
||||
for j in 0..d {
|
||||
let mut s = dt_bias[j];
|
||||
for r in 0..dt_rank {
|
||||
s += xdbl[li * dbc + r] * dt_proj[r * d + j];
|
||||
}
|
||||
dt_pre[li * d + j] = s;
|
||||
delta[li * d + j] = softplus_f32(s);
|
||||
}
|
||||
}
|
||||
// scan, caching da and the h trajectory (h after update at each li).
|
||||
let mut da = vec![0.0f32; l * d * n];
|
||||
let mut h_tr = vec![0.0f32; l * d * n];
|
||||
let mut y = vec![0.0f32; l * d];
|
||||
let mut h = vec![0.0f32; d * n];
|
||||
for li in 0..l {
|
||||
for j in 0..d {
|
||||
let dj = delta[li * d + j];
|
||||
let uj = u[li * d + j];
|
||||
let mut yy = d_skip[j] * uj;
|
||||
for nn in 0..n {
|
||||
let dav = (dj * a_mat[j * n + nn]).exp();
|
||||
let dbu = dj * bmat[li * n + nn] * uj;
|
||||
let hv = dav * h[j * n + nn] + dbu;
|
||||
h[j * n + nn] = hv;
|
||||
da[(li * d + j) * n + nn] = dav;
|
||||
h_tr[(li * d + j) * n + nn] = hv;
|
||||
yy += cmat[li * n + nn] * hv;
|
||||
}
|
||||
y[li * d + j] = yy;
|
||||
}
|
||||
}
|
||||
|
||||
// ---- reverse pass ----
|
||||
let mut d_u = vec![0.0f32; l * d];
|
||||
let mut d_z = vec![0.0f32; l * d];
|
||||
let mut d_delta = vec![0.0f32; l * d];
|
||||
let mut d_bmat = vec![0.0f32; l * n];
|
||||
let mut d_cmat = vec![0.0f32; l * n];
|
||||
let mut dy = vec![0.0f32; l * d];
|
||||
|
||||
// (9) out_proj + (8) gate yg = y·silu(z)
|
||||
for li in 0..l {
|
||||
let obase = (bi * l + li) * d_model;
|
||||
for j in 0..d {
|
||||
let sz = silu_f32(z[li * d + j]);
|
||||
let yg = y[li * d + j] * sz;
|
||||
let mut dyg = 0.0f32;
|
||||
for m in 0..d_model {
|
||||
let go = dov[obase + m];
|
||||
g_op[j * d_model + m] += yg * go;
|
||||
dyg += go * out_proj[j * d_model + m];
|
||||
}
|
||||
dy[li * d + j] = dyg * sz;
|
||||
d_z[li * d + j] = dyg * y[li * d + j] * silu_grad_f32(z[li * d + j]);
|
||||
}
|
||||
}
|
||||
// (7) y = D·u + Σ_nn C·h → d_skip, d_u(D), d_cmat, dh_from_y
|
||||
let mut dh_from_y = vec![0.0f32; l * d * n];
|
||||
for li in 0..l {
|
||||
for j in 0..d {
|
||||
let dyj = dy[li * d + j];
|
||||
g_dsk[j] += dyj * u[li * d + j];
|
||||
d_u[li * d + j] += dyj * d_skip[j];
|
||||
for nn in 0..n {
|
||||
d_cmat[li * n + nn] += dyj * h_tr[(li * d + j) * n + nn];
|
||||
dh_from_y[(li * d + j) * n + nn] = dyj * cmat[li * n + nn];
|
||||
}
|
||||
}
|
||||
}
|
||||
// (6) scan reverse: dh[li] = dh_from_y[li] + carry from li+1.
|
||||
let mut dh_next = vec![0.0f32; d * n];
|
||||
for li in (0..l).rev() {
|
||||
let mut dh_prev = vec![0.0f32; d * n];
|
||||
for j in 0..d {
|
||||
let dj = delta[li * d + j];
|
||||
let uj = u[li * d + j];
|
||||
let mut dd_from_da = 0.0f32;
|
||||
let mut dd_from_dbu = 0.0f32;
|
||||
let mut du_dbu = 0.0f32;
|
||||
for nn in 0..n {
|
||||
let idx = (li * d + j) * n + nn;
|
||||
let dh = dh_from_y[idx] + dh_next[j * n + nn];
|
||||
let h_prev = if li > 0 { h_tr[((li - 1) * d + j) * n + nn] } else { 0.0 };
|
||||
let dav = da[idx];
|
||||
// da[li] = exp(delta·A): d(delta·A) = (dh·h_prev)·da
|
||||
let d_deltaA = (dh * h_prev) * dav;
|
||||
dd_from_da += d_deltaA * a_mat[j * n + nn];
|
||||
g_a[j * n + nn] += d_deltaA * dj;
|
||||
// dbu = delta·B·u
|
||||
let ddbu = dh; // ∂h/∂dbu = 1
|
||||
dd_from_dbu += ddbu * bmat[li * n + nn] * uj;
|
||||
d_bmat[li * n + nn] += ddbu * dj * uj;
|
||||
du_dbu += ddbu * dj * bmat[li * n + nn];
|
||||
// carry to h[li-1]: dh·da
|
||||
dh_prev[j * n + nn] = dh * dav;
|
||||
}
|
||||
d_delta[li * d + j] += dd_from_da + dd_from_dbu;
|
||||
d_u[li * d + j] += du_dbu;
|
||||
}
|
||||
dh_next = dh_prev;
|
||||
}
|
||||
// (5) delta = softplus(dt_pre); dt_pre = dt_bias + Σ_r xdbl[r]·dt_proj[r,j]
|
||||
let mut d_xdbl = vec![0.0f32; l * dbc];
|
||||
for li in 0..l {
|
||||
for j in 0..d {
|
||||
let ddt = d_delta[li * d + j] * sigmoid_f32(dt_pre[li * d + j]);
|
||||
g_dtb[j] += ddt;
|
||||
for r in 0..dt_rank {
|
||||
g_dtp[r * d + j] += ddt * xdbl[li * dbc + r];
|
||||
d_xdbl[li * dbc + r] += ddt * dt_proj[r * d + j];
|
||||
}
|
||||
}
|
||||
// (4b) B,C slices of xdbl
|
||||
for nn in 0..n {
|
||||
d_xdbl[li * dbc + dt_rank + nn] += d_bmat[li * n + nn];
|
||||
d_xdbl[li * dbc + dt_rank + n + nn] += d_cmat[li * n + nn];
|
||||
}
|
||||
}
|
||||
// (4a) xdbl = u @ x_proj
|
||||
for li in 0..l {
|
||||
for q in 0..dbc {
|
||||
let dq = d_xdbl[li * dbc + q];
|
||||
for j in 0..d {
|
||||
g_xp[j * dbc + q] += u[li * d + j] * dq;
|
||||
d_u[li * d + j] += dq * x_proj[j * dbc + q];
|
||||
}
|
||||
}
|
||||
}
|
||||
// (3) u = silu(conv_pre) → (2) conv → (1) in_proj
|
||||
let mut d_x_in = vec![0.0f32; l * d];
|
||||
for li in 0..l {
|
||||
for j in 0..d {
|
||||
let dcp = d_u[li * d + j] * silu_grad_f32(conv_pre[li * d + j]);
|
||||
g_cb[j] += dcp;
|
||||
for kk in 0..kc {
|
||||
let src = li as isize - (kc as isize - 1) + kk as isize;
|
||||
if src >= 0 {
|
||||
let s = src as usize;
|
||||
g_cw[j * kc + kk] += dcp * x_in[s * d + j];
|
||||
d_x_in[s * d + j] += dcp * conv_w[j * kc + kk];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
for li in 0..l {
|
||||
let xr = &xv[(bi * l + li) * d_model..(bi * l + li) * d_model + d_model];
|
||||
for j in 0..d {
|
||||
let dxi = d_x_in[li * d + j];
|
||||
let dzj = d_z[li * d + j];
|
||||
for m in 0..d_model {
|
||||
g_in[m * (2 * d) + j] += xr[m] * dxi;
|
||||
g_in[m * (2 * d) + d + j] += xr[m] * dzj;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// A = -exp(A_log) ⇒ ∂L/∂A_log = ∂L/∂A · A.
|
||||
let mut g_alog = vec![0.0f32; d * n];
|
||||
for i in 0..d * n {
|
||||
g_alog[i] = g_a[i] * a_mat[i];
|
||||
}
|
||||
|
||||
let dev = &self.device;
|
||||
let mut grads: HashMap<String, Tensor> = HashMap::new();
|
||||
grads.insert("in_proj".into(), Tensor::from_vec(g_in, &[d_model, 2 * d], dev)?);
|
||||
grads.insert("conv1d_weight".into(), Tensor::from_vec(g_cw, &[d, 1, kc], dev)?);
|
||||
if self.conv1d_bias.is_some() {
|
||||
grads.insert("conv1d_bias".into(), Tensor::from_vec(g_cb, &[d], dev)?);
|
||||
}
|
||||
grads.insert("A_log".into(), Tensor::from_vec(g_alog, &[d, n], dev)?);
|
||||
grads.insert("x_proj".into(), Tensor::from_vec(g_xp, &[d, dbc], dev)?);
|
||||
grads.insert("dt_proj".into(), Tensor::from_vec(g_dtp, &[dt_rank, d], dev)?);
|
||||
grads.insert("dt_bias".into(), Tensor::from_vec(g_dtb, &[d], dev)?);
|
||||
grads.insert("D".into(), Tensor::from_vec(g_dsk, &[d], dev)?);
|
||||
grads.insert("out_proj".into(), Tensor::from_vec(g_op, &[d, d_model], dev)?);
|
||||
Ok(grads)
|
||||
}
|
||||
}
|
||||
|
||||
/// Numerically-stable SiLU (a.k.a. swish): `x · σ(x)`.
|
||||
@@ -644,6 +922,24 @@ fn softplus_f32(x: f32) -> f32 {
|
||||
x.max(0.0) + (1.0 + (-x.abs()).exp()).ln()
|
||||
}
|
||||
|
||||
/// Numerically-stable logistic sigmoid `σ(x)`.
|
||||
#[inline]
|
||||
fn sigmoid_f32(x: f32) -> f32 {
|
||||
if x >= 0.0 {
|
||||
1.0 / (1.0 + (-x).exp())
|
||||
} else {
|
||||
let e = x.exp();
|
||||
e / (1.0 + e)
|
||||
}
|
||||
}
|
||||
|
||||
/// Derivative of SiLU: `σ(x) + x·σ(x)·(1−σ(x))`.
|
||||
#[inline]
|
||||
fn silu_grad_f32(x: f32) -> f32 {
|
||||
let s = sigmoid_f32(x);
|
||||
s + x * s * (1.0 - s)
|
||||
}
|
||||
|
||||
impl Layer for MambaBlock {
|
||||
fn forward(&self, input: &Tensor) -> Result<Tensor> {
|
||||
let output = self.forward(input)?;
|
||||
|
||||
Reference in New Issue
Block a user