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