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,
|
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)`.
|
/// 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()
|
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 {
|
impl Layer for MambaBlock {
|
||||||
fn forward(&self, input: &Tensor) -> Result<Tensor> {
|
fn forward(&self, input: &Tensor) -> Result<Tensor> {
|
||||||
let output = self.forward(input)?;
|
let output = self.forward(input)?;
|
||||||
|
|||||||
@@ -112,6 +112,83 @@ fn seeded_forward_is_deterministic() {
|
|||||||
assert_eq!(fwd(&a, &x), fwd(&b, &x), "same (config, seed) must be bit-exact");
|
assert_eq!(fwd(&a, &x), fwd(&b, &x), "same (config, seed) must be bit-exact");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Build a seeded block whose every weight is scaled by `scale`, so the
|
||||||
|
/// SSM operates in a well-conditioned regime (small `A_log` ⇒ `A≈-1` ⇒
|
||||||
|
/// the scan neither saturates nor underflows), giving a meaningful
|
||||||
|
/// finite-difference signal.
|
||||||
|
fn scaled_block(cfg: &MambaConfig, dev: &Device, seed: u64, scale: f32) -> MambaBlock {
|
||||||
|
let base = MambaBlock::new_seeded(cfg.clone(), dev, seed).expect("base");
|
||||||
|
let mut map = HashMap::new();
|
||||||
|
for (name, t) in base.persistence_tensors() {
|
||||||
|
let mut v = t.to_vec().expect("to_vec");
|
||||||
|
for x in v.iter_mut() {
|
||||||
|
*x *= scale;
|
||||||
|
}
|
||||||
|
map.insert(name.to_string(), Tensor::from_vec(v, t.shape().dims(), dev).expect("from_vec"));
|
||||||
|
}
|
||||||
|
MambaBlock::from_persistence_tensors(cfg.clone(), dev, map).expect("rebuild")
|
||||||
|
}
|
||||||
|
|
||||||
|
fn loss(block: &MambaBlock, x: &Tensor, dov: &[f32]) -> f32 {
|
||||||
|
let out = block.forward(x).expect("fwd").output.to_vec().expect("to_vec");
|
||||||
|
out.iter().zip(dov).map(|(o, g)| o * g).sum::<f32>()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn perturbed_one(block: &MambaBlock, name: &str, idx: usize, eps: f32, cfg: &MambaConfig, dev: &Device) -> MambaBlock {
|
||||||
|
let mut map = HashMap::new();
|
||||||
|
for (nm, t) in block.persistence_tensors() {
|
||||||
|
let mut v = t.to_vec().expect("to_vec");
|
||||||
|
if nm == name {
|
||||||
|
v[idx] += eps;
|
||||||
|
}
|
||||||
|
map.insert(nm.to_string(), Tensor::from_vec(v, t.shape().dims(), dev).expect("from_vec"));
|
||||||
|
}
|
||||||
|
MambaBlock::from_persistence_tensors(cfg.clone(), dev, map).expect("rebuild")
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn analytic_gradients_match_finite_differences() {
|
||||||
|
let dev = Device::cpu();
|
||||||
|
// Small, well-conditioned instance for a clean f32 FD signal.
|
||||||
|
let cfg = MambaConfig::new(4, 4, 3);
|
||||||
|
let block = scaled_block(&cfg, &dev, 21, 0.2);
|
||||||
|
let ll = 4usize;
|
||||||
|
let dmodel = 4usize;
|
||||||
|
|
||||||
|
// Input and a fixed upstream gradient d_out (so L = Σ d_out·out).
|
||||||
|
let xvec: Vec<f32> = (0..ll * dmodel).map(|i| 0.5 * ((i as f32 * 0.37).sin())).collect();
|
||||||
|
let x = Tensor::from_vec(xvec, &[1, ll, dmodel], &dev).expect("x");
|
||||||
|
let dov: Vec<f32> = (0..ll * dmodel).map(|i| ((i as f32 * 0.91).cos())).collect();
|
||||||
|
|
||||||
|
let grads = block.backward(&x, &Tensor::from_vec(dov.clone(), &[1, ll, dmodel], &dev).expect("dov")).expect("backward");
|
||||||
|
|
||||||
|
let eps = 1e-2f32;
|
||||||
|
let mut checked = 0;
|
||||||
|
for name in ["in_proj", "conv1d_weight", "conv1d_bias", "A_log", "x_proj", "dt_proj", "dt_bias", "D", "out_proj"] {
|
||||||
|
let g = grads.get(name).unwrap_or_else(|| panic!("missing grad {name}")).to_vec().expect("g");
|
||||||
|
let len = g.len();
|
||||||
|
// Sample up to 4 spread-out indices per parameter.
|
||||||
|
let idxs: Vec<usize> = if len <= 4 {
|
||||||
|
(0..len).collect()
|
||||||
|
} else {
|
||||||
|
vec![0, len / 4, len / 2, (3 * len) / 4]
|
||||||
|
};
|
||||||
|
for &idx in &idxs {
|
||||||
|
let bp = perturbed_one(&block, name, idx, eps, &cfg, &dev);
|
||||||
|
let bm = perturbed_one(&block, name, idx, -eps, &cfg, &dev);
|
||||||
|
let fd = (loss(&bp, &x, &dov) - loss(&bm, &x, &dov)) / (2.0 * eps);
|
||||||
|
let an = g[idx];
|
||||||
|
let tol = 5e-3 + 5e-2 * fd.abs();
|
||||||
|
assert!(
|
||||||
|
(an - fd).abs() <= tol,
|
||||||
|
"grad mismatch for `{name}`[{idx}]: analytic={an:.6} finite-diff={fd:.6} (tol {tol:.4})"
|
||||||
|
);
|
||||||
|
checked += 1;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
assert!(checked >= 30, "expected to check ≥30 grad elements, got {checked}");
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn output_is_finite_and_non_constant() {
|
fn output_is_finite_and_non_constant() {
|
||||||
let dev = Device::cpu();
|
let dev = Device::cpu();
|
||||||
|
|||||||
Reference in New Issue
Block a user