1use crate::pool::Pool;
22
23pub trait Fp:
28 Copy
29 + PartialOrd
30 + core::ops::Add<Output = Self>
31 + core::ops::Sub<Output = Self>
32 + core::ops::Mul<Output = Self>
33 + core::ops::Div<Output = Self>
34 + core::ops::Neg<Output = Self>
35 + core::ops::AddAssign
36 + core::ops::MulAssign
37 + Send
38 + Sync
39 + 'static
40{
41 const ZERO: Self;
42 const ONE: Self;
43 fn exp(self) -> Self;
44 fn sqrt(self) -> Self;
45 fn maxf(self, o: Self) -> Self;
46 fn fromf(x: f64) -> Self;
47 fn f64(self) -> f64;
48}
49
50impl Fp for f32 {
51 const ZERO: Self = 0.0;
52 const ONE: Self = 1.0;
53 #[inline]
54 fn exp(self) -> Self {
55 f32::exp(self)
56 }
57 #[inline]
58 fn sqrt(self) -> Self {
59 f32::sqrt(self)
60 }
61 #[inline]
62 fn maxf(self, o: Self) -> Self {
63 f32::max(self, o)
64 }
65 #[inline]
66 fn fromf(x: f64) -> Self {
67 x as f32
68 }
69 #[inline]
70 fn f64(self) -> f64 {
71 self as f64
72 }
73}
74
75impl Fp for f64 {
76 const ZERO: Self = 0.0;
77 const ONE: Self = 1.0;
78 #[inline]
79 fn exp(self) -> Self {
80 f64::exp(self)
81 }
82 #[inline]
83 fn sqrt(self) -> Self {
84 f64::sqrt(self)
85 }
86 #[inline]
87 fn maxf(self, o: Self) -> Self {
88 f64::max(self, o)
89 }
90 #[inline]
91 fn fromf(x: f64) -> Self {
92 x
93 }
94 #[inline]
95 fn f64(self) -> f64 {
96 self
97 }
98}
99
100#[inline]
101fn dot<F: Fp>(a: &[F], b: &[F]) -> F {
102 let mut s = F::ZERO;
103 for (x, y) in a.iter().zip(b) {
104 s += *x * *y;
105 }
106 s
107}
108
109pub fn matmul_nt<F: Fp>(x: &[F], w: &[F], y: &mut [F], n: usize, k: usize, m: usize) {
114 for i in 0..n {
115 let xr = &x[i * k..(i + 1) * k];
116 for o in 0..m {
117 y[i * m + o] = dot(xr, &w[o * k..(o + 1) * k]);
118 }
119 }
120}
121
122pub fn matmul_nt_dx<F: Fp>(dy: &[F], w: &[F], dx: &mut [F], n: usize, k: usize, m: usize) {
124 for i in 0..n {
125 let dxr = &mut dx[i * k..(i + 1) * k];
126 for o in 0..m {
127 let g = dy[i * m + o];
128 for (d, wv) in dxr.iter_mut().zip(&w[o * k..(o + 1) * k]) {
129 *d += g * *wv;
130 }
131 }
132 }
133}
134
135pub fn matmul_nt_dw<F: Fp>(dy: &[F], x: &[F], dw: &mut [F], n: usize, k: usize, m: usize) {
137 for i in 0..n {
138 let xr = &x[i * k..(i + 1) * k];
139 for o in 0..m {
140 let g = dy[i * m + o];
141 for (d, xv) in dw[o * k..(o + 1) * k].iter_mut().zip(xr) {
142 *d += g * *xv;
143 }
144 }
145 }
146}
147
148const GEMM_BLOCK: usize = 128;
153
154struct SendMut<T>(*mut T);
157unsafe impl<T> Send for SendMut<T> {}
158unsafe impl<T> Sync for SendMut<T> {}
159impl<T> SendMut<T> {
160 #[inline]
161 #[allow(clippy::mut_from_ref)]
165 unsafe fn slice(&self, off: usize, len: usize) -> &mut [T] {
166 unsafe { std::slice::from_raw_parts_mut(self.0.add(off), len) }
167 }
168}
169
170#[cfg(target_os = "macos")]
176mod accel {
177 #[link(name = "Accelerate", kind = "framework")]
178 unsafe extern "C" {
179 pub fn cblas_sgemm(
180 order: i32,
181 ta: i32,
182 tb: i32,
183 m: i32,
184 n: i32,
185 k: i32,
186 alpha: f32,
187 a: *const f32,
188 lda: i32,
189 b: *const f32,
190 ldb: i32,
191 beta: f32,
192 c: *mut f32,
193 ldc: i32,
194 );
195 }
196
197 pub fn on() -> bool {
198 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
199 *ON.get_or_init(|| std::env::var("CMF_ACCEL").map(|v| v != "0").unwrap_or(true))
200 }
201}
202
203#[cfg(target_arch = "x86_64")]
205fn gemm_avx512() -> bool {
206 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
207 *ON.get_or_init(|| {
208 std::env::var("CMF_GEMM_BLOCK4").map(|v| v != "0").unwrap_or(true)
209 && std::arch::is_x86_feature_detected!("avx512f")
210 })
211}
212
213#[cfg(target_arch = "x86_64")]
218#[target_feature(enable = "avx512f")]
219unsafe fn gemm_block4x4_avx512(
220 x: &[f32],
221 w: &[f32],
222 i0: usize,
223 o0: usize,
224 k: usize,
225) -> [[f32; 4]; 4] {
226 unsafe {
228 use core::arch::x86_64::*;
229 let mut acc = [[_mm512_setzero_ps(); 4]; 4];
230 let xp = x.as_ptr().add(i0 * k);
231 let wp = w.as_ptr().add(o0 * k);
232 let mut c = 0usize;
233 while c + 16 <= k {
234 let xv = [
235 _mm512_loadu_ps(xp.add(c)),
236 _mm512_loadu_ps(xp.add(k + c)),
237 _mm512_loadu_ps(xp.add(2 * k + c)),
238 _mm512_loadu_ps(xp.add(3 * k + c)),
239 ];
240 let wv = [
241 _mm512_loadu_ps(wp.add(c)),
242 _mm512_loadu_ps(wp.add(k + c)),
243 _mm512_loadu_ps(wp.add(2 * k + c)),
244 _mm512_loadu_ps(wp.add(3 * k + c)),
245 ];
246 for a in 0..4 {
247 for b in 0..4 {
248 acc[a][b] = _mm512_fmadd_ps(xv[a], wv[b], acc[a][b]);
249 }
250 }
251 c += 16;
252 }
253 let mut out = [[0f32; 4]; 4];
254 for a in 0..4 {
255 for b in 0..4 {
256 out[a][b] = _mm512_reduce_add_ps(acc[a][b]);
257 }
258 }
259 while c < k {
262 for a in 0..4 {
263 let xe = *xp.add(a * k + c);
264 for b in 0..4 {
265 out[a][b] += xe * *wp.add(b * k + c);
266 }
267 }
268 c += 1;
269 }
270 out
271 }
272}
273
274struct GemmGuard(std::time::Instant, usize, usize, usize);
290impl Drop for GemmGuard {
291 fn drop(&mut self) {
292 crate::fcd::prof::add(&crate::fcd::prof::GEMM, self.0);
293 crate::fcd::prof::gemm_shape(self.1, self.2, self.3, self.0);
294 }
295}
296fn scopeguard_gemm(t: std::time::Instant, n: usize, k: usize, m: usize) -> GemmGuard {
297 GemmGuard(t, n, k, m)
298}
299
300pub fn gemm_nt(
301 x: &[f32],
302 w: &[f32],
303 y: &mut [f32],
304 n: usize,
305 k: usize,
306 m: usize,
307 pool: Option<&Pool>,
308) {
309 let _t_gemm = std::time::Instant::now();
310 crate::fcd::prof::GEMM_CALLS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
311 let _guard = scopeguard_gemm(_t_gemm, n, k, m);
312 #[cfg(feature = "gpu")]
313 if n * k * m >= (1 << 22) && crate::gpu::enabled_here() {
314 let t0 = std::time::Instant::now();
315 match crate::gpu::probe_arm(crate::gpu::OpClass::GemmNt) {
316 crate::gpu::ProbeArm::Gpu => {
317 if crate::gpu_wgpu::gemm_nt_f32(x, w, y, n, k, m) {
318 crate::gpu::probe_record(crate::gpu::OpClass::GemmNt, true, t0.elapsed());
319 return;
320 }
321 }
322 crate::gpu::ProbeArm::CpuTimed => {
323 crate::gpu::cpu_scope(|| gemm_nt_cpu(x, w, y, n, k, m, pool));
324 crate::gpu::probe_record(crate::gpu::OpClass::GemmNt, false, t0.elapsed());
325 return;
326 }
327 crate::gpu::ProbeArm::Cpu => {}
328 }
329 }
330 gemm_nt_cpu(x, w, y, n, k, m, pool)
331}
332
333fn gemm_nt_cpu(
334 x: &[f32],
335 w: &[f32],
336 y: &mut [f32],
337 n: usize,
338 k: usize,
339 m: usize,
340 pool: Option<&Pool>,
341) {
342 debug_assert_eq!(x.len(), n * k);
343 debug_assert_eq!(w.len(), m * k);
344 debug_assert_eq!(y.len(), n * m);
345 #[cfg(target_os = "macos")]
346 if accel::on() && n * k * m >= 1 << 18 {
347 unsafe {
349 accel::cblas_sgemm(
350 101,
351 111,
352 112,
353 n as i32,
354 m as i32,
355 k as i32,
356 1.0,
357 x.as_ptr(),
358 k as i32,
359 w.as_ptr(),
360 k as i32,
361 0.0,
362 y.as_mut_ptr(),
363 m as i32,
364 );
365 }
366 return;
367 }
368 let nb = n.div_ceil(GEMM_BLOCK);
369 let block = |r0: usize, r1: usize, y: &mut [f32]| {
370 let mut o = 0usize;
382 #[cfg(target_arch = "x86_64")]
383 if gemm_avx512() {
384 while o + 4 <= m {
385 let mut i = r0;
386 while i + 4 <= r1 {
387 let acc = unsafe { gemm_block4x4_avx512(x, w, i, o, k) };
390 for (di, ai) in acc.iter().enumerate() {
391 for (dj, v) in ai.iter().enumerate() {
392 y[(i - r0 + di) * m + o + dj] = *v;
393 }
394 }
395 i += 4;
396 }
397 for i in i..r1 {
398 let xr = &x[i * k..(i + 1) * k];
399 for dj in 0..4 {
400 let wr = &w[(o + dj) * k..(o + dj + 1) * k];
401 y[(i - r0) * m + o + dj] =
402 crate::attention::dot_f32(xr, wr);
403 }
404 }
405 o += 4;
406 }
407 }
408 for o in o..m {
409 let wr = &w[o * k..(o + 1) * k];
410 for i in r0..r1 {
411 y[(i - r0) * m + o] = crate::attention::dot_f32(&x[i * k..(i + 1) * k], wr);
412 }
413 }
414 };
415 match pool {
416 Some(p) if nb > 1 => {
417 let yp = SendMut(y.as_mut_ptr());
418 p.run(&|widx, nw| {
419 for bi in (widx..nb).step_by(nw) {
420 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
421 let ys = unsafe { yp.slice(r0 * m, (r1 - r0) * m) };
423 block(r0, r1, ys);
424 }
425 });
426 }
427 _ => {
428 for bi in 0..nb {
429 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
430 block(r0, r1, &mut y[r0 * m..r1 * m]);
431 }
432 }
433 }
434}
435
436pub fn gemm_dx(
438 dy: &[f32],
439 w: &[f32],
440 dx: &mut [f32],
441 n: usize,
442 k: usize,
443 m: usize,
444 pool: Option<&Pool>,
445) {
446 let _t = std::time::Instant::now();
447 crate::fcd::prof::GEMM_CALLS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
448 let _guard = scopeguard_gemm(_t, n, k, m);
449 debug_assert_eq!(dy.len(), n * m);
450 debug_assert_eq!(w.len(), m * k);
451 debug_assert_eq!(dx.len(), n * k);
452 #[cfg(feature = "gpu")]
456 if crate::gpu_wgpu::gemm_dx_f32(dy, w, dx, n, k, m) {
457 return;
458 }
459 #[cfg(target_os = "macos")]
460 if accel::on() && n * k * m >= 1 << 18 {
461 unsafe {
463 accel::cblas_sgemm(
464 101,
465 111,
466 111,
467 n as i32,
468 k as i32,
469 m as i32,
470 1.0,
471 dy.as_ptr(),
472 m as i32,
473 w.as_ptr(),
474 k as i32,
475 1.0,
476 dx.as_mut_ptr(),
477 k as i32,
478 );
479 }
480 return;
481 }
482 let nb = n.div_ceil(GEMM_BLOCK);
483 let block = |r0: usize, r1: usize, dxs: &mut [f32]| {
484 for o in 0..m {
485 let wr = &w[o * k..(o + 1) * k];
486 for i in r0..r1 {
487 let g = dy[i * m + o];
488 if g != 0.0 {
489 crate::attention::axpy_f32(&mut dxs[(i - r0) * k..(i - r0 + 1) * k], wr, g);
490 }
491 }
492 }
493 };
494 match pool {
495 Some(p) if nb > 1 => {
496 let dxp = SendMut(dx.as_mut_ptr());
497 p.run(&|widx, nw| {
498 for bi in (widx..nb).step_by(nw) {
499 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
500 let dxs = unsafe { dxp.slice(r0 * k, (r1 - r0) * k) };
502 block(r0, r1, dxs);
503 }
504 });
505 }
506 _ => {
507 for bi in 0..nb {
508 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
509 block(r0, r1, &mut dx[r0 * k..r1 * k]);
510 }
511 }
512 }
513}
514
515pub fn gemm_dw(
518 dy: &[f32],
519 x: &[f32],
520 dw: &mut [f32],
521 n: usize,
522 k: usize,
523 m: usize,
524 pool: Option<&Pool>,
525) {
526 let _t = std::time::Instant::now();
527 crate::fcd::prof::GEMM_CALLS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
528 let _guard = scopeguard_gemm(_t, n, k, m);
529 debug_assert_eq!(dy.len(), n * m);
530 debug_assert_eq!(x.len(), n * k);
531 debug_assert_eq!(dw.len(), m * k);
532 #[cfg(target_os = "macos")]
533 if accel::on() && n * k * m >= 1 << 18 {
534 unsafe {
536 accel::cblas_sgemm(
537 101,
538 112,
539 111,
540 m as i32,
541 k as i32,
542 n as i32,
543 1.0,
544 dy.as_ptr(),
545 m as i32,
546 x.as_ptr(),
547 k as i32,
548 1.0,
549 dw.as_mut_ptr(),
550 k as i32,
551 );
552 }
553 return;
554 }
555 let range = |o0: usize, o1: usize, dws: &mut [f32]| {
556 let nb = n.div_ceil(GEMM_BLOCK);
558 for bi in 0..nb {
559 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
560 for o in o0..o1 {
561 let dwr = &mut dws[(o - o0) * k..(o - o0 + 1) * k];
562 for i in r0..r1 {
563 let g = dy[i * m + o];
564 if g != 0.0 {
565 crate::attention::axpy_f32(dwr, &x[i * k..(i + 1) * k], g);
566 }
567 }
568 }
569 }
570 };
571 match pool {
572 Some(p) if m >= 8 => {
573 let dwp = SendMut(dw.as_mut_ptr());
574 p.run(&|widx, nw| {
575 let (o0, o1) = (widx * m / nw, (widx + 1) * m / nw);
576 if o0 < o1 {
577 let dws = unsafe { dwp.slice(o0 * k, (o1 - o0) * k) };
579 range(o0, o1, dws);
580 }
581 });
582 }
583 _ => range(0, m, dw),
584 }
585}
586
587#[inline]
590pub fn silu<F: Fp>(x: F) -> F {
591 x / (F::ONE + (-x).exp())
592}
593
594#[inline]
596pub fn silu_bwd<F: Fp>(x: F) -> F {
597 let s = F::ONE / (F::ONE + (-x).exp());
598 s * (F::ONE + x * (F::ONE - s))
599}
600
601pub fn rmsnorm_fwd<F: Fp>(x: &[F], w: &[F], eps: f64, gemma: bool, y: &mut [F], inv_out: &mut [F]) {
608 let d = w.len();
609 let n = x.len() / d;
610 for r in 0..n {
611 let xr = &x[r * d..(r + 1) * d];
612 let mut ss = 0f64;
613 for v in xr {
614 ss += v.f64() * v.f64();
615 }
616 let inv = F::fromf(1.0 / (ss / d as f64 + eps).sqrt());
617 inv_out[r] = inv;
618 let yr = &mut y[r * d..(r + 1) * d];
619 for j in 0..d {
620 let weff = if gemma { F::ONE + w[j] } else { w[j] };
621 yr[j] = xr[j] * inv * weff;
622 }
623 }
624}
625
626pub fn rmsnorm_bwd<F: Fp>(
632 x: &[F],
633 w: &[F],
634 inv: &[F],
635 dy: &[F],
636 gemma: bool,
637 dx: &mut [F],
638 mut dw: Option<&mut [F]>,
639) {
640 let d = w.len();
641 let n = x.len() / d;
642 for r in 0..n {
643 let xr = &x[r * d..(r + 1) * d];
644 let dyr = &dy[r * d..(r + 1) * d];
645 let iv = inv[r];
646 let mut s = 0f64;
647 for j in 0..d {
648 let weff = if gemma { F::ONE + w[j] } else { w[j] };
649 s += (dyr[j] * weff * xr[j]).f64();
650 }
651 let coef = F::fromf(s / d as f64) * iv * iv * iv;
652 let dxr = &mut dx[r * d..(r + 1) * d];
653 for j in 0..d {
654 let weff = if gemma { F::ONE + w[j] } else { w[j] };
655 dxr[j] += iv * weff * dyr[j] - xr[j] * coef;
656 }
657 if let Some(dwv) = dw.as_deref_mut() {
658 for j in 0..d {
659 dwv[j] += dyr[j] * xr[j] * iv;
660 }
661 }
662 }
663}
664
665pub fn rope_fwd<F: Fp>(x: &mut [F], position: usize, inv_freq: &[f64]) {
670 let half = inv_freq.len();
671 for (i, &freq) in inv_freq.iter().enumerate() {
672 let angle = position as f64 * freq;
673 let (sin, cos) = (F::fromf(angle.sin()), F::fromf(angle.cos()));
674 let x0 = x[i];
675 let x1 = x[i + half];
676 x[i] = x0 * cos - x1 * sin;
677 x[i + half] = x0 * sin + x1 * cos;
678 }
679}
680
681pub fn rope_bwd<F: Fp>(dy: &mut [F], position: usize, inv_freq: &[f64]) {
684 let half = inv_freq.len();
685 for (i, &freq) in inv_freq.iter().enumerate() {
686 let angle = position as f64 * freq;
687 let (sin, cos) = (F::fromf(angle.sin()), F::fromf(angle.cos()));
688 let g0 = dy[i];
689 let g1 = dy[i + half];
690 dy[i] = g0 * cos + g1 * sin;
691 dy[i + half] = -g0 * sin + g1 * cos;
692 }
693}
694
695pub fn seg_means<F: Fp>(x: &[F], t: usize, d: usize, m: usize, out: &mut [F]) {
700 for i in 0..m {
701 let (lo, hi) = (i * t / m, (i + 1) * t / m);
702 let or = &mut out[i * d..(i + 1) * d];
703 for v in or.iter_mut() {
704 *v = F::ZERO;
705 }
706 for j in lo..hi {
707 for c in 0..d {
708 or[c] += x[j * d + c];
709 }
710 }
711 let inv = F::fromf(1.0 / (hi - lo) as f64);
712 for v in or.iter_mut() {
713 *v *= inv;
714 }
715 }
716}
717
718pub fn seg_means_bwd<F: Fp>(dl: &[F], t: usize, d: usize, m: usize, dx: &mut [F]) {
721 for i in 0..m {
722 let (lo, hi) = (i * t / m, (i + 1) * t / m);
723 let inv = F::fromf(1.0 / (hi - lo) as f64);
724 let dlr = &dl[i * d..(i + 1) * d];
725 for j in lo..hi {
726 for c in 0..d {
727 dx[j * d + c] += dlr[c] * inv;
728 }
729 }
730 }
731}
732
733#[allow(clippy::needless_range_loop)] pub fn attn_head_fwd<F: Fp>(
739 q: &[F],
740 k: &[F],
741 v: &[F],
742 t: usize,
743 d: usize,
744 dv: usize,
745 out: &mut [F],
746) {
747 let scale = F::fromf(1.0 / (d as f64).sqrt());
748 let mut row = vec![F::ZERO; t];
749 for ti in 0..t {
750 let qr = &q[ti * d..(ti + 1) * d];
751 let mut mx = F::fromf(f64::NEG_INFINITY);
752 for j in 0..=ti {
753 let s = dot(qr, &k[j * d..(j + 1) * d]) * scale;
754 row[j] = s;
755 mx = mx.maxf(s);
756 }
757 let mut den = F::ZERO;
758 for j in 0..=ti {
759 row[j] = (row[j] - mx).exp();
760 den += row[j];
761 }
762 let or = &mut out[ti * dv..(ti + 1) * dv];
763 for o in or.iter_mut() {
764 *o = F::ZERO;
765 }
766 for j in 0..=ti {
767 let p = row[j] / den;
768 for (o, vv) in or.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
769 *o += p * *vv;
770 }
771 }
772 }
773}
774
775#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
779pub fn attn_head_bwd<F: Fp>(
780 q: &[F],
781 k: &[F],
782 v: &[F],
783 dout: &[F],
784 t: usize,
785 d: usize,
786 dv: usize,
787 dq: &mut [F],
788 dk: &mut [F],
789 dvv: &mut [F],
790) {
791 let scale = F::fromf(1.0 / (d as f64).sqrt());
792 let mut row = vec![F::ZERO; t];
793 for ti in 0..t {
794 let qr = &q[ti * d..(ti + 1) * d];
795 let mut mx = F::fromf(f64::NEG_INFINITY);
796 for j in 0..=ti {
797 let s = dot(qr, &k[j * d..(j + 1) * d]) * scale;
798 row[j] = s;
799 mx = mx.maxf(s);
800 }
801 let mut den = F::ZERO;
802 for j in 0..=ti {
803 row[j] = (row[j] - mx).exp();
804 den += row[j];
805 }
806 let dor = &dout[ti * dv..(ti + 1) * dv];
807 let mut pdp = F::ZERO;
809 let mut dp = vec![F::ZERO; ti + 1];
810 for j in 0..=ti {
811 let p = row[j] / den;
812 row[j] = p; dp[j] = dot(dor, &v[j * dv..(j + 1) * dv]);
814 pdp += p * dp[j];
815 }
816 let dqr = &mut dq[ti * d..(ti + 1) * d];
817 for j in 0..=ti {
818 let p = row[j];
819 for (dvo, o) in dvv[j * dv..(j + 1) * dv].iter_mut().zip(dor) {
821 *dvo += p * *o;
822 }
823 let ds = p * (dp[j] - pdp) * scale;
825 let kr = &k[j * d..(j + 1) * d];
826 for c in 0..d {
827 dqr[c] += ds * kr[c];
828 }
829 let dkr = &mut dk[j * d..(j + 1) * d];
830 for c in 0..d {
831 dkr[c] += ds * qr[c];
832 }
833 }
834 }
835}
836
837#[derive(Clone, Copy, Debug)]
841pub struct NysCfg {
842 pub m: usize,
844 pub w: usize,
846 pub sink: usize,
848 pub prefill: Option<usize>,
860}
861
862impl NysCfg {
863 #[inline]
866 pub fn prefill_len(&self, t: usize) -> usize {
867 self.prefill.unwrap_or(t / 2).clamp(1, t)
868 }
869}
870
871const NYS_DEN_EPS: f64 = 1e-30;
873
874struct NysGraph<F: Fp> {
878 m_eff: usize,
879 tp: usize,
882 q_l: Vec<F>,
883 k_l: Vec<F>,
884 mu: Vec<F>,
885 fu: Vec<F>,
886 e: Vec<F>,
887 fumu: Vec<F>,
888 far_keep: Vec<bool>,
897 wmat: Vec<F>,
900 c_row: Vec<F>,
901 den: Vec<F>,
902}
903
904#[inline]
909fn nys_exact_row(ti: usize, tp: usize) -> bool {
910 ti < tp
911}
912
913#[inline]
915fn nys_near(ti: usize, j: usize, w: usize, sink: usize) -> bool {
916 ti - j < w || j < sink
917}
918
919fn nys_graph<F: Fp>(
924 q: &[F],
925 k: &[F],
926 t: usize,
927 d: usize,
928 cfg: &NysCfg,
929 mu_override: Option<&[F]>,
930) -> NysGraph<F> {
931 let scale = 1.0 / (d as f64).sqrt();
932 let fscale = F::fromf(scale);
933 let tp = cfg.prefill_len(t);
938 let m_eff = (tp / 8).clamp(4, cfg.m);
939 let mut q_l = vec![F::ZERO; m_eff * d];
940 let mut k_l = vec![F::ZERO; m_eff * d];
941 seg_means(&q[..tp * d], tp, d, m_eff, &mut q_l);
942 seg_means(&k[..tp * d], tp, d, m_eff, &mut k_l);
943
944 let mut au = vec![0f64; m_eff * m_eff];
946 for i in 0..m_eff {
947 for j in 0..m_eff {
948 let mut s = 0f64;
949 for c in 0..d {
950 s += q_l[i * d + c].f64() * k_l[j * d + c].f64();
951 }
952 au[i * m_eff + j] = (s * scale).exp();
953 }
954 }
955 let mu: Vec<F> = match mu_override {
956 Some(m) => m.to_vec(),
957 None => crate::nystrom::ridge_pinv(&au, m_eff)
958 .iter()
959 .map(|&x| F::fromf(x))
960 .collect(),
961 };
962
963 let mut fu = vec![F::ZERO; t * m_eff];
965 for ti in 0..t {
966 for i in 0..m_eff {
967 fu[ti * m_eff + i] =
968 (dot(&q[ti * d..(ti + 1) * d], &k_l[i * d..(i + 1) * d]) * fscale).exp();
969 }
970 }
971 let mut e = vec![F::ZERO; m_eff * t];
972 for i in 0..m_eff {
973 for j in 0..t {
974 e[i * t + j] = (dot(&q_l[i * d..(i + 1) * d], &k[j * d..(j + 1) * d]) * fscale).exp();
975 }
976 }
977 let mut fumu = vec![F::ZERO; t * m_eff];
978 matmul_nt(
979 &fu,
980 &transpose(&mu, m_eff, m_eff),
983 &mut fumu,
984 t,
985 m_eff,
986 m_eff,
987 );
988 let mut a = vec![F::ZERO; t * t];
993 for ti in tp..t {
994 let fr = &fumu[ti * m_eff..(ti + 1) * m_eff];
995 let ar = &mut a[ti * t..ti * t + ti + 1]; for (i, &f) in fr.iter().enumerate() {
997 let er = &e[i * t..i * t + ti + 1];
998 for (av, ev) in ar.iter_mut().zip(er) {
999 *av += f * *ev;
1000 }
1001 }
1002 }
1003
1004 let mut wmat = vec![F::ZERO; t * t];
1007 let mut c_row = vec![F::ZERO; t];
1008 let mut den = vec![F::ZERO; t];
1009 let mut far_keep = vec![false; t];
1010 let mut lg_row = vec![F::ZERO; t];
1011 for ti in 0..t {
1012 let qr = &q[ti * d..(ti + 1) * d];
1013 let mut c = F::fromf(f64::NEG_INFINITY);
1014 for j in 0..=ti {
1015 let s = dot(qr, &k[j * d..(j + 1) * d]) * fscale;
1016 lg_row[j] = s;
1017 c = c.maxf(s);
1018 }
1019 c_row[ti] = c;
1020 let emc = (-c).exp();
1021 let exact_row = nys_exact_row(ti, tp);
1022
1023 let mut far_sum = F::ZERO;
1038 if !exact_row {
1039 for j in 0..=ti {
1040 if !nys_near(ti, j, cfg.w, cfg.sink) {
1041 far_sum += a[ti * t + j];
1042 }
1043 }
1044 }
1045 let keep = !exact_row && far_sum.f64() >= 0.0;
1046 far_keep[ti] = keep;
1047
1048 let wr = &mut wmat[ti * t..(ti + 1) * t];
1049 let mut dsum = F::ZERO;
1050 for j in 0..=ti {
1051 let wv = if exact_row || nys_near(ti, j, cfg.w, cfg.sink) {
1052 (lg_row[j] - c).exp()
1053 } else if keep {
1054 a[ti * t + j] * emc
1055 } else {
1056 F::ZERO
1057 };
1058 wr[j] = wv;
1059 dsum += wv;
1060 }
1061 den[ti] = dsum.maxf(F::fromf(NYS_DEN_EPS));
1062 }
1063 NysGraph {
1064 m_eff,
1065 tp,
1066 q_l,
1067 k_l,
1068 mu,
1069 fu,
1070 e,
1071 fumu,
1072 far_keep,
1073 wmat,
1074 c_row,
1075 den,
1076 }
1077}
1078
1079fn transpose<F: Fp>(x: &[F], rows: usize, cols: usize) -> Vec<F> {
1080 let mut out = vec![F::ZERO; rows * cols];
1081 for r in 0..rows {
1082 for c in 0..cols {
1083 out[c * rows + r] = x[r * cols + c];
1084 }
1085 }
1086 out
1087}
1088
1089#[allow(clippy::too_many_arguments)]
1092pub fn nystrom_head_fwd<F: Fp>(
1093 q: &[F],
1094 k: &[F],
1095 v: &[F],
1096 t: usize,
1097 d: usize,
1098 dv: usize,
1099 cfg: &NysCfg,
1100 out: &mut [F],
1101) {
1102 if nys_degenerate(t, cfg) {
1103 attn_head_fwd(q, k, v, t, d, dv, out);
1104 return;
1105 }
1106 nystrom_head_fwd_mu(q, k, v, t, d, dv, cfg, None, out);
1107}
1108
1109#[inline]
1118fn nys_degenerate(t: usize, cfg: &NysCfg) -> bool {
1119 cfg.prefill_len(t) <= cfg.w + cfg.sink + 8
1120}
1121
1122#[doc(hidden)]
1125pub fn nystrom_mu_for_test<F: Fp>(q: &[F], k: &[F], t: usize, d: usize, cfg: &NysCfg) -> Vec<F> {
1126 nys_graph(q, k, t, d, cfg, None).mu
1127}
1128
1129#[doc(hidden)]
1132#[allow(clippy::too_many_arguments)]
1133pub fn nystrom_head_fwd_mu<F: Fp>(
1134 q: &[F],
1135 k: &[F],
1136 v: &[F],
1137 t: usize,
1138 d: usize,
1139 dv: usize,
1140 cfg: &NysCfg,
1141 mu_override: Option<&[F]>,
1142 out: &mut [F],
1143) {
1144 let g = nys_graph(q, k, t, d, cfg, mu_override);
1145 for ti in 0..t {
1146 let wr = &g.wmat[ti * t..(ti + 1) * t];
1147 let den = g.den[ti];
1148 let or = &mut out[ti * dv..(ti + 1) * dv];
1149 for o in or.iter_mut() {
1150 *o = F::ZERO;
1151 }
1152 for j in 0..=ti {
1153 let p = wr[j] / den;
1154 if p.f64() != 0.0 {
1155 for (o, vv) in or.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
1156 *o += p * *vv;
1157 }
1158 }
1159 }
1160 }
1161}
1162
1163#[allow(clippy::too_many_arguments)]
1174pub fn nystrom_head_bwd<F: Fp>(
1175 q: &[F],
1176 k: &[F],
1177 v: &[F],
1178 dout: &[F],
1179 t: usize,
1180 d: usize,
1181 dv: usize,
1182 cfg: &NysCfg,
1183 dq: &mut [F],
1184 dk: &mut [F],
1185 dvv: &mut [F],
1186) {
1187 if nys_degenerate(t, cfg) {
1188 attn_head_bwd(q, k, v, dout, t, d, dv, dq, dk, dvv);
1189 return;
1190 }
1191 nystrom_head_bwd_mu(q, k, v, dout, t, d, dv, cfg, None, dq, dk, dvv);
1192}
1193
1194#[doc(hidden)]
1196#[allow(clippy::too_many_arguments)]
1197pub fn nystrom_head_bwd_mu<F: Fp>(
1198 q: &[F],
1199 k: &[F],
1200 v: &[F],
1201 dout: &[F],
1202 t: usize,
1203 d: usize,
1204 dv: usize,
1205 cfg: &NysCfg,
1206 mu_override: Option<&[F]>,
1207 dq: &mut [F],
1208 dk: &mut [F],
1209 dvv: &mut [F],
1210) {
1211 let scale = F::fromf(1.0 / (d as f64).sqrt());
1212 let g = nys_graph(q, k, t, d, cfg, mu_override);
1213 let m_eff = g.m_eff;
1214
1215 let mut dwmat = vec![F::ZERO; t * t];
1217 let mut out_row = vec![F::ZERO; dv];
1218 for ti in 0..t {
1219 let wr = &g.wmat[ti * t..(ti + 1) * t];
1220 let den = g.den[ti];
1221 for o in out_row.iter_mut() {
1222 *o = F::ZERO;
1223 }
1224 for j in 0..=ti {
1225 let p = wr[j] / den;
1226 if p.f64() != 0.0 {
1227 for (o, vv) in out_row.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
1228 *o += p * *vv;
1229 }
1230 }
1231 }
1232 let dor = &dout[ti * dv..(ti + 1) * dv];
1233 let dwr = &mut dwmat[ti * t..(ti + 1) * t];
1234 for j in 0..=ti {
1235 let p = wr[j] / den;
1237 if p.f64() != 0.0 {
1238 for (dvo, o) in dvv[j * dv..(j + 1) * dv].iter_mut().zip(dor) {
1239 *dvo += p * *o;
1240 }
1241 }
1242 let mut s = F::ZERO;
1244 for c in 0..dv {
1245 s += dor[c] * (v[j * dv + c] - out_row[c]);
1246 }
1247 dwr[j] = s / den;
1248 }
1249 }
1250
1251 for ti in 0..t {
1254 let qr = &q[ti * d..(ti + 1) * d];
1255 let dqr_base = ti * d;
1256 for j in 0..=ti {
1257 if !(nys_exact_row(ti, g.tp) || nys_near(ti, j, cfg.w, cfg.sink)) {
1258 continue;
1259 }
1260 let dlg = dwmat[ti * t + j] * g.wmat[ti * t + j] * scale;
1261 if dlg.f64() == 0.0 {
1262 continue;
1263 }
1264 let kr = &k[j * d..(j + 1) * d];
1265 for c in 0..d {
1266 dq[dqr_base + c] += dlg * kr[c];
1267 }
1268 let dkr = &mut dk[j * d..(j + 1) * d];
1269 for c in 0..d {
1270 dkr[c] += dlg * qr[c];
1271 }
1272 }
1273 }
1274
1275 let mut da = vec![F::ZERO; t * t];
1290 for ti in g.tp..t {
1291 if !g.far_keep[ti] {
1292 continue; }
1294 let emc = (-g.c_row[ti]).exp();
1295 for j in 0..=ti {
1296 if nys_near(ti, j, cfg.w, cfg.sink) {
1297 continue;
1298 }
1299 da[ti * t + j] = dwmat[ti * t + j] * emc;
1300 }
1301 }
1302 let mut dfumu = vec![F::ZERO; t * m_eff];
1304 for ti in 0..t {
1305 let dar = &da[ti * t..ti * t + ti + 1];
1306 let dfr = &mut dfumu[ti * m_eff..(ti + 1) * m_eff];
1307 for (i, df) in dfr.iter_mut().enumerate() {
1308 let er = &g.e[i * t..i * t + ti + 1];
1309 let mut s = F::ZERO;
1310 for (av, ev) in dar.iter().zip(er) {
1311 s += *av * *ev;
1312 }
1313 *df = s;
1314 }
1315 }
1316 let mut dfu = vec![F::ZERO; t * m_eff];
1318 matmul_nt(&dfumu, &g.mu, &mut dfu, t, m_eff, m_eff);
1319 let mut de = vec![F::ZERO; m_eff * t];
1321 for ti in 0..t {
1322 let dar = &da[ti * t..ti * t + ti + 1];
1323 let fr = &g.fumu[ti * m_eff..(ti + 1) * m_eff];
1324 for (i, &f) in fr.iter().enumerate() {
1325 if f.f64() == 0.0 {
1326 continue;
1327 }
1328 let der = &mut de[i * t..i * t + ti + 1];
1329 for (dev, av) in der.iter_mut().zip(dar) {
1330 *dev += f * *av;
1331 }
1332 }
1333 }
1334 let mut dq_l = vec![F::ZERO; m_eff * d];
1337 let mut dk_l = vec![F::ZERO; m_eff * d];
1338 for ti in 0..t {
1339 let qr = &q[ti * d..(ti + 1) * d];
1340 for i in 0..m_eff {
1341 let dlg = dfu[ti * m_eff + i] * g.fu[ti * m_eff + i] * scale;
1342 if dlg.f64() == 0.0 {
1343 continue;
1344 }
1345 let klr = &g.k_l[i * d..(i + 1) * d];
1346 for c in 0..d {
1347 dq[ti * d + c] += dlg * klr[c];
1348 }
1349 let dklr = &mut dk_l[i * d..(i + 1) * d];
1350 for c in 0..d {
1351 dklr[c] += dlg * qr[c];
1352 }
1353 }
1354 }
1355 for i in 0..m_eff {
1356 let qlr = &g.q_l[i * d..(i + 1) * d];
1357 for j in 0..t {
1358 let dlg = de[i * t + j] * g.e[i * t + j] * scale;
1359 if dlg.f64() == 0.0 {
1360 continue;
1361 }
1362 let kr = &k[j * d..(j + 1) * d];
1363 let dqlr = &mut dq_l[i * d..(i + 1) * d];
1364 for c in 0..d {
1365 dqlr[c] += dlg * kr[c];
1366 }
1367 for c in 0..d {
1368 dk[j * d + c] += dlg * qlr[c];
1369 }
1370 }
1371 }
1372 let tp = g.tp;
1376 seg_means_bwd(&dq_l, tp, d, m_eff, &mut dq[..tp * d]);
1377 seg_means_bwd(&dk_l, tp, d, m_eff, &mut dk[..tp * d]);
1378}
1379
1380pub fn ce_kl_position<F: Fp>(
1388 s_logits: &[F],
1389 t_logits: &[F],
1390 target: usize,
1391 kl_w: f64,
1392 inv_n: f64,
1393 dlogits: &mut [F],
1394) -> (f64, f64) {
1395 let vsz = s_logits.len();
1396 debug_assert_eq!(t_logits.len(), vsz);
1397 let mut smax = f64::NEG_INFINITY;
1399 let mut tmax = f64::NEG_INFINITY;
1400 for i in 0..vsz {
1401 smax = smax.max(s_logits[i].f64());
1402 tmax = tmax.max(t_logits[i].f64());
1403 }
1404 let mut ssum = 0f64;
1405 let mut tsum = 0f64;
1406 for i in 0..vsz {
1407 ssum += (s_logits[i].f64() - smax).exp();
1408 tsum += (t_logits[i].f64() - tmax).exp();
1409 }
1410 let slz = smax + ssum.ln();
1411 let tlz = tmax + tsum.ln();
1412 let ce = slz - s_logits[target].f64();
1413 let mut kl = 0f64;
1414 for i in 0..vsz {
1415 let ls = s_logits[i].f64() - slz;
1416 let lt = t_logits[i].f64() - tlz;
1417 let pt = lt.exp();
1418 let ps = ls.exp();
1419 if pt > 0.0 {
1420 kl += pt * (lt - ls);
1421 }
1422 let mut gd = (1.0 - kl_w) * ps + kl_w * (ps - pt);
1423 if i == target {
1424 gd -= 1.0 - kl_w;
1425 }
1426 dlogits[i] = F::fromf(gd * inv_n);
1427 }
1428 (ce, kl)
1429}
1430
1431pub struct GdnSeqCfg<'a> {
1448 pub nv: usize,
1449 pub nk: usize,
1450 pub dk: usize,
1451 pub dv: usize,
1452 pub kk: usize,
1453 pub rms_eps: f64,
1454 pub conv: &'a [f32],
1457 pub a_log: &'a [f32],
1459 pub dt_bias: &'a [f32],
1461 pub norm: &'a [f32],
1463}
1464
1465impl GdnSeqCfg<'_> {
1466 pub fn c_dim(&self) -> usize {
1467 2 * self.nk * self.dk + self.nv * self.dv
1468 }
1469}
1470
1471#[inline]
1472fn softplus_f<F: Fp>(x: F) -> F {
1473 if x.f64() > 20.0 {
1476 x
1477 } else {
1478 F::fromf(x.f64().exp().ln_1p())
1479 }
1480}
1481
1482#[inline]
1483fn sigmoid_f<F: Fp>(x: F) -> F {
1484 F::ONE / (F::ONE + (-x).exp())
1485}
1486
1487pub fn gdn_conv_fwd<F: Fp>(
1491 raw: &[F],
1492 t: usize,
1493 c_dim: usize,
1494 kk: usize,
1495 conv: &[f32],
1496 pre: &mut [F],
1497 cq: &mut [F],
1498) {
1499 for ti in 0..t {
1500 for c in 0..c_dim {
1501 let taps = &conv[c * kk..(c + 1) * kk];
1502 let mut acc = F::ZERO;
1503 for (j, &tap) in taps.iter().enumerate() {
1504 let p = ti as isize - (kk as isize - 1) + j as isize;
1506 if p >= 0 {
1507 acc += raw[p as usize * c_dim + c] * F::fromf(tap as f64);
1508 }
1509 }
1510 pre[ti * c_dim + c] = acc;
1511 cq[ti * c_dim + c] = silu(acc);
1512 }
1513 }
1514}
1515
1516pub fn gdn_conv_bwd<F: Fp>(
1518 pre: &[F],
1519 t: usize,
1520 c_dim: usize,
1521 kk: usize,
1522 conv: &[f32],
1523 dcq: &[F],
1524 draw: &mut [F],
1525) {
1526 for ti in 0..t {
1527 for c in 0..c_dim {
1528 let g = dcq[ti * c_dim + c];
1529 if g.f64() == 0.0 {
1530 continue;
1531 }
1532 let dp = g * silu_bwd(pre[ti * c_dim + c]);
1533 let taps = &conv[c * kk..(c + 1) * kk];
1534 for (j, &tap) in taps.iter().enumerate() {
1535 let p = ti as isize - (kk as isize - 1) + j as isize;
1536 if p >= 0 {
1537 draw[p as usize * c_dim + c] += dp * F::fromf(tap as f64);
1538 }
1539 }
1540 }
1541 }
1542}
1543
1544#[inline]
1547fn gdn_inv<F: Fp>(x: &[F], extra_scale: f64) -> (F, F) {
1548 let mut n2 = F::ZERO;
1549 for v in x {
1550 n2 += *v * *v;
1551 }
1552 let n2e = n2 + F::fromf(1e-6);
1553 let inv = F::ONE / (n2e.sqrt() * F::fromf(extra_scale));
1554 (inv, n2e)
1555}
1556
1557#[allow(clippy::too_many_arguments)]
1560pub fn gdn_group_fwd<F: Fp>(
1561 cq: &[F],
1562 z: &[F],
1563 a: &[F],
1564 b: &[F],
1565 t: usize,
1566 cfg: &GdnSeqCfg,
1567 ko: usize,
1568 out: &mut [F],
1569) {
1570 let (nv, nk, dk, dv) = (cfg.nv, cfg.nk, cfg.dk, cfg.dv);
1571 let c_dim = cfg.c_dim();
1572 let kd = nk * dk;
1573 let rep = nv / nk;
1574 let vd = nv * dv;
1575 let sqdk = (dk as f64).sqrt();
1576 for hh in 0..rep {
1577 let h = ko * rep + hh;
1578 let ea = F::fromf((cfg.a_log[h] as f64).exp());
1579 let mut s = vec![F::ZERO; dk * dv];
1580 let mut kv = vec![F::ZERO; dv];
1581 let mut o = vec![F::ZERO; dv];
1582 for ti in 0..t {
1583 let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
1584 let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
1585 let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
1586 let (invq, _) = gdn_inv(qrow, sqdk);
1587 let (invk, _) = gdn_inv(krow, 1.0);
1588 let g = (-ea * softplus_f(a[ti * nv + h] + F::fromf(cfg.dt_bias[h] as f64))).exp();
1589 let beta = sigmoid_f(b[ti * nv + h]);
1590 for x in kv.iter_mut() {
1592 *x = F::ZERO;
1593 }
1594 for di in 0..dk {
1595 let kf = krow[di] * invk;
1596 let row = &mut s[di * dv..(di + 1) * dv];
1597 for dj in 0..dv {
1598 row[dj] *= g;
1599 kv[dj] += row[dj] * kf;
1600 }
1601 }
1602 for x in o.iter_mut() {
1603 *x = F::ZERO;
1604 }
1605 for di in 0..dk {
1606 let kf = krow[di] * invk;
1607 let qf = qrow[di] * invq;
1608 let row = &mut s[di * dv..(di + 1) * dv];
1609 for dj in 0..dv {
1610 row[dj] += kf * (vrow[dj] - kv[dj]) * beta;
1611 o[dj] += qf * row[dj];
1612 }
1613 }
1614 let mut ss = 0f64;
1616 for v in &o {
1617 ss += v.f64() * v.f64();
1618 }
1619 let inv = F::fromf(1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt());
1620 for dj in 0..dv {
1621 let zv = z[ti * vd + h * dv + dj];
1622 out[ti * vd + h * dv + dj] = o[dj] * inv * F::fromf(cfg.norm[dj] as f64) * silu(zv);
1623 }
1624 }
1625 }
1626}
1627
1628#[allow(clippy::too_many_arguments)]
1632pub fn gdn_group_bwd<F: Fp>(
1633 cq: &[F],
1634 z: &[F],
1635 a: &[F],
1636 b: &[F],
1637 t: usize,
1638 cfg: &GdnSeqCfg,
1639 ko: usize,
1640 dout: &[F],
1641 dcq: &mut [F],
1642 dz: &mut [F],
1643 da: &mut [F],
1644 db: &mut [F],
1645) {
1646 let (nv, nk, dk, dv) = (cfg.nv, cfg.nk, cfg.dk, cfg.dv);
1647 let c_dim = cfg.c_dim();
1648 let kd = nk * dk;
1649 let rep = nv / nk;
1650 let vd = nv * dv;
1651 let sqdk = (dk as f64).sqrt();
1652 for hh in 0..rep {
1653 let h = ko * rep + hh;
1654 let ea = F::fromf((cfg.a_log[h] as f64).exp());
1655
1656 let mut s_hist = vec![F::ZERO; (t + 1) * dk * dv]; let mut kv_hist = vec![F::ZERO; t * dv];
1659 let mut o_hist = vec![F::ZERO; t * dv];
1660 let mut g_v = vec![F::ZERO; t];
1661 let mut beta_v = vec![F::ZERO; t];
1662 let mut sp_arg = vec![F::ZERO; t]; for ti in 0..t {
1664 let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
1665 let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
1666 let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
1667 let (invq, _) = gdn_inv(qrow, sqdk);
1668 let (invk, _) = gdn_inv(krow, 1.0);
1669 let arg = a[ti * nv + h] + F::fromf(cfg.dt_bias[h] as f64);
1670 let g = (-ea * softplus_f(arg)).exp();
1671 let beta = sigmoid_f(b[ti * nv + h]);
1672 sp_arg[ti] = arg;
1673 g_v[ti] = g;
1674 beta_v[ti] = beta;
1675 let (prev, cur) = s_hist.split_at_mut((ti + 1) * dk * dv);
1676 let sp = &prev[ti * dk * dv..];
1677 let sn = &mut cur[..dk * dv];
1678 let kvr = &mut kv_hist[ti * dv..(ti + 1) * dv];
1679 for di in 0..dk {
1680 let kf = krow[di] * invk;
1681 for dj in 0..dv {
1682 let dec = sp[di * dv + dj] * g;
1683 sn[di * dv + dj] = dec;
1684 kvr[dj] += dec * kf;
1685 }
1686 }
1687 let or = &mut o_hist[ti * dv..(ti + 1) * dv];
1688 for di in 0..dk {
1689 let kf = krow[di] * invk;
1690 let qf = qrow[di] * invq;
1691 for dj in 0..dv {
1692 let sv = sn[di * dv + dj] + kf * (vrow[dj] - kvr[dj]) * beta_v[ti];
1693 sn[di * dv + dj] = sv;
1694 or[dj] += qf * sv;
1695 }
1696 }
1697 }
1698
1699 let mut ds = vec![F::ZERO; dk * dv];
1701 let mut do_o = vec![F::ZERO; dv];
1702 let mut du = vec![F::ZERO; dv];
1703 let mut dkv = vec![F::ZERO; dv];
1704 let mut dqh = vec![F::ZERO; dk];
1705 let mut dkh = vec![F::ZERO; dk];
1706 for ti in (0..t).rev() {
1707 let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
1708 let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
1709 let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
1710 let (invq, nq2) = gdn_inv(qrow, sqdk);
1711 let (invk, nk2) = gdn_inv(krow, 1.0);
1712 let g = g_v[ti];
1713 let beta = beta_v[ti];
1714 let s_t = &s_hist[(ti + 1) * dk * dv..(ti + 2) * dk * dv];
1715 let s_prev = &s_hist[ti * dk * dv..(ti + 1) * dk * dv];
1716 let kvr = &kv_hist[ti * dv..(ti + 1) * dv];
1717 let or = &o_hist[ti * dv..(ti + 1) * dv];
1718
1719 let mut ss = 0f64;
1721 for v in or {
1722 ss += v.f64() * v.f64();
1723 }
1724 let inv = F::fromf(1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt());
1725 let dofr = &dout[ti * vd + h * dv..ti * vd + (h + 1) * dv];
1726 let mut sdot = 0f64; for dj in 0..dv {
1729 let zv = z[ti * vd + h * dv + dj];
1730 let w = F::fromf(cfg.norm[dj] as f64);
1731 let weff = w * silu(zv);
1732 sdot += (dofr[dj] * weff * or[dj]).f64();
1733 dz[ti * vd + h * dv + dj] += dofr[dj] * or[dj] * inv * w * silu_bwd(zv);
1734 }
1735 let coef = F::fromf(sdot / dv as f64) * inv * inv * inv;
1736 for dj in 0..dv {
1737 let zv = z[ti * vd + h * dv + dj];
1738 let weff = F::fromf(cfg.norm[dj] as f64) * silu(zv);
1739 do_o[dj] = inv * weff * dofr[dj] - or[dj] * coef;
1740 }
1741
1742 for x in dqh.iter_mut() {
1744 *x = F::ZERO;
1745 }
1746 for di in 0..dk {
1747 let qf = qrow[di] * invq;
1748 let row = &s_t[di * dv..(di + 1) * dv];
1749 let dsr = &mut ds[di * dv..(di + 1) * dv];
1750 let mut acc = F::ZERO;
1751 for dj in 0..dv {
1752 dsr[dj] += qf * do_o[dj];
1753 acc += row[dj] * do_o[dj];
1754 }
1755 dqh[di] = acc;
1756 }
1757
1758 for x in du.iter_mut() {
1761 *x = F::ZERO;
1762 }
1763 for x in dkh.iter_mut() {
1764 *x = F::ZERO;
1765 }
1766 for di in 0..dk {
1767 let kf = krow[di] * invk;
1768 let dsr = &ds[di * dv..(di + 1) * dv];
1769 let mut acc = F::ZERO;
1770 for dj in 0..dv {
1771 du[dj] += dsr[dj] * kf;
1772 acc += dsr[dj] * (vrow[dj] - kvr[dj]) * beta;
1773 }
1774 dkh[di] = acc;
1775 }
1776 let mut dbeta = F::ZERO;
1777 for dj in 0..dv {
1778 dbeta += du[dj] * (vrow[dj] - kvr[dj]);
1779 dcq[ti * c_dim + 2 * kd + h * dv + dj] += beta * du[dj];
1781 dkv[dj] = -(beta * du[dj]);
1782 }
1783
1784 let mut dg = F::ZERO;
1789 for di in 0..dk {
1790 let kf = krow[di] * invk;
1791 let spr = &s_prev[di * dv..(di + 1) * dv];
1792 let dsr = &mut ds[di * dv..(di + 1) * dv];
1793 let mut acc = F::ZERO;
1794 for dj in 0..dv {
1795 let dspre = dsr[dj] + kf * dkv[dj];
1796 acc += (spr[dj] * g) * dkv[dj];
1797 dg += dspre * spr[dj];
1798 dsr[dj] = g * dspre;
1799 }
1800 dkh[di] += acc;
1801 }
1802
1803 let sig = sigmoid_f(sp_arg[ti]);
1805 da[ti * nv + h] += dg * g * (-ea) * sig;
1806 db[ti * nv + h] += dbeta * beta * (F::ONE - beta);
1807
1808 let mut qdot = F::ZERO;
1812 let mut kdot = F::ZERO;
1813 for di in 0..dk {
1814 qdot += dqh[di] * qrow[di];
1815 kdot += dkh[di] * krow[di];
1816 }
1817 for di in 0..dk {
1818 dcq[ti * c_dim + ko * dk + di] += invq * dqh[di] - qrow[di] * qdot * invq / nq2;
1819 dcq[ti * c_dim + kd + ko * dk + di] +=
1820 invk * dkh[di] - krow[di] * kdot * invk / nk2;
1821 }
1822 }
1823 }
1824}
1825
1826pub fn gdn_seq_fwd<F: Fp>(
1830 qkv: &[F],
1831 z: &[F],
1832 a: &[F],
1833 b: &[F],
1834 t: usize,
1835 cfg: &GdnSeqCfg,
1836 out: &mut [F],
1837) {
1838 let c_dim = cfg.c_dim();
1839 let mut pre = vec![F::ZERO; t * c_dim];
1840 let mut cq = vec![F::ZERO; t * c_dim];
1841 gdn_conv_fwd(qkv, t, c_dim, cfg.kk, cfg.conv, &mut pre, &mut cq);
1842 for ko in 0..cfg.nk {
1843 gdn_group_fwd(&cq, z, a, b, t, cfg, ko, out);
1844 }
1845}
1846
1847#[allow(clippy::too_many_arguments)]
1850pub fn gdn_seq_bwd<F: Fp>(
1851 qkv: &[F],
1852 z: &[F],
1853 a: &[F],
1854 b: &[F],
1855 t: usize,
1856 cfg: &GdnSeqCfg,
1857 dout: &[F],
1858 dqkv: &mut [F],
1859 dz: &mut [F],
1860 da: &mut [F],
1861 db: &mut [F],
1862) {
1863 let c_dim = cfg.c_dim();
1864 let mut pre = vec![F::ZERO; t * c_dim];
1865 let mut cq = vec![F::ZERO; t * c_dim];
1866 gdn_conv_fwd(qkv, t, c_dim, cfg.kk, cfg.conv, &mut pre, &mut cq);
1867 let mut dcq = vec![F::ZERO; t * c_dim];
1868 for ko in 0..cfg.nk {
1869 gdn_group_bwd(&cq, z, a, b, t, cfg, ko, dout, &mut dcq, dz, da, db);
1870 }
1871 gdn_conv_bwd(&pre, t, c_dim, cfg.kk, cfg.conv, &dcq, dqkv);
1872}
1873
1874
1875#[cfg(test)]
1876mod gpu_bake_tests {
1877 #[test]
1881 #[cfg(feature = "gpu")]
1882 fn gemm_dx_gpu_matches_cpu() {
1883 if std::env::var("CMF_GPU").is_err() {
1884 eprintln!("no backend requested — skip");
1885 return;
1886 }
1887 let (n, k, m) = (8usize, 512usize, 1024usize);
1888 let dy: Vec<f32> = (0..n * m).map(|i| ((i * 37 % 101) as f32 - 50.0) / 50.0).collect();
1889 let w: Vec<f32> = (0..m * k).map(|i| ((i * 17 % 97) as f32 - 48.0) / 48.0).collect();
1890 let mut want = vec![0f32; n * k];
1891 super::gemm_dx(&dy, &w, &mut want, n, k, m, None);
1892 let mut got = vec![0f32; n * k];
1893 if !crate::gpu_wgpu::gemm_dx_f32(&dy, &w, &mut got, n, k, m) {
1894 eprintln!("device declined — skip");
1895 return;
1896 }
1897 let num: f64 = want.iter().zip(&got).map(|(a, b)| ((a - b) as f64).powi(2)).sum();
1898 let den: f64 = want.iter().map(|a| (*a as f64).powi(2)).sum::<f64>().max(1e-30);
1899 let rel = (num / den).sqrt();
1900 assert!(rel < 1e-4, "gemm_dx GPU vs CPU rel {rel:e}");
1901 }
1902}