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
335pub(crate) fn gemm_nt_host(x: &[f32], w: &[f32], y: &mut [f32], n: usize, k: usize, m: usize, pool: Option<&Pool>) {
338 gemm_nt_cpu(x, w, y, n, k, m, pool)
339}
340
341fn gemm_nt_cpu(
342 x: &[f32],
343 w: &[f32],
344 y: &mut [f32],
345 n: usize,
346 k: usize,
347 m: usize,
348 pool: Option<&Pool>,
349) {
350 debug_assert_eq!(x.len(), n * k);
351 debug_assert_eq!(w.len(), m * k);
352 debug_assert_eq!(y.len(), n * m);
353 #[cfg(target_os = "macos")]
354 if accel::on() && n * k * m >= 1 << 18 {
355 unsafe {
357 accel::cblas_sgemm(
358 101,
359 111,
360 112,
361 n as i32,
362 m as i32,
363 k as i32,
364 1.0,
365 x.as_ptr(),
366 k as i32,
367 w.as_ptr(),
368 k as i32,
369 0.0,
370 y.as_mut_ptr(),
371 m as i32,
372 );
373 }
374 return;
375 }
376 let nb = n.div_ceil(GEMM_BLOCK);
377 let block = |r0: usize, r1: usize, y: &mut [f32]| {
378 let mut o = 0usize;
390 #[cfg(target_arch = "x86_64")]
391 if gemm_avx512() {
392 while o + 4 <= m {
393 let mut i = r0;
394 while i + 4 <= r1 {
395 let acc = unsafe { gemm_block4x4_avx512(x, w, i, o, k) };
398 for (di, ai) in acc.iter().enumerate() {
399 for (dj, v) in ai.iter().enumerate() {
400 y[(i - r0 + di) * m + o + dj] = *v;
401 }
402 }
403 i += 4;
404 }
405 for i in i..r1 {
406 let xr = &x[i * k..(i + 1) * k];
407 for dj in 0..4 {
408 let wr = &w[(o + dj) * k..(o + dj + 1) * k];
409 y[(i - r0) * m + o + dj] = crate::attention::dot_f32(xr, wr);
410 }
411 }
412 o += 4;
413 }
414 }
415 for o in o..m {
416 let wr = &w[o * k..(o + 1) * k];
417 for i in r0..r1 {
418 y[(i - r0) * m + o] = crate::attention::dot_f32(&x[i * k..(i + 1) * k], wr);
419 }
420 }
421 };
422 match pool {
423 Some(p) if nb > 1 => {
424 let yp = SendMut(y.as_mut_ptr());
425 p.run(&|widx, nw| {
426 for bi in (widx..nb).step_by(nw) {
427 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
428 let ys = unsafe { yp.slice(r0 * m, (r1 - r0) * m) };
430 block(r0, r1, ys);
431 }
432 });
433 }
434 _ => {
435 for bi in 0..nb {
436 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
437 block(r0, r1, &mut y[r0 * m..r1 * m]);
438 }
439 }
440 }
441}
442
443pub fn gemm_dx(
445 dy: &[f32],
446 w: &[f32],
447 dx: &mut [f32],
448 n: usize,
449 k: usize,
450 m: usize,
451 pool: Option<&Pool>,
452) {
453 let _t = std::time::Instant::now();
454 crate::fcd::prof::GEMM_CALLS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
455 let _guard = scopeguard_gemm(_t, n, k, m);
456 debug_assert_eq!(dy.len(), n * m);
457 debug_assert_eq!(w.len(), m * k);
458 debug_assert_eq!(dx.len(), n * k);
459 #[cfg(feature = "gpu")]
463 if crate::gpu_wgpu::gemm_dx_f32(dy, w, dx, n, k, m) {
464 return;
465 }
466 #[cfg(target_os = "macos")]
467 if accel::on() && n * k * m >= 1 << 18 {
468 unsafe {
470 accel::cblas_sgemm(
471 101,
472 111,
473 111,
474 n as i32,
475 k as i32,
476 m as i32,
477 1.0,
478 dy.as_ptr(),
479 m as i32,
480 w.as_ptr(),
481 k as i32,
482 1.0,
483 dx.as_mut_ptr(),
484 k as i32,
485 );
486 }
487 return;
488 }
489 let nb = n.div_ceil(GEMM_BLOCK);
490 let block = |r0: usize, r1: usize, dxs: &mut [f32]| {
491 for o in 0..m {
492 let wr = &w[o * k..(o + 1) * k];
493 for i in r0..r1 {
494 let g = dy[i * m + o];
495 if g != 0.0 {
496 crate::attention::axpy_f32(&mut dxs[(i - r0) * k..(i - r0 + 1) * k], wr, g);
497 }
498 }
499 }
500 };
501 match pool {
502 Some(p) if nb > 1 => {
503 let dxp = SendMut(dx.as_mut_ptr());
504 p.run(&|widx, nw| {
505 for bi in (widx..nb).step_by(nw) {
506 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
507 let dxs = unsafe { dxp.slice(r0 * k, (r1 - r0) * k) };
509 block(r0, r1, dxs);
510 }
511 });
512 }
513 _ => {
514 for bi in 0..nb {
515 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
516 block(r0, r1, &mut dx[r0 * k..r1 * k]);
517 }
518 }
519 }
520}
521
522pub fn gemm_dw(
525 dy: &[f32],
526 x: &[f32],
527 dw: &mut [f32],
528 n: usize,
529 k: usize,
530 m: usize,
531 pool: Option<&Pool>,
532) {
533 let _t = std::time::Instant::now();
534 crate::fcd::prof::GEMM_CALLS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
535 let _guard = scopeguard_gemm(_t, n, k, m);
536 debug_assert_eq!(dy.len(), n * m);
537 debug_assert_eq!(x.len(), n * k);
538 debug_assert_eq!(dw.len(), m * k);
539 #[cfg(target_os = "macos")]
540 if accel::on() && n * k * m >= 1 << 18 {
541 unsafe {
543 accel::cblas_sgemm(
544 101,
545 112,
546 111,
547 m as i32,
548 k as i32,
549 n as i32,
550 1.0,
551 dy.as_ptr(),
552 m as i32,
553 x.as_ptr(),
554 k as i32,
555 1.0,
556 dw.as_mut_ptr(),
557 k as i32,
558 );
559 }
560 return;
561 }
562 let range = |o0: usize, o1: usize, dws: &mut [f32]| {
563 let nb = n.div_ceil(GEMM_BLOCK);
565 for bi in 0..nb {
566 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
567 for o in o0..o1 {
568 let dwr = &mut dws[(o - o0) * k..(o - o0 + 1) * k];
569 for i in r0..r1 {
570 let g = dy[i * m + o];
571 if g != 0.0 {
572 crate::attention::axpy_f32(dwr, &x[i * k..(i + 1) * k], g);
573 }
574 }
575 }
576 }
577 };
578 match pool {
579 Some(p) if m >= 8 => {
580 let dwp = SendMut(dw.as_mut_ptr());
581 p.run(&|widx, nw| {
582 let (o0, o1) = (widx * m / nw, (widx + 1) * m / nw);
583 if o0 < o1 {
584 let dws = unsafe { dwp.slice(o0 * k, (o1 - o0) * k) };
586 range(o0, o1, dws);
587 }
588 });
589 }
590 _ => range(0, m, dw),
591 }
592}
593
594#[inline]
597pub fn silu<F: Fp>(x: F) -> F {
598 x / (F::ONE + (-x).exp())
599}
600
601#[inline]
603pub fn silu_bwd<F: Fp>(x: F) -> F {
604 let s = F::ONE / (F::ONE + (-x).exp());
605 s * (F::ONE + x * (F::ONE - s))
606}
607
608pub fn rmsnorm_fwd<F: Fp>(x: &[F], w: &[F], eps: f64, gemma: bool, y: &mut [F], inv_out: &mut [F]) {
615 let d = w.len();
616 let n = x.len() / d;
617 for r in 0..n {
618 let xr = &x[r * d..(r + 1) * d];
619 let mut ss = 0f64;
620 for v in xr {
621 ss += v.f64() * v.f64();
622 }
623 let inv = F::fromf(1.0 / (ss / d as f64 + eps).sqrt());
624 inv_out[r] = inv;
625 let yr = &mut y[r * d..(r + 1) * d];
626 for j in 0..d {
627 let weff = if gemma { F::ONE + w[j] } else { w[j] };
628 yr[j] = xr[j] * inv * weff;
629 }
630 }
631}
632
633pub fn rmsnorm_bwd<F: Fp>(
639 x: &[F],
640 w: &[F],
641 inv: &[F],
642 dy: &[F],
643 gemma: bool,
644 dx: &mut [F],
645 mut dw: Option<&mut [F]>,
646) {
647 let d = w.len();
648 let n = x.len() / d;
649 for r in 0..n {
650 let xr = &x[r * d..(r + 1) * d];
651 let dyr = &dy[r * d..(r + 1) * d];
652 let iv = inv[r];
653 let mut s = 0f64;
654 for j in 0..d {
655 let weff = if gemma { F::ONE + w[j] } else { w[j] };
656 s += (dyr[j] * weff * xr[j]).f64();
657 }
658 let coef = F::fromf(s / d as f64) * iv * iv * iv;
659 let dxr = &mut dx[r * d..(r + 1) * d];
660 for j in 0..d {
661 let weff = if gemma { F::ONE + w[j] } else { w[j] };
662 dxr[j] += iv * weff * dyr[j] - xr[j] * coef;
663 }
664 if let Some(dwv) = dw.as_deref_mut() {
665 for j in 0..d {
666 dwv[j] += dyr[j] * xr[j] * iv;
667 }
668 }
669 }
670}
671
672pub fn rope_fwd<F: Fp>(x: &mut [F], position: usize, inv_freq: &[f64]) {
677 let half = inv_freq.len();
678 for (i, &freq) in inv_freq.iter().enumerate() {
679 let angle = position as f64 * freq;
680 let (sin, cos) = (F::fromf(angle.sin()), F::fromf(angle.cos()));
681 let x0 = x[i];
682 let x1 = x[i + half];
683 x[i] = x0 * cos - x1 * sin;
684 x[i + half] = x0 * sin + x1 * cos;
685 }
686}
687
688pub fn rope_bwd<F: Fp>(dy: &mut [F], position: usize, inv_freq: &[f64]) {
691 let half = inv_freq.len();
692 for (i, &freq) in inv_freq.iter().enumerate() {
693 let angle = position as f64 * freq;
694 let (sin, cos) = (F::fromf(angle.sin()), F::fromf(angle.cos()));
695 let g0 = dy[i];
696 let g1 = dy[i + half];
697 dy[i] = g0 * cos + g1 * sin;
698 dy[i + half] = -g0 * sin + g1 * cos;
699 }
700}
701
702pub fn seg_means<F: Fp>(x: &[F], t: usize, d: usize, m: usize, out: &mut [F]) {
707 for i in 0..m {
708 let (lo, hi) = (i * t / m, (i + 1) * t / m);
709 let or = &mut out[i * d..(i + 1) * d];
710 for v in or.iter_mut() {
711 *v = F::ZERO;
712 }
713 for j in lo..hi {
714 for c in 0..d {
715 or[c] += x[j * d + c];
716 }
717 }
718 let inv = F::fromf(1.0 / (hi - lo) as f64);
719 for v in or.iter_mut() {
720 *v *= inv;
721 }
722 }
723}
724
725pub fn seg_means_bwd<F: Fp>(dl: &[F], t: usize, d: usize, m: usize, dx: &mut [F]) {
728 for i in 0..m {
729 let (lo, hi) = (i * t / m, (i + 1) * t / m);
730 let inv = F::fromf(1.0 / (hi - lo) as f64);
731 let dlr = &dl[i * d..(i + 1) * d];
732 for j in lo..hi {
733 for c in 0..d {
734 dx[j * d + c] += dlr[c] * inv;
735 }
736 }
737 }
738}
739
740#[allow(clippy::needless_range_loop)] pub fn attn_head_fwd<F: Fp>(
746 q: &[F],
747 k: &[F],
748 v: &[F],
749 t: usize,
750 d: usize,
751 dv: usize,
752 out: &mut [F],
753) {
754 let scale = F::fromf(1.0 / (d as f64).sqrt());
755 let mut row = vec![F::ZERO; t];
756 for ti in 0..t {
757 let qr = &q[ti * d..(ti + 1) * d];
758 let mut mx = F::fromf(f64::NEG_INFINITY);
759 for j in 0..=ti {
760 let s = dot(qr, &k[j * d..(j + 1) * d]) * scale;
761 row[j] = s;
762 mx = mx.maxf(s);
763 }
764 let mut den = F::ZERO;
765 for j in 0..=ti {
766 row[j] = (row[j] - mx).exp();
767 den += row[j];
768 }
769 let or = &mut out[ti * dv..(ti + 1) * dv];
770 for o in or.iter_mut() {
771 *o = F::ZERO;
772 }
773 for j in 0..=ti {
774 let p = row[j] / den;
775 for (o, vv) in or.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
776 *o += p * *vv;
777 }
778 }
779 }
780}
781
782#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
786pub fn attn_head_bwd<F: Fp>(
787 q: &[F],
788 k: &[F],
789 v: &[F],
790 dout: &[F],
791 t: usize,
792 d: usize,
793 dv: usize,
794 dq: &mut [F],
795 dk: &mut [F],
796 dvv: &mut [F],
797) {
798 let scale = F::fromf(1.0 / (d as f64).sqrt());
799 let mut row = vec![F::ZERO; t];
800 for ti in 0..t {
801 let qr = &q[ti * d..(ti + 1) * d];
802 let mut mx = F::fromf(f64::NEG_INFINITY);
803 for j in 0..=ti {
804 let s = dot(qr, &k[j * d..(j + 1) * d]) * scale;
805 row[j] = s;
806 mx = mx.maxf(s);
807 }
808 let mut den = F::ZERO;
809 for j in 0..=ti {
810 row[j] = (row[j] - mx).exp();
811 den += row[j];
812 }
813 let dor = &dout[ti * dv..(ti + 1) * dv];
814 let mut pdp = F::ZERO;
816 let mut dp = vec![F::ZERO; ti + 1];
817 for j in 0..=ti {
818 let p = row[j] / den;
819 row[j] = p; dp[j] = dot(dor, &v[j * dv..(j + 1) * dv]);
821 pdp += p * dp[j];
822 }
823 let dqr = &mut dq[ti * d..(ti + 1) * d];
824 for j in 0..=ti {
825 let p = row[j];
826 for (dvo, o) in dvv[j * dv..(j + 1) * dv].iter_mut().zip(dor) {
828 *dvo += p * *o;
829 }
830 let ds = p * (dp[j] - pdp) * scale;
832 let kr = &k[j * d..(j + 1) * d];
833 for c in 0..d {
834 dqr[c] += ds * kr[c];
835 }
836 let dkr = &mut dk[j * d..(j + 1) * d];
837 for c in 0..d {
838 dkr[c] += ds * qr[c];
839 }
840 }
841 }
842}
843
844#[derive(Clone, Copy, Debug)]
848pub struct NysCfg {
849 pub m: usize,
851 pub w: usize,
853 pub sink: usize,
855 pub prefill: Option<usize>,
867}
868
869impl NysCfg {
870 #[inline]
873 pub fn prefill_len(&self, t: usize) -> usize {
874 self.prefill.unwrap_or(t / 2).clamp(1, t)
875 }
876}
877
878const NYS_DEN_EPS: f64 = 1e-30;
880
881struct NysGraph<F: Fp> {
885 m_eff: usize,
886 tp: usize,
889 q_l: Vec<F>,
890 k_l: Vec<F>,
891 mu: Vec<F>,
892 fu: Vec<F>,
893 e: Vec<F>,
894 fumu: Vec<F>,
895 far_keep: Vec<bool>,
904 wmat: Vec<F>,
907 c_row: Vec<F>,
908 den: Vec<F>,
909}
910
911#[inline]
916fn nys_exact_row(ti: usize, tp: usize) -> bool {
917 ti < tp
918}
919
920#[inline]
922fn nys_near(ti: usize, j: usize, w: usize, sink: usize) -> bool {
923 ti - j < w || j < sink
924}
925
926fn nys_graph<F: Fp>(
931 q: &[F],
932 k: &[F],
933 t: usize,
934 d: usize,
935 cfg: &NysCfg,
936 mu_override: Option<&[F]>,
937) -> NysGraph<F> {
938 let scale = 1.0 / (d as f64).sqrt();
939 let fscale = F::fromf(scale);
940 let tp = cfg.prefill_len(t);
945 let m_eff = (tp / 8).clamp(4, cfg.m);
946 let mut q_l = vec![F::ZERO; m_eff * d];
947 let mut k_l = vec![F::ZERO; m_eff * d];
948 seg_means(&q[..tp * d], tp, d, m_eff, &mut q_l);
949 seg_means(&k[..tp * d], tp, d, m_eff, &mut k_l);
950
951 let mut au = vec![0f64; m_eff * m_eff];
953 for i in 0..m_eff {
954 for j in 0..m_eff {
955 let mut s = 0f64;
956 for c in 0..d {
957 s += q_l[i * d + c].f64() * k_l[j * d + c].f64();
958 }
959 au[i * m_eff + j] = (s * scale).exp();
960 }
961 }
962 let mu: Vec<F> = match mu_override {
963 Some(m) => m.to_vec(),
964 None => crate::nystrom::ridge_pinv(&au, m_eff)
965 .iter()
966 .map(|&x| F::fromf(x))
967 .collect(),
968 };
969
970 let mut fu = vec![F::ZERO; t * m_eff];
972 for ti in 0..t {
973 for i in 0..m_eff {
974 fu[ti * m_eff + i] =
975 (dot(&q[ti * d..(ti + 1) * d], &k_l[i * d..(i + 1) * d]) * fscale).exp();
976 }
977 }
978 let mut e = vec![F::ZERO; m_eff * t];
979 for i in 0..m_eff {
980 for j in 0..t {
981 e[i * t + j] = (dot(&q_l[i * d..(i + 1) * d], &k[j * d..(j + 1) * d]) * fscale).exp();
982 }
983 }
984 let mut fumu = vec![F::ZERO; t * m_eff];
985 matmul_nt(
986 &fu,
987 &transpose(&mu, m_eff, m_eff),
990 &mut fumu,
991 t,
992 m_eff,
993 m_eff,
994 );
995 let mut a = vec![F::ZERO; t * t];
1000 for ti in tp..t {
1001 let fr = &fumu[ti * m_eff..(ti + 1) * m_eff];
1002 let ar = &mut a[ti * t..ti * t + ti + 1]; for (i, &f) in fr.iter().enumerate() {
1004 let er = &e[i * t..i * t + ti + 1];
1005 for (av, ev) in ar.iter_mut().zip(er) {
1006 *av += f * *ev;
1007 }
1008 }
1009 }
1010
1011 let mut wmat = vec![F::ZERO; t * t];
1014 let mut c_row = vec![F::ZERO; t];
1015 let mut den = vec![F::ZERO; t];
1016 let mut far_keep = vec![false; t];
1017 let mut lg_row = vec![F::ZERO; t];
1018 for ti in 0..t {
1019 let qr = &q[ti * d..(ti + 1) * d];
1020 let mut c = F::fromf(f64::NEG_INFINITY);
1021 for j in 0..=ti {
1022 let s = dot(qr, &k[j * d..(j + 1) * d]) * fscale;
1023 lg_row[j] = s;
1024 c = c.maxf(s);
1025 }
1026 c_row[ti] = c;
1027 let emc = (-c).exp();
1028 let exact_row = nys_exact_row(ti, tp);
1029
1030 let mut far_sum = F::ZERO;
1045 if !exact_row {
1046 for j in 0..=ti {
1047 if !nys_near(ti, j, cfg.w, cfg.sink) {
1048 far_sum += a[ti * t + j];
1049 }
1050 }
1051 }
1052 let keep = !exact_row && far_sum.f64() >= 0.0;
1053 far_keep[ti] = keep;
1054
1055 let wr = &mut wmat[ti * t..(ti + 1) * t];
1056 let mut dsum = F::ZERO;
1057 for j in 0..=ti {
1058 let wv = if exact_row || nys_near(ti, j, cfg.w, cfg.sink) {
1059 (lg_row[j] - c).exp()
1060 } else if keep {
1061 a[ti * t + j] * emc
1062 } else {
1063 F::ZERO
1064 };
1065 wr[j] = wv;
1066 dsum += wv;
1067 }
1068 den[ti] = dsum.maxf(F::fromf(NYS_DEN_EPS));
1069 }
1070 NysGraph {
1071 m_eff,
1072 tp,
1073 q_l,
1074 k_l,
1075 mu,
1076 fu,
1077 e,
1078 fumu,
1079 far_keep,
1080 wmat,
1081 c_row,
1082 den,
1083 }
1084}
1085
1086fn transpose<F: Fp>(x: &[F], rows: usize, cols: usize) -> Vec<F> {
1087 let mut out = vec![F::ZERO; rows * cols];
1088 for r in 0..rows {
1089 for c in 0..cols {
1090 out[c * rows + r] = x[r * cols + c];
1091 }
1092 }
1093 out
1094}
1095
1096#[allow(clippy::too_many_arguments)]
1099pub fn nystrom_head_fwd<F: Fp>(
1100 q: &[F],
1101 k: &[F],
1102 v: &[F],
1103 t: usize,
1104 d: usize,
1105 dv: usize,
1106 cfg: &NysCfg,
1107 out: &mut [F],
1108) {
1109 if nys_degenerate(t, cfg) {
1110 attn_head_fwd(q, k, v, t, d, dv, out);
1111 return;
1112 }
1113 nystrom_head_fwd_mu(q, k, v, t, d, dv, cfg, None, out);
1114}
1115
1116#[inline]
1125fn nys_degenerate(t: usize, cfg: &NysCfg) -> bool {
1126 cfg.prefill_len(t) <= cfg.w + cfg.sink + 8
1127}
1128
1129#[doc(hidden)]
1132pub fn nystrom_mu_for_test<F: Fp>(q: &[F], k: &[F], t: usize, d: usize, cfg: &NysCfg) -> Vec<F> {
1133 nys_graph(q, k, t, d, cfg, None).mu
1134}
1135
1136#[doc(hidden)]
1139#[allow(clippy::too_many_arguments)]
1140pub fn nystrom_head_fwd_mu<F: Fp>(
1141 q: &[F],
1142 k: &[F],
1143 v: &[F],
1144 t: usize,
1145 d: usize,
1146 dv: usize,
1147 cfg: &NysCfg,
1148 mu_override: Option<&[F]>,
1149 out: &mut [F],
1150) {
1151 let g = nys_graph(q, k, t, d, cfg, mu_override);
1152 for ti in 0..t {
1153 let wr = &g.wmat[ti * t..(ti + 1) * t];
1154 let den = g.den[ti];
1155 let or = &mut out[ti * dv..(ti + 1) * dv];
1156 for o in or.iter_mut() {
1157 *o = F::ZERO;
1158 }
1159 for j in 0..=ti {
1160 let p = wr[j] / den;
1161 if p.f64() != 0.0 {
1162 for (o, vv) in or.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
1163 *o += p * *vv;
1164 }
1165 }
1166 }
1167 }
1168}
1169
1170#[allow(clippy::too_many_arguments)]
1181pub fn nystrom_head_bwd<F: Fp>(
1182 q: &[F],
1183 k: &[F],
1184 v: &[F],
1185 dout: &[F],
1186 t: usize,
1187 d: usize,
1188 dv: usize,
1189 cfg: &NysCfg,
1190 dq: &mut [F],
1191 dk: &mut [F],
1192 dvv: &mut [F],
1193) {
1194 if nys_degenerate(t, cfg) {
1195 attn_head_bwd(q, k, v, dout, t, d, dv, dq, dk, dvv);
1196 return;
1197 }
1198 nystrom_head_bwd_mu(q, k, v, dout, t, d, dv, cfg, None, dq, dk, dvv);
1199}
1200
1201#[doc(hidden)]
1203#[allow(clippy::too_many_arguments)]
1204pub fn nystrom_head_bwd_mu<F: Fp>(
1205 q: &[F],
1206 k: &[F],
1207 v: &[F],
1208 dout: &[F],
1209 t: usize,
1210 d: usize,
1211 dv: usize,
1212 cfg: &NysCfg,
1213 mu_override: Option<&[F]>,
1214 dq: &mut [F],
1215 dk: &mut [F],
1216 dvv: &mut [F],
1217) {
1218 let scale = F::fromf(1.0 / (d as f64).sqrt());
1219 let g = nys_graph(q, k, t, d, cfg, mu_override);
1220 let m_eff = g.m_eff;
1221
1222 let mut dwmat = vec![F::ZERO; t * t];
1224 let mut out_row = vec![F::ZERO; dv];
1225 for ti in 0..t {
1226 let wr = &g.wmat[ti * t..(ti + 1) * t];
1227 let den = g.den[ti];
1228 for o in out_row.iter_mut() {
1229 *o = F::ZERO;
1230 }
1231 for j in 0..=ti {
1232 let p = wr[j] / den;
1233 if p.f64() != 0.0 {
1234 for (o, vv) in out_row.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
1235 *o += p * *vv;
1236 }
1237 }
1238 }
1239 let dor = &dout[ti * dv..(ti + 1) * dv];
1240 let dwr = &mut dwmat[ti * t..(ti + 1) * t];
1241 for j in 0..=ti {
1242 let p = wr[j] / den;
1244 if p.f64() != 0.0 {
1245 for (dvo, o) in dvv[j * dv..(j + 1) * dv].iter_mut().zip(dor) {
1246 *dvo += p * *o;
1247 }
1248 }
1249 let mut s = F::ZERO;
1251 for c in 0..dv {
1252 s += dor[c] * (v[j * dv + c] - out_row[c]);
1253 }
1254 dwr[j] = s / den;
1255 }
1256 }
1257
1258 for ti in 0..t {
1261 let qr = &q[ti * d..(ti + 1) * d];
1262 let dqr_base = ti * d;
1263 for j in 0..=ti {
1264 if !(nys_exact_row(ti, g.tp) || nys_near(ti, j, cfg.w, cfg.sink)) {
1265 continue;
1266 }
1267 let dlg = dwmat[ti * t + j] * g.wmat[ti * t + j] * scale;
1268 if dlg.f64() == 0.0 {
1269 continue;
1270 }
1271 let kr = &k[j * d..(j + 1) * d];
1272 for c in 0..d {
1273 dq[dqr_base + c] += dlg * kr[c];
1274 }
1275 let dkr = &mut dk[j * d..(j + 1) * d];
1276 for c in 0..d {
1277 dkr[c] += dlg * qr[c];
1278 }
1279 }
1280 }
1281
1282 let mut da = vec![F::ZERO; t * t];
1297 for ti in g.tp..t {
1298 if !g.far_keep[ti] {
1299 continue; }
1301 let emc = (-g.c_row[ti]).exp();
1302 for j in 0..=ti {
1303 if nys_near(ti, j, cfg.w, cfg.sink) {
1304 continue;
1305 }
1306 da[ti * t + j] = dwmat[ti * t + j] * emc;
1307 }
1308 }
1309 let mut dfumu = vec![F::ZERO; t * m_eff];
1311 for ti in 0..t {
1312 let dar = &da[ti * t..ti * t + ti + 1];
1313 let dfr = &mut dfumu[ti * m_eff..(ti + 1) * m_eff];
1314 for (i, df) in dfr.iter_mut().enumerate() {
1315 let er = &g.e[i * t..i * t + ti + 1];
1316 let mut s = F::ZERO;
1317 for (av, ev) in dar.iter().zip(er) {
1318 s += *av * *ev;
1319 }
1320 *df = s;
1321 }
1322 }
1323 let mut dfu = vec![F::ZERO; t * m_eff];
1325 matmul_nt(&dfumu, &g.mu, &mut dfu, t, m_eff, m_eff);
1326 let mut de = vec![F::ZERO; m_eff * t];
1328 for ti in 0..t {
1329 let dar = &da[ti * t..ti * t + ti + 1];
1330 let fr = &g.fumu[ti * m_eff..(ti + 1) * m_eff];
1331 for (i, &f) in fr.iter().enumerate() {
1332 if f.f64() == 0.0 {
1333 continue;
1334 }
1335 let der = &mut de[i * t..i * t + ti + 1];
1336 for (dev, av) in der.iter_mut().zip(dar) {
1337 *dev += f * *av;
1338 }
1339 }
1340 }
1341 let mut dq_l = vec![F::ZERO; m_eff * d];
1344 let mut dk_l = vec![F::ZERO; m_eff * d];
1345 for ti in 0..t {
1346 let qr = &q[ti * d..(ti + 1) * d];
1347 for i in 0..m_eff {
1348 let dlg = dfu[ti * m_eff + i] * g.fu[ti * m_eff + i] * scale;
1349 if dlg.f64() == 0.0 {
1350 continue;
1351 }
1352 let klr = &g.k_l[i * d..(i + 1) * d];
1353 for c in 0..d {
1354 dq[ti * d + c] += dlg * klr[c];
1355 }
1356 let dklr = &mut dk_l[i * d..(i + 1) * d];
1357 for c in 0..d {
1358 dklr[c] += dlg * qr[c];
1359 }
1360 }
1361 }
1362 for i in 0..m_eff {
1363 let qlr = &g.q_l[i * d..(i + 1) * d];
1364 for j in 0..t {
1365 let dlg = de[i * t + j] * g.e[i * t + j] * scale;
1366 if dlg.f64() == 0.0 {
1367 continue;
1368 }
1369 let kr = &k[j * d..(j + 1) * d];
1370 let dqlr = &mut dq_l[i * d..(i + 1) * d];
1371 for c in 0..d {
1372 dqlr[c] += dlg * kr[c];
1373 }
1374 for c in 0..d {
1375 dk[j * d + c] += dlg * qlr[c];
1376 }
1377 }
1378 }
1379 let tp = g.tp;
1383 seg_means_bwd(&dq_l, tp, d, m_eff, &mut dq[..tp * d]);
1384 seg_means_bwd(&dk_l, tp, d, m_eff, &mut dk[..tp * d]);
1385}
1386
1387pub fn ce_kl_position<F: Fp>(
1395 s_logits: &[F],
1396 t_logits: &[F],
1397 target: usize,
1398 kl_w: f64,
1399 inv_n: f64,
1400 dlogits: &mut [F],
1401) -> (f64, f64) {
1402 let vsz = s_logits.len();
1403 debug_assert_eq!(t_logits.len(), vsz);
1404 let mut smax = f64::NEG_INFINITY;
1406 let mut tmax = f64::NEG_INFINITY;
1407 for i in 0..vsz {
1408 smax = smax.max(s_logits[i].f64());
1409 tmax = tmax.max(t_logits[i].f64());
1410 }
1411 let mut ssum = 0f64;
1412 let mut tsum = 0f64;
1413 for i in 0..vsz {
1414 ssum += (s_logits[i].f64() - smax).exp();
1415 tsum += (t_logits[i].f64() - tmax).exp();
1416 }
1417 let slz = smax + ssum.ln();
1418 let tlz = tmax + tsum.ln();
1419 let ce = slz - s_logits[target].f64();
1420 let mut kl = 0f64;
1421 for i in 0..vsz {
1422 let ls = s_logits[i].f64() - slz;
1423 let lt = t_logits[i].f64() - tlz;
1424 let pt = lt.exp();
1425 let ps = ls.exp();
1426 if pt > 0.0 {
1427 kl += pt * (lt - ls);
1428 }
1429 let mut gd = (1.0 - kl_w) * ps + kl_w * (ps - pt);
1430 if i == target {
1431 gd -= 1.0 - kl_w;
1432 }
1433 dlogits[i] = F::fromf(gd * inv_n);
1434 }
1435 (ce, kl)
1436}
1437
1438pub struct GdnSeqCfg<'a> {
1455 pub nv: usize,
1456 pub nk: usize,
1457 pub dk: usize,
1458 pub dv: usize,
1459 pub kk: usize,
1460 pub rms_eps: f64,
1461 pub conv: &'a [f32],
1464 pub a_log: &'a [f32],
1466 pub dt_bias: &'a [f32],
1468 pub norm: &'a [f32],
1470}
1471
1472impl GdnSeqCfg<'_> {
1473 pub fn c_dim(&self) -> usize {
1474 2 * self.nk * self.dk + self.nv * self.dv
1475 }
1476}
1477
1478#[inline]
1479fn softplus_f<F: Fp>(x: F) -> F {
1480 if x.f64() > 20.0 {
1483 x
1484 } else {
1485 F::fromf(x.f64().exp().ln_1p())
1486 }
1487}
1488
1489#[inline]
1490fn sigmoid_f<F: Fp>(x: F) -> F {
1491 F::ONE / (F::ONE + (-x).exp())
1492}
1493
1494pub fn gdn_conv_fwd<F: Fp>(
1498 raw: &[F],
1499 t: usize,
1500 c_dim: usize,
1501 kk: usize,
1502 conv: &[f32],
1503 pre: &mut [F],
1504 cq: &mut [F],
1505) {
1506 for ti in 0..t {
1507 for c in 0..c_dim {
1508 let taps = &conv[c * kk..(c + 1) * kk];
1509 let mut acc = F::ZERO;
1510 for (j, &tap) in taps.iter().enumerate() {
1511 let p = ti as isize - (kk as isize - 1) + j as isize;
1513 if p >= 0 {
1514 acc += raw[p as usize * c_dim + c] * F::fromf(tap as f64);
1515 }
1516 }
1517 pre[ti * c_dim + c] = acc;
1518 cq[ti * c_dim + c] = silu(acc);
1519 }
1520 }
1521}
1522
1523pub fn gdn_conv_bwd<F: Fp>(
1525 pre: &[F],
1526 t: usize,
1527 c_dim: usize,
1528 kk: usize,
1529 conv: &[f32],
1530 dcq: &[F],
1531 draw: &mut [F],
1532) {
1533 for ti in 0..t {
1534 for c in 0..c_dim {
1535 let g = dcq[ti * c_dim + c];
1536 if g.f64() == 0.0 {
1537 continue;
1538 }
1539 let dp = g * silu_bwd(pre[ti * c_dim + c]);
1540 let taps = &conv[c * kk..(c + 1) * kk];
1541 for (j, &tap) in taps.iter().enumerate() {
1542 let p = ti as isize - (kk as isize - 1) + j as isize;
1543 if p >= 0 {
1544 draw[p as usize * c_dim + c] += dp * F::fromf(tap as f64);
1545 }
1546 }
1547 }
1548 }
1549}
1550
1551#[inline]
1554fn gdn_inv<F: Fp>(x: &[F], extra_scale: f64) -> (F, F) {
1555 let mut n2 = F::ZERO;
1556 for v in x {
1557 n2 += *v * *v;
1558 }
1559 let n2e = n2 + F::fromf(1e-6);
1560 let inv = F::ONE / (n2e.sqrt() * F::fromf(extra_scale));
1561 (inv, n2e)
1562}
1563
1564#[allow(clippy::too_many_arguments)]
1567pub fn gdn_group_fwd<F: Fp>(
1568 cq: &[F],
1569 z: &[F],
1570 a: &[F],
1571 b: &[F],
1572 t: usize,
1573 cfg: &GdnSeqCfg,
1574 ko: usize,
1575 out: &mut [F],
1576) {
1577 let (nv, nk, dk, dv) = (cfg.nv, cfg.nk, cfg.dk, cfg.dv);
1578 let c_dim = cfg.c_dim();
1579 let kd = nk * dk;
1580 let rep = nv / nk;
1581 let vd = nv * dv;
1582 let sqdk = (dk as f64).sqrt();
1583 for hh in 0..rep {
1584 let h = ko * rep + hh;
1585 let ea = F::fromf((cfg.a_log[h] as f64).exp());
1586 let mut s = vec![F::ZERO; dk * dv];
1587 let mut kv = vec![F::ZERO; dv];
1588 let mut o = vec![F::ZERO; dv];
1589 for ti in 0..t {
1590 let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
1591 let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
1592 let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
1593 let (invq, _) = gdn_inv(qrow, sqdk);
1594 let (invk, _) = gdn_inv(krow, 1.0);
1595 let g = (-ea * softplus_f(a[ti * nv + h] + F::fromf(cfg.dt_bias[h] as f64))).exp();
1596 let beta = sigmoid_f(b[ti * nv + h]);
1597 for x in kv.iter_mut() {
1599 *x = F::ZERO;
1600 }
1601 for di in 0..dk {
1602 let kf = krow[di] * invk;
1603 let row = &mut s[di * dv..(di + 1) * dv];
1604 for dj in 0..dv {
1605 row[dj] *= g;
1606 kv[dj] += row[dj] * kf;
1607 }
1608 }
1609 for x in o.iter_mut() {
1610 *x = F::ZERO;
1611 }
1612 for di in 0..dk {
1613 let kf = krow[di] * invk;
1614 let qf = qrow[di] * invq;
1615 let row = &mut s[di * dv..(di + 1) * dv];
1616 for dj in 0..dv {
1617 row[dj] += kf * (vrow[dj] - kv[dj]) * beta;
1618 o[dj] += qf * row[dj];
1619 }
1620 }
1621 let mut ss = 0f64;
1623 for v in &o {
1624 ss += v.f64() * v.f64();
1625 }
1626 let inv = F::fromf(1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt());
1627 for dj in 0..dv {
1628 let zv = z[ti * vd + h * dv + dj];
1629 out[ti * vd + h * dv + dj] = o[dj] * inv * F::fromf(cfg.norm[dj] as f64) * silu(zv);
1630 }
1631 }
1632 }
1633}
1634
1635#[allow(clippy::too_many_arguments)]
1639pub fn gdn_group_bwd<F: Fp>(
1640 cq: &[F],
1641 z: &[F],
1642 a: &[F],
1643 b: &[F],
1644 t: usize,
1645 cfg: &GdnSeqCfg,
1646 ko: usize,
1647 dout: &[F],
1648 dcq: &mut [F],
1649 dz: &mut [F],
1650 da: &mut [F],
1651 db: &mut [F],
1652) {
1653 let (nv, nk, dk, dv) = (cfg.nv, cfg.nk, cfg.dk, cfg.dv);
1654 let c_dim = cfg.c_dim();
1655 let kd = nk * dk;
1656 let rep = nv / nk;
1657 let vd = nv * dv;
1658 let sqdk = (dk as f64).sqrt();
1659 for hh in 0..rep {
1660 let h = ko * rep + hh;
1661 let ea = F::fromf((cfg.a_log[h] as f64).exp());
1662
1663 let mut s_hist = vec![F::ZERO; (t + 1) * dk * dv]; let mut kv_hist = vec![F::ZERO; t * dv];
1666 let mut o_hist = vec![F::ZERO; t * dv];
1667 let mut g_v = vec![F::ZERO; t];
1668 let mut beta_v = vec![F::ZERO; t];
1669 let mut sp_arg = vec![F::ZERO; t]; for ti in 0..t {
1671 let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
1672 let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
1673 let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
1674 let (invq, _) = gdn_inv(qrow, sqdk);
1675 let (invk, _) = gdn_inv(krow, 1.0);
1676 let arg = a[ti * nv + h] + F::fromf(cfg.dt_bias[h] as f64);
1677 let g = (-ea * softplus_f(arg)).exp();
1678 let beta = sigmoid_f(b[ti * nv + h]);
1679 sp_arg[ti] = arg;
1680 g_v[ti] = g;
1681 beta_v[ti] = beta;
1682 let (prev, cur) = s_hist.split_at_mut((ti + 1) * dk * dv);
1683 let sp = &prev[ti * dk * dv..];
1684 let sn = &mut cur[..dk * dv];
1685 let kvr = &mut kv_hist[ti * dv..(ti + 1) * dv];
1686 for di in 0..dk {
1687 let kf = krow[di] * invk;
1688 for dj in 0..dv {
1689 let dec = sp[di * dv + dj] * g;
1690 sn[di * dv + dj] = dec;
1691 kvr[dj] += dec * kf;
1692 }
1693 }
1694 let or = &mut o_hist[ti * dv..(ti + 1) * dv];
1695 for di in 0..dk {
1696 let kf = krow[di] * invk;
1697 let qf = qrow[di] * invq;
1698 for dj in 0..dv {
1699 let sv = sn[di * dv + dj] + kf * (vrow[dj] - kvr[dj]) * beta_v[ti];
1700 sn[di * dv + dj] = sv;
1701 or[dj] += qf * sv;
1702 }
1703 }
1704 }
1705
1706 let mut ds = vec![F::ZERO; dk * dv];
1708 let mut do_o = vec![F::ZERO; dv];
1709 let mut du = vec![F::ZERO; dv];
1710 let mut dkv = vec![F::ZERO; dv];
1711 let mut dqh = vec![F::ZERO; dk];
1712 let mut dkh = vec![F::ZERO; dk];
1713 for ti in (0..t).rev() {
1714 let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
1715 let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
1716 let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
1717 let (invq, nq2) = gdn_inv(qrow, sqdk);
1718 let (invk, nk2) = gdn_inv(krow, 1.0);
1719 let g = g_v[ti];
1720 let beta = beta_v[ti];
1721 let s_t = &s_hist[(ti + 1) * dk * dv..(ti + 2) * dk * dv];
1722 let s_prev = &s_hist[ti * dk * dv..(ti + 1) * dk * dv];
1723 let kvr = &kv_hist[ti * dv..(ti + 1) * dv];
1724 let or = &o_hist[ti * dv..(ti + 1) * dv];
1725
1726 let mut ss = 0f64;
1728 for v in or {
1729 ss += v.f64() * v.f64();
1730 }
1731 let inv = F::fromf(1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt());
1732 let dofr = &dout[ti * vd + h * dv..ti * vd + (h + 1) * dv];
1733 let mut sdot = 0f64; for dj in 0..dv {
1736 let zv = z[ti * vd + h * dv + dj];
1737 let w = F::fromf(cfg.norm[dj] as f64);
1738 let weff = w * silu(zv);
1739 sdot += (dofr[dj] * weff * or[dj]).f64();
1740 dz[ti * vd + h * dv + dj] += dofr[dj] * or[dj] * inv * w * silu_bwd(zv);
1741 }
1742 let coef = F::fromf(sdot / dv as f64) * inv * inv * inv;
1743 for dj in 0..dv {
1744 let zv = z[ti * vd + h * dv + dj];
1745 let weff = F::fromf(cfg.norm[dj] as f64) * silu(zv);
1746 do_o[dj] = inv * weff * dofr[dj] - or[dj] * coef;
1747 }
1748
1749 for x in dqh.iter_mut() {
1751 *x = F::ZERO;
1752 }
1753 for di in 0..dk {
1754 let qf = qrow[di] * invq;
1755 let row = &s_t[di * dv..(di + 1) * dv];
1756 let dsr = &mut ds[di * dv..(di + 1) * dv];
1757 let mut acc = F::ZERO;
1758 for dj in 0..dv {
1759 dsr[dj] += qf * do_o[dj];
1760 acc += row[dj] * do_o[dj];
1761 }
1762 dqh[di] = acc;
1763 }
1764
1765 for x in du.iter_mut() {
1768 *x = F::ZERO;
1769 }
1770 for x in dkh.iter_mut() {
1771 *x = F::ZERO;
1772 }
1773 for di in 0..dk {
1774 let kf = krow[di] * invk;
1775 let dsr = &ds[di * dv..(di + 1) * dv];
1776 let mut acc = F::ZERO;
1777 for dj in 0..dv {
1778 du[dj] += dsr[dj] * kf;
1779 acc += dsr[dj] * (vrow[dj] - kvr[dj]) * beta;
1780 }
1781 dkh[di] = acc;
1782 }
1783 let mut dbeta = F::ZERO;
1784 for dj in 0..dv {
1785 dbeta += du[dj] * (vrow[dj] - kvr[dj]);
1786 dcq[ti * c_dim + 2 * kd + h * dv + dj] += beta * du[dj];
1788 dkv[dj] = -(beta * du[dj]);
1789 }
1790
1791 let mut dg = F::ZERO;
1796 for di in 0..dk {
1797 let kf = krow[di] * invk;
1798 let spr = &s_prev[di * dv..(di + 1) * dv];
1799 let dsr = &mut ds[di * dv..(di + 1) * dv];
1800 let mut acc = F::ZERO;
1801 for dj in 0..dv {
1802 let dspre = dsr[dj] + kf * dkv[dj];
1803 acc += (spr[dj] * g) * dkv[dj];
1804 dg += dspre * spr[dj];
1805 dsr[dj] = g * dspre;
1806 }
1807 dkh[di] += acc;
1808 }
1809
1810 let sig = sigmoid_f(sp_arg[ti]);
1812 da[ti * nv + h] += dg * g * (-ea) * sig;
1813 db[ti * nv + h] += dbeta * beta * (F::ONE - beta);
1814
1815 let mut qdot = F::ZERO;
1819 let mut kdot = F::ZERO;
1820 for di in 0..dk {
1821 qdot += dqh[di] * qrow[di];
1822 kdot += dkh[di] * krow[di];
1823 }
1824 for di in 0..dk {
1825 dcq[ti * c_dim + ko * dk + di] += invq * dqh[di] - qrow[di] * qdot * invq / nq2;
1826 dcq[ti * c_dim + kd + ko * dk + di] +=
1827 invk * dkh[di] - krow[di] * kdot * invk / nk2;
1828 }
1829 }
1830 }
1831}
1832
1833pub fn gdn_seq_fwd<F: Fp>(
1837 qkv: &[F],
1838 z: &[F],
1839 a: &[F],
1840 b: &[F],
1841 t: usize,
1842 cfg: &GdnSeqCfg,
1843 out: &mut [F],
1844) {
1845 let c_dim = cfg.c_dim();
1846 let mut pre = vec![F::ZERO; t * c_dim];
1847 let mut cq = vec![F::ZERO; t * c_dim];
1848 gdn_conv_fwd(qkv, t, c_dim, cfg.kk, cfg.conv, &mut pre, &mut cq);
1849 for ko in 0..cfg.nk {
1850 gdn_group_fwd(&cq, z, a, b, t, cfg, ko, out);
1851 }
1852}
1853
1854#[allow(clippy::too_many_arguments)]
1857pub fn gdn_seq_bwd<F: Fp>(
1858 qkv: &[F],
1859 z: &[F],
1860 a: &[F],
1861 b: &[F],
1862 t: usize,
1863 cfg: &GdnSeqCfg,
1864 dout: &[F],
1865 dqkv: &mut [F],
1866 dz: &mut [F],
1867 da: &mut [F],
1868 db: &mut [F],
1869) {
1870 let c_dim = cfg.c_dim();
1871 let mut pre = vec![F::ZERO; t * c_dim];
1872 let mut cq = vec![F::ZERO; t * c_dim];
1873 gdn_conv_fwd(qkv, t, c_dim, cfg.kk, cfg.conv, &mut pre, &mut cq);
1874 let mut dcq = vec![F::ZERO; t * c_dim];
1875 for ko in 0..cfg.nk {
1876 gdn_group_bwd(&cq, z, a, b, t, cfg, ko, dout, &mut dcq, dz, da, db);
1877 }
1878 gdn_conv_bwd(&pre, t, c_dim, cfg.kk, cfg.conv, &dcq, dqkv);
1879}
1880
1881#[cfg(test)]
1882mod gpu_bake_tests {
1883 #[test]
1887 #[cfg(feature = "gpu")]
1888 fn gemm_dx_gpu_matches_cpu() {
1889 if std::env::var("CMF_GPU").is_err() {
1890 eprintln!("no backend requested — skip");
1891 return;
1892 }
1893 let (n, k, m) = (8usize, 512usize, 1024usize);
1894 let dy: Vec<f32> = (0..n * m)
1895 .map(|i| ((i * 37 % 101) as f32 - 50.0) / 50.0)
1896 .collect();
1897 let w: Vec<f32> = (0..m * k)
1898 .map(|i| ((i * 17 % 97) as f32 - 48.0) / 48.0)
1899 .collect();
1900 let mut want = vec![0f32; n * k];
1901 super::gemm_dx(&dy, &w, &mut want, n, k, m, None);
1902 let mut got = vec![0f32; n * k];
1903 if !crate::gpu_wgpu::gemm_dx_f32(&dy, &w, &mut got, n, k, m) {
1904 eprintln!("device declined — skip");
1905 return;
1906 }
1907 let num: f64 = want
1908 .iter()
1909 .zip(&got)
1910 .map(|(a, b)| ((a - b) as f64).powi(2))
1911 .sum();
1912 let den: f64 = want
1913 .iter()
1914 .map(|a| (*a as f64).powi(2))
1915 .sum::<f64>()
1916 .max(1e-30);
1917 let rel = (num / den).sqrt();
1918 assert!(rel < 1e-4, "gemm_dx GPU vs CPU rel {rel:e}");
1919 }
1920}