1use crate::{EngineError, SimdMode};
2use rayon::prelude::*;
3
4pub fn matmul(
7 a: &[f32],
8 a_rows: usize,
9 a_cols: usize,
10 b: &[f32],
11 b_rows: usize,
12 b_cols: usize,
13 _mode: SimdMode,
14) -> Result<Vec<f32>, EngineError> {
15 if a_cols != b_rows {
16 return Err(EngineError::ShapeMismatch(format!(
17 "matmul inner dim {a_cols} != {b_rows}"
18 )));
19 }
20 if a.len() != a_rows * a_cols || b.len() != b_rows * b_cols {
21 return Err(EngineError::ShapeMismatch(
22 "matmul buffer length does not match shape".into(),
23 ));
24 }
25 let mut out = vec![0.0f32; a_rows * b_cols];
26 for i in 0..a_rows {
27 for j in 0..b_cols {
28 let mut s = 0.0f32;
29 for k in 0..a_cols {
30 s += a[i * a_cols + k] * b[k * b_cols + j];
31 }
32 out[i * b_cols + j] = s;
33 }
34 }
35 Ok(out)
36}
37
38pub fn linear(x: &[f32], w: &[f32], out_f: usize, in_f: usize) -> Result<Vec<f32>, EngineError> {
40 if !x.len().is_multiple_of(in_f) {
41 return Err(EngineError::ShapeMismatch(format!(
42 "linear x len {} not divisible by in_f {in_f}",
43 x.len()
44 )));
45 }
46 if w.len() != out_f * in_f {
47 return Err(EngineError::ShapeMismatch(format!(
48 "linear weight length mismatch: got {} want out_f*in_f={out_f}*{in_f}={}",
49 w.len(),
50 out_f * in_f
51 )));
52 }
53 let batch = x.len() / in_f;
54 let mut out = vec![0.0f32; batch * out_f];
55 for b in 0..batch {
56 for o in 0..out_f {
57 let mut s = 0.0f32;
58 let wr = &w[o * in_f..(o + 1) * in_f];
59 let xr = &x[b * in_f..(b + 1) * in_f];
60 for i in 0..in_f {
61 s += xr[i] * wr[i];
62 }
63 out[b * out_f + o] = s;
64 }
65 }
66 Ok(out)
67}
68
69fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
70 debug_assert_eq!(a.len(), b.len());
71 #[cfg(target_arch = "x86_64")]
72 {
73 if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
74 return unsafe { dot_avx2(a, b) };
75 }
76 }
77 #[cfg(target_arch = "aarch64")]
78 {
79 unsafe { dot_neon(a, b) }
80 }
81 #[cfg(not(target_arch = "aarch64"))]
82 {
83 let mut s = 0.0f32;
84 for i in 0..a.len() {
85 s += a[i] * b[i];
86 }
87 s
88 }
89}
90
91#[cfg(target_arch = "x86_64")]
92#[target_feature(enable = "avx2,fma")]
93unsafe fn dot_avx2(a: &[f32], b: &[f32]) -> f32 {
94 use std::arch::x86_64::*;
95 let n = a.len();
96 let mut i = 0usize;
97 let mut acc = _mm256_setzero_ps();
98 while i + 8 <= n {
99 let va = _mm256_loadu_ps(a.as_ptr().add(i));
100 let vb = _mm256_loadu_ps(b.as_ptr().add(i));
101 acc = _mm256_fmadd_ps(va, vb, acc);
102 i += 8;
103 }
104 let mut tmp = [0.0f32; 8];
105 _mm256_storeu_ps(tmp.as_mut_ptr(), acc);
106 let mut s = tmp.iter().sum::<f32>();
107 while i < n {
108 s += a[i] * b[i];
109 i += 1;
110 }
111 s
112}
113
114#[cfg(target_arch = "aarch64")]
115unsafe fn dot_neon(a: &[f32], b: &[f32]) -> f32 {
116 use std::arch::aarch64::*;
117 let n = a.len();
118 let mut i = 0usize;
119 let mut acc = vdupq_n_f32(0.0);
120 while i + 4 <= n {
121 let va = vld1q_f32(a.as_ptr().add(i));
122 let vb = vld1q_f32(b.as_ptr().add(i));
123 acc = vfmaq_f32(acc, va, vb);
124 i += 4;
125 }
126 let mut tmp = [0.0f32; 4];
127 vst1q_f32(tmp.as_mut_ptr(), acc);
128 let mut s = tmp[0] + tmp[1] + tmp[2] + tmp[3];
129 while i < n {
130 s += a[i] * b[i];
131 i += 1;
132 }
133 s
134}
135
136pub fn linear_cpu(x: &[f32], w: &[f32], out_f: usize, in_f: usize) -> Result<Vec<f32>, EngineError> {
138 if !x.len().is_multiple_of(in_f) {
139 return Err(EngineError::ShapeMismatch(format!(
140 "linear x len {} not divisible by in_f {in_f}",
141 x.len()
142 )));
143 }
144 if w.len() != out_f * in_f {
145 return Err(EngineError::ShapeMismatch(format!(
146 "linear weight length mismatch: got {} want out_f*in_f={out_f}*{in_f}={}",
147 w.len(),
148 out_f * in_f
149 )));
150 }
151 let batch = x.len() / in_f;
152 let mut out = vec![0.0f32; batch * out_f];
153 if batch == 0 || out_f == 0 {
154 return Ok(out);
155 }
156 out.par_iter_mut().enumerate().for_each(|(idx, slot)| {
159 let b = idx / out_f;
160 let o = idx % out_f;
161 *slot = dot_f32(
162 &w[o * in_f..(o + 1) * in_f],
163 &x[b * in_f..(b + 1) * in_f],
164 );
165 });
166 Ok(out)
167}
168
169pub fn rms_norm(x: &[f32], weight: &[f32], eps: f32) -> Result<Vec<f32>, EngineError> {
170 let n = weight.len();
171 if n == 0 || !x.len().is_multiple_of(n) {
172 return Err(EngineError::ShapeMismatch(
173 "rms_norm: x length must be multiple of weight len".into(),
174 ));
175 }
176 let batch = x.len() / n;
177 let mut out = vec![0.0f32; x.len()];
178 for b in 0..batch {
179 let base = b * n;
180 let mut ms = 0.0f32;
181 for i in 0..n {
182 ms += x[base + i] * x[base + i];
183 }
184 let scale = (ms / n as f32 + eps).sqrt().recip();
185 for i in 0..n {
186 out[base + i] = x[base + i] * scale * weight[i];
187 }
188 }
189 Ok(out)
190}
191
192pub fn softmax_inplace(logits: &mut [f32]) {
193 if logits.is_empty() {
194 return;
195 }
196 let m = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
197 let mut sum = 0.0f32;
198 for v in logits.iter_mut() {
199 *v = (*v - m).exp();
200 sum += *v;
201 }
202 let inv = if sum > 0.0 { 1.0 / sum } else { 0.0 };
203 for v in logits.iter_mut() {
204 *v *= inv;
205 }
206}
207
208pub fn softmax(logits: &[f32]) -> Vec<f32> {
209 let mut o = logits.to_vec();
210 softmax_inplace(&mut o);
211 o
212}
213
214pub fn rope(x: &mut [f32], head_dim: usize, pos: usize, theta: f32) -> Result<(), EngineError> {
216 if head_dim == 0 || !head_dim.is_multiple_of(2) {
217 return Err(EngineError::ShapeMismatch(
218 "rope head_dim must be positive even".into(),
219 ));
220 }
221 if !x.len().is_multiple_of(head_dim) {
222 return Err(EngineError::ShapeMismatch(
223 "rope x len not divisible by head_dim".into(),
224 ));
225 }
226 let n_heads = x.len() / head_dim;
227 for h in 0..n_heads {
228 let base = h * head_dim;
229 for i in 0..(head_dim / 2) {
230 let freq = 1.0 / theta.powf((2 * i) as f32 / head_dim as f32);
231 let angle = pos as f32 * freq;
232 let (c, s) = (angle.cos(), angle.sin());
233 let u = x[base + 2 * i];
234 let v = x[base + 2 * i + 1];
235 x[base + 2 * i] = u * c - v * s;
236 x[base + 2 * i + 1] = u * s + v * c;
237 }
238 }
239 Ok(())
240}
241
242pub fn attention(
245 q: &[f32],
246 k_cache: &[f32],
247 v_cache: &[f32],
248 n_heads: usize,
249 n_kv_heads: usize,
250 head_dim: usize,
251) -> Result<Vec<f32>, EngineError> {
252 let scale = 1.0 / (head_dim as f32).sqrt();
253 attention_with_scale(q, k_cache, v_cache, n_heads, n_kv_heads, head_dim, scale)
254}
255
256fn sliding_kv_start(seq: usize, window: Option<usize>) -> usize {
257 match window {
258 Some(w) if w > 0 => seq.saturating_sub(w),
259 _ => 0,
260 }
261}
262
263pub fn kv_sliding_view<'a>(
265 k_cache: &'a [f32],
266 v_cache: &'a [f32],
267 kv_dim: usize,
268 window: Option<usize>,
269) -> Result<(&'a [f32], &'a [f32]), EngineError> {
270 if kv_dim == 0
271 || k_cache.len() != v_cache.len()
272 || !k_cache.len().is_multiple_of(kv_dim)
273 {
274 return Err(EngineError::ShapeMismatch(
275 "sliding kv view shape".into(),
276 ));
277 }
278 let seq = k_cache.len() / kv_dim;
279 let start = sliding_kv_start(seq, window);
280 Ok((
281 &k_cache[start * kv_dim..],
282 &v_cache[start * kv_dim..],
283 ))
284}
285
286pub fn attention_with_scale(
288 q: &[f32],
289 k_cache: &[f32],
290 v_cache: &[f32],
291 n_heads: usize,
292 n_kv_heads: usize,
293 head_dim: usize,
294 scale: f32,
295) -> Result<Vec<f32>, EngineError> {
296 if n_heads == 0 || head_dim == 0 || n_kv_heads == 0 || !n_heads.is_multiple_of(n_kv_heads) {
297 return Err(EngineError::ShapeMismatch(
298 "attention invalid head configuration".into(),
299 ));
300 }
301 let kv_dim = n_kv_heads * head_dim;
302 if q.len() != n_heads * head_dim {
303 return Err(EngineError::ShapeMismatch("attention q shape".into()));
304 }
305 if k_cache.len() != v_cache.len() || !k_cache.len().is_multiple_of(kv_dim) {
306 return Err(EngineError::ShapeMismatch(
307 "attention kv cache shape".into(),
308 ));
309 }
310 let seq = k_cache.len() / kv_dim;
311 let rep = n_heads / n_kv_heads;
312 let mut out = vec![0.0f32; n_heads * head_dim];
313 for h in 0..n_heads {
314 let kv_h = h / rep;
315 let qh = &q[h * head_dim..(h + 1) * head_dim];
316 let mut scores = vec![0.0f32; seq];
317 for t in 0..seq {
318 let kh = &k_cache[t * kv_dim + kv_h * head_dim..t * kv_dim + (kv_h + 1) * head_dim];
319 let mut dot = 0.0f32;
320 for i in 0..head_dim {
321 dot += qh[i] * kh[i];
322 }
323 scores[t] = dot * scale;
324 }
325 softmax_inplace(&mut scores);
326 let oh = &mut out[h * head_dim..(h + 1) * head_dim];
327 for t in 0..seq {
328 let vh = &v_cache[t * kv_dim + kv_h * head_dim..t * kv_dim + (kv_h + 1) * head_dim];
329 for i in 0..head_dim {
330 oh[i] += scores[t] * vh[i];
331 }
332 }
333 }
334 Ok(out)
335}
336
337pub fn attention_causal(
341 q: &[f32],
342 k_cache: &[f32],
343 v_cache: &[f32],
344 n_heads: usize,
345 n_kv_heads: usize,
346 head_dim: usize,
347) -> Result<Vec<f32>, EngineError> {
348 let q_dim = n_heads * head_dim;
349 if q_dim == 0 || !q.len().is_multiple_of(q_dim) {
350 return Err(EngineError::ShapeMismatch("attention_causal q shape".into()));
351 }
352 let seq_q = q.len() / q_dim;
353 if seq_q == 1 {
354 return attention(q, k_cache, v_cache, n_heads, n_kv_heads, head_dim);
355 }
356 let kv_dim = n_kv_heads * head_dim;
357 let seq_kv = k_cache.len() / kv_dim;
358 let causal = seq_q == seq_kv;
359 let mut out = vec![0.0f32; q.len()];
360 for tq in 0..seq_q {
361 let q_tok = &q[tq * q_dim..(tq + 1) * q_dim];
362 let k_end = if causal { tq + 1 } else { seq_kv };
363 let attn = attention(
364 q_tok,
365 &k_cache[..k_end * kv_dim],
366 &v_cache[..k_end * kv_dim],
367 n_heads,
368 n_kv_heads,
369 head_dim,
370 )?;
371 out[tq * q_dim..(tq + 1) * q_dim].copy_from_slice(&attn);
372 }
373 Ok(out)
374}
375
376#[allow(clippy::too_many_arguments)]
379pub fn attention_causal_with_scale(
380 q: &[f32],
381 k_cache: &[f32],
382 v_cache: &[f32],
383 n_heads: usize,
384 n_kv_heads: usize,
385 head_dim: usize,
386 scale: f32,
387 window: Option<usize>,
388) -> Result<Vec<f32>, EngineError> {
389 let q_dim = n_heads * head_dim;
390 if q_dim == 0 || !q.len().is_multiple_of(q_dim) {
391 return Err(EngineError::ShapeMismatch("attention_causal q shape".into()));
392 }
393 let seq_q = q.len() / q_dim;
394 let kv_dim = n_kv_heads * head_dim;
395 if seq_q == 1 {
396 let (k, v) = kv_sliding_view(k_cache, v_cache, kv_dim, window)?;
397 return attention_with_scale(q, k, v, n_heads, n_kv_heads, head_dim, scale);
398 }
399 let seq_kv = k_cache.len() / kv_dim;
400 let causal = seq_q == seq_kv;
401 let mut out = vec![0.0f32; q.len()];
402 for tq in 0..seq_q {
403 let q_tok = &q[tq * q_dim..(tq + 1) * q_dim];
404 let k_end = if causal { tq + 1 } else { seq_kv };
405 let (k, v) = kv_sliding_view(
406 &k_cache[..k_end * kv_dim],
407 &v_cache[..k_end * kv_dim],
408 kv_dim,
409 window,
410 )?;
411 let attn = attention_with_scale(q_tok, k, v, n_heads, n_kv_heads, head_dim, scale)?;
412 out[tq * q_dim..(tq + 1) * q_dim].copy_from_slice(&attn);
413 }
414 Ok(out)
415}
416
417pub fn swiglu(gate: &[f32], up: &[f32]) -> Result<Vec<f32>, EngineError> {
418 if gate.len() != up.len() {
419 return Err(EngineError::ShapeMismatch("swiglu length mismatch".into()));
420 }
421 Ok(gate
422 .iter()
423 .zip(up.iter())
424 .map(|(g, u)| {
425 let s = 1.0 / (1.0 + (-g).exp());
426 s * g * u
427 })
428 .collect())
429}
430
431pub fn short_conv_step(
436 x: &[f32],
437 w: &[f32],
438 state: &mut [f32],
439 hidden: usize,
440 kernel: usize,
441) -> Result<Vec<f32>, EngineError> {
442 if hidden == 0 || kernel == 0 {
443 return Err(EngineError::ShapeMismatch(
444 "short_conv_step: hidden and kernel must be > 0".into(),
445 ));
446 }
447 if x.len() != hidden {
448 return Err(EngineError::ShapeMismatch(
449 "short_conv_step: x len != hidden".into(),
450 ));
451 }
452 if w.len() != hidden * kernel {
453 return Err(EngineError::ShapeMismatch(
454 "short_conv_step: weight len != hidden*kernel".into(),
455 ));
456 }
457 let hist = kernel.saturating_sub(1);
458 if state.len() != hidden * hist {
459 return Err(EngineError::ShapeMismatch(
460 "short_conv_step: state len != hidden*(kernel-1)".into(),
461 ));
462 }
463 let mut out = vec![0.0f32; hidden];
464 for c in 0..hidden {
465 let mut acc = 0.0f32;
466 let wbase = c * kernel;
467 let sbase = c * hist;
468 for k in 0..hist {
469 acc += w[wbase + k] * state[sbase + k];
470 }
471 acc += w[wbase + hist] * x[c];
472 out[c] = acc;
473 if hist > 0 {
474 for k in 0..(hist - 1) {
475 state[sbase + k] = state[sbase + k + 1];
476 }
477 state[sbase + hist - 1] = x[c];
478 }
479 }
480 Ok(out)
481}
482
483pub fn moe_topk_route(
485 logits: &[f32],
486 top_k: usize,
487 use_sigmoid: bool,
488) -> Result<(Vec<usize>, Vec<f32>), EngineError> {
489 let n = logits.len();
490 if n == 0 || top_k == 0 {
491 return Err(EngineError::InvalidParam(
492 "moe_topk_route: num_experts and top_k must be > 0".into(),
493 ));
494 }
495 let k = top_k.min(n);
496 let scores: Vec<f32> = if use_sigmoid {
497 logits.iter().map(|x| 1.0 / (1.0 + (-x).exp())).collect()
498 } else {
499 let m = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
500 let mut exps: Vec<f32> = logits.iter().map(|x| (x - m).exp()).collect();
501 let s: f32 = exps.iter().sum();
502 if s > 0.0 {
503 for e in &mut exps {
504 *e /= s;
505 }
506 }
507 exps
508 };
509 let mut idx: Vec<usize> = (0..n).collect();
510 idx.sort_by(|&a, &b| {
511 scores[b]
512 .partial_cmp(&scores[a])
513 .unwrap_or(std::cmp::Ordering::Equal)
514 });
515 idx.truncate(k);
516 let mut weights: Vec<f32> = idx.iter().map(|&i| scores[i]).collect();
517 let sum: f32 = weights.iter().sum();
518 if sum > 0.0 {
519 for w in &mut weights {
520 *w /= sum;
521 }
522 }
523 Ok((idx, weights))
524}
525
526fn silu(x: f32) -> f32 {
527 x / (1.0 + (-x).exp())
528}
529
530pub fn silu_vec(x: &mut [f32]) {
532 for v in x.iter_mut() {
533 *v = silu(*v);
534 }
535}
536
537pub fn softplus(x: f32) -> f32 {
539 if x > 20.0 {
540 x
541 } else {
542 (1.0 + x.exp()).ln()
543 }
544}
545
546fn l2_normalize_inplace(x: &mut [f32]) {
547 let mut ss = 0.0f32;
548 for v in x.iter() {
549 ss += *v * *v;
550 }
551 let n = (ss + 1e-6).sqrt();
552 if n > 0.0 {
553 for v in x.iter_mut() {
554 *v /= n;
555 }
556 }
557}
558
559pub struct GatedDeltaStep<'a> {
561 pub q: &'a [f32],
562 pub k: &'a [f32],
563 pub v: &'a [f32],
564 pub g: &'a [f32],
565 pub beta: &'a [f32],
566 pub state: &'a mut [f32],
567 pub n_heads: usize,
568 pub dk: usize,
569 pub dv: usize,
570}
571
572pub fn gated_delta_step(p: GatedDeltaStep<'_>) -> Result<Vec<f32>, EngineError> {
577 let GatedDeltaStep {
578 q,
579 k,
580 v,
581 g,
582 beta,
583 state,
584 n_heads,
585 dk,
586 dv,
587 } = p;
588 if n_heads == 0 || dk == 0 || dv == 0 {
589 return Err(EngineError::ShapeMismatch(
590 "gated_delta_step: heads/dk/dv must be > 0".into(),
591 ));
592 }
593 if q.len() != n_heads * dk
594 || k.len() != n_heads * dk
595 || v.len() != n_heads * dv
596 || g.len() != n_heads
597 || beta.len() != n_heads
598 || state.len() != n_heads * dk * dv
599 {
600 return Err(EngineError::ShapeMismatch(
601 "gated_delta_step: q/k/v/g/beta/state shape mismatch".into(),
602 ));
603 }
604 let mut qq = q.to_vec();
605 let mut kk = k.to_vec();
606 for h in 0..n_heads {
607 l2_normalize_inplace(&mut qq[h * dk..(h + 1) * dk]);
608 l2_normalize_inplace(&mut kk[h * dk..(h + 1) * dk]);
609 }
610 let scale = (dk as f32).sqrt().recip();
611 let mut out = vec![0.0f32; n_heads * dv];
612 for h in 0..n_heads {
613 let sbase = h * dk * dv;
614 let gh = g[h];
615 for i in 0..dk * dv {
616 state[sbase + i] *= gh;
617 }
618 let mut kv_mem = vec![0.0f32; dv];
619 for i in 0..dk {
620 let kv = kk[h * dk + i];
621 for j in 0..dv {
622 kv_mem[j] += state[sbase + i * dv + j] * kv;
623 }
624 }
625 for j in 0..dv {
626 let delta = (v[h * dv + j] - kv_mem[j]) * beta[h];
627 for i in 0..dk {
628 state[sbase + i * dv + j] += kk[h * dk + i] * delta;
629 }
630 }
631 for j in 0..dv {
632 let mut o = 0.0f32;
633 for i in 0..dk {
634 o += state[sbase + i * dv + j] * qq[h * dk + i];
635 }
636 out[h * dv + j] = o * scale;
637 }
638 }
639 Ok(out)
640}
641
642pub fn geglu(gate: &[f32], up: &[f32]) -> Result<Vec<f32>, EngineError> {
644 if gate.len() != up.len() {
645 return Err(EngineError::ShapeMismatch("geglu length mismatch".into()));
646 }
647 Ok(gate
648 .iter()
649 .zip(up.iter())
650 .map(|(g, u)| gelu_pytorch_tanh(*g) * u)
651 .collect())
652}
653
654pub fn gelu_pytorch_tanh(x: f32) -> f32 {
656 const SQRT_2_OVER_PI: f32 = 0.797_884_6;
658 const COEFF: f32 = 0.044_715;
659 let inner = SQRT_2_OVER_PI * (x + COEFF * x * x * x);
660 0.5 * x * (1.0 + inner.tanh())
661}
662
663pub fn rms_norm_gemma(x: &[f32], weight: &[f32], eps: f32) -> Result<Vec<f32>, EngineError> {
665 let n = weight.len();
666 if n == 0 || !x.len().is_multiple_of(n) {
667 return Err(EngineError::ShapeMismatch(
668 "rms_norm_gemma: x length must be multiple of weight len".into(),
669 ));
670 }
671 let batch = x.len() / n;
672 let mut out = vec![0.0f32; x.len()];
673 for b in 0..batch {
674 let base = b * n;
675 let mut ms = 0.0f32;
676 for i in 0..n {
677 ms += x[base + i] * x[base + i];
678 }
679 let rrms = (ms / n as f32 + eps).sqrt().recip();
680 for i in 0..n {
681 out[base + i] = x[base + i] * rrms * (1.0 + weight[i]);
682 }
683 }
684 Ok(out)
685}
686
687pub fn rope_half(
689 x: &mut [f32],
690 head_dim: usize,
691 pos: usize,
692 theta: f32,
693) -> Result<(), EngineError> {
694 if head_dim == 0 || !head_dim.is_multiple_of(2) {
695 return Err(EngineError::ShapeMismatch(
696 "rope_half head_dim must be positive even".into(),
697 ));
698 }
699 if !x.len().is_multiple_of(head_dim) {
700 return Err(EngineError::ShapeMismatch(
701 "rope_half x len not divisible by head_dim".into(),
702 ));
703 }
704 let half = head_dim / 2;
705 let n_heads = x.len() / head_dim;
706 for h in 0..n_heads {
707 let base = h * head_dim;
708 for i in 0..half {
709 let freq = 1.0 / theta.powf((2 * i) as f32 / head_dim as f32);
710 let angle = pos as f32 * freq;
711 let (c, s) = (angle.cos(), angle.sin());
712 let u = x[base + i];
713 let v = x[base + i + half];
714 x[base + i] = u * c - v * s;
715 x[base + i + half] = u * s + v * c;
716 }
717 }
718 Ok(())
719}
720
721pub fn rope_half_partial(
723 x: &mut [f32],
724 head_dim: usize,
725 rotary_dim: usize,
726 pos: usize,
727 theta: f32,
728) -> Result<(), EngineError> {
729 if rotary_dim == 0 || rotary_dim > head_dim || !rotary_dim.is_multiple_of(2) {
730 return Err(EngineError::ShapeMismatch(
731 "rope_half_partial rotary_dim must be positive even and <= head_dim".into(),
732 ));
733 }
734 if rotary_dim == head_dim {
735 return rope_half(x, head_dim, pos, theta);
736 }
737 if !x.len().is_multiple_of(head_dim) {
738 return Err(EngineError::ShapeMismatch(
739 "rope_half_partial x len not divisible by head_dim".into(),
740 ));
741 }
742 let n_heads = x.len() / head_dim;
743 for h in 0..n_heads {
744 let sl = &mut x[h * head_dim..h * head_dim + rotary_dim];
745 rope_half(sl, rotary_dim, pos, theta)?;
746 }
747 Ok(())
748}
749
750pub fn rope_half_proportional(
754 x: &mut [f32],
755 head_dim: usize,
756 factor: f32,
757 pos: usize,
758 theta: f32,
759) -> Result<(), EngineError> {
760 if !(0.0..=1.0).contains(&factor) {
761 return Err(EngineError::ShapeMismatch(
762 "rope_half_proportional factor must be in [0, 1]".into(),
763 ));
764 }
765 if (factor - 1.0).abs() < 1e-6 {
766 return rope_half(x, head_dim, pos, theta);
767 }
768 if head_dim == 0 || !head_dim.is_multiple_of(2) {
769 return Err(EngineError::ShapeMismatch(
770 "rope_half_proportional head_dim must be positive even".into(),
771 ));
772 }
773 if !x.len().is_multiple_of(head_dim) {
774 return Err(EngineError::ShapeMismatch(
775 "rope_half_proportional x len not divisible by head_dim".into(),
776 ));
777 }
778 let half = head_dim / 2;
779 let rope_angles = (factor * head_dim as f32 / 2.0) as usize;
780 if rope_angles == 0 {
781 return Ok(());
782 }
783 let n_heads = x.len() / head_dim;
784 for h in 0..n_heads {
785 let base = h * head_dim;
786 for i in 0..rope_angles.min(half) {
787 let freq = 1.0 / theta.powf((2 * i) as f32 / head_dim as f32);
788 let angle = pos as f32 * freq;
789 let (c, s) = (angle.cos(), angle.sin());
790 let u = x[base + i];
791 let v = x[base + i + half];
792 x[base + i] = u * c - v * s;
793 x[base + i + half] = u * s + v * c;
794 }
795 }
796 Ok(())
797}
798
799pub fn hdm_linear(
805 x: &[f32],
806 w_rot: &[f32],
807 out_f: usize,
808 in_f: usize,
809 hadamard_seed: Option<i64>,
810) -> Result<Vec<f32>, EngineError> {
811 let mut y = linear(x, w_rot, out_f, in_f)?;
812 if out_f == 0 || !y.len().is_multiple_of(out_f) {
813 return Err(EngineError::ShapeMismatch(
814 "hdm_linear output not divisible by out_f".into(),
815 ));
816 }
817 let batch = y.len() / out_f;
818 for b in 0..batch {
819 let sl = b * out_f..(b + 1) * out_f;
820 hadamard_blocked_vec(&mut y[sl], hadamard_seed, true)?;
821 }
822 Ok(y)
823}
824
825pub fn fwht(x: &mut [f32]) -> Result<(), EngineError> {
827 let n = x.len();
828 if n == 0 || !n.is_power_of_two() {
829 return Err(EngineError::ShapeMismatch(
830 "fwht length must be power of two".into(),
831 ));
832 }
833 if n == 1 {
834 return Ok(());
835 }
836 let mut h = 1usize;
837 while h < n {
838 for i in (0..n).step_by(h * 2) {
839 for j in i..(i + h) {
840 let a = x[j];
841 let b = x[j + h];
842 x[j] = a + b;
843 x[j + h] = a - b;
844 }
845 }
846 h *= 2;
847 }
848 let scale = 1.0 / (n as f32).sqrt();
849 for v in x.iter_mut() {
850 *v *= scale;
851 }
852 Ok(())
853}
854
855pub fn pow2_tile_sizes(k: usize) -> Result<Vec<usize>, EngineError> {
857 if k == 0 {
858 return Err(EngineError::ShapeMismatch(
859 "pow2_tile_sizes expects k>=1".into(),
860 ));
861 }
862 let mut sizes = Vec::new();
863 let mut rem = k;
864 while rem > 0 {
865 let mut b = 1usize;
866 while (b << 1) <= rem {
867 b <<= 1;
868 }
869 sizes.push(b);
870 rem -= b;
871 }
872 Ok(sizes)
873}
874
875pub fn portable_block_signs(seed: i64, start: usize, size: usize) -> Vec<f32> {
877 let mut signs = vec![0.0f32; size];
878 let mut state = (seed as u64) ^ ((start as u64).wrapping_mul(0x9E3779B97F4A7C15));
879 for s in signs.iter_mut() {
880 state = state.wrapping_add(0x9E3779B97F4A7C15);
881 let mut z = state;
882 z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
883 z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
884 z ^= z >> 31;
885 *s = if (z & 1) == 0 { 1.0 } else { -1.0 };
886 }
887 signs
888}
889
890pub fn hadamard_blocked_rows(
893 data: &mut [f32],
894 rows: usize,
895 cols: usize,
896 seed: Option<i64>,
897 inverse: bool,
898) -> Result<(), EngineError> {
899 if rows == 0 || cols == 0 || data.len() != rows * cols {
900 return Err(EngineError::ShapeMismatch(
901 "hadamard_blocked_rows shape mismatch".into(),
902 ));
903 }
904 let sizes = pow2_tile_sizes(rows)?;
905 let mut start = 0usize;
906 for &sz in &sizes {
907 let signs = seed.map(|s| portable_block_signs(s, start, sz));
908 let mut work = vec![0.0f32; sz * cols];
910 for r in 0..sz {
911 let src = (start + r) * cols;
912 work[r * cols..(r + 1) * cols].copy_from_slice(&data[src..src + cols]);
913 }
914 let col_out: Result<Vec<Vec<f32>>, EngineError> = (0..cols)
915 .into_par_iter()
916 .map(|c| {
917 let mut colbuf = vec![0.0f32; sz];
918 for r in 0..sz {
919 colbuf[r] = work[r * cols + c];
920 }
921 if let Some(ref sg) = signs {
922 if !inverse {
923 for r in 0..sz {
924 colbuf[r] *= sg[r];
925 }
926 fwht(&mut colbuf)?;
927 } else {
928 fwht(&mut colbuf)?;
929 for r in 0..sz {
930 colbuf[r] *= sg[r];
931 }
932 }
933 } else if sz > 1 {
934 fwht(&mut colbuf)?;
935 }
936 Ok(colbuf)
937 })
938 .collect();
939 let col_out = col_out?;
940 for c in 0..cols {
941 for r in 0..sz {
942 work[r * cols + c] = col_out[c][r];
943 }
944 }
945 for r in 0..sz {
946 let dst = (start + r) * cols;
947 data[dst..dst + cols].copy_from_slice(&work[r * cols..(r + 1) * cols]);
948 }
949 start += sz;
950 }
951 Ok(())
952}
953
954pub fn hadamard_blocked_vec(
956 x: &mut [f32],
957 seed: Option<i64>,
958 inverse: bool,
959) -> Result<(), EngineError> {
960 let rows = x.len();
961 hadamard_blocked_rows(x, rows, 1, seed, inverse)
962}
963
964pub fn dequant_lookup_group(
966 indices: &[u8],
967 codebook: &[f32],
968 num_groups: usize,
969 group_size: usize,
970 n: usize,
971 kc: usize,
972 k0: usize,
973) -> Result<Vec<f32>, EngineError> {
974 let k_work = num_groups * group_size;
975 if indices.len() != k_work * n {
976 return Err(EngineError::ShapeMismatch(
977 "dequant indices length mismatch".into(),
978 ));
979 }
980 if codebook.len() != num_groups * kc {
981 return Err(EngineError::Quant(
982 "dequant codebook length mismatch".into(),
983 ));
984 }
985 let mut out = vec![0.0f32; k_work * n];
986 for g in 0..num_groups {
987 let cb = &codebook[g * kc..(g + 1) * kc];
988 for r in 0..group_size {
989 let row = g * group_size + r;
990 for j in 0..n {
991 let idx = indices[row * n + j] as usize;
992 if idx >= kc {
993 return Err(EngineError::Quant(format!("index {idx} >= kc {kc}")));
994 }
995 out[row * n + j] = cb[idx];
996 }
997 }
998 }
999 out.truncate(k0 * n);
1000 Ok(out)
1001}
1002
1003pub fn matmul_blocked(
1005 a: &[f32],
1006 a_rows: usize,
1007 a_cols: usize,
1008 b: &[f32],
1009 b_rows: usize,
1010 b_cols: usize,
1011 block: usize,
1012) -> Result<Vec<f32>, EngineError> {
1013 if a_cols != b_rows {
1014 return Err(EngineError::ShapeMismatch(format!(
1015 "matmul inner dim {a_cols} != {b_rows}"
1016 )));
1017 }
1018 if a.len() != a_rows * a_cols || b.len() != b_rows * b_cols {
1019 return Err(EngineError::ShapeMismatch(
1020 "matmul buffer length does not match shape".into(),
1021 ));
1022 }
1023 let block = block.max(1);
1024 let mut out = vec![0.0f32; a_rows * b_cols];
1025 for i0 in (0..a_rows).step_by(block) {
1026 for j0 in (0..b_cols).step_by(block) {
1027 for k0 in (0..a_cols).step_by(block) {
1028 let i_max = (i0 + block).min(a_rows);
1029 let j_max = (j0 + block).min(b_cols);
1030 let k_max = (k0 + block).min(a_cols);
1031 for i in i0..i_max {
1032 for j in j0..j_max {
1033 let mut s = out[i * b_cols + j];
1034 for k in k0..k_max {
1035 s += a[i * a_cols + k] * b[k * b_cols + j];
1036 }
1037 out[i * b_cols + j] = s;
1038 }
1039 }
1040 }
1041 }
1042 }
1043 Ok(out)
1044}
1045
1046pub fn matmul_dispatch(
1049 a: &[f32],
1050 a_rows: usize,
1051 a_cols: usize,
1052 b: &[f32],
1053 b_rows: usize,
1054 b_cols: usize,
1055 mode: SimdMode,
1056) -> Result<Vec<f32>, EngineError> {
1057 match mode {
1058 SimdMode::Scalar => matmul(a, a_rows, a_cols, b, b_rows, b_cols, SimdMode::Scalar),
1059 SimdMode::Neon | SimdMode::Avx2 => matmul_blocked(a, a_rows, a_cols, b, b_rows, b_cols, 8),
1060 }
1061}
1062
1063#[cfg(test)]
1064mod tests {
1065 use super::*;
1066
1067 #[test]
1068 fn matmul_ok() {
1069 let a = [1.0f32, 2.0, 3.0, 4.0]; let b = [1.0f32, 0.0, 0.0, 1.0];
1071 let c = matmul(&a, 2, 2, &b, 2, 2, SimdMode::Scalar).unwrap();
1072 assert_eq!(c, vec![1.0, 2.0, 3.0, 4.0]);
1073 }
1074
1075 #[test]
1076 fn matmul_shape_err() {
1077 let err = matmul(&[1.0], 1, 1, &[1.0, 2.0], 2, 1, SimdMode::Scalar).unwrap_err();
1078 assert!(matches!(err, EngineError::ShapeMismatch(_)));
1079 }
1080
1081 #[test]
1082 fn rms_and_softmax() {
1083 let w = [1.0f32, 1.0];
1084 let y = rms_norm(&[3.0, 4.0], &w, 1e-6).unwrap();
1085 let rms = (12.5f32).sqrt();
1087 assert!((y[0] - 3.0 / rms).abs() < 1e-4);
1088 let s = softmax(&[1.0, 2.0, 3.0]);
1089 let sum: f32 = s.iter().sum();
1090 assert!((sum - 1.0).abs() < 1e-5);
1091 }
1092
1093 #[test]
1094 fn fwht_roundtrip_ish() {
1095 let mut x = [1.0f32, 2.0, 3.0, 4.0];
1096 let orig = x;
1097 fwht(&mut x).unwrap();
1098 fwht(&mut x).unwrap();
1099 for (a, b) in x.iter().zip(orig.iter()) {
1100 assert!((a - b).abs() < 1e-4);
1101 }
1102 }
1103
1104 #[test]
1105 fn pow2_tiles() {
1106 assert_eq!(pow2_tile_sizes(10).unwrap(), vec![8, 2]);
1107 assert_eq!(pow2_tile_sizes(3072).unwrap(), vec![2048, 1024]);
1108 assert_eq!(pow2_tile_sizes(64).unwrap(), vec![64]);
1109 assert_eq!(pow2_tile_sizes(1).unwrap(), vec![1]);
1110 assert_eq!(pow2_tile_sizes(151936).unwrap()[0], 131072);
1111 assert!(matches!(
1112 pow2_tile_sizes(0),
1113 Err(EngineError::ShapeMismatch(_))
1114 ));
1115 }
1116
1117 #[test]
1118 fn blocked_roundtrip_non_pow2() {
1119 let rows = 10usize;
1120 let cols = 3usize;
1121 let mut w: Vec<f32> = (0..rows * cols).map(|i| (i as f32) * 0.1 - 0.5).collect();
1122 let orig = w.clone();
1123 hadamard_blocked_rows(&mut w, rows, cols, Some(7), false).unwrap();
1124 let mut twice = w.clone();
1126 hadamard_blocked_rows(&mut twice, rows, cols, Some(7), false).unwrap();
1127 let mut err_wrong = 0.0f32;
1128 for (a, b) in twice.iter().zip(orig.iter()) {
1129 err_wrong += (a - b).abs();
1130 }
1131 assert!(err_wrong > 1.0, "second forward unexpectedly near identity");
1132 hadamard_blocked_rows(&mut w, rows, cols, Some(7), true).unwrap();
1133 for (a, b) in w.iter().zip(orig.iter()) {
1134 assert!((a - b).abs() < 1e-4, "{a} vs {b}");
1135 }
1136 }
1137
1138 #[test]
1139 fn blocked_roundtrip_unsigned_and_pow2() {
1140 let rows = 16usize;
1141 let cols = 5usize;
1142 let mut w: Vec<f32> = (0..rows * cols).map(|i| (i as f32) * 0.03).collect();
1143 let orig = w.clone();
1144 hadamard_blocked_rows(&mut w, rows, cols, None, false).unwrap();
1145 hadamard_blocked_rows(&mut w, rows, cols, None, true).unwrap();
1146 for (a, b) in w.iter().zip(orig.iter()) {
1147 assert!((a - b).abs() < 1e-4);
1148 }
1149 }
1150
1151 #[test]
1152 fn blocked_matches_python_golden() {
1153 let rows = 10usize;
1155 let cols = 3usize;
1156 let mut w: Vec<f32> = (0..rows * cols).map(|i| (i as f32) * 0.1 - 0.5).collect();
1157 hadamard_blocked_rows(&mut w, rows, cols, Some(7), false).unwrap();
1158 let golden: [f32; 30] = [
1160 0.919239,
1161 0.989949,
1162 1.060_66,
1163 0.919239,
1164 0.989950,
1165 1.060_66,
1166 -0.919239,
1167 -0.989949,
1168 -1.060_66,
1169 0.777817,
1170 std::f32::consts::FRAC_1_SQRT_2,
1171 0.636396,
1172 -0.919239,
1173 -0.989949,
1174 -1.060_66,
1175 -0.070711,
1176 -0.141421,
1177 -0.212132,
1178 1.343503,
1179 std::f32::consts::SQRT_2,
1180 1.484924,
1181 -0.636396,
1182 -0.848528,
1183 -1.060_66,
1184 -2.899138,
1185 -3.040559,
1186 -3.181_98,
1187 0.212132,
1188 0.212132,
1189 0.212132,
1190 ];
1191 assert_eq!(w.len(), golden.len());
1192 for (a, b) in w.iter().zip(golden.iter()) {
1193 assert!((a - b).abs() < 1e-4, "{a} vs {b}");
1194 }
1195 assert_eq!(
1196 portable_block_signs(7, 0, 8),
1197 vec![-1.0, 1.0, 1.0, -1.0, 1.0, -1.0, 1.0, 1.0]
1198 );
1199 assert_eq!(portable_block_signs(7, 8, 2), vec![-1.0, -1.0]);
1200 }
1201
1202 #[test]
1203 fn portable_signs_stable() {
1204 let a = portable_block_signs(0, 0, 8);
1205 let b = portable_block_signs(0, 0, 8);
1206 assert_eq!(a, b);
1207 let golden = [-1.0f32, 1.0, -1.0, 1.0, -1.0, 1.0, -1.0, 1.0];
1209 assert_eq!(a, golden);
1210 }
1211
1212 #[test]
1213 fn hadamard_blocked_shape_errors() {
1214 let mut w = [1.0f32, 2.0];
1215 assert!(matches!(
1216 hadamard_blocked_rows(&mut w, 2, 2, None, false),
1217 Err(EngineError::ShapeMismatch(_))
1218 ));
1219 assert!(matches!(
1220 hadamard_blocked_rows(&mut [], 0, 1, None, false),
1221 Err(EngineError::ShapeMismatch(_))
1222 ));
1223 }
1224
1225 #[test]
1226 fn hadamard_blocked_vec_roundtrip() {
1227 let mut x: Vec<f32> = (0..10).map(|i| (i as f32) * 0.2 - 1.0).collect();
1228 let orig = x.clone();
1229 hadamard_blocked_vec(&mut x, Some(11), false).unwrap();
1230 hadamard_blocked_vec(&mut x, Some(11), true).unwrap();
1231 for (a, b) in x.iter().zip(orig.iter()) {
1232 assert!((a - b).abs() < 1e-4);
1233 }
1234 }
1235
1236 #[test]
1237 fn dequant_group() {
1238 let indices = [0u8, 1, 1, 0];
1240 let codebook = [10.0f32, 20.0];
1241 let out = dequant_lookup_group(&indices, &codebook, 1, 2, 2, 2, 2).unwrap();
1242 assert_eq!(out, vec![10.0, 20.0, 20.0, 10.0]);
1243 }
1244
1245 #[test]
1246 fn neon_scalar_matmul_parity() {
1247 let a: Vec<f32> = (0..64).map(|i| (i as f32) * 0.01).collect();
1248 let b: Vec<f32> = (0..64).map(|i| (i as f32) * 0.02 - 0.5).collect();
1249 let s = matmul_dispatch(&a, 8, 8, &b, 8, 8, SimdMode::Scalar).unwrap();
1250 let n = matmul_dispatch(&a, 8, 8, &b, 8, 8, SimdMode::Neon).unwrap();
1251 assert_eq!(s.len(), n.len());
1252 for (x, y) in s.iter().zip(n.iter()) {
1253 assert!((x - y).abs() < 1e-5, "{x} vs {y}");
1254 }
1255 }
1256
1257 #[test]
1258 fn short_conv_step_and_moe_route() {
1259 let hidden = 2usize;
1260 let kernel = 3usize;
1261 let w = vec![0.0f32, 0.0, 1.0, 0.0, 0.0, 1.0];
1263 let mut state = vec![0.0f32; hidden * (kernel - 1)];
1264 let x = [3.0f32, 5.0];
1265 let y = short_conv_step(&x, &w, &mut state, hidden, kernel).unwrap();
1266 assert!((y[0] - 3.0).abs() < 1e-5);
1267 assert!((y[1] - 5.0).abs() < 1e-5);
1268 let w_old = vec![1.0f32, 0.0, 0.0, 1.0, 0.0, 0.0];
1270 let y2 = short_conv_step(&[1.0, 2.0], &w_old, &mut state, hidden, kernel).unwrap();
1271 assert!((y2[0] - 0.0).abs() < 1e-5);
1272 assert!((y2[1] - 0.0).abs() < 1e-5);
1273
1274 let (ids, ws) = moe_topk_route(&[0.1, 2.0, 0.5, -1.0], 2, false).unwrap();
1275 assert_eq!(ids.len(), 2);
1276 assert_eq!(ids[0], 1);
1277 assert!((ws.iter().sum::<f32>() - 1.0).abs() < 1e-5);
1278 let (ids_s, _) = moe_topk_route(&[0.0, 10.0, 0.0], 1, true).unwrap();
1279 assert_eq!(ids_s, vec![1]);
1280 }
1281
1282 #[test]
1283 fn gated_delta_step_updates_state() {
1284 let n_heads = 1usize;
1285 let dk = 2usize;
1286 let dv = 2usize;
1287 let q = [1.0f32, 0.0];
1288 let k = [1.0f32, 0.0];
1289 let v = [0.5f32, -0.25];
1290 let g = [0.9f32];
1291 let beta = [1.0f32];
1292 let mut s = vec![0.0f32; n_heads * dk * dv];
1293 let o1 = gated_delta_step(GatedDeltaStep {
1294 q: &q,
1295 k: &k,
1296 v: &v,
1297 g: &g,
1298 beta: &beta,
1299 state: &mut s,
1300 n_heads,
1301 dk,
1302 dv,
1303 })
1304 .unwrap();
1305 assert_eq!(o1.len(), dv);
1306 let s_after = s.clone();
1307 let o2 = gated_delta_step(GatedDeltaStep {
1308 q: &q,
1309 k: &k,
1310 v: &v,
1311 g: &g,
1312 beta: &beta,
1313 state: &mut s,
1314 n_heads,
1315 dk,
1316 dv,
1317 })
1318 .unwrap();
1319 assert!(s != s_after || (o2[0] - o1[0]).abs() > 0.0);
1320 }
1321
1322 #[test]
1323 fn geglu_and_rms_norm_gemma() {
1324 let gate = [0.5f32, -1.0];
1325 let up = [2.0f32, 3.0];
1326 let y = geglu(&gate, &up).unwrap();
1327 assert!((y[0] - gelu_pytorch_tanh(0.5) * 2.0).abs() < 1e-5);
1328 assert!((y[1] - gelu_pytorch_tanh(-1.0) * 3.0).abs() < 1e-5);
1329 let x = [1.0f32, -1.0, 2.0, 0.0];
1330 let w = [0.1f32, 0.2];
1331 let n = rms_norm_gemma(&x, &w, 1e-6).unwrap();
1332 assert_eq!(n.len(), 4);
1333 let plain = rms_norm(&x, &w, 1e-6).unwrap();
1335 assert!((n[0] - plain[0]).abs() > 1e-4 || (n[1] - plain[1]).abs() > 1e-4);
1336 }
1337
1338 #[test]
1339 fn rope_half_layout() {
1340 let mut x = [1.0f32, 2.0, 3.0, 4.0];
1341 rope_half(&mut x, 4, 1, 10000.0).unwrap();
1342 let freq0 = 1.0 / 10000f32.powf(0.0);
1344 let (c0, s0) = (freq0.cos(), freq0.sin());
1345 let freq1 = 1.0 / 10000f32.powf(2.0 / 4.0);
1346 let (c1, s1) = (freq1.cos(), freq1.sin());
1347 assert!((x[0] - (1.0 * c0 - 3.0 * s0)).abs() < 1e-5);
1348 assert!((x[2] - (1.0 * s0 + 3.0 * c0)).abs() < 1e-5);
1349 assert!((x[1] - (2.0 * c1 - 4.0 * s1)).abs() < 1e-5);
1350 assert!((x[3] - (2.0 * s1 + 4.0 * c1)).abs() < 1e-5);
1351 assert!(matches!(
1352 rope_half(&mut [1.0, 2.0, 3.0], 3, 0, 10000.0),
1353 Err(EngineError::ShapeMismatch(_))
1354 ));
1355 }
1356
1357 #[test]
1358 fn rope_half_partial_full_matches_rope_half() {
1359 let mut a = [1.0f32, 2.0, 3.0, 4.0];
1360 let mut b = a;
1361 rope_half(&mut a, 4, 2, 10000.0).unwrap();
1362 rope_half_partial(&mut b, 4, 4, 2, 10000.0).unwrap();
1363 for (x, y) in a.iter().zip(b.iter()) {
1364 assert!((x - y).abs() < 1e-6);
1365 }
1366 let mut c = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
1367 let tail = [5.0f32, 6.0, 7.0, 8.0];
1368 rope_half_partial(&mut c, 8, 4, 1, 10000.0).unwrap();
1369 assert_eq!(&c[4..], &tail);
1370 }
1371
1372 #[test]
1373 fn rope_half_proportional_identity_tail() {
1374 let mut full = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
1375 let orig = full;
1376 rope_half(&mut full, 8, 3, 10000.0).unwrap();
1377 let mut prop = orig;
1378 rope_half_proportional(&mut prop, 8, 1.0, 3, 10000.0).unwrap();
1379 for (a, b) in full.iter().zip(prop.iter()) {
1380 assert!((a - b).abs() < 1e-6);
1381 }
1382 let mut half = orig;
1383 rope_half_proportional(&mut half, 8, 0.25, 3, 10000.0).unwrap();
1384 assert!((half[0] - orig[0]).abs() > 1e-6);
1386 assert!((half[4] - orig[4]).abs() > 1e-6);
1387 assert_eq!(&half[1..4], &orig[1..4]);
1388 assert_eq!(&half[5..], &orig[5..]);
1389 }
1390
1391 #[test]
1392 fn hdm_linear_matches_unrotated_weight() {
1393 let out_f = 8usize;
1394 let in_f = 4usize;
1395 let seed = Some(7i64);
1396 let mut w_orig: Vec<f32> = (0..out_f * in_f).map(|i| (i as f32) * 0.05 - 0.2).collect();
1397 let x: Vec<f32> = (0..in_f).map(|i| (i as f32) * 0.1).collect();
1398 let y_ref = linear(&x, &w_orig, out_f, in_f).unwrap();
1399 hadamard_blocked_rows(&mut w_orig, out_f, in_f, seed, false).unwrap();
1400 let y = hdm_linear(&x, &w_orig, out_f, in_f, seed).unwrap();
1401 for (a, b) in y.iter().zip(y_ref.iter()) {
1402 assert!((a - b).abs() < 1e-4, "{a} vs {b}");
1403 }
1404 }
1405
1406 #[test]
1407 fn linear_cpu_matches_scalar_linear() {
1408 let out_f = 7usize;
1409 let in_f = 5usize;
1410 let w: Vec<f32> = (0..out_f * in_f).map(|i| (i as f32) * 0.02 - 0.1).collect();
1411 let x: Vec<f32> = (0..in_f * 3).map(|i| (i as f32) * 0.03).collect();
1412 let a = linear(&x, &w, out_f, in_f).unwrap();
1413 let b = linear_cpu(&x, &w, out_f, in_f).unwrap();
1414 for (x, y) in a.iter().zip(b.iter()) {
1415 assert!((x - y).abs() < 1e-4, "{x} vs {y}");
1416 }
1417 }
1418
1419 #[test]
1420 fn attention_causal_matches_stepwise() {
1421 let n_heads = 2usize;
1422 let n_kv = 1usize;
1423 let head_dim = 4usize;
1424 let seq = 3usize;
1425 let q_dim = n_heads * head_dim;
1426 let kv_dim = n_kv * head_dim;
1427 let q: Vec<f32> = (0..seq * q_dim).map(|i| (i as f32) * 0.01).collect();
1428 let k: Vec<f32> = (0..seq * kv_dim).map(|i| (i as f32) * 0.02).collect();
1429 let v: Vec<f32> = (0..seq * kv_dim).map(|i| (i as f32) * 0.03).collect();
1430 let batched = attention_causal(&q, &k, &v, n_heads, n_kv, head_dim).unwrap();
1431 let mut step = Vec::new();
1432 for t in 0..seq {
1433 let qi = &q[t * q_dim..(t + 1) * q_dim];
1434 let a = attention(
1435 qi,
1436 &k[..(t + 1) * kv_dim],
1437 &v[..(t + 1) * kv_dim],
1438 n_heads,
1439 n_kv,
1440 head_dim,
1441 )
1442 .unwrap();
1443 step.extend_from_slice(&a);
1444 }
1445 for (a, b) in batched.iter().zip(step.iter()) {
1446 assert!((a - b).abs() < 1e-5, "{a} vs {b}");
1447 }
1448 }
1449
1450 #[test]
1451 fn sliding_window_matches_truncated_kv() {
1452 let n_heads = 2usize;
1453 let n_kv = 1usize;
1454 let head_dim = 4usize;
1455 let seq = 6usize;
1456 let window = 2usize;
1457 let q_dim = n_heads * head_dim;
1458 let kv_dim = n_kv * head_dim;
1459 let scale = 1.0 / (head_dim as f32).sqrt();
1460 let q: Vec<f32> = (0..seq * q_dim).map(|i| (i as f32) * 0.01).collect();
1461 let k: Vec<f32> = (0..seq * kv_dim).map(|i| (i as f32) * 0.02).collect();
1462 let v: Vec<f32> = (0..seq * kv_dim).map(|i| (i as f32) * 0.03).collect();
1463
1464 let wide = attention_causal_with_scale(&q, &k, &v, n_heads, n_kv, head_dim, scale, None)
1465 .unwrap();
1466 let win = attention_causal_with_scale(
1467 &q,
1468 &k,
1469 &v,
1470 n_heads,
1471 n_kv,
1472 head_dim,
1473 scale,
1474 Some(window),
1475 )
1476 .unwrap();
1477 for i in 0..window * q_dim {
1479 assert!((wide[i] - win[i]).abs() < 1e-6, "prefix {i}");
1480 }
1481 assert!(
1482 wide.iter()
1483 .zip(win.iter())
1484 .any(|(a, b)| (a - b).abs() > 1e-5),
1485 "window must change scores once seq > window"
1486 );
1487
1488 let last_q = &q[(seq - 1) * q_dim..];
1489 let start = seq - window;
1490 let sliced = attention_with_scale(
1491 last_q,
1492 &k[start * kv_dim..],
1493 &v[start * kv_dim..],
1494 n_heads,
1495 n_kv,
1496 head_dim,
1497 scale,
1498 )
1499 .unwrap();
1500 let (k_win, v_win) = kv_sliding_view(&k, &v, kv_dim, Some(window)).unwrap();
1501 let via_window =
1502 attention_with_scale(last_q, k_win, v_win, n_heads, n_kv, head_dim, scale).unwrap();
1503 for (a, b) in sliced.iter().zip(via_window.iter()) {
1504 assert!((a - b).abs() < 1e-6, "{a} vs {b}");
1505 }
1506 let (k_noop, v_noop) = kv_sliding_view(&k, &v, kv_dim, Some(seq + 8)).unwrap();
1507 let noop =
1508 attention_with_scale(last_q, k_noop, v_noop, n_heads, n_kv, head_dim, scale).unwrap();
1509 let full = attention_with_scale(last_q, &k, &v, n_heads, n_kv, head_dim, scale).unwrap();
1510 for (a, b) in noop.iter().zip(full.iter()) {
1511 assert!((a - b).abs() < 1e-6);
1512 }
1513 }
1514
1515 #[test]
1516 fn embedding_row_gather_needs_full_matrix_unrotate() {
1517 let vocab = 8usize;
1518 let hidden = 4usize;
1519 let seed = Some(7i64);
1520 let tid = 3usize;
1521 let mut w: Vec<f32> = (0..vocab * hidden)
1522 .map(|i| (i as f32) * 0.05 - 0.2)
1523 .collect();
1524 let orig = w[tid * hidden..(tid + 1) * hidden].to_vec();
1525 hadamard_blocked_rows(&mut w, vocab, hidden, seed, false).unwrap();
1526 let rotated_row = &w[tid * hidden..(tid + 1) * hidden];
1527 let drift: f32 = orig
1528 .iter()
1529 .zip(rotated_row.iter())
1530 .map(|(a, b)| (a - b).abs())
1531 .sum();
1532 assert!(
1533 drift > 1e-3,
1534 "axis-0 Hadamard must mix vocab rows; gather of W_rot[token] is wrong"
1535 );
1536 hadamard_blocked_rows(&mut w, vocab, hidden, seed, true).unwrap();
1537 for (a, b) in orig.iter().zip(w[tid * hidden..(tid + 1) * hidden].iter()) {
1538 assert!((a - b).abs() < 1e-4, "{a} vs {b}");
1539 }
1540 }
1541
1542 #[test]
1543 fn shape_errors() {
1544 assert!(matches!(
1545 rms_norm(&[1.0, 2.0], &[1.0, 1.0, 1.0], 1e-6),
1546 Err(EngineError::ShapeMismatch(_))
1547 ));
1548 assert!(matches!(
1549 linear(&[1.0], &[1.0, 2.0], 1, 1),
1550 Err(EngineError::ShapeMismatch(_))
1551 ));
1552 assert!(matches!(
1553 rope(&mut [1.0, 2.0, 3.0], 3, 0, 10000.0),
1554 Err(EngineError::ShapeMismatch(_))
1555 ));
1556 assert!(matches!(
1557 attention(&[1.0], &[1.0], &[1.0], 1, 1, 2),
1558 Err(EngineError::ShapeMismatch(_))
1559 ));
1560 assert!(matches!(
1561 fwht(&mut [1.0, 2.0, 3.0]),
1562 Err(EngineError::ShapeMismatch(_))
1563 ));
1564 }
1565}