pub fn cross_entropy_backward(
logits: &[f32],
targets: &[u32],
seq_len: usize,
vocab_size: usize,
completion_start: usize,
) -> Vec<f32> {
assert_eq!(logits.len(), seq_len * vocab_size);
assert_eq!(targets.len(), seq_len);
let n_comp = (seq_len - completion_start) as f32;
let mut dlogits = vec![0.0f32; seq_len * vocab_size];
for t in completion_start..seq_len {
let logit_row = &logits[t * vocab_size..(t + 1) * vocab_size];
let target = targets[t] as usize;
let max = logit_row.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut sum_exp = 0.0f64;
for &l in logit_row {
sum_exp += ((l - max) as f64).exp();
}
let inv_sum = (1.0 / sum_exp) as f32;
let drow = &mut dlogits[t * vocab_size..(t + 1) * vocab_size];
for (v, d) in drow.iter_mut().enumerate() {
let p = ((logit_row[v] - max).exp()) * inv_sum;
let indicator = if v == target { 1.0f32 } else { 0.0 };
*d = (p - indicator) / n_comp;
}
}
dlogits
}
pub fn linear_vjp(w: &[f32], g: &[f32], d_in: usize, d_out: usize) -> Vec<f32> {
assert_eq!(w.len(), d_out * d_in);
assert_eq!(g.len(), d_out);
let mut dx = vec![0.0f32; d_in];
crate::forward::cpu::matmul_into(g, w, &mut dx, 1, d_out, d_in);
dx
}
pub fn lora_vjp(
g: &[f32],
x: &[f32],
h: &[f32],
a: &[f32],
b: &[f32],
rank: usize,
d_in: usize,
d_out: usize,
scale: f32,
) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
assert_eq!(a.len(), rank * d_in);
assert_eq!(b.len(), d_out * rank);
assert_eq!(g.len(), d_out);
assert_eq!(x.len(), d_in);
assert_eq!(h.len(), rank);
let mut grad_b = vec![0.0f32; d_out * rank];
for i in 0..d_out {
for r in 0..rank {
grad_b[i * rank + r] = scale * g[i] * h[r];
}
}
let mut bt_g = vec![0.0f32; rank];
crate::forward::cpu::matmul_into(g, b, &mut bt_g, 1, d_out, rank);
let mut grad_a = vec![0.0f32; rank * d_in];
for r in 0..rank {
for j in 0..d_in {
grad_a[r * d_in + j] = scale * bt_g[r] * x[j];
}
}
for v in bt_g.iter_mut() {
*v *= scale;
}
let mut dx = vec![0.0f32; d_in];
crate::forward::cpu::matmul_into(&bt_g, a, &mut dx, 1, rank, d_in);
(grad_b, grad_a, dx)
}
pub fn rmsnorm_backward(x: &[f32], w: &[f32], inv_rms: f32, g: &[f32]) -> Vec<f32> {
let d = x.len();
assert_eq!(w.len(), d);
assert_eq!(g.len(), d);
let sum_xwg: f32 = (0..d).map(|j| x[j] * w[j] * g[j]).sum();
let inv_rms3_over_d = inv_rms * inv_rms * inv_rms / d as f32;
let mut dx = vec![0.0f32; d];
for i in 0..d {
dx[i] = w[i] * g[i] * inv_rms - x[i] * inv_rms3_over_d * sum_xwg;
}
dx
}
pub fn rope_backward(g: &[f32], cos_vals: &[f32], sin_vals: &[f32], rope_dim: usize) -> Vec<f32> {
let head_dim = g.len();
assert!(rope_dim <= head_dim);
assert_eq!(rope_dim % 2, 0);
let half = rope_dim / 2;
assert_eq!(cos_vals.len(), half);
assert_eq!(sin_vals.len(), half);
let mut dx = g.to_vec();
for i in 0..half {
let c = cos_vals[i];
let s = sin_vals[i];
let g0 = g[i];
let g1 = g[half + i];
dx[i] = g0 * c + g1 * s;
dx[half + i] = -g0 * s + g1 * c;
}
dx
}
pub fn swiglu_backward(
dy: &[f32],
gate_pre: &[f32],
up_pre: &[f32],
w_down: &[f32],
w_gate: &[f32],
w_up: &[f32],
hidden: usize,
inter: usize,
) -> (Vec<f32>, Vec<f32>) {
assert_eq!(dy.len(), hidden);
assert_eq!(gate_pre.len(), inter);
assert_eq!(up_pre.len(), inter);
assert_eq!(w_down.len(), hidden * inter);
assert_eq!(w_gate.len(), inter * hidden);
assert_eq!(w_up.len(), inter * hidden);
let mut dm = vec![0.0f32; inter];
crate::forward::cpu::matmul_into(dy, w_down, &mut dm, 1, hidden, inter);
let mut d_up = vec![0.0f32; inter];
let mut d_gate_pre = vec![0.0f32; inter];
for j in 0..inter {
let a = gate_pre[j];
let sigma_a = 1.0 / (1.0 + (-a).exp());
let silu_a = a * sigma_a;
let up_j = up_pre[j];
let dm_j = dm[j];
d_up[j] = dm_j * silu_a;
let d_s = dm_j * up_j;
let silu_prime = sigma_a + a * sigma_a * (1.0 - sigma_a);
d_gate_pre[j] = d_s * silu_prime;
}
let mut dx = vec![0.0f32; hidden];
crate::forward::cpu::matmul_into(&d_up, w_up, &mut dx, 1, inter, hidden);
let mut dx_gate = vec![0.0f32; hidden];
crate::forward::cpu::matmul_into(&d_gate_pre, w_gate, &mut dx_gate, 1, inter, hidden);
for (dxi, dgi) in dx.iter_mut().zip(dx_gate.iter()) {
*dxi += *dgi;
}
(dx, dm)
}
#[cfg(test)]
fn linear_vjp_scalar(w: &[f32], g: &[f32], d_in: usize, d_out: usize) -> Vec<f32> {
let mut dx = vec![0.0f32; d_in];
for i in 0..d_out {
let gi = g[i];
let row = &w[i * d_in..(i + 1) * d_in];
for (j, &wij) in row.iter().enumerate() {
dx[j] += wij * gi;
}
}
dx
}
#[cfg(test)]
fn lora_bt_g_dx_scalar(
g: &[f32],
a: &[f32],
b: &[f32],
rank: usize,
d_in: usize,
d_out: usize,
scale: f32,
) -> (Vec<f32>, Vec<f32>) {
let mut bt_g = vec![0.0f32; rank];
for r in 0..rank {
let mut acc = 0.0f32;
for i in 0..d_out {
acc += b[i * rank + r] * g[i];
}
bt_g[r] = acc;
}
let mut dx = vec![0.0f32; d_in];
for r in 0..rank {
let scaled_bt_g = scale * bt_g[r];
let row = &a[r * d_in..(r + 1) * d_in];
for (j, &aij) in row.iter().enumerate() {
dx[j] += aij * scaled_bt_g;
}
}
(bt_g, dx)
}
#[cfg(test)]
fn swiglu_backward_scalar(
dy: &[f32],
gate_pre: &[f32],
up_pre: &[f32],
w_down: &[f32],
w_gate: &[f32],
w_up: &[f32],
hidden: usize,
inter: usize,
) -> (Vec<f32>, Vec<f32>) {
let mut dm = vec![0.0f32; inter];
for j in 0..inter {
let mut acc = 0.0f32;
for i in 0..hidden {
acc += w_down[i * inter + j] * dy[i];
}
dm[j] = acc;
}
let mut dx = vec![0.0f32; hidden];
let mut d_gate_pre = vec![0.0f32; inter];
for j in 0..inter {
let a = gate_pre[j];
let sigma_a = 1.0 / (1.0 + (-a).exp());
let silu_a = a * sigma_a;
let up_j = up_pre[j];
let dm_j = dm[j];
let d_up = dm_j * silu_a;
let d_s = dm_j * up_j;
let silu_prime = sigma_a + a * sigma_a * (1.0 - sigma_a);
d_gate_pre[j] = d_s * silu_prime;
let up_row = &w_up[j * hidden..(j + 1) * hidden];
for (i, &wu) in up_row.iter().enumerate() {
dx[i] += wu * d_up;
}
}
for j in 0..inter {
let wg_row = &w_gate[j * hidden..(j + 1) * hidden];
let dg = d_gate_pre[j];
for (i, &wg) in wg_row.iter().enumerate() {
dx[i] += wg * dg;
}
}
(dx, dm)
}
#[cfg(test)]
fn rel_err(analytic: &[f32], fd: &[f32]) -> f64 {
let diff_sq: f64 = analytic
.iter()
.zip(fd.iter())
.map(|(&a, &b)| ((a - b) as f64).powi(2))
.sum();
let norm_sq: f64 = analytic.iter().map(|&a| (a as f64).powi(2)).sum();
(diff_sq / norm_sq.max(1e-30)).sqrt()
}
#[cfg(test)]
mod tests {
use super::*;
const EPS: f32 = 1e-3;
const TOL: f64 = 1e-3;
#[test]
fn cross_entropy_backward_gradcheck() {
let seq_len = 4;
let vocab = 8;
let completion_start = 2;
let logits: Vec<f32> = (0..seq_len * vocab)
.map(|i| (i as f32) * 0.1 - 2.0)
.collect();
let targets: Vec<u32> = (0..seq_len as u32).map(|i| i % vocab as u32).collect();
let loss_fn = |logits: &[f32]| -> f32 {
let mut total = 0.0f64;
let n_comp = (seq_len - completion_start) as f64;
for t in completion_start..seq_len {
let row = &logits[t * vocab..(t + 1) * vocab];
let max = row.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let sum_exp: f64 = row.iter().map(|&x| ((x - max) as f64).exp()).sum();
let lse = (max as f64) + sum_exp.ln();
total += lse - logits[t * vocab + targets[t] as usize] as f64;
}
(total / n_comp) as f32
};
let analytic = cross_entropy_backward(&logits, &targets, seq_len, vocab, completion_start);
let mut fd = vec![0.0f32; seq_len * vocab];
for k in 0..seq_len * vocab {
let mut lp = logits.clone();
let mut lm = logits.clone();
lp[k] += EPS;
lm[k] -= EPS;
fd[k] = (loss_fn(&lp) - loss_fn(&lm)) / (2.0 * EPS);
}
let err = rel_err(&analytic, &fd);
eprintln!("cross_entropy_backward rel_err={err:.2e}");
assert!(err < TOL, "cross_entropy rel_err {err:.2e} >= {TOL:.2e}");
}
#[test]
fn linear_vjp_gradcheck() {
let d_in = 6;
let d_out = 4;
let w: Vec<f32> = (0..d_in * d_out).map(|i| (i as f32 + 1.0) * 0.1).collect();
let x: Vec<f32> = (0..d_in).map(|i| i as f32 * 0.3 - 0.5).collect();
let g: Vec<f32> = (0..d_out).map(|i| if i == 0 { 1.0 } else { 0.0 }).collect();
let loss_fn = |x: &[f32]| -> f32 {
(0..d_out)
.map(|i| {
let row = &w[i * d_in..(i + 1) * d_in];
let wx_i: f32 = row.iter().zip(x.iter()).map(|(a, b)| a * b).sum();
g[i] * wx_i
})
.sum()
};
let analytic = linear_vjp(&w, &g, d_in, d_out);
let mut fd = vec![0.0f32; d_in];
for j in 0..d_in {
let mut xp = x.clone();
let mut xm = x.clone();
xp[j] += EPS;
xm[j] -= EPS;
fd[j] = (loss_fn(&xp) - loss_fn(&xm)) / (2.0 * EPS);
}
let err = rel_err(&analytic, &fd);
eprintln!("linear_vjp rel_err={err:.2e}");
assert!(err < TOL, "linear_vjp rel_err {err:.2e} >= {TOL:.2e}");
}
#[test]
fn lora_vjp_gradcheck() {
let rank = 2;
let d_in = 3;
let d_out = 4;
let scale = 0.5f32;
let a: Vec<f32> = (0..rank * d_in)
.map(|i| (i as f32 + 1.0) * 0.2 - 0.5)
.collect();
let b: Vec<f32> = (0..d_out * rank)
.map(|i| (i as f32 + 1.0) * 0.15 - 0.4)
.collect();
let x: Vec<f32> = (0..d_in).map(|i| (i as f32 + 1.0) * 0.4).collect();
let g: Vec<f32> = (0..d_out).map(|i| (i as f32 + 1.0) * 0.3 - 0.5).collect();
let forward = |a: &[f32], b: &[f32], x: &[f32]| -> f32 {
let mut h = vec![0.0f32; rank];
for r in 0..rank {
h[r] = a[r * d_in..(r + 1) * d_in]
.iter()
.zip(x.iter())
.map(|(ai, xi)| ai * xi)
.sum();
}
let mut y = vec![0.0f32; d_out];
for i in 0..d_out {
y[i] = scale
* b[i * rank..(i + 1) * rank]
.iter()
.zip(h.iter())
.map(|(bi, hi)| bi * hi)
.sum::<f32>();
}
g.iter().zip(y.iter()).map(|(gi, yi)| gi * yi).sum()
};
let h: Vec<f32> = (0..rank)
.map(|r| {
a[r * d_in..(r + 1) * d_in]
.iter()
.zip(x.iter())
.map(|(ai, xi)| ai * xi)
.sum()
})
.collect();
let (grad_b, grad_a, dx) = lora_vjp(&g, &x, &h, &a, &b, rank, d_in, d_out, scale);
let mut fd_a = vec![0.0f32; rank * d_in];
for k in 0..rank * d_in {
let mut ap = a.clone();
let mut am = a.clone();
ap[k] += EPS;
am[k] -= EPS;
fd_a[k] = (forward(&ap, &b, &x) - forward(&am, &b, &x)) / (2.0 * EPS);
}
let mut fd_b = vec![0.0f32; d_out * rank];
for k in 0..d_out * rank {
let mut bp = b.clone();
let mut bm = b.clone();
bp[k] += EPS;
bm[k] -= EPS;
fd_b[k] = (forward(&a, &bp, &x) - forward(&a, &bm, &x)) / (2.0 * EPS);
}
let mut fd_x = vec![0.0f32; d_in];
for k in 0..d_in {
let mut xp = x.clone();
let mut xm = x.clone();
xp[k] += EPS;
xm[k] -= EPS;
fd_x[k] = (forward(&a, &b, &xp) - forward(&a, &b, &xm)) / (2.0 * EPS);
}
let err_a = rel_err(&grad_a, &fd_a);
let err_b = rel_err(&grad_b, &fd_b);
let err_x = rel_err(&dx, &fd_x);
eprintln!("lora_vjp: grad_A rel_err={err_a:.2e} grad_B={err_b:.2e} dx={err_x:.2e}");
assert!(err_a < TOL, "lora grad_A rel_err {err_a:.2e} >= {TOL:.2e}");
assert!(err_b < TOL, "lora grad_B rel_err {err_b:.2e} >= {TOL:.2e}");
assert!(err_x < TOL, "lora dx rel_err {err_x:.2e} >= {TOL:.2e}");
}
#[test]
fn rmsnorm_backward_gradcheck() {
let d = 8;
let eps = 1e-6f32;
let x: Vec<f32> = (0..d).map(|i| (i as f32 + 1.0) * 0.3 - 1.0).collect();
let w: Vec<f32> = (0..d).map(|i| 1.0 + i as f32 * 0.05).collect();
let g: Vec<f32> = (0..d)
.map(|i| if i % 3 == 0 { 1.0 } else { -0.5 })
.collect();
let rms = |x: &[f32]| -> f32 {
let mean_sq: f32 = x.iter().map(|xi| xi * xi).sum::<f32>() / d as f32;
(mean_sq + eps).sqrt()
};
let loss_fn = |x: &[f32]| -> f32 {
let r = rms(x);
x.iter()
.zip(w.iter())
.zip(g.iter())
.map(|((&xi, &wi), &gi)| gi * xi * wi / r)
.sum()
};
let inv_rms_val = 1.0 / rms(&x);
let analytic = rmsnorm_backward(&x, &w, inv_rms_val, &g);
let mut fd = vec![0.0f32; d];
for k in 0..d {
let mut xp = x.clone();
let mut xm = x.clone();
xp[k] += EPS;
xm[k] -= EPS;
fd[k] = (loss_fn(&xp) - loss_fn(&xm)) / (2.0 * EPS);
}
let err = rel_err(&analytic, &fd);
eprintln!("rmsnorm_backward rel_err={err:.2e}");
assert!(err < TOL, "rmsnorm rel_err {err:.2e} >= {TOL:.2e}");
}
#[test]
fn rope_backward_gradcheck() {
let head_dim = 8;
let rope_dim = 4;
let half = rope_dim / 2;
let cos_vals: Vec<f32> = (0..half).map(|i| (i as f32 * 0.3).cos()).collect();
let sin_vals: Vec<f32> = (0..half).map(|i| (i as f32 * 0.3).sin()).collect();
let x: Vec<f32> = (0..head_dim).map(|i| (i as f32 + 1.0) * 0.2).collect();
let g: Vec<f32> = (0..head_dim)
.map(|i| (i as f32 + 1.0) * 0.15 - 1.0)
.collect();
let rope_fwd = |x: &[f32]| -> Vec<f32> {
let mut y = x.to_vec();
for i in 0..half {
let c = cos_vals[i];
let s = sin_vals[i];
let x0 = x[i];
let x1 = x[half + i];
y[i] = x0 * c - x1 * s;
y[half + i] = x0 * s + x1 * c;
}
y
};
let loss_fn = |x: &[f32]| -> f32 {
let y = rope_fwd(x);
y.iter().zip(g.iter()).map(|(yi, gi)| yi * gi).sum()
};
let analytic = rope_backward(&g, &cos_vals, &sin_vals, rope_dim);
let mut fd = vec![0.0f32; head_dim];
for k in 0..head_dim {
let mut xp = x.clone();
let mut xm = x.clone();
xp[k] += EPS;
xm[k] -= EPS;
fd[k] = (loss_fn(&xp) - loss_fn(&xm)) / (2.0 * EPS);
}
let err = rel_err(&analytic, &fd);
eprintln!("rope_backward rel_err={err:.2e}");
assert!(err < TOL, "rope rel_err {err:.2e} >= {TOL:.2e}");
}
#[test]
fn swiglu_backward_gradcheck() {
let hidden = 4;
let inter = 6;
let w_gate: Vec<f32> = (0..inter * hidden)
.map(|i| (i as f32 + 1.0) * 0.1 - 1.5)
.collect();
let w_up: Vec<f32> = (0..inter * hidden)
.map(|i| (i as f32 + 1.0) * 0.12 - 1.2)
.collect();
let w_down: Vec<f32> = (0..hidden * inter)
.map(|i| (i as f32 + 1.0) * 0.08 - 0.9)
.collect();
let x: Vec<f32> = (0..hidden).map(|i| (i as f32 + 1.0) * 0.25).collect();
let dy: Vec<f32> = (0..hidden).map(|i| (i as f32 + 1.0) * 0.2 - 1.0).collect();
let swiglu_fwd = |x: &[f32]| -> f32 {
let mut gate = vec![0.0f32; inter];
let mut up = vec![0.0f32; inter];
for j in 0..inter {
gate[j] = w_gate[j * hidden..(j + 1) * hidden]
.iter()
.zip(x.iter())
.map(|(a, b)| a * b)
.sum();
up[j] = w_up[j * hidden..(j + 1) * hidden]
.iter()
.zip(x.iter())
.map(|(a, b)| a * b)
.sum();
}
let mut m = vec![0.0f32; inter];
for j in 0..inter {
let s = gate[j] / (1.0 + (-gate[j]).exp());
m[j] = s * up[j];
}
let mut y = vec![0.0f32; hidden];
for i in 0..hidden {
y[i] = w_down[i * inter..(i + 1) * inter]
.iter()
.zip(m.iter())
.map(|(a, b)| a * b)
.sum();
}
dy.iter().zip(y.iter()).map(|(d, yi)| d * yi).sum()
};
let gate_pre: Vec<f32> = (0..inter)
.map(|j| {
w_gate[j * hidden..(j + 1) * hidden]
.iter()
.zip(x.iter())
.map(|(a, b)| a * b)
.sum()
})
.collect();
let up_pre: Vec<f32> = (0..inter)
.map(|j| {
w_up[j * hidden..(j + 1) * hidden]
.iter()
.zip(x.iter())
.map(|(a, b)| a * b)
.sum()
})
.collect();
let (analytic_dx, _dm) = swiglu_backward(
&dy, &gate_pre, &up_pre, &w_down, &w_gate, &w_up, hidden, inter,
);
let mut fd = vec![0.0f32; hidden];
for k in 0..hidden {
let mut xp = x.clone();
let mut xm = x.clone();
xp[k] += EPS;
xm[k] -= EPS;
fd[k] = (swiglu_fwd(&xp) - swiglu_fwd(&xm)) / (2.0 * EPS);
}
let err = rel_err(&analytic_dx, &fd);
eprintln!("swiglu_backward rel_err={err:.2e}");
assert!(err < TOL, "swiglu rel_err {err:.2e} >= {TOL:.2e}");
}
const PARITY_TOL: f64 = 1e-4;
fn xorshift_fill(seed: u64, n: usize, amp: f32) -> Vec<f32> {
let mut state = seed | 1;
(0..n)
.map(|_| {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
((state >> 32) as u32 as f32 / u32::MAX as f32 * 2.0 - 1.0) * amp
})
.collect()
}
#[test]
fn linear_vjp_parity_vs_scalar() {
let d_in = 517;
let d_out = 251;
let w = xorshift_fill(1, d_out * d_in, 0.5);
let g = xorshift_fill(2, d_out, 1.0);
let vectorized = linear_vjp(&w, &g, d_in, d_out);
let scalar = linear_vjp_scalar(&w, &g, d_in, d_out);
let err = rel_err(&vectorized, &scalar);
eprintln!("linear_vjp parity rel_err={err:.2e}");
assert!(
err < PARITY_TOL,
"linear_vjp vectorized vs scalar rel_err {err:.2e} >= {PARITY_TOL:.2e}"
);
}
#[test]
fn lora_vjp_parity_vs_scalar() {
let rank = 11;
let d_in = 199;
let d_out = 257;
let scale = 0.37f32;
let g = xorshift_fill(3, d_out, 1.0);
let x = xorshift_fill(4, d_in, 0.5);
let a = xorshift_fill(5, rank * d_in, 0.3);
let b = xorshift_fill(6, d_out * rank, 0.3);
let h = xorshift_fill(7, rank, 0.4);
let (_grad_b, _grad_a, dx_vec) = lora_vjp(&g, &x, &h, &a, &b, rank, d_in, d_out, scale);
let (bt_g_scalar, dx_scalar) = lora_bt_g_dx_scalar(&g, &a, &b, rank, d_in, d_out, scale);
let mut bt_g_vec = vec![0.0f32; rank];
crate::forward::cpu::matmul_into(&g, &b, &mut bt_g_vec, 1, d_out, rank);
let err_btg = rel_err(&bt_g_vec, &bt_g_scalar);
let err_dx = rel_err(&dx_vec, &dx_scalar);
eprintln!("lora_vjp parity bt_g rel_err={err_btg:.2e} dx rel_err={err_dx:.2e}");
assert!(
err_btg < PARITY_TOL,
"lora bt_g vectorized vs scalar rel_err {err_btg:.2e} >= {PARITY_TOL:.2e}"
);
assert!(
err_dx < PARITY_TOL,
"lora dx vectorized vs scalar rel_err {err_dx:.2e} >= {PARITY_TOL:.2e}"
);
}
#[test]
fn swiglu_backward_parity_vs_scalar() {
let hidden = 131;
let inter = 347;
let dy = xorshift_fill(8, hidden, 0.8);
let gate_pre = xorshift_fill(9, inter, 1.2);
let up_pre = xorshift_fill(10, inter, 1.2);
let w_down = xorshift_fill(11, hidden * inter, 0.2);
let w_gate = xorshift_fill(12, inter * hidden, 0.2);
let w_up = xorshift_fill(13, inter * hidden, 0.2);
let (dx_vec, dm_vec) = swiglu_backward(
&dy, &gate_pre, &up_pre, &w_down, &w_gate, &w_up, hidden, inter,
);
let (dx_scalar, dm_scalar) = swiglu_backward_scalar(
&dy, &gate_pre, &up_pre, &w_down, &w_gate, &w_up, hidden, inter,
);
let err_dm = rel_err(&dm_vec, &dm_scalar);
let err_dx = rel_err(&dx_vec, &dx_scalar);
eprintln!("swiglu_backward parity dm rel_err={err_dm:.2e} dx rel_err={err_dx:.2e}");
assert!(
err_dm < PARITY_TOL,
"swiglu dm vectorized vs scalar rel_err {err_dm:.2e} >= {PARITY_TOL:.2e}"
);
assert!(
err_dx < PARITY_TOL,
"swiglu dx vectorized vs scalar rel_err {err_dx:.2e} >= {PARITY_TOL:.2e}"
);
}
}