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
203pub fn gemm_nt(
204 x: &[f32],
205 w: &[f32],
206 y: &mut [f32],
207 n: usize,
208 k: usize,
209 m: usize,
210 pool: Option<&Pool>,
211) {
212 debug_assert_eq!(x.len(), n * k);
213 debug_assert_eq!(w.len(), m * k);
214 debug_assert_eq!(y.len(), n * m);
215 #[cfg(target_os = "macos")]
216 if accel::on() && n * k * m >= 1 << 18 {
217 unsafe {
219 accel::cblas_sgemm(
220 101,
221 111,
222 112,
223 n as i32,
224 m as i32,
225 k as i32,
226 1.0,
227 x.as_ptr(),
228 k as i32,
229 w.as_ptr(),
230 k as i32,
231 0.0,
232 y.as_mut_ptr(),
233 m as i32,
234 );
235 }
236 return;
237 }
238 let nb = n.div_ceil(GEMM_BLOCK);
239 let block = |r0: usize, r1: usize, y: &mut [f32]| {
240 for o in 0..m {
241 let wr = &w[o * k..(o + 1) * k];
242 for i in r0..r1 {
243 y[(i - r0) * m + o] = crate::attention::dot_f32(&x[i * k..(i + 1) * k], wr);
244 }
245 }
246 };
247 match pool {
248 Some(p) if nb > 1 => {
249 let yp = SendMut(y.as_mut_ptr());
250 p.run(&|widx, nw| {
251 for bi in (widx..nb).step_by(nw) {
252 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
253 let ys = unsafe { yp.slice(r0 * m, (r1 - r0) * m) };
255 block(r0, r1, ys);
256 }
257 });
258 }
259 _ => {
260 for bi in 0..nb {
261 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
262 block(r0, r1, &mut y[r0 * m..r1 * m]);
263 }
264 }
265 }
266}
267
268pub fn gemm_dx(
270 dy: &[f32],
271 w: &[f32],
272 dx: &mut [f32],
273 n: usize,
274 k: usize,
275 m: usize,
276 pool: Option<&Pool>,
277) {
278 debug_assert_eq!(dy.len(), n * m);
279 debug_assert_eq!(w.len(), m * k);
280 debug_assert_eq!(dx.len(), n * k);
281 #[cfg(target_os = "macos")]
282 if accel::on() && n * k * m >= 1 << 18 {
283 unsafe {
285 accel::cblas_sgemm(
286 101,
287 111,
288 111,
289 n as i32,
290 k as i32,
291 m as i32,
292 1.0,
293 dy.as_ptr(),
294 m as i32,
295 w.as_ptr(),
296 k as i32,
297 1.0,
298 dx.as_mut_ptr(),
299 k as i32,
300 );
301 }
302 return;
303 }
304 let nb = n.div_ceil(GEMM_BLOCK);
305 let block = |r0: usize, r1: usize, dxs: &mut [f32]| {
306 for o in 0..m {
307 let wr = &w[o * k..(o + 1) * k];
308 for i in r0..r1 {
309 let g = dy[i * m + o];
310 if g != 0.0 {
311 crate::attention::axpy_f32(&mut dxs[(i - r0) * k..(i - r0 + 1) * k], wr, g);
312 }
313 }
314 }
315 };
316 match pool {
317 Some(p) if nb > 1 => {
318 let dxp = SendMut(dx.as_mut_ptr());
319 p.run(&|widx, nw| {
320 for bi in (widx..nb).step_by(nw) {
321 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
322 let dxs = unsafe { dxp.slice(r0 * k, (r1 - r0) * k) };
324 block(r0, r1, dxs);
325 }
326 });
327 }
328 _ => {
329 for bi in 0..nb {
330 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
331 block(r0, r1, &mut dx[r0 * k..r1 * k]);
332 }
333 }
334 }
335}
336
337pub fn gemm_dw(
340 dy: &[f32],
341 x: &[f32],
342 dw: &mut [f32],
343 n: usize,
344 k: usize,
345 m: usize,
346 pool: Option<&Pool>,
347) {
348 debug_assert_eq!(dy.len(), n * m);
349 debug_assert_eq!(x.len(), n * k);
350 debug_assert_eq!(dw.len(), m * k);
351 #[cfg(target_os = "macos")]
352 if accel::on() && n * k * m >= 1 << 18 {
353 unsafe {
355 accel::cblas_sgemm(
356 101,
357 112,
358 111,
359 m as i32,
360 k as i32,
361 n as i32,
362 1.0,
363 dy.as_ptr(),
364 m as i32,
365 x.as_ptr(),
366 k as i32,
367 1.0,
368 dw.as_mut_ptr(),
369 k as i32,
370 );
371 }
372 return;
373 }
374 let range = |o0: usize, o1: usize, dws: &mut [f32]| {
375 let nb = n.div_ceil(GEMM_BLOCK);
377 for bi in 0..nb {
378 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
379 for o in o0..o1 {
380 let dwr = &mut dws[(o - o0) * k..(o - o0 + 1) * k];
381 for i in r0..r1 {
382 let g = dy[i * m + o];
383 if g != 0.0 {
384 crate::attention::axpy_f32(dwr, &x[i * k..(i + 1) * k], g);
385 }
386 }
387 }
388 }
389 };
390 match pool {
391 Some(p) if m >= 8 => {
392 let dwp = SendMut(dw.as_mut_ptr());
393 p.run(&|widx, nw| {
394 let (o0, o1) = (widx * m / nw, (widx + 1) * m / nw);
395 if o0 < o1 {
396 let dws = unsafe { dwp.slice(o0 * k, (o1 - o0) * k) };
398 range(o0, o1, dws);
399 }
400 });
401 }
402 _ => range(0, m, dw),
403 }
404}
405
406#[inline]
409pub fn silu<F: Fp>(x: F) -> F {
410 x / (F::ONE + (-x).exp())
411}
412
413#[inline]
415pub fn silu_bwd<F: Fp>(x: F) -> F {
416 let s = F::ONE / (F::ONE + (-x).exp());
417 s * (F::ONE + x * (F::ONE - s))
418}
419
420pub fn rmsnorm_fwd<F: Fp>(x: &[F], w: &[F], eps: f64, gemma: bool, y: &mut [F], inv_out: &mut [F]) {
427 let d = w.len();
428 let n = x.len() / d;
429 for r in 0..n {
430 let xr = &x[r * d..(r + 1) * d];
431 let mut ss = 0f64;
432 for v in xr {
433 ss += v.f64() * v.f64();
434 }
435 let inv = F::fromf(1.0 / (ss / d as f64 + eps).sqrt());
436 inv_out[r] = inv;
437 let yr = &mut y[r * d..(r + 1) * d];
438 for j in 0..d {
439 let weff = if gemma { F::ONE + w[j] } else { w[j] };
440 yr[j] = xr[j] * inv * weff;
441 }
442 }
443}
444
445pub fn rmsnorm_bwd<F: Fp>(
451 x: &[F],
452 w: &[F],
453 inv: &[F],
454 dy: &[F],
455 gemma: bool,
456 dx: &mut [F],
457 mut dw: Option<&mut [F]>,
458) {
459 let d = w.len();
460 let n = x.len() / d;
461 for r in 0..n {
462 let xr = &x[r * d..(r + 1) * d];
463 let dyr = &dy[r * d..(r + 1) * d];
464 let iv = inv[r];
465 let mut s = 0f64;
466 for j in 0..d {
467 let weff = if gemma { F::ONE + w[j] } else { w[j] };
468 s += (dyr[j] * weff * xr[j]).f64();
469 }
470 let coef = F::fromf(s / d as f64) * iv * iv * iv;
471 let dxr = &mut dx[r * d..(r + 1) * d];
472 for j in 0..d {
473 let weff = if gemma { F::ONE + w[j] } else { w[j] };
474 dxr[j] += iv * weff * dyr[j] - xr[j] * coef;
475 }
476 if let Some(dwv) = dw.as_deref_mut() {
477 for j in 0..d {
478 dwv[j] += dyr[j] * xr[j] * iv;
479 }
480 }
481 }
482}
483
484pub fn rope_fwd<F: Fp>(x: &mut [F], position: usize, inv_freq: &[f64]) {
489 let half = inv_freq.len();
490 for (i, &freq) in inv_freq.iter().enumerate() {
491 let angle = position as f64 * freq;
492 let (sin, cos) = (F::fromf(angle.sin()), F::fromf(angle.cos()));
493 let x0 = x[i];
494 let x1 = x[i + half];
495 x[i] = x0 * cos - x1 * sin;
496 x[i + half] = x0 * sin + x1 * cos;
497 }
498}
499
500pub fn rope_bwd<F: Fp>(dy: &mut [F], position: usize, inv_freq: &[f64]) {
503 let half = inv_freq.len();
504 for (i, &freq) in inv_freq.iter().enumerate() {
505 let angle = position as f64 * freq;
506 let (sin, cos) = (F::fromf(angle.sin()), F::fromf(angle.cos()));
507 let g0 = dy[i];
508 let g1 = dy[i + half];
509 dy[i] = g0 * cos + g1 * sin;
510 dy[i + half] = -g0 * sin + g1 * cos;
511 }
512}
513
514pub fn seg_means<F: Fp>(x: &[F], t: usize, d: usize, m: usize, out: &mut [F]) {
519 for i in 0..m {
520 let (lo, hi) = (i * t / m, (i + 1) * t / m);
521 let or = &mut out[i * d..(i + 1) * d];
522 for v in or.iter_mut() {
523 *v = F::ZERO;
524 }
525 for j in lo..hi {
526 for c in 0..d {
527 or[c] += x[j * d + c];
528 }
529 }
530 let inv = F::fromf(1.0 / (hi - lo) as f64);
531 for v in or.iter_mut() {
532 *v *= inv;
533 }
534 }
535}
536
537pub fn seg_means_bwd<F: Fp>(dl: &[F], t: usize, d: usize, m: usize, dx: &mut [F]) {
540 for i in 0..m {
541 let (lo, hi) = (i * t / m, (i + 1) * t / m);
542 let inv = F::fromf(1.0 / (hi - lo) as f64);
543 let dlr = &dl[i * d..(i + 1) * d];
544 for j in lo..hi {
545 for c in 0..d {
546 dx[j * d + c] += dlr[c] * inv;
547 }
548 }
549 }
550}
551
552#[allow(clippy::needless_range_loop)] pub fn attn_head_fwd<F: Fp>(
558 q: &[F],
559 k: &[F],
560 v: &[F],
561 t: usize,
562 d: usize,
563 dv: usize,
564 out: &mut [F],
565) {
566 let scale = F::fromf(1.0 / (d as f64).sqrt());
567 let mut row = vec![F::ZERO; t];
568 for ti in 0..t {
569 let qr = &q[ti * d..(ti + 1) * d];
570 let mut mx = F::fromf(f64::NEG_INFINITY);
571 for j in 0..=ti {
572 let s = dot(qr, &k[j * d..(j + 1) * d]) * scale;
573 row[j] = s;
574 mx = mx.maxf(s);
575 }
576 let mut den = F::ZERO;
577 for j in 0..=ti {
578 row[j] = (row[j] - mx).exp();
579 den += row[j];
580 }
581 let or = &mut out[ti * dv..(ti + 1) * dv];
582 for o in or.iter_mut() {
583 *o = F::ZERO;
584 }
585 for j in 0..=ti {
586 let p = row[j] / den;
587 for (o, vv) in or.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
588 *o += p * *vv;
589 }
590 }
591 }
592}
593
594#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
598pub fn attn_head_bwd<F: Fp>(
599 q: &[F],
600 k: &[F],
601 v: &[F],
602 dout: &[F],
603 t: usize,
604 d: usize,
605 dv: usize,
606 dq: &mut [F],
607 dk: &mut [F],
608 dvv: &mut [F],
609) {
610 let scale = F::fromf(1.0 / (d as f64).sqrt());
611 let mut row = vec![F::ZERO; t];
612 for ti in 0..t {
613 let qr = &q[ti * d..(ti + 1) * d];
614 let mut mx = F::fromf(f64::NEG_INFINITY);
615 for j in 0..=ti {
616 let s = dot(qr, &k[j * d..(j + 1) * d]) * scale;
617 row[j] = s;
618 mx = mx.maxf(s);
619 }
620 let mut den = F::ZERO;
621 for j in 0..=ti {
622 row[j] = (row[j] - mx).exp();
623 den += row[j];
624 }
625 let dor = &dout[ti * dv..(ti + 1) * dv];
626 let mut pdp = F::ZERO;
628 let mut dp = vec![F::ZERO; ti + 1];
629 for j in 0..=ti {
630 let p = row[j] / den;
631 row[j] = p; dp[j] = dot(dor, &v[j * dv..(j + 1) * dv]);
633 pdp += p * dp[j];
634 }
635 let dqr = &mut dq[ti * d..(ti + 1) * d];
636 for j in 0..=ti {
637 let p = row[j];
638 for (dvo, o) in dvv[j * dv..(j + 1) * dv].iter_mut().zip(dor) {
640 *dvo += p * *o;
641 }
642 let ds = p * (dp[j] - pdp) * scale;
644 let kr = &k[j * d..(j + 1) * d];
645 for c in 0..d {
646 dqr[c] += ds * kr[c];
647 }
648 let dkr = &mut dk[j * d..(j + 1) * d];
649 for c in 0..d {
650 dkr[c] += ds * qr[c];
651 }
652 }
653 }
654}
655
656#[derive(Clone, Copy, Debug)]
660pub struct NysCfg {
661 pub m: usize,
663 pub w: usize,
665 pub sink: usize,
667 pub prefill: Option<usize>,
679}
680
681impl NysCfg {
682 #[inline]
685 pub fn prefill_len(&self, t: usize) -> usize {
686 self.prefill.unwrap_or(t / 2).clamp(1, t)
687 }
688}
689
690const NYS_DEN_EPS: f64 = 1e-30;
692
693struct NysGraph<F: Fp> {
697 m_eff: usize,
698 tp: usize,
701 q_l: Vec<F>,
702 k_l: Vec<F>,
703 mu: Vec<F>,
704 fu: Vec<F>,
705 e: Vec<F>,
706 fumu: Vec<F>,
707 far_keep: Vec<bool>,
716 wmat: Vec<F>,
719 c_row: Vec<F>,
720 den: Vec<F>,
721}
722
723#[inline]
728fn nys_exact_row(ti: usize, tp: usize) -> bool {
729 ti < tp
730}
731
732#[inline]
734fn nys_near(ti: usize, j: usize, w: usize, sink: usize) -> bool {
735 ti - j < w || j < sink
736}
737
738fn nys_graph<F: Fp>(
743 q: &[F],
744 k: &[F],
745 t: usize,
746 d: usize,
747 cfg: &NysCfg,
748 mu_override: Option<&[F]>,
749) -> NysGraph<F> {
750 let scale = 1.0 / (d as f64).sqrt();
751 let fscale = F::fromf(scale);
752 let tp = cfg.prefill_len(t);
757 let m_eff = (tp / 8).clamp(4, cfg.m);
758 let mut q_l = vec![F::ZERO; m_eff * d];
759 let mut k_l = vec![F::ZERO; m_eff * d];
760 seg_means(&q[..tp * d], tp, d, m_eff, &mut q_l);
761 seg_means(&k[..tp * d], tp, d, m_eff, &mut k_l);
762
763 let mut au = vec![0f64; m_eff * m_eff];
765 for i in 0..m_eff {
766 for j in 0..m_eff {
767 let mut s = 0f64;
768 for c in 0..d {
769 s += q_l[i * d + c].f64() * k_l[j * d + c].f64();
770 }
771 au[i * m_eff + j] = (s * scale).exp();
772 }
773 }
774 let mu: Vec<F> = match mu_override {
775 Some(m) => m.to_vec(),
776 None => crate::nystrom::ridge_pinv(&au, m_eff)
777 .iter()
778 .map(|&x| F::fromf(x))
779 .collect(),
780 };
781
782 let mut fu = vec![F::ZERO; t * m_eff];
784 for ti in 0..t {
785 for i in 0..m_eff {
786 fu[ti * m_eff + i] =
787 (dot(&q[ti * d..(ti + 1) * d], &k_l[i * d..(i + 1) * d]) * fscale).exp();
788 }
789 }
790 let mut e = vec![F::ZERO; m_eff * t];
791 for i in 0..m_eff {
792 for j in 0..t {
793 e[i * t + j] = (dot(&q_l[i * d..(i + 1) * d], &k[j * d..(j + 1) * d]) * fscale).exp();
794 }
795 }
796 let mut fumu = vec![F::ZERO; t * m_eff];
797 matmul_nt(
798 &fu,
799 &transpose(&mu, m_eff, m_eff),
802 &mut fumu,
803 t,
804 m_eff,
805 m_eff,
806 );
807 let mut a = vec![F::ZERO; t * t];
812 for ti in tp..t {
813 let fr = &fumu[ti * m_eff..(ti + 1) * m_eff];
814 let ar = &mut a[ti * t..ti * t + ti + 1]; for (i, &f) in fr.iter().enumerate() {
816 let er = &e[i * t..i * t + ti + 1];
817 for (av, ev) in ar.iter_mut().zip(er) {
818 *av += f * *ev;
819 }
820 }
821 }
822
823 let mut wmat = vec![F::ZERO; t * t];
826 let mut c_row = vec![F::ZERO; t];
827 let mut den = vec![F::ZERO; t];
828 let mut far_keep = vec![false; t];
829 let mut lg_row = vec![F::ZERO; t];
830 for ti in 0..t {
831 let qr = &q[ti * d..(ti + 1) * d];
832 let mut c = F::fromf(f64::NEG_INFINITY);
833 for j in 0..=ti {
834 let s = dot(qr, &k[j * d..(j + 1) * d]) * fscale;
835 lg_row[j] = s;
836 c = c.maxf(s);
837 }
838 c_row[ti] = c;
839 let emc = (-c).exp();
840 let exact_row = nys_exact_row(ti, tp);
841
842 let mut far_sum = F::ZERO;
857 if !exact_row {
858 for j in 0..=ti {
859 if !nys_near(ti, j, cfg.w, cfg.sink) {
860 far_sum += a[ti * t + j];
861 }
862 }
863 }
864 let keep = !exact_row && far_sum.f64() >= 0.0;
865 far_keep[ti] = keep;
866
867 let wr = &mut wmat[ti * t..(ti + 1) * t];
868 let mut dsum = F::ZERO;
869 for j in 0..=ti {
870 let wv = if exact_row || nys_near(ti, j, cfg.w, cfg.sink) {
871 (lg_row[j] - c).exp()
872 } else if keep {
873 a[ti * t + j] * emc
874 } else {
875 F::ZERO
876 };
877 wr[j] = wv;
878 dsum += wv;
879 }
880 den[ti] = dsum.maxf(F::fromf(NYS_DEN_EPS));
881 }
882 NysGraph {
883 m_eff,
884 tp,
885 q_l,
886 k_l,
887 mu,
888 fu,
889 e,
890 fumu,
891 far_keep,
892 wmat,
893 c_row,
894 den,
895 }
896}
897
898fn transpose<F: Fp>(x: &[F], rows: usize, cols: usize) -> Vec<F> {
899 let mut out = vec![F::ZERO; rows * cols];
900 for r in 0..rows {
901 for c in 0..cols {
902 out[c * rows + r] = x[r * cols + c];
903 }
904 }
905 out
906}
907
908#[allow(clippy::too_many_arguments)]
911pub fn nystrom_head_fwd<F: Fp>(
912 q: &[F],
913 k: &[F],
914 v: &[F],
915 t: usize,
916 d: usize,
917 dv: usize,
918 cfg: &NysCfg,
919 out: &mut [F],
920) {
921 if nys_degenerate(t, cfg) {
922 attn_head_fwd(q, k, v, t, d, dv, out);
923 return;
924 }
925 nystrom_head_fwd_mu(q, k, v, t, d, dv, cfg, None, out);
926}
927
928#[inline]
937fn nys_degenerate(t: usize, cfg: &NysCfg) -> bool {
938 cfg.prefill_len(t) <= cfg.w + cfg.sink + 8
939}
940
941#[doc(hidden)]
944pub fn nystrom_mu_for_test<F: Fp>(q: &[F], k: &[F], t: usize, d: usize, cfg: &NysCfg) -> Vec<F> {
945 nys_graph(q, k, t, d, cfg, None).mu
946}
947
948#[doc(hidden)]
951#[allow(clippy::too_many_arguments)]
952pub fn nystrom_head_fwd_mu<F: Fp>(
953 q: &[F],
954 k: &[F],
955 v: &[F],
956 t: usize,
957 d: usize,
958 dv: usize,
959 cfg: &NysCfg,
960 mu_override: Option<&[F]>,
961 out: &mut [F],
962) {
963 let g = nys_graph(q, k, t, d, cfg, mu_override);
964 for ti in 0..t {
965 let wr = &g.wmat[ti * t..(ti + 1) * t];
966 let den = g.den[ti];
967 let or = &mut out[ti * dv..(ti + 1) * dv];
968 for o in or.iter_mut() {
969 *o = F::ZERO;
970 }
971 for j in 0..=ti {
972 let p = wr[j] / den;
973 if p.f64() != 0.0 {
974 for (o, vv) in or.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
975 *o += p * *vv;
976 }
977 }
978 }
979 }
980}
981
982#[allow(clippy::too_many_arguments)]
993pub fn nystrom_head_bwd<F: Fp>(
994 q: &[F],
995 k: &[F],
996 v: &[F],
997 dout: &[F],
998 t: usize,
999 d: usize,
1000 dv: usize,
1001 cfg: &NysCfg,
1002 dq: &mut [F],
1003 dk: &mut [F],
1004 dvv: &mut [F],
1005) {
1006 if nys_degenerate(t, cfg) {
1007 attn_head_bwd(q, k, v, dout, t, d, dv, dq, dk, dvv);
1008 return;
1009 }
1010 nystrom_head_bwd_mu(q, k, v, dout, t, d, dv, cfg, None, dq, dk, dvv);
1011}
1012
1013#[doc(hidden)]
1015#[allow(clippy::too_many_arguments)]
1016pub fn nystrom_head_bwd_mu<F: Fp>(
1017 q: &[F],
1018 k: &[F],
1019 v: &[F],
1020 dout: &[F],
1021 t: usize,
1022 d: usize,
1023 dv: usize,
1024 cfg: &NysCfg,
1025 mu_override: Option<&[F]>,
1026 dq: &mut [F],
1027 dk: &mut [F],
1028 dvv: &mut [F],
1029) {
1030 let scale = F::fromf(1.0 / (d as f64).sqrt());
1031 let g = nys_graph(q, k, t, d, cfg, mu_override);
1032 let m_eff = g.m_eff;
1033
1034 let mut dwmat = vec![F::ZERO; t * t];
1036 let mut out_row = vec![F::ZERO; dv];
1037 for ti in 0..t {
1038 let wr = &g.wmat[ti * t..(ti + 1) * t];
1039 let den = g.den[ti];
1040 for o in out_row.iter_mut() {
1041 *o = F::ZERO;
1042 }
1043 for j in 0..=ti {
1044 let p = wr[j] / den;
1045 if p.f64() != 0.0 {
1046 for (o, vv) in out_row.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
1047 *o += p * *vv;
1048 }
1049 }
1050 }
1051 let dor = &dout[ti * dv..(ti + 1) * dv];
1052 let dwr = &mut dwmat[ti * t..(ti + 1) * t];
1053 for j in 0..=ti {
1054 let p = wr[j] / den;
1056 if p.f64() != 0.0 {
1057 for (dvo, o) in dvv[j * dv..(j + 1) * dv].iter_mut().zip(dor) {
1058 *dvo += p * *o;
1059 }
1060 }
1061 let mut s = F::ZERO;
1063 for c in 0..dv {
1064 s += dor[c] * (v[j * dv + c] - out_row[c]);
1065 }
1066 dwr[j] = s / den;
1067 }
1068 }
1069
1070 for ti in 0..t {
1073 let qr = &q[ti * d..(ti + 1) * d];
1074 let dqr_base = ti * d;
1075 for j in 0..=ti {
1076 if !(nys_exact_row(ti, g.tp) || nys_near(ti, j, cfg.w, cfg.sink)) {
1077 continue;
1078 }
1079 let dlg = dwmat[ti * t + j] * g.wmat[ti * t + j] * scale;
1080 if dlg.f64() == 0.0 {
1081 continue;
1082 }
1083 let kr = &k[j * d..(j + 1) * d];
1084 for c in 0..d {
1085 dq[dqr_base + c] += dlg * kr[c];
1086 }
1087 let dkr = &mut dk[j * d..(j + 1) * d];
1088 for c in 0..d {
1089 dkr[c] += dlg * qr[c];
1090 }
1091 }
1092 }
1093
1094 let mut da = vec![F::ZERO; t * t];
1109 for ti in g.tp..t {
1110 if !g.far_keep[ti] {
1111 continue; }
1113 let emc = (-g.c_row[ti]).exp();
1114 for j in 0..=ti {
1115 if nys_near(ti, j, cfg.w, cfg.sink) {
1116 continue;
1117 }
1118 da[ti * t + j] = dwmat[ti * t + j] * emc;
1119 }
1120 }
1121 let mut dfumu = vec![F::ZERO; t * m_eff];
1123 for ti in 0..t {
1124 let dar = &da[ti * t..ti * t + ti + 1];
1125 let dfr = &mut dfumu[ti * m_eff..(ti + 1) * m_eff];
1126 for (i, df) in dfr.iter_mut().enumerate() {
1127 let er = &g.e[i * t..i * t + ti + 1];
1128 let mut s = F::ZERO;
1129 for (av, ev) in dar.iter().zip(er) {
1130 s += *av * *ev;
1131 }
1132 *df = s;
1133 }
1134 }
1135 let mut dfu = vec![F::ZERO; t * m_eff];
1137 matmul_nt(&dfumu, &g.mu, &mut dfu, t, m_eff, m_eff);
1138 let mut de = vec![F::ZERO; m_eff * t];
1140 for ti in 0..t {
1141 let dar = &da[ti * t..ti * t + ti + 1];
1142 let fr = &g.fumu[ti * m_eff..(ti + 1) * m_eff];
1143 for (i, &f) in fr.iter().enumerate() {
1144 if f.f64() == 0.0 {
1145 continue;
1146 }
1147 let der = &mut de[i * t..i * t + ti + 1];
1148 for (dev, av) in der.iter_mut().zip(dar) {
1149 *dev += f * *av;
1150 }
1151 }
1152 }
1153 let mut dq_l = vec![F::ZERO; m_eff * d];
1156 let mut dk_l = vec![F::ZERO; m_eff * d];
1157 for ti in 0..t {
1158 let qr = &q[ti * d..(ti + 1) * d];
1159 for i in 0..m_eff {
1160 let dlg = dfu[ti * m_eff + i] * g.fu[ti * m_eff + i] * scale;
1161 if dlg.f64() == 0.0 {
1162 continue;
1163 }
1164 let klr = &g.k_l[i * d..(i + 1) * d];
1165 for c in 0..d {
1166 dq[ti * d + c] += dlg * klr[c];
1167 }
1168 let dklr = &mut dk_l[i * d..(i + 1) * d];
1169 for c in 0..d {
1170 dklr[c] += dlg * qr[c];
1171 }
1172 }
1173 }
1174 for i in 0..m_eff {
1175 let qlr = &g.q_l[i * d..(i + 1) * d];
1176 for j in 0..t {
1177 let dlg = de[i * t + j] * g.e[i * t + j] * scale;
1178 if dlg.f64() == 0.0 {
1179 continue;
1180 }
1181 let kr = &k[j * d..(j + 1) * d];
1182 let dqlr = &mut dq_l[i * d..(i + 1) * d];
1183 for c in 0..d {
1184 dqlr[c] += dlg * kr[c];
1185 }
1186 for c in 0..d {
1187 dk[j * d + c] += dlg * qlr[c];
1188 }
1189 }
1190 }
1191 let tp = g.tp;
1195 seg_means_bwd(&dq_l, tp, d, m_eff, &mut dq[..tp * d]);
1196 seg_means_bwd(&dk_l, tp, d, m_eff, &mut dk[..tp * d]);
1197}
1198
1199pub fn ce_kl_position<F: Fp>(
1207 s_logits: &[F],
1208 t_logits: &[F],
1209 target: usize,
1210 kl_w: f64,
1211 inv_n: f64,
1212 dlogits: &mut [F],
1213) -> (f64, f64) {
1214 let vsz = s_logits.len();
1215 debug_assert_eq!(t_logits.len(), vsz);
1216 let mut smax = f64::NEG_INFINITY;
1218 let mut tmax = f64::NEG_INFINITY;
1219 for i in 0..vsz {
1220 smax = smax.max(s_logits[i].f64());
1221 tmax = tmax.max(t_logits[i].f64());
1222 }
1223 let mut ssum = 0f64;
1224 let mut tsum = 0f64;
1225 for i in 0..vsz {
1226 ssum += (s_logits[i].f64() - smax).exp();
1227 tsum += (t_logits[i].f64() - tmax).exp();
1228 }
1229 let slz = smax + ssum.ln();
1230 let tlz = tmax + tsum.ln();
1231 let ce = slz - s_logits[target].f64();
1232 let mut kl = 0f64;
1233 for i in 0..vsz {
1234 let ls = s_logits[i].f64() - slz;
1235 let lt = t_logits[i].f64() - tlz;
1236 let pt = lt.exp();
1237 let ps = ls.exp();
1238 if pt > 0.0 {
1239 kl += pt * (lt - ls);
1240 }
1241 let mut gd = (1.0 - kl_w) * ps + kl_w * (ps - pt);
1242 if i == target {
1243 gd -= 1.0 - kl_w;
1244 }
1245 dlogits[i] = F::fromf(gd * inv_n);
1246 }
1247 (ce, kl)
1248}
1249
1250pub struct GdnSeqCfg<'a> {
1267 pub nv: usize,
1268 pub nk: usize,
1269 pub dk: usize,
1270 pub dv: usize,
1271 pub kk: usize,
1272 pub rms_eps: f64,
1273 pub conv: &'a [f32],
1276 pub a_log: &'a [f32],
1278 pub dt_bias: &'a [f32],
1280 pub norm: &'a [f32],
1282}
1283
1284impl GdnSeqCfg<'_> {
1285 pub fn c_dim(&self) -> usize {
1286 2 * self.nk * self.dk + self.nv * self.dv
1287 }
1288}
1289
1290#[inline]
1291fn softplus_f<F: Fp>(x: F) -> F {
1292 if x.f64() > 20.0 {
1295 x
1296 } else {
1297 F::fromf(x.f64().exp().ln_1p())
1298 }
1299}
1300
1301#[inline]
1302fn sigmoid_f<F: Fp>(x: F) -> F {
1303 F::ONE / (F::ONE + (-x).exp())
1304}
1305
1306pub fn gdn_conv_fwd<F: Fp>(
1310 raw: &[F],
1311 t: usize,
1312 c_dim: usize,
1313 kk: usize,
1314 conv: &[f32],
1315 pre: &mut [F],
1316 cq: &mut [F],
1317) {
1318 for ti in 0..t {
1319 for c in 0..c_dim {
1320 let taps = &conv[c * kk..(c + 1) * kk];
1321 let mut acc = F::ZERO;
1322 for (j, &tap) in taps.iter().enumerate() {
1323 let p = ti as isize - (kk as isize - 1) + j as isize;
1325 if p >= 0 {
1326 acc += raw[p as usize * c_dim + c] * F::fromf(tap as f64);
1327 }
1328 }
1329 pre[ti * c_dim + c] = acc;
1330 cq[ti * c_dim + c] = silu(acc);
1331 }
1332 }
1333}
1334
1335pub fn gdn_conv_bwd<F: Fp>(
1337 pre: &[F],
1338 t: usize,
1339 c_dim: usize,
1340 kk: usize,
1341 conv: &[f32],
1342 dcq: &[F],
1343 draw: &mut [F],
1344) {
1345 for ti in 0..t {
1346 for c in 0..c_dim {
1347 let g = dcq[ti * c_dim + c];
1348 if g.f64() == 0.0 {
1349 continue;
1350 }
1351 let dp = g * silu_bwd(pre[ti * c_dim + c]);
1352 let taps = &conv[c * kk..(c + 1) * kk];
1353 for (j, &tap) in taps.iter().enumerate() {
1354 let p = ti as isize - (kk as isize - 1) + j as isize;
1355 if p >= 0 {
1356 draw[p as usize * c_dim + c] += dp * F::fromf(tap as f64);
1357 }
1358 }
1359 }
1360 }
1361}
1362
1363#[inline]
1366fn gdn_inv<F: Fp>(x: &[F], extra_scale: f64) -> (F, F) {
1367 let mut n2 = F::ZERO;
1368 for v in x {
1369 n2 += *v * *v;
1370 }
1371 let n2e = n2 + F::fromf(1e-6);
1372 let inv = F::ONE / (n2e.sqrt() * F::fromf(extra_scale));
1373 (inv, n2e)
1374}
1375
1376#[allow(clippy::too_many_arguments)]
1379pub fn gdn_group_fwd<F: Fp>(
1380 cq: &[F],
1381 z: &[F],
1382 a: &[F],
1383 b: &[F],
1384 t: usize,
1385 cfg: &GdnSeqCfg,
1386 ko: usize,
1387 out: &mut [F],
1388) {
1389 let (nv, nk, dk, dv) = (cfg.nv, cfg.nk, cfg.dk, cfg.dv);
1390 let c_dim = cfg.c_dim();
1391 let kd = nk * dk;
1392 let rep = nv / nk;
1393 let vd = nv * dv;
1394 let sqdk = (dk as f64).sqrt();
1395 for hh in 0..rep {
1396 let h = ko * rep + hh;
1397 let ea = F::fromf((cfg.a_log[h] as f64).exp());
1398 let mut s = vec![F::ZERO; dk * dv];
1399 let mut kv = vec![F::ZERO; dv];
1400 let mut o = vec![F::ZERO; dv];
1401 for ti in 0..t {
1402 let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
1403 let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
1404 let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
1405 let (invq, _) = gdn_inv(qrow, sqdk);
1406 let (invk, _) = gdn_inv(krow, 1.0);
1407 let g = (-ea * softplus_f(a[ti * nv + h] + F::fromf(cfg.dt_bias[h] as f64))).exp();
1408 let beta = sigmoid_f(b[ti * nv + h]);
1409 for x in kv.iter_mut() {
1411 *x = F::ZERO;
1412 }
1413 for di in 0..dk {
1414 let kf = krow[di] * invk;
1415 let row = &mut s[di * dv..(di + 1) * dv];
1416 for dj in 0..dv {
1417 row[dj] *= g;
1418 kv[dj] += row[dj] * kf;
1419 }
1420 }
1421 for x in o.iter_mut() {
1422 *x = F::ZERO;
1423 }
1424 for di in 0..dk {
1425 let kf = krow[di] * invk;
1426 let qf = qrow[di] * invq;
1427 let row = &mut s[di * dv..(di + 1) * dv];
1428 for dj in 0..dv {
1429 row[dj] += kf * (vrow[dj] - kv[dj]) * beta;
1430 o[dj] += qf * row[dj];
1431 }
1432 }
1433 let mut ss = 0f64;
1435 for v in &o {
1436 ss += v.f64() * v.f64();
1437 }
1438 let inv = F::fromf(1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt());
1439 for dj in 0..dv {
1440 let zv = z[ti * vd + h * dv + dj];
1441 out[ti * vd + h * dv + dj] = o[dj] * inv * F::fromf(cfg.norm[dj] as f64) * silu(zv);
1442 }
1443 }
1444 }
1445}
1446
1447#[allow(clippy::too_many_arguments)]
1451pub fn gdn_group_bwd<F: Fp>(
1452 cq: &[F],
1453 z: &[F],
1454 a: &[F],
1455 b: &[F],
1456 t: usize,
1457 cfg: &GdnSeqCfg,
1458 ko: usize,
1459 dout: &[F],
1460 dcq: &mut [F],
1461 dz: &mut [F],
1462 da: &mut [F],
1463 db: &mut [F],
1464) {
1465 let (nv, nk, dk, dv) = (cfg.nv, cfg.nk, cfg.dk, cfg.dv);
1466 let c_dim = cfg.c_dim();
1467 let kd = nk * dk;
1468 let rep = nv / nk;
1469 let vd = nv * dv;
1470 let sqdk = (dk as f64).sqrt();
1471 for hh in 0..rep {
1472 let h = ko * rep + hh;
1473 let ea = F::fromf((cfg.a_log[h] as f64).exp());
1474
1475 let mut s_hist = vec![F::ZERO; (t + 1) * dk * dv]; let mut kv_hist = vec![F::ZERO; t * dv];
1478 let mut o_hist = vec![F::ZERO; t * dv];
1479 let mut g_v = vec![F::ZERO; t];
1480 let mut beta_v = vec![F::ZERO; t];
1481 let mut sp_arg = vec![F::ZERO; t]; for ti in 0..t {
1483 let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
1484 let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
1485 let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
1486 let (invq, _) = gdn_inv(qrow, sqdk);
1487 let (invk, _) = gdn_inv(krow, 1.0);
1488 let arg = a[ti * nv + h] + F::fromf(cfg.dt_bias[h] as f64);
1489 let g = (-ea * softplus_f(arg)).exp();
1490 let beta = sigmoid_f(b[ti * nv + h]);
1491 sp_arg[ti] = arg;
1492 g_v[ti] = g;
1493 beta_v[ti] = beta;
1494 let (prev, cur) = s_hist.split_at_mut((ti + 1) * dk * dv);
1495 let sp = &prev[ti * dk * dv..];
1496 let sn = &mut cur[..dk * dv];
1497 let kvr = &mut kv_hist[ti * dv..(ti + 1) * dv];
1498 for di in 0..dk {
1499 let kf = krow[di] * invk;
1500 for dj in 0..dv {
1501 let dec = sp[di * dv + dj] * g;
1502 sn[di * dv + dj] = dec;
1503 kvr[dj] += dec * kf;
1504 }
1505 }
1506 let or = &mut o_hist[ti * dv..(ti + 1) * dv];
1507 for di in 0..dk {
1508 let kf = krow[di] * invk;
1509 let qf = qrow[di] * invq;
1510 for dj in 0..dv {
1511 let sv = sn[di * dv + dj] + kf * (vrow[dj] - kvr[dj]) * beta_v[ti];
1512 sn[di * dv + dj] = sv;
1513 or[dj] += qf * sv;
1514 }
1515 }
1516 }
1517
1518 let mut ds = vec![F::ZERO; dk * dv];
1520 let mut do_o = vec![F::ZERO; dv];
1521 let mut du = vec![F::ZERO; dv];
1522 let mut dkv = vec![F::ZERO; dv];
1523 let mut dqh = vec![F::ZERO; dk];
1524 let mut dkh = vec![F::ZERO; dk];
1525 for ti in (0..t).rev() {
1526 let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
1527 let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
1528 let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
1529 let (invq, nq2) = gdn_inv(qrow, sqdk);
1530 let (invk, nk2) = gdn_inv(krow, 1.0);
1531 let g = g_v[ti];
1532 let beta = beta_v[ti];
1533 let s_t = &s_hist[(ti + 1) * dk * dv..(ti + 2) * dk * dv];
1534 let s_prev = &s_hist[ti * dk * dv..(ti + 1) * dk * dv];
1535 let kvr = &kv_hist[ti * dv..(ti + 1) * dv];
1536 let or = &o_hist[ti * dv..(ti + 1) * dv];
1537
1538 let mut ss = 0f64;
1540 for v in or {
1541 ss += v.f64() * v.f64();
1542 }
1543 let inv = F::fromf(1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt());
1544 let dofr = &dout[ti * vd + h * dv..ti * vd + (h + 1) * dv];
1545 let mut sdot = 0f64; for dj in 0..dv {
1548 let zv = z[ti * vd + h * dv + dj];
1549 let w = F::fromf(cfg.norm[dj] as f64);
1550 let weff = w * silu(zv);
1551 sdot += (dofr[dj] * weff * or[dj]).f64();
1552 dz[ti * vd + h * dv + dj] += dofr[dj] * or[dj] * inv * w * silu_bwd(zv);
1553 }
1554 let coef = F::fromf(sdot / dv as f64) * inv * inv * inv;
1555 for dj in 0..dv {
1556 let zv = z[ti * vd + h * dv + dj];
1557 let weff = F::fromf(cfg.norm[dj] as f64) * silu(zv);
1558 do_o[dj] = inv * weff * dofr[dj] - or[dj] * coef;
1559 }
1560
1561 for x in dqh.iter_mut() {
1563 *x = F::ZERO;
1564 }
1565 for di in 0..dk {
1566 let qf = qrow[di] * invq;
1567 let row = &s_t[di * dv..(di + 1) * dv];
1568 let dsr = &mut ds[di * dv..(di + 1) * dv];
1569 let mut acc = F::ZERO;
1570 for dj in 0..dv {
1571 dsr[dj] += qf * do_o[dj];
1572 acc += row[dj] * do_o[dj];
1573 }
1574 dqh[di] = acc;
1575 }
1576
1577 for x in du.iter_mut() {
1580 *x = F::ZERO;
1581 }
1582 for x in dkh.iter_mut() {
1583 *x = F::ZERO;
1584 }
1585 for di in 0..dk {
1586 let kf = krow[di] * invk;
1587 let dsr = &ds[di * dv..(di + 1) * dv];
1588 let mut acc = F::ZERO;
1589 for dj in 0..dv {
1590 du[dj] += dsr[dj] * kf;
1591 acc += dsr[dj] * (vrow[dj] - kvr[dj]) * beta;
1592 }
1593 dkh[di] = acc;
1594 }
1595 let mut dbeta = F::ZERO;
1596 for dj in 0..dv {
1597 dbeta += du[dj] * (vrow[dj] - kvr[dj]);
1598 dcq[ti * c_dim + 2 * kd + h * dv + dj] += beta * du[dj];
1600 dkv[dj] = -(beta * du[dj]);
1601 }
1602
1603 let mut dg = F::ZERO;
1608 for di in 0..dk {
1609 let kf = krow[di] * invk;
1610 let spr = &s_prev[di * dv..(di + 1) * dv];
1611 let dsr = &mut ds[di * dv..(di + 1) * dv];
1612 let mut acc = F::ZERO;
1613 for dj in 0..dv {
1614 let dspre = dsr[dj] + kf * dkv[dj];
1615 acc += (spr[dj] * g) * dkv[dj];
1616 dg += dspre * spr[dj];
1617 dsr[dj] = g * dspre;
1618 }
1619 dkh[di] += acc;
1620 }
1621
1622 let sig = sigmoid_f(sp_arg[ti]);
1624 da[ti * nv + h] += dg * g * (-ea) * sig;
1625 db[ti * nv + h] += dbeta * beta * (F::ONE - beta);
1626
1627 let mut qdot = F::ZERO;
1631 let mut kdot = F::ZERO;
1632 for di in 0..dk {
1633 qdot += dqh[di] * qrow[di];
1634 kdot += dkh[di] * krow[di];
1635 }
1636 for di in 0..dk {
1637 dcq[ti * c_dim + ko * dk + di] += invq * dqh[di] - qrow[di] * qdot * invq / nq2;
1638 dcq[ti * c_dim + kd + ko * dk + di] +=
1639 invk * dkh[di] - krow[di] * kdot * invk / nk2;
1640 }
1641 }
1642 }
1643}
1644
1645pub fn gdn_seq_fwd<F: Fp>(
1649 qkv: &[F],
1650 z: &[F],
1651 a: &[F],
1652 b: &[F],
1653 t: usize,
1654 cfg: &GdnSeqCfg,
1655 out: &mut [F],
1656) {
1657 let c_dim = cfg.c_dim();
1658 let mut pre = vec![F::ZERO; t * c_dim];
1659 let mut cq = vec![F::ZERO; t * c_dim];
1660 gdn_conv_fwd(qkv, t, c_dim, cfg.kk, cfg.conv, &mut pre, &mut cq);
1661 for ko in 0..cfg.nk {
1662 gdn_group_fwd(&cq, z, a, b, t, cfg, ko, out);
1663 }
1664}
1665
1666#[allow(clippy::too_many_arguments)]
1669pub fn gdn_seq_bwd<F: Fp>(
1670 qkv: &[F],
1671 z: &[F],
1672 a: &[F],
1673 b: &[F],
1674 t: usize,
1675 cfg: &GdnSeqCfg,
1676 dout: &[F],
1677 dqkv: &mut [F],
1678 dz: &mut [F],
1679 da: &mut [F],
1680 db: &mut [F],
1681) {
1682 let c_dim = cfg.c_dim();
1683 let mut pre = vec![F::ZERO; t * c_dim];
1684 let mut cq = vec![F::ZERO; t * c_dim];
1685 gdn_conv_fwd(qkv, t, c_dim, cfg.kk, cfg.conv, &mut pre, &mut cq);
1686 let mut dcq = vec![F::ZERO; t * c_dim];
1687 for ko in 0..cfg.nk {
1688 gdn_group_bwd(&cq, z, a, b, t, cfg, ko, dout, &mut dcq, dz, da, db);
1689 }
1690 gdn_conv_bwd(&pre, t, c_dim, cfg.kk, cfg.conv, &dcq, dqkv);
1691}