#[derive(Clone, Copy, Debug)]
pub struct MlaDims {
pub n_head: usize,
pub d_nope: usize,
pub d_rope: usize,
pub d_v: usize,
pub kv_rank: usize,
}
impl MlaDims {
pub const GLM52: MlaDims = MlaDims {
n_head: 64,
d_nope: 192,
d_rope: 64,
d_v: 256,
kv_rank: 512,
};
pub fn scale(&self) -> f32 {
1.0 / ((self.d_nope + self.d_rope) as f32).sqrt()
}
}
pub struct MlaInputs<'a> {
pub q_nope: &'a [f32],
pub q_pe: &'a [f32],
pub c_kv: &'a [f32],
pub k_pe: &'a [f32],
pub w_uk: &'a [f32],
pub w_uv: &'a [f32],
pub t_q: usize,
pub t_kv: usize,
}
fn check_shapes(d: &MlaDims, x: &MlaInputs) {
assert_eq!(x.q_nope.len(), x.t_q * d.n_head * d.d_nope, "q_nope shape");
assert_eq!(x.q_pe.len(), x.t_q * d.n_head * d.d_rope, "q_pe shape");
assert_eq!(x.c_kv.len(), x.t_kv * d.kv_rank, "c_kv shape");
assert_eq!(x.k_pe.len(), x.t_kv * d.d_rope, "k_pe shape");
assert_eq!(x.w_uk.len(), d.n_head * d.d_nope * d.kv_rank, "w_uk shape");
assert_eq!(x.w_uv.len(), d.n_head * d.d_v * d.kv_rank, "w_uv shape");
assert!(x.t_q <= x.t_kv, "queries must be a suffix of the cache");
}
fn softmax(s: &mut [f32]) {
let m = s.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f32;
for v in s.iter_mut() {
*v = (*v - m).exp();
sum += *v;
}
let inv = 1.0 / sum;
for v in s.iter_mut() {
*v *= inv;
}
}
pub fn mla_attend_naive(d: &MlaDims, x: &MlaInputs) -> Vec<f32> {
check_shapes(d, x);
let (nh, dn, dr, dv, r) = (d.n_head, d.d_nope, d.d_rope, d.d_v, d.kv_rank);
let scale = d.scale();
let mut out = vec![0.0f32; x.t_q * nh * dv];
let mut k_nope = vec![0.0f32; x.t_kv * dn];
let mut v = vec![0.0f32; x.t_kv * dv];
let mut scores = vec![0.0f32; x.t_kv];
for h in 0..nh {
let wuk = &x.w_uk[h * dn * r..(h + 1) * dn * r];
let wuv = &x.w_uv[h * dv * r..(h + 1) * dv * r];
for t in 0..x.t_kv {
let c = &x.c_kv[t * r..(t + 1) * r];
for p in 0..dn {
let row = &wuk[p * r..(p + 1) * r];
let mut acc = 0.0f32;
for l in 0..r {
acc += row[l] * c[l];
}
k_nope[t * dn + p] = acc;
}
for j in 0..dv {
let row = &wuv[j * r..(j + 1) * r];
let mut acc = 0.0f32;
for l in 0..r {
acc += row[l] * c[l];
}
v[t * dv + j] = acc;
}
}
for i in 0..x.t_q {
let visible = x.t_kv - x.t_q + i + 1; let qn = &x.q_nope[(i * nh + h) * dn..(i * nh + h + 1) * dn];
let qp = &x.q_pe[(i * nh + h) * dr..(i * nh + h + 1) * dr];
for t in 0..visible {
let mut s = 0.0f32;
let kn = &k_nope[t * dn..(t + 1) * dn];
for p in 0..dn {
s += qn[p] * kn[p];
}
let kp = &x.k_pe[t * dr..(t + 1) * dr];
for p in 0..dr {
s += qp[p] * kp[p];
}
scores[t] = s * scale;
}
softmax(&mut scores[..visible]);
let o = &mut out[(i * nh + h) * dv..(i * nh + h + 1) * dv];
for t in 0..visible {
let p = scores[t];
let vt = &v[t * dv..(t + 1) * dv];
for j in 0..dv {
o[j] += p * vt[j];
}
}
}
}
out
}
pub fn mla_attend_absorbed(d: &MlaDims, x: &MlaInputs) -> Vec<f32> {
check_shapes(d, x);
let (nh, dn, dr, dv, r) = (d.n_head, d.d_nope, d.d_rope, d.d_v, d.kv_rank);
let scale = d.scale();
let mut out = vec![0.0f32; x.t_q * nh * dv];
let mut q_lat = vec![0.0f32; r]; let mut o_lat = vec![0.0f32; r]; let mut scores = vec![0.0f32; x.t_kv];
for h in 0..nh {
let wuk = &x.w_uk[h * dn * r..(h + 1) * dn * r];
let wuv = &x.w_uv[h * dv * r..(h + 1) * dv * r];
for i in 0..x.t_q {
let visible = x.t_kv - x.t_q + i + 1;
let qn = &x.q_nope[(i * nh + h) * dn..(i * nh + h + 1) * dn];
let qp = &x.q_pe[(i * nh + h) * dr..(i * nh + h + 1) * dr];
q_lat.iter_mut().for_each(|v| *v = 0.0);
for p in 0..dn {
let row = &wuk[p * r..(p + 1) * r];
let qv = qn[p];
for l in 0..r {
q_lat[l] += qv * row[l];
}
}
for t in 0..visible {
let c = &x.c_kv[t * r..(t + 1) * r];
let mut s = 0.0f32;
for l in 0..r {
s += q_lat[l] * c[l];
}
let kp = &x.k_pe[t * dr..(t + 1) * dr];
for p in 0..dr {
s += qp[p] * kp[p];
}
scores[t] = s * scale;
}
softmax(&mut scores[..visible]);
o_lat.iter_mut().for_each(|v| *v = 0.0);
for t in 0..visible {
let p = scores[t];
let c = &x.c_kv[t * r..(t + 1) * r];
for l in 0..r {
o_lat[l] += p * c[l];
}
}
let o = &mut out[(i * nh + h) * dv..(i * nh + h + 1) * dv];
for j in 0..dv {
let row = &wuv[j * r..(j + 1) * r];
let mut acc = 0.0f32;
for l in 0..r {
acc += row[l] * o_lat[l];
}
o[j] = acc;
}
}
}
out
}
pub fn rope_interleaved(x: &mut [f32], n_dims: usize, pos: f32, base: f32) {
let half = n_dims / 2;
let theta_scale = base.powf(-2.0 / n_dims as f32);
let mut theta = pos;
for j in 0..half {
let (sin, cos) = theta.sin_cos();
let a = x[2 * j];
let b = x[2 * j + 1];
x[2 * j] = a * cos - b * sin;
x[2 * j + 1] = a * sin + b * cos;
theta *= theta_scale;
}
}
pub fn rope_neox(x: &mut [f32], n_dims: usize, pos: f32, base: f32) {
let half = n_dims / 2;
let theta_scale = base.powf(-2.0 / n_dims as f32);
let mut theta = pos;
for j in 0..half {
let (sin, cos) = theta.sin_cos();
let a = x[j];
let b = x[j + half];
x[j] = a * cos - b * sin;
x[j + half] = a * sin + b * cos;
theta *= theta_scale;
}
}
pub fn norm_to_neox_perm(n_dims: usize) -> Vec<usize> {
let half = n_dims / 2;
let mut p = vec![0usize; n_dims];
for j in 0..half {
p[2 * j] = j;
p[2 * j + 1] = j + half;
}
p
}
#[cfg(test)]
mod tests {
use super::*;
struct Rng(u64);
impl Rng {
fn next_f32(&mut self) -> f32 {
self.0 ^= self.0 << 13;
self.0 ^= self.0 >> 7;
self.0 ^= self.0 << 17;
let v = (self.0.wrapping_mul(0x2545F4914F6CDD1D) >> 40) as u32;
(v as f32 / (1u32 << 24) as f32) * 2.0 - 1.0 }
fn fill(&mut self, n: usize, scale: f32) -> Vec<f32> {
(0..n).map(|_| self.next_f32() * scale).collect()
}
}
fn maxdiff(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len());
a.iter()
.zip(b)
.map(|(x, y)| (x - y).abs())
.fold(0.0f32, f32::max)
}
fn maxabs(a: &[f32]) -> f32 {
a.iter().map(|x| x.abs()).fold(0.0f32, f32::max)
}
fn random_case(
d: &MlaDims,
t_q: usize,
t_kv: usize,
seed: u64,
) -> (Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>) {
let mut rng = Rng(seed | 1);
let ws = 1.0 / (d.kv_rank as f32).sqrt();
(
rng.fill(t_q * d.n_head * d.d_nope, 1.0),
rng.fill(t_q * d.n_head * d.d_rope, 1.0),
rng.fill(t_kv * d.kv_rank, 1.0),
rng.fill(t_kv * d.d_rope, 1.0),
rng.fill(d.n_head * d.d_nope * d.kv_rank, ws),
rng.fill(d.n_head * d.d_v * d.kv_rank, ws),
)
}
fn run_case(d: &MlaDims, t_q: usize, t_kv: usize, seed: u64, tol: f32) {
let (q_nope, q_pe, c_kv, k_pe, w_uk, w_uv) = random_case(d, t_q, t_kv, seed);
let x = MlaInputs {
q_nope: &q_nope,
q_pe: &q_pe,
c_kv: &c_kv,
k_pe: &k_pe,
w_uk: &w_uk,
w_uv: &w_uv,
t_q,
t_kv,
};
let naive = mla_attend_naive(d, &x);
let absorbed = mla_attend_absorbed(d, &x);
let md = maxdiff(&naive, &absorbed);
let scale = maxabs(&naive).max(1.0);
assert!(
md <= tol * scale,
"naive vs absorbed disagree: maxdiff {md:.3e} (scale {scale:.3e}, rel {:.3e}) \
dims {d:?} t_q {t_q} t_kv {t_kv} seed {seed}",
md / scale
);
assert!(naive.iter().all(|v| v.is_finite()));
assert!(maxabs(&naive) > 1e-6);
}
#[test]
fn naive_equals_absorbed_decode_t1() {
let shapes = [
MlaDims {
n_head: 4,
d_nope: 24,
d_rope: 8,
d_v: 32,
kv_rank: 64,
},
MlaDims {
n_head: 2,
d_nope: 16,
d_rope: 16,
d_v: 16,
kv_rank: 32,
},
MlaDims {
n_head: 3,
d_nope: 12,
d_rope: 4,
d_v: 20,
kv_rank: 48,
},
];
for (i, d) in shapes.iter().enumerate() {
for seed in [7, 1234, 0xB1E55ED] {
run_case(d, 1, 17, seed + i as u64, 1e-5);
}
}
}
#[test]
fn naive_equals_absorbed_prefill_causal() {
let d = MlaDims {
n_head: 4,
d_nope: 24,
d_rope: 8,
d_v: 32,
kv_rank: 64,
};
run_case(&d, 5, 9, 42, 1e-5);
run_case(&d, 8, 8, 43, 1e-5); let d2 = MlaDims {
n_head: 2,
d_nope: 16,
d_rope: 16,
d_v: 16,
kv_rank: 32,
};
run_case(&d2, 3, 11, 44, 1e-5);
}
#[test]
fn naive_equals_absorbed_glm52_full_dims() {
run_case(&MlaDims::GLM52, 1, 8, 20260801, 1e-4);
}
#[test]
fn rope_norm_equals_permuted_neox() {
let n_dims = 64;
let base = 8_000_000.0f32; let perm = norm_to_neox_perm(n_dims);
let mut rng = Rng(99);
for pos in [0.0f32, 1.0, 17.0, 4096.0, 1_000_000.0] {
let x0: Vec<f32> = (0..n_dims).map(|_| rng.next_f32()).collect();
let y0: Vec<f32> = (0..n_dims).map(|_| rng.next_f32()).collect();
let mut xa = x0.clone();
rope_interleaved(&mut xa, n_dims, pos, base);
let mut xa_p = vec![0.0f32; n_dims];
for (src, &dst) in perm.iter().enumerate() {
xa_p[dst] = xa[src];
}
let mut xb = vec![0.0f32; n_dims];
for (src, &dst) in perm.iter().enumerate() {
xb[dst] = x0[src];
}
rope_neox(&mut xb, n_dims, pos, base);
assert!(
maxdiff(&xa_p, &xb) <= 1e-6,
"perm/rope orders disagree at pos {pos}"
);
let mut ya = y0.clone();
rope_interleaved(&mut ya, n_dims, pos, base);
let dot_norm: f32 = xa.iter().zip(&ya).map(|(a, b)| a * b).sum();
let mut yb = vec![0.0f32; n_dims];
for (src, &dst) in perm.iter().enumerate() {
yb[dst] = y0[src];
}
rope_neox(&mut yb, n_dims, pos, base);
let dot_neox: f32 = xb.iter().zip(&yb).map(|(a, b)| a * b).sum();
assert!(
(dot_norm - dot_neox).abs() <= 1e-4 * dot_norm.abs().max(1.0),
"roped dot products diverge at pos {pos}: {dot_norm} vs {dot_neox}"
);
}
}
}