1#[derive(Clone, Copy, Debug)]
20pub struct MlaDims {
21 pub n_head: usize,
22 pub d_nope: usize,
24 pub d_rope: usize,
26 pub d_v: usize,
28 pub kv_rank: usize,
30}
31
32impl MlaDims {
33 pub const GLM52: MlaDims = MlaDims {
34 n_head: 64,
35 d_nope: 192,
36 d_rope: 64,
37 d_v: 256,
38 kv_rank: 512,
39 };
40
41 pub fn scale(&self) -> f32 {
44 1.0 / ((self.d_nope + self.d_rope) as f32).sqrt()
45 }
46}
47
48pub struct MlaInputs<'a> {
63 pub q_nope: &'a [f32],
64 pub q_pe: &'a [f32],
65 pub c_kv: &'a [f32],
66 pub k_pe: &'a [f32],
67 pub w_uk: &'a [f32],
68 pub w_uv: &'a [f32],
69 pub t_q: usize,
70 pub t_kv: usize,
71}
72
73fn check_shapes(d: &MlaDims, x: &MlaInputs) {
74 assert_eq!(x.q_nope.len(), x.t_q * d.n_head * d.d_nope, "q_nope shape");
75 assert_eq!(x.q_pe.len(), x.t_q * d.n_head * d.d_rope, "q_pe shape");
76 assert_eq!(x.c_kv.len(), x.t_kv * d.kv_rank, "c_kv shape");
77 assert_eq!(x.k_pe.len(), x.t_kv * d.d_rope, "k_pe shape");
78 assert_eq!(x.w_uk.len(), d.n_head * d.d_nope * d.kv_rank, "w_uk shape");
79 assert_eq!(x.w_uv.len(), d.n_head * d.d_v * d.kv_rank, "w_uv shape");
80 assert!(x.t_q <= x.t_kv, "queries must be a suffix of the cache");
81}
82
83fn softmax(s: &mut [f32]) {
85 let m = s.iter().copied().fold(f32::NEG_INFINITY, f32::max);
86 let mut sum = 0.0f32;
87 for v in s.iter_mut() {
88 *v = (*v - m).exp();
89 sum += *v;
90 }
91 let inv = 1.0 / sum;
92 for v in s.iter_mut() {
93 *v *= inv;
94 }
95}
96
97pub fn mla_attend_naive(d: &MlaDims, x: &MlaInputs) -> Vec<f32> {
101 check_shapes(d, x);
102 let (nh, dn, dr, dv, r) = (d.n_head, d.d_nope, d.d_rope, d.d_v, d.kv_rank);
103 let scale = d.scale();
104 let mut out = vec![0.0f32; x.t_q * nh * dv];
105
106 let mut k_nope = vec![0.0f32; x.t_kv * dn];
108 let mut v = vec![0.0f32; x.t_kv * dv];
109 let mut scores = vec![0.0f32; x.t_kv];
110 for h in 0..nh {
111 let wuk = &x.w_uk[h * dn * r..(h + 1) * dn * r];
112 let wuv = &x.w_uv[h * dv * r..(h + 1) * dv * r];
113 for t in 0..x.t_kv {
114 let c = &x.c_kv[t * r..(t + 1) * r];
115 for p in 0..dn {
116 let row = &wuk[p * r..(p + 1) * r];
117 let mut acc = 0.0f32;
118 for l in 0..r {
119 acc += row[l] * c[l];
120 }
121 k_nope[t * dn + p] = acc;
122 }
123 for j in 0..dv {
124 let row = &wuv[j * r..(j + 1) * r];
125 let mut acc = 0.0f32;
126 for l in 0..r {
127 acc += row[l] * c[l];
128 }
129 v[t * dv + j] = acc;
130 }
131 }
132 for i in 0..x.t_q {
133 let visible = x.t_kv - x.t_q + i + 1; let qn = &x.q_nope[(i * nh + h) * dn..(i * nh + h + 1) * dn];
135 let qp = &x.q_pe[(i * nh + h) * dr..(i * nh + h + 1) * dr];
136 for t in 0..visible {
137 let mut s = 0.0f32;
138 let kn = &k_nope[t * dn..(t + 1) * dn];
139 for p in 0..dn {
140 s += qn[p] * kn[p];
141 }
142 let kp = &x.k_pe[t * dr..(t + 1) * dr];
143 for p in 0..dr {
144 s += qp[p] * kp[p];
145 }
146 scores[t] = s * scale;
147 }
148 softmax(&mut scores[..visible]);
149 let o = &mut out[(i * nh + h) * dv..(i * nh + h + 1) * dv];
150 for t in 0..visible {
151 let p = scores[t];
152 let vt = &v[t * dv..(t + 1) * dv];
153 for j in 0..dv {
154 o[j] += p * vt[j];
155 }
156 }
157 }
158 }
159 out
160}
161
162pub fn mla_attend_absorbed(d: &MlaDims, x: &MlaInputs) -> Vec<f32> {
167 check_shapes(d, x);
168 let (nh, dn, dr, dv, r) = (d.n_head, d.d_nope, d.d_rope, d.d_v, d.kv_rank);
169 let scale = d.scale();
170 let mut out = vec![0.0f32; x.t_q * nh * dv];
171
172 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];
175 for h in 0..nh {
176 let wuk = &x.w_uk[h * dn * r..(h + 1) * dn * r];
177 let wuv = &x.w_uv[h * dv * r..(h + 1) * dv * r];
178 for i in 0..x.t_q {
179 let visible = x.t_kv - x.t_q + i + 1;
180 let qn = &x.q_nope[(i * nh + h) * dn..(i * nh + h + 1) * dn];
181 let qp = &x.q_pe[(i * nh + h) * dr..(i * nh + h + 1) * dr];
182 q_lat.iter_mut().for_each(|v| *v = 0.0);
184 for p in 0..dn {
185 let row = &wuk[p * r..(p + 1) * r];
186 let qv = qn[p];
187 for l in 0..r {
188 q_lat[l] += qv * row[l];
189 }
190 }
191 for t in 0..visible {
193 let c = &x.c_kv[t * r..(t + 1) * r];
194 let mut s = 0.0f32;
195 for l in 0..r {
196 s += q_lat[l] * c[l];
197 }
198 let kp = &x.k_pe[t * dr..(t + 1) * dr];
199 for p in 0..dr {
200 s += qp[p] * kp[p];
201 }
202 scores[t] = s * scale;
203 }
204 softmax(&mut scores[..visible]);
205 o_lat.iter_mut().for_each(|v| *v = 0.0);
207 for t in 0..visible {
208 let p = scores[t];
209 let c = &x.c_kv[t * r..(t + 1) * r];
210 for l in 0..r {
211 o_lat[l] += p * c[l];
212 }
213 }
214 let o = &mut out[(i * nh + h) * dv..(i * nh + h + 1) * dv];
216 for j in 0..dv {
217 let row = &wuv[j * r..(j + 1) * r];
218 let mut acc = 0.0f32;
219 for l in 0..r {
220 acc += row[l] * o_lat[l];
221 }
222 o[j] = acc;
223 }
224 }
225 }
226 out
227}
228
229pub fn rope_interleaved(x: &mut [f32], n_dims: usize, pos: f32, base: f32) {
239 let half = n_dims / 2;
240 let theta_scale = base.powf(-2.0 / n_dims as f32);
241 let mut theta = pos;
242 for j in 0..half {
243 let (sin, cos) = theta.sin_cos();
244 let a = x[2 * j];
245 let b = x[2 * j + 1];
246 x[2 * j] = a * cos - b * sin;
247 x[2 * j + 1] = a * sin + b * cos;
248 theta *= theta_scale;
249 }
250}
251
252pub fn rope_neox(x: &mut [f32], n_dims: usize, pos: f32, base: f32) {
255 let half = n_dims / 2;
256 let theta_scale = base.powf(-2.0 / n_dims as f32);
257 let mut theta = pos;
258 for j in 0..half {
259 let (sin, cos) = theta.sin_cos();
260 let a = x[j];
261 let b = x[j + half];
262 x[j] = a * cos - b * sin;
263 x[j + half] = a * sin + b * cos;
264 theta *= theta_scale;
265 }
266}
267
268pub fn norm_to_neox_perm(n_dims: usize) -> Vec<usize> {
273 let half = n_dims / 2;
274 let mut p = vec![0usize; n_dims];
275 for j in 0..half {
276 p[2 * j] = j;
277 p[2 * j + 1] = j + half;
278 }
279 p
280}
281
282#[cfg(test)]
283mod tests {
284 use super::*;
285
286 struct Rng(u64);
288 impl Rng {
289 fn next_f32(&mut self) -> f32 {
290 self.0 ^= self.0 << 13;
291 self.0 ^= self.0 >> 7;
292 self.0 ^= self.0 << 17;
293 let v = (self.0.wrapping_mul(0x2545F4914F6CDD1D) >> 40) as u32;
294 (v as f32 / (1u32 << 24) as f32) * 2.0 - 1.0 }
296 fn fill(&mut self, n: usize, scale: f32) -> Vec<f32> {
297 (0..n).map(|_| self.next_f32() * scale).collect()
298 }
299 }
300
301 fn maxdiff(a: &[f32], b: &[f32]) -> f32 {
302 assert_eq!(a.len(), b.len());
303 a.iter()
304 .zip(b)
305 .map(|(x, y)| (x - y).abs())
306 .fold(0.0f32, f32::max)
307 }
308 fn maxabs(a: &[f32]) -> f32 {
309 a.iter().map(|x| x.abs()).fold(0.0f32, f32::max)
310 }
311
312 fn random_case(
315 d: &MlaDims,
316 t_q: usize,
317 t_kv: usize,
318 seed: u64,
319 ) -> (Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>) {
320 let mut rng = Rng(seed | 1);
321 let ws = 1.0 / (d.kv_rank as f32).sqrt();
322 (
323 rng.fill(t_q * d.n_head * d.d_nope, 1.0),
324 rng.fill(t_q * d.n_head * d.d_rope, 1.0),
325 rng.fill(t_kv * d.kv_rank, 1.0),
326 rng.fill(t_kv * d.d_rope, 1.0),
327 rng.fill(d.n_head * d.d_nope * d.kv_rank, ws),
328 rng.fill(d.n_head * d.d_v * d.kv_rank, ws),
329 )
330 }
331
332 fn run_case(d: &MlaDims, t_q: usize, t_kv: usize, seed: u64, tol: f32) {
333 let (q_nope, q_pe, c_kv, k_pe, w_uk, w_uv) = random_case(d, t_q, t_kv, seed);
334 let x = MlaInputs {
335 q_nope: &q_nope,
336 q_pe: &q_pe,
337 c_kv: &c_kv,
338 k_pe: &k_pe,
339 w_uk: &w_uk,
340 w_uv: &w_uv,
341 t_q,
342 t_kv,
343 };
344 let naive = mla_attend_naive(d, &x);
345 let absorbed = mla_attend_absorbed(d, &x);
346 let md = maxdiff(&naive, &absorbed);
347 let scale = maxabs(&naive).max(1.0);
348 assert!(
349 md <= tol * scale,
350 "naive vs absorbed disagree: maxdiff {md:.3e} (scale {scale:.3e}, rel {:.3e}) \
351 dims {d:?} t_q {t_q} t_kv {t_kv} seed {seed}",
352 md / scale
353 );
354 assert!(naive.iter().all(|v| v.is_finite()));
356 assert!(maxabs(&naive) > 1e-6);
357 }
358
359 #[test]
360 fn naive_equals_absorbed_decode_t1() {
361 let shapes = [
363 MlaDims {
364 n_head: 4,
365 d_nope: 24,
366 d_rope: 8,
367 d_v: 32,
368 kv_rank: 64,
369 },
370 MlaDims {
371 n_head: 2,
372 d_nope: 16,
373 d_rope: 16,
374 d_v: 16,
375 kv_rank: 32,
376 },
377 MlaDims {
380 n_head: 3,
381 d_nope: 12,
382 d_rope: 4,
383 d_v: 20,
384 kv_rank: 48,
385 },
386 ];
387 for (i, d) in shapes.iter().enumerate() {
388 for seed in [7, 1234, 0xB1E55ED] {
389 run_case(d, 1, 17, seed + i as u64, 1e-5);
390 }
391 }
392 }
393
394 #[test]
395 fn naive_equals_absorbed_prefill_causal() {
396 let d = MlaDims {
398 n_head: 4,
399 d_nope: 24,
400 d_rope: 8,
401 d_v: 32,
402 kv_rank: 64,
403 };
404 run_case(&d, 5, 9, 42, 1e-5);
405 run_case(&d, 8, 8, 43, 1e-5); let d2 = MlaDims {
407 n_head: 2,
408 d_nope: 16,
409 d_rope: 16,
410 d_v: 16,
411 kv_rank: 32,
412 };
413 run_case(&d2, 3, 11, 44, 1e-5);
414 }
415
416 #[test]
417 fn naive_equals_absorbed_glm52_full_dims() {
418 run_case(&MlaDims::GLM52, 1, 8, 20260801, 1e-4);
421 }
422
423 #[test]
424 fn rope_norm_equals_permuted_neox() {
425 let n_dims = 64;
431 let base = 8_000_000.0f32; let perm = norm_to_neox_perm(n_dims);
433 let mut rng = Rng(99);
434 for pos in [0.0f32, 1.0, 17.0, 4096.0, 1_000_000.0] {
435 let x0: Vec<f32> = (0..n_dims).map(|_| rng.next_f32()).collect();
436 let y0: Vec<f32> = (0..n_dims).map(|_| rng.next_f32()).collect();
437
438 let mut xa = x0.clone();
440 rope_interleaved(&mut xa, n_dims, pos, base);
441 let mut xa_p = vec![0.0f32; n_dims];
442 for (src, &dst) in perm.iter().enumerate() {
443 xa_p[dst] = xa[src];
444 }
445 let mut xb = vec![0.0f32; n_dims];
447 for (src, &dst) in perm.iter().enumerate() {
448 xb[dst] = x0[src];
449 }
450 rope_neox(&mut xb, n_dims, pos, base);
451
452 assert!(
453 maxdiff(&xa_p, &xb) <= 1e-6,
454 "perm/rope orders disagree at pos {pos}"
455 );
456
457 let mut ya = y0.clone();
459 rope_interleaved(&mut ya, n_dims, pos, base);
460 let dot_norm: f32 = xa.iter().zip(&ya).map(|(a, b)| a * b).sum();
461
462 let mut yb = vec![0.0f32; n_dims];
463 for (src, &dst) in perm.iter().enumerate() {
464 yb[dst] = y0[src];
465 }
466 rope_neox(&mut yb, n_dims, pos, base);
467 let dot_neox: f32 = xb.iter().zip(&yb).map(|(a, b)| a * b).sum();
468
469 assert!(
470 (dot_norm - dot_neox).abs() <= 1e-4 * dot_norm.abs().max(1.0),
471 "roped dot products diverge at pos {pos}: {dot_norm} vs {dot_neox}"
472 );
473 }
474 }
475}