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(x: &[f32], w: &[f32], y: &mut [f32], n: usize, k: usize, m: usize, pool: Option<&Pool>) {
204 debug_assert_eq!(x.len(), n * k);
205 debug_assert_eq!(w.len(), m * k);
206 debug_assert_eq!(y.len(), n * m);
207 #[cfg(target_os = "macos")]
208 if accel::on() && n * k * m >= 1 << 18 {
209 unsafe {
211 accel::cblas_sgemm(
212 101, 111, 112, n as i32, m as i32, k as i32, 1.0, x.as_ptr(), k as i32,
213 w.as_ptr(), k as i32, 0.0, y.as_mut_ptr(), m as i32,
214 );
215 }
216 return;
217 }
218 let nb = n.div_ceil(GEMM_BLOCK);
219 let block = |r0: usize, r1: usize, y: &mut [f32]| {
220 for o in 0..m {
221 let wr = &w[o * k..(o + 1) * k];
222 for i in r0..r1 {
223 y[(i - r0) * m + o] = crate::attention::dot_f32(&x[i * k..(i + 1) * k], wr);
224 }
225 }
226 };
227 match pool {
228 Some(p) if nb > 1 => {
229 let yp = SendMut(y.as_mut_ptr());
230 p.run(&|widx, nw| {
231 for bi in (widx..nb).step_by(nw) {
232 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
233 let ys = unsafe { yp.slice(r0 * m, (r1 - r0) * m) };
235 block(r0, r1, ys);
236 }
237 });
238 }
239 _ => {
240 for bi in 0..nb {
241 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
242 block(r0, r1, &mut y[r0 * m..r1 * m]);
243 }
244 }
245 }
246}
247
248pub fn gemm_dx(dy: &[f32], w: &[f32], dx: &mut [f32], n: usize, k: usize, m: usize, pool: Option<&Pool>) {
250 debug_assert_eq!(dy.len(), n * m);
251 debug_assert_eq!(w.len(), m * k);
252 debug_assert_eq!(dx.len(), n * k);
253 #[cfg(target_os = "macos")]
254 if accel::on() && n * k * m >= 1 << 18 {
255 unsafe {
257 accel::cblas_sgemm(
258 101, 111, 111, n as i32, k as i32, m as i32, 1.0, dy.as_ptr(), m as i32,
259 w.as_ptr(), k as i32, 1.0, dx.as_mut_ptr(), k as i32,
260 );
261 }
262 return;
263 }
264 let nb = n.div_ceil(GEMM_BLOCK);
265 let block = |r0: usize, r1: usize, dxs: &mut [f32]| {
266 for o in 0..m {
267 let wr = &w[o * k..(o + 1) * k];
268 for i in r0..r1 {
269 let g = dy[i * m + o];
270 if g != 0.0 {
271 crate::attention::axpy_f32(&mut dxs[(i - r0) * k..(i - r0 + 1) * k], wr, g);
272 }
273 }
274 }
275 };
276 match pool {
277 Some(p) if nb > 1 => {
278 let dxp = SendMut(dx.as_mut_ptr());
279 p.run(&|widx, nw| {
280 for bi in (widx..nb).step_by(nw) {
281 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
282 let dxs = unsafe { dxp.slice(r0 * k, (r1 - r0) * k) };
284 block(r0, r1, dxs);
285 }
286 });
287 }
288 _ => {
289 for bi in 0..nb {
290 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
291 block(r0, r1, &mut dx[r0 * k..r1 * k]);
292 }
293 }
294 }
295}
296
297pub fn gemm_dw(dy: &[f32], x: &[f32], dw: &mut [f32], n: usize, k: usize, m: usize, pool: Option<&Pool>) {
300 debug_assert_eq!(dy.len(), n * m);
301 debug_assert_eq!(x.len(), n * k);
302 debug_assert_eq!(dw.len(), m * k);
303 #[cfg(target_os = "macos")]
304 if accel::on() && n * k * m >= 1 << 18 {
305 unsafe {
307 accel::cblas_sgemm(
308 101, 112, 111, m as i32, k as i32, n as i32, 1.0, dy.as_ptr(), m as i32,
309 x.as_ptr(), k as i32, 1.0, dw.as_mut_ptr(), k as i32,
310 );
311 }
312 return;
313 }
314 let range = |o0: usize, o1: usize, dws: &mut [f32]| {
315 let nb = n.div_ceil(GEMM_BLOCK);
317 for bi in 0..nb {
318 let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
319 for o in o0..o1 {
320 let dwr = &mut dws[(o - o0) * k..(o - o0 + 1) * k];
321 for i in r0..r1 {
322 let g = dy[i * m + o];
323 if g != 0.0 {
324 crate::attention::axpy_f32(dwr, &x[i * k..(i + 1) * k], g);
325 }
326 }
327 }
328 }
329 };
330 match pool {
331 Some(p) if m >= 8 => {
332 let dwp = SendMut(dw.as_mut_ptr());
333 p.run(&|widx, nw| {
334 let (o0, o1) = (widx * m / nw, (widx + 1) * m / nw);
335 if o0 < o1 {
336 let dws = unsafe { dwp.slice(o0 * k, (o1 - o0) * k) };
338 range(o0, o1, dws);
339 }
340 });
341 }
342 _ => range(0, m, dw),
343 }
344}
345
346#[inline]
349pub fn silu<F: Fp>(x: F) -> F {
350 x / (F::ONE + (-x).exp())
351}
352
353#[inline]
355pub fn silu_bwd<F: Fp>(x: F) -> F {
356 let s = F::ONE / (F::ONE + (-x).exp());
357 s * (F::ONE + x * (F::ONE - s))
358}
359
360pub fn rmsnorm_fwd<F: Fp>(x: &[F], w: &[F], eps: f64, gemma: bool, y: &mut [F], inv_out: &mut [F]) {
367 let d = w.len();
368 let n = x.len() / d;
369 for r in 0..n {
370 let xr = &x[r * d..(r + 1) * d];
371 let mut ss = 0f64;
372 for v in xr {
373 ss += v.f64() * v.f64();
374 }
375 let inv = F::fromf(1.0 / (ss / d as f64 + eps).sqrt());
376 inv_out[r] = inv;
377 let yr = &mut y[r * d..(r + 1) * d];
378 for j in 0..d {
379 let weff = if gemma { F::ONE + w[j] } else { w[j] };
380 yr[j] = xr[j] * inv * weff;
381 }
382 }
383}
384
385pub fn rmsnorm_bwd<F: Fp>(
391 x: &[F],
392 w: &[F],
393 inv: &[F],
394 dy: &[F],
395 gemma: bool,
396 dx: &mut [F],
397 mut dw: Option<&mut [F]>,
398) {
399 let d = w.len();
400 let n = x.len() / d;
401 for r in 0..n {
402 let xr = &x[r * d..(r + 1) * d];
403 let dyr = &dy[r * d..(r + 1) * d];
404 let iv = inv[r];
405 let mut s = 0f64;
406 for j in 0..d {
407 let weff = if gemma { F::ONE + w[j] } else { w[j] };
408 s += (dyr[j] * weff * xr[j]).f64();
409 }
410 let coef = F::fromf(s / d as f64) * iv * iv * iv;
411 let dxr = &mut dx[r * d..(r + 1) * d];
412 for j in 0..d {
413 let weff = if gemma { F::ONE + w[j] } else { w[j] };
414 dxr[j] += iv * weff * dyr[j] - xr[j] * coef;
415 }
416 if let Some(dwv) = dw.as_deref_mut() {
417 for j in 0..d {
418 dwv[j] += dyr[j] * xr[j] * iv;
419 }
420 }
421 }
422}
423
424pub fn rope_fwd<F: Fp>(x: &mut [F], position: usize, inv_freq: &[f64]) {
429 let half = inv_freq.len();
430 for (i, &freq) in inv_freq.iter().enumerate() {
431 let angle = position as f64 * freq;
432 let (sin, cos) = (F::fromf(angle.sin()), F::fromf(angle.cos()));
433 let x0 = x[i];
434 let x1 = x[i + half];
435 x[i] = x0 * cos - x1 * sin;
436 x[i + half] = x0 * sin + x1 * cos;
437 }
438}
439
440pub fn rope_bwd<F: Fp>(dy: &mut [F], position: usize, inv_freq: &[f64]) {
443 let half = inv_freq.len();
444 for (i, &freq) in inv_freq.iter().enumerate() {
445 let angle = position as f64 * freq;
446 let (sin, cos) = (F::fromf(angle.sin()), F::fromf(angle.cos()));
447 let g0 = dy[i];
448 let g1 = dy[i + half];
449 dy[i] = g0 * cos + g1 * sin;
450 dy[i + half] = -g0 * sin + g1 * cos;
451 }
452}
453
454pub fn seg_means<F: Fp>(x: &[F], t: usize, d: usize, m: usize, out: &mut [F]) {
459 for i in 0..m {
460 let (lo, hi) = (i * t / m, (i + 1) * t / m);
461 let or = &mut out[i * d..(i + 1) * d];
462 for v in or.iter_mut() {
463 *v = F::ZERO;
464 }
465 for j in lo..hi {
466 for c in 0..d {
467 or[c] += x[j * d + c];
468 }
469 }
470 let inv = F::fromf(1.0 / (hi - lo) as f64);
471 for v in or.iter_mut() {
472 *v *= inv;
473 }
474 }
475}
476
477pub fn seg_means_bwd<F: Fp>(dl: &[F], t: usize, d: usize, m: usize, dx: &mut [F]) {
480 for i in 0..m {
481 let (lo, hi) = (i * t / m, (i + 1) * t / m);
482 let inv = F::fromf(1.0 / (hi - lo) as f64);
483 let dlr = &dl[i * d..(i + 1) * d];
484 for j in lo..hi {
485 for c in 0..d {
486 dx[j * d + c] += dlr[c] * inv;
487 }
488 }
489 }
490}
491
492#[allow(clippy::needless_range_loop)] pub fn attn_head_fwd<F: Fp>(q: &[F], k: &[F], v: &[F], t: usize, d: usize, dv: usize, out: &mut [F]) {
498 let scale = F::fromf(1.0 / (d as f64).sqrt());
499 let mut row = vec![F::ZERO; t];
500 for ti in 0..t {
501 let qr = &q[ti * d..(ti + 1) * d];
502 let mut mx = F::fromf(f64::NEG_INFINITY);
503 for j in 0..=ti {
504 let s = dot(qr, &k[j * d..(j + 1) * d]) * scale;
505 row[j] = s;
506 mx = mx.maxf(s);
507 }
508 let mut den = F::ZERO;
509 for j in 0..=ti {
510 row[j] = (row[j] - mx).exp();
511 den += row[j];
512 }
513 let or = &mut out[ti * dv..(ti + 1) * dv];
514 for o in or.iter_mut() {
515 *o = F::ZERO;
516 }
517 for j in 0..=ti {
518 let p = row[j] / den;
519 for (o, vv) in or.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
520 *o += p * *vv;
521 }
522 }
523 }
524}
525
526#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
530pub fn attn_head_bwd<F: Fp>(
531 q: &[F],
532 k: &[F],
533 v: &[F],
534 dout: &[F],
535 t: usize,
536 d: usize,
537 dv: usize,
538 dq: &mut [F],
539 dk: &mut [F],
540 dvv: &mut [F],
541) {
542 let scale = F::fromf(1.0 / (d as f64).sqrt());
543 let mut row = vec![F::ZERO; t];
544 for ti in 0..t {
545 let qr = &q[ti * d..(ti + 1) * d];
546 let mut mx = F::fromf(f64::NEG_INFINITY);
547 for j in 0..=ti {
548 let s = dot(qr, &k[j * d..(j + 1) * d]) * scale;
549 row[j] = s;
550 mx = mx.maxf(s);
551 }
552 let mut den = F::ZERO;
553 for j in 0..=ti {
554 row[j] = (row[j] - mx).exp();
555 den += row[j];
556 }
557 let dor = &dout[ti * dv..(ti + 1) * dv];
558 let mut pdp = F::ZERO;
560 let mut dp = vec![F::ZERO; ti + 1];
561 for j in 0..=ti {
562 let p = row[j] / den;
563 row[j] = p; dp[j] = dot(dor, &v[j * dv..(j + 1) * dv]);
565 pdp += p * dp[j];
566 }
567 let dqr = &mut dq[ti * d..(ti + 1) * d];
568 for j in 0..=ti {
569 let p = row[j];
570 for (dvo, o) in dvv[j * dv..(j + 1) * dv].iter_mut().zip(dor) {
572 *dvo += p * *o;
573 }
574 let ds = p * (dp[j] - pdp) * scale;
576 let kr = &k[j * d..(j + 1) * d];
577 for c in 0..d {
578 dqr[c] += ds * kr[c];
579 }
580 let dkr = &mut dk[j * d..(j + 1) * d];
581 for c in 0..d {
582 dkr[c] += ds * qr[c];
583 }
584 }
585 }
586}
587
588#[derive(Clone, Copy, Debug)]
592pub struct NysCfg {
593 pub m: usize,
595 pub w: usize,
597 pub sink: usize,
599 pub prefill: Option<usize>,
611}
612
613impl NysCfg {
614 #[inline]
617 pub fn prefill_len(&self, t: usize) -> usize {
618 self.prefill.unwrap_or(t / 2).clamp(1, t)
619 }
620}
621
622const NYS_DEN_EPS: f64 = 1e-30;
624
625struct NysGraph<F: Fp> {
629 m_eff: usize,
630 tp: usize,
633 q_l: Vec<F>,
634 k_l: Vec<F>,
635 mu: Vec<F>,
636 fu: Vec<F>,
637 e: Vec<F>,
638 fumu: Vec<F>,
639 far_keep: Vec<bool>,
648 wmat: Vec<F>,
651 c_row: Vec<F>,
652 den: Vec<F>,
653}
654
655#[inline]
660fn nys_exact_row(ti: usize, tp: usize) -> bool {
661 ti < tp
662}
663
664#[inline]
666fn nys_near(ti: usize, j: usize, w: usize, sink: usize) -> bool {
667 ti - j < w || j < sink
668}
669
670fn nys_graph<F: Fp>(
675 q: &[F],
676 k: &[F],
677 t: usize,
678 d: usize,
679 cfg: &NysCfg,
680 mu_override: Option<&[F]>,
681) -> NysGraph<F> {
682 let scale = 1.0 / (d as f64).sqrt();
683 let fscale = F::fromf(scale);
684 let tp = cfg.prefill_len(t);
689 let m_eff = (tp / 8).clamp(4, cfg.m);
690 let mut q_l = vec![F::ZERO; m_eff * d];
691 let mut k_l = vec![F::ZERO; m_eff * d];
692 seg_means(&q[..tp * d], tp, d, m_eff, &mut q_l);
693 seg_means(&k[..tp * d], tp, d, m_eff, &mut k_l);
694
695 let mut au = vec![0f64; m_eff * m_eff];
697 for i in 0..m_eff {
698 for j in 0..m_eff {
699 let mut s = 0f64;
700 for c in 0..d {
701 s += q_l[i * d + c].f64() * k_l[j * d + c].f64();
702 }
703 au[i * m_eff + j] = (s * scale).exp();
704 }
705 }
706 let mu: Vec<F> = match mu_override {
707 Some(m) => m.to_vec(),
708 None => crate::nystrom::ridge_pinv(&au, m_eff)
709 .iter()
710 .map(|&x| F::fromf(x))
711 .collect(),
712 };
713
714 let mut fu = vec![F::ZERO; t * m_eff];
716 for ti in 0..t {
717 for i in 0..m_eff {
718 fu[ti * m_eff + i] =
719 (dot(&q[ti * d..(ti + 1) * d], &k_l[i * d..(i + 1) * d]) * fscale).exp();
720 }
721 }
722 let mut e = vec![F::ZERO; m_eff * t];
723 for i in 0..m_eff {
724 for j in 0..t {
725 e[i * t + j] =
726 (dot(&q_l[i * d..(i + 1) * d], &k[j * d..(j + 1) * d]) * fscale).exp();
727 }
728 }
729 let mut fumu = vec![F::ZERO; t * m_eff];
730 matmul_nt(
731 &fu,
732 &transpose(&mu, m_eff, m_eff),
735 &mut fumu,
736 t,
737 m_eff,
738 m_eff,
739 );
740 let mut a = vec![F::ZERO; t * t];
745 for ti in tp..t {
746 let fr = &fumu[ti * m_eff..(ti + 1) * m_eff];
747 let ar = &mut a[ti * t..ti * t + ti + 1]; for (i, &f) in fr.iter().enumerate() {
749 let er = &e[i * t..i * t + ti + 1];
750 for (av, ev) in ar.iter_mut().zip(er) {
751 *av += f * *ev;
752 }
753 }
754 }
755
756 let mut wmat = vec![F::ZERO; t * t];
759 let mut c_row = vec![F::ZERO; t];
760 let mut den = vec![F::ZERO; t];
761 let mut far_keep = vec![false; t];
762 let mut lg_row = vec![F::ZERO; t];
763 for ti in 0..t {
764 let qr = &q[ti * d..(ti + 1) * d];
765 let mut c = F::fromf(f64::NEG_INFINITY);
766 for j in 0..=ti {
767 let s = dot(qr, &k[j * d..(j + 1) * d]) * fscale;
768 lg_row[j] = s;
769 c = c.maxf(s);
770 }
771 c_row[ti] = c;
772 let emc = (-c).exp();
773 let exact_row = nys_exact_row(ti, tp);
774
775 let mut far_sum = F::ZERO;
790 if !exact_row {
791 for j in 0..=ti {
792 if !nys_near(ti, j, cfg.w, cfg.sink) {
793 far_sum += a[ti * t + j];
794 }
795 }
796 }
797 let keep = !exact_row && far_sum.f64() >= 0.0;
798 far_keep[ti] = keep;
799
800 let wr = &mut wmat[ti * t..(ti + 1) * t];
801 let mut dsum = F::ZERO;
802 for j in 0..=ti {
803 let wv = if exact_row || nys_near(ti, j, cfg.w, cfg.sink) {
804 (lg_row[j] - c).exp()
805 } else if keep {
806 a[ti * t + j] * emc
807 } else {
808 F::ZERO
809 };
810 wr[j] = wv;
811 dsum += wv;
812 }
813 den[ti] = dsum.maxf(F::fromf(NYS_DEN_EPS));
814 }
815 NysGraph { m_eff, tp, q_l, k_l, mu, fu, e, fumu, far_keep, wmat, c_row, den }
816}
817
818fn transpose<F: Fp>(x: &[F], rows: usize, cols: usize) -> Vec<F> {
819 let mut out = vec![F::ZERO; rows * cols];
820 for r in 0..rows {
821 for c in 0..cols {
822 out[c * rows + r] = x[r * cols + c];
823 }
824 }
825 out
826}
827
828#[allow(clippy::too_many_arguments)]
831pub fn nystrom_head_fwd<F: Fp>(
832 q: &[F],
833 k: &[F],
834 v: &[F],
835 t: usize,
836 d: usize,
837 dv: usize,
838 cfg: &NysCfg,
839 out: &mut [F],
840) {
841 if nys_degenerate(t, cfg) {
842 attn_head_fwd(q, k, v, t, d, dv, out);
843 return;
844 }
845 nystrom_head_fwd_mu(q, k, v, t, d, dv, cfg, None, out);
846}
847
848#[inline]
857fn nys_degenerate(t: usize, cfg: &NysCfg) -> bool {
858 cfg.prefill_len(t) <= cfg.w + cfg.sink + 8
859}
860
861#[doc(hidden)]
864pub fn nystrom_mu_for_test<F: Fp>(q: &[F], k: &[F], t: usize, d: usize, cfg: &NysCfg) -> Vec<F> {
865 nys_graph(q, k, t, d, cfg, None).mu
866}
867
868#[doc(hidden)]
871#[allow(clippy::too_many_arguments)]
872pub fn nystrom_head_fwd_mu<F: Fp>(
873 q: &[F],
874 k: &[F],
875 v: &[F],
876 t: usize,
877 d: usize,
878 dv: usize,
879 cfg: &NysCfg,
880 mu_override: Option<&[F]>,
881 out: &mut [F],
882) {
883 let g = nys_graph(q, k, t, d, cfg, mu_override);
884 for ti in 0..t {
885 let wr = &g.wmat[ti * t..(ti + 1) * t];
886 let den = g.den[ti];
887 let or = &mut out[ti * dv..(ti + 1) * dv];
888 for o in or.iter_mut() {
889 *o = F::ZERO;
890 }
891 for j in 0..=ti {
892 let p = wr[j] / den;
893 if p.f64() != 0.0 {
894 for (o, vv) in or.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
895 *o += p * *vv;
896 }
897 }
898 }
899 }
900}
901
902#[allow(clippy::too_many_arguments)]
913pub fn nystrom_head_bwd<F: Fp>(
914 q: &[F],
915 k: &[F],
916 v: &[F],
917 dout: &[F],
918 t: usize,
919 d: usize,
920 dv: usize,
921 cfg: &NysCfg,
922 dq: &mut [F],
923 dk: &mut [F],
924 dvv: &mut [F],
925) {
926 if nys_degenerate(t, cfg) {
927 attn_head_bwd(q, k, v, dout, t, d, dv, dq, dk, dvv);
928 return;
929 }
930 nystrom_head_bwd_mu(q, k, v, dout, t, d, dv, cfg, None, dq, dk, dvv);
931}
932
933#[doc(hidden)]
935#[allow(clippy::too_many_arguments)]
936pub fn nystrom_head_bwd_mu<F: Fp>(
937 q: &[F],
938 k: &[F],
939 v: &[F],
940 dout: &[F],
941 t: usize,
942 d: usize,
943 dv: usize,
944 cfg: &NysCfg,
945 mu_override: Option<&[F]>,
946 dq: &mut [F],
947 dk: &mut [F],
948 dvv: &mut [F],
949) {
950 let scale = F::fromf(1.0 / (d as f64).sqrt());
951 let g = nys_graph(q, k, t, d, cfg, mu_override);
952 let m_eff = g.m_eff;
953
954 let mut dwmat = vec![F::ZERO; t * t];
956 let mut out_row = vec![F::ZERO; dv];
957 for ti in 0..t {
958 let wr = &g.wmat[ti * t..(ti + 1) * t];
959 let den = g.den[ti];
960 for o in out_row.iter_mut() {
961 *o = F::ZERO;
962 }
963 for j in 0..=ti {
964 let p = wr[j] / den;
965 if p.f64() != 0.0 {
966 for (o, vv) in out_row.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
967 *o += p * *vv;
968 }
969 }
970 }
971 let dor = &dout[ti * dv..(ti + 1) * dv];
972 let dwr = &mut dwmat[ti * t..(ti + 1) * t];
973 for j in 0..=ti {
974 let p = wr[j] / den;
976 if p.f64() != 0.0 {
977 for (dvo, o) in dvv[j * dv..(j + 1) * dv].iter_mut().zip(dor) {
978 *dvo += p * *o;
979 }
980 }
981 let mut s = F::ZERO;
983 for c in 0..dv {
984 s += dor[c] * (v[j * dv + c] - out_row[c]);
985 }
986 dwr[j] = s / den;
987 }
988 }
989
990 for ti in 0..t {
993 let qr = &q[ti * d..(ti + 1) * d];
994 let dqr_base = ti * d;
995 for j in 0..=ti {
996 if !(nys_exact_row(ti, g.tp) || nys_near(ti, j, cfg.w, cfg.sink)) {
997 continue;
998 }
999 let dlg = dwmat[ti * t + j] * g.wmat[ti * t + j] * scale;
1000 if dlg.f64() == 0.0 {
1001 continue;
1002 }
1003 let kr = &k[j * d..(j + 1) * d];
1004 for c in 0..d {
1005 dq[dqr_base + c] += dlg * kr[c];
1006 }
1007 let dkr = &mut dk[j * d..(j + 1) * d];
1008 for c in 0..d {
1009 dkr[c] += dlg * qr[c];
1010 }
1011 }
1012 }
1013
1014 let mut da = vec![F::ZERO; t * t];
1029 for ti in g.tp..t {
1030 if !g.far_keep[ti] {
1031 continue; }
1033 let emc = (-g.c_row[ti]).exp();
1034 for j in 0..=ti {
1035 if nys_near(ti, j, cfg.w, cfg.sink) {
1036 continue;
1037 }
1038 da[ti * t + j] = dwmat[ti * t + j] * emc;
1039 }
1040 }
1041 let mut dfumu = vec![F::ZERO; t * m_eff];
1043 for ti in 0..t {
1044 let dar = &da[ti * t..ti * t + ti + 1];
1045 let dfr = &mut dfumu[ti * m_eff..(ti + 1) * m_eff];
1046 for (i, df) in dfr.iter_mut().enumerate() {
1047 let er = &g.e[i * t..i * t + ti + 1];
1048 let mut s = F::ZERO;
1049 for (av, ev) in dar.iter().zip(er) {
1050 s += *av * *ev;
1051 }
1052 *df = s;
1053 }
1054 }
1055 let mut dfu = vec![F::ZERO; t * m_eff];
1057 matmul_nt(&dfumu, &g.mu, &mut dfu, t, m_eff, m_eff);
1058 let mut de = vec![F::ZERO; m_eff * t];
1060 for ti in 0..t {
1061 let dar = &da[ti * t..ti * t + ti + 1];
1062 let fr = &g.fumu[ti * m_eff..(ti + 1) * m_eff];
1063 for (i, &f) in fr.iter().enumerate() {
1064 if f.f64() == 0.0 {
1065 continue;
1066 }
1067 let der = &mut de[i * t..i * t + ti + 1];
1068 for (dev, av) in der.iter_mut().zip(dar) {
1069 *dev += f * *av;
1070 }
1071 }
1072 }
1073 let mut dq_l = vec![F::ZERO; m_eff * d];
1076 let mut dk_l = vec![F::ZERO; m_eff * d];
1077 for ti in 0..t {
1078 let qr = &q[ti * d..(ti + 1) * d];
1079 for i in 0..m_eff {
1080 let dlg = dfu[ti * m_eff + i] * g.fu[ti * m_eff + i] * scale;
1081 if dlg.f64() == 0.0 {
1082 continue;
1083 }
1084 let klr = &g.k_l[i * d..(i + 1) * d];
1085 for c in 0..d {
1086 dq[ti * d + c] += dlg * klr[c];
1087 }
1088 let dklr = &mut dk_l[i * d..(i + 1) * d];
1089 for c in 0..d {
1090 dklr[c] += dlg * qr[c];
1091 }
1092 }
1093 }
1094 for i in 0..m_eff {
1095 let qlr = &g.q_l[i * d..(i + 1) * d];
1096 for j in 0..t {
1097 let dlg = de[i * t + j] * g.e[i * t + j] * scale;
1098 if dlg.f64() == 0.0 {
1099 continue;
1100 }
1101 let kr = &k[j * d..(j + 1) * d];
1102 let dqlr = &mut dq_l[i * d..(i + 1) * d];
1103 for c in 0..d {
1104 dqlr[c] += dlg * kr[c];
1105 }
1106 for c in 0..d {
1107 dk[j * d + c] += dlg * qlr[c];
1108 }
1109 }
1110 }
1111 let tp = g.tp;
1115 seg_means_bwd(&dq_l, tp, d, m_eff, &mut dq[..tp * d]);
1116 seg_means_bwd(&dk_l, tp, d, m_eff, &mut dk[..tp * d]);
1117}
1118
1119pub fn ce_kl_position<F: Fp>(
1127 s_logits: &[F],
1128 t_logits: &[F],
1129 target: usize,
1130 kl_w: f64,
1131 inv_n: f64,
1132 dlogits: &mut [F],
1133) -> (f64, f64) {
1134 let vsz = s_logits.len();
1135 debug_assert_eq!(t_logits.len(), vsz);
1136 let mut smax = f64::NEG_INFINITY;
1138 let mut tmax = f64::NEG_INFINITY;
1139 for i in 0..vsz {
1140 smax = smax.max(s_logits[i].f64());
1141 tmax = tmax.max(t_logits[i].f64());
1142 }
1143 let mut ssum = 0f64;
1144 let mut tsum = 0f64;
1145 for i in 0..vsz {
1146 ssum += (s_logits[i].f64() - smax).exp();
1147 tsum += (t_logits[i].f64() - tmax).exp();
1148 }
1149 let slz = smax + ssum.ln();
1150 let tlz = tmax + tsum.ln();
1151 let ce = slz - s_logits[target].f64();
1152 let mut kl = 0f64;
1153 for i in 0..vsz {
1154 let ls = s_logits[i].f64() - slz;
1155 let lt = t_logits[i].f64() - tlz;
1156 let pt = lt.exp();
1157 let ps = ls.exp();
1158 if pt > 0.0 {
1159 kl += pt * (lt - ls);
1160 }
1161 let mut gd = (1.0 - kl_w) * ps + kl_w * (ps - pt);
1162 if i == target {
1163 gd -= 1.0 - kl_w;
1164 }
1165 dlogits[i] = F::fromf(gd * inv_n);
1166 }
1167 (ce, kl)
1168}
1169
1170pub struct GdnSeqCfg<'a> {
1187 pub nv: usize,
1188 pub nk: usize,
1189 pub dk: usize,
1190 pub dv: usize,
1191 pub kk: usize,
1192 pub rms_eps: f64,
1193 pub conv: &'a [f32],
1196 pub a_log: &'a [f32],
1198 pub dt_bias: &'a [f32],
1200 pub norm: &'a [f32],
1202}
1203
1204impl GdnSeqCfg<'_> {
1205 pub fn c_dim(&self) -> usize {
1206 2 * self.nk * self.dk + self.nv * self.dv
1207 }
1208}
1209
1210#[inline]
1211fn softplus_f<F: Fp>(x: F) -> F {
1212 if x.f64() > 20.0 {
1215 x
1216 } else {
1217 F::fromf(x.f64().exp().ln_1p())
1218 }
1219}
1220
1221#[inline]
1222fn sigmoid_f<F: Fp>(x: F) -> F {
1223 F::ONE / (F::ONE + (-x).exp())
1224}
1225
1226pub fn gdn_conv_fwd<F: Fp>(
1230 raw: &[F],
1231 t: usize,
1232 c_dim: usize,
1233 kk: usize,
1234 conv: &[f32],
1235 pre: &mut [F],
1236 cq: &mut [F],
1237) {
1238 for ti in 0..t {
1239 for c in 0..c_dim {
1240 let taps = &conv[c * kk..(c + 1) * kk];
1241 let mut acc = F::ZERO;
1242 for (j, &tap) in taps.iter().enumerate() {
1243 let p = ti as isize - (kk as isize - 1) + j as isize;
1245 if p >= 0 {
1246 acc += raw[p as usize * c_dim + c] * F::fromf(tap as f64);
1247 }
1248 }
1249 pre[ti * c_dim + c] = acc;
1250 cq[ti * c_dim + c] = silu(acc);
1251 }
1252 }
1253}
1254
1255pub fn gdn_conv_bwd<F: Fp>(
1257 pre: &[F],
1258 t: usize,
1259 c_dim: usize,
1260 kk: usize,
1261 conv: &[f32],
1262 dcq: &[F],
1263 draw: &mut [F],
1264) {
1265 for ti in 0..t {
1266 for c in 0..c_dim {
1267 let g = dcq[ti * c_dim + c];
1268 if g.f64() == 0.0 {
1269 continue;
1270 }
1271 let dp = g * silu_bwd(pre[ti * c_dim + c]);
1272 let taps = &conv[c * kk..(c + 1) * kk];
1273 for (j, &tap) in taps.iter().enumerate() {
1274 let p = ti as isize - (kk as isize - 1) + j as isize;
1275 if p >= 0 {
1276 draw[p as usize * c_dim + c] += dp * F::fromf(tap as f64);
1277 }
1278 }
1279 }
1280 }
1281}
1282
1283#[inline]
1286fn gdn_inv<F: Fp>(x: &[F], extra_scale: f64) -> (F, F) {
1287 let mut n2 = F::ZERO;
1288 for v in x {
1289 n2 += *v * *v;
1290 }
1291 let n2e = n2 + F::fromf(1e-6);
1292 let inv = F::ONE / (n2e.sqrt() * F::fromf(extra_scale));
1293 (inv, n2e)
1294}
1295
1296#[allow(clippy::too_many_arguments)]
1299pub fn gdn_group_fwd<F: Fp>(
1300 cq: &[F],
1301 z: &[F],
1302 a: &[F],
1303 b: &[F],
1304 t: usize,
1305 cfg: &GdnSeqCfg,
1306 ko: usize,
1307 out: &mut [F],
1308) {
1309 let (nv, nk, dk, dv) = (cfg.nv, cfg.nk, cfg.dk, cfg.dv);
1310 let c_dim = cfg.c_dim();
1311 let kd = nk * dk;
1312 let rep = nv / nk;
1313 let vd = nv * dv;
1314 let sqdk = (dk as f64).sqrt();
1315 for hh in 0..rep {
1316 let h = ko * rep + hh;
1317 let ea = F::fromf((cfg.a_log[h] as f64).exp());
1318 let mut s = vec![F::ZERO; dk * dv];
1319 let mut kv = vec![F::ZERO; dv];
1320 let mut o = vec![F::ZERO; dv];
1321 for ti in 0..t {
1322 let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
1323 let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
1324 let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
1325 let (invq, _) = gdn_inv(qrow, sqdk);
1326 let (invk, _) = gdn_inv(krow, 1.0);
1327 let g = (-ea * softplus_f(a[ti * nv + h] + F::fromf(cfg.dt_bias[h] as f64))).exp();
1328 let beta = sigmoid_f(b[ti * nv + h]);
1329 for x in kv.iter_mut() {
1331 *x = F::ZERO;
1332 }
1333 for di in 0..dk {
1334 let kf = krow[di] * invk;
1335 let row = &mut s[di * dv..(di + 1) * dv];
1336 for dj in 0..dv {
1337 row[dj] *= g;
1338 kv[dj] += row[dj] * kf;
1339 }
1340 }
1341 for x in o.iter_mut() {
1342 *x = F::ZERO;
1343 }
1344 for di in 0..dk {
1345 let kf = krow[di] * invk;
1346 let qf = qrow[di] * invq;
1347 let row = &mut s[di * dv..(di + 1) * dv];
1348 for dj in 0..dv {
1349 row[dj] += kf * (vrow[dj] - kv[dj]) * beta;
1350 o[dj] += qf * row[dj];
1351 }
1352 }
1353 let mut ss = 0f64;
1355 for v in &o {
1356 ss += v.f64() * v.f64();
1357 }
1358 let inv = F::fromf(1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt());
1359 for dj in 0..dv {
1360 let zv = z[ti * vd + h * dv + dj];
1361 out[ti * vd + h * dv + dj] =
1362 o[dj] * inv * F::fromf(cfg.norm[dj] as f64) * silu(zv);
1363 }
1364 }
1365 }
1366}
1367
1368#[allow(clippy::too_many_arguments)]
1372pub fn gdn_group_bwd<F: Fp>(
1373 cq: &[F],
1374 z: &[F],
1375 a: &[F],
1376 b: &[F],
1377 t: usize,
1378 cfg: &GdnSeqCfg,
1379 ko: usize,
1380 dout: &[F],
1381 dcq: &mut [F],
1382 dz: &mut [F],
1383 da: &mut [F],
1384 db: &mut [F],
1385) {
1386 let (nv, nk, dk, dv) = (cfg.nv, cfg.nk, cfg.dk, cfg.dv);
1387 let c_dim = cfg.c_dim();
1388 let kd = nk * dk;
1389 let rep = nv / nk;
1390 let vd = nv * dv;
1391 let sqdk = (dk as f64).sqrt();
1392 for hh in 0..rep {
1393 let h = ko * rep + hh;
1394 let ea = F::fromf((cfg.a_log[h] as f64).exp());
1395
1396 let mut s_hist = vec![F::ZERO; (t + 1) * dk * dv]; let mut kv_hist = vec![F::ZERO; t * dv];
1399 let mut o_hist = vec![F::ZERO; t * dv];
1400 let mut g_v = vec![F::ZERO; t];
1401 let mut beta_v = vec![F::ZERO; t];
1402 let mut sp_arg = vec![F::ZERO; t]; for ti in 0..t {
1404 let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
1405 let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
1406 let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
1407 let (invq, _) = gdn_inv(qrow, sqdk);
1408 let (invk, _) = gdn_inv(krow, 1.0);
1409 let arg = a[ti * nv + h] + F::fromf(cfg.dt_bias[h] as f64);
1410 let g = (-ea * softplus_f(arg)).exp();
1411 let beta = sigmoid_f(b[ti * nv + h]);
1412 sp_arg[ti] = arg;
1413 g_v[ti] = g;
1414 beta_v[ti] = beta;
1415 let (prev, cur) = s_hist.split_at_mut((ti + 1) * dk * dv);
1416 let sp = &prev[ti * dk * dv..];
1417 let sn = &mut cur[..dk * dv];
1418 let kvr = &mut kv_hist[ti * dv..(ti + 1) * dv];
1419 for di in 0..dk {
1420 let kf = krow[di] * invk;
1421 for dj in 0..dv {
1422 let dec = sp[di * dv + dj] * g;
1423 sn[di * dv + dj] = dec;
1424 kvr[dj] += dec * kf;
1425 }
1426 }
1427 let or = &mut o_hist[ti * dv..(ti + 1) * dv];
1428 for di in 0..dk {
1429 let kf = krow[di] * invk;
1430 let qf = qrow[di] * invq;
1431 for dj in 0..dv {
1432 let sv = sn[di * dv + dj] + kf * (vrow[dj] - kvr[dj]) * beta_v[ti];
1433 sn[di * dv + dj] = sv;
1434 or[dj] += qf * sv;
1435 }
1436 }
1437 }
1438
1439 let mut ds = vec![F::ZERO; dk * dv];
1441 let mut do_o = vec![F::ZERO; dv];
1442 let mut du = vec![F::ZERO; dv];
1443 let mut dkv = vec![F::ZERO; dv];
1444 let mut dqh = vec![F::ZERO; dk];
1445 let mut dkh = vec![F::ZERO; dk];
1446 for ti in (0..t).rev() {
1447 let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
1448 let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
1449 let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
1450 let (invq, nq2) = gdn_inv(qrow, sqdk);
1451 let (invk, nk2) = gdn_inv(krow, 1.0);
1452 let g = g_v[ti];
1453 let beta = beta_v[ti];
1454 let s_t = &s_hist[(ti + 1) * dk * dv..(ti + 2) * dk * dv];
1455 let s_prev = &s_hist[ti * dk * dv..(ti + 1) * dk * dv];
1456 let kvr = &kv_hist[ti * dv..(ti + 1) * dv];
1457 let or = &o_hist[ti * dv..(ti + 1) * dv];
1458
1459 let mut ss = 0f64;
1461 for v in or {
1462 ss += v.f64() * v.f64();
1463 }
1464 let inv = F::fromf(1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt());
1465 let dofr = &dout[ti * vd + h * dv..ti * vd + (h + 1) * dv];
1466 let mut sdot = 0f64; for dj in 0..dv {
1469 let zv = z[ti * vd + h * dv + dj];
1470 let w = F::fromf(cfg.norm[dj] as f64);
1471 let weff = w * silu(zv);
1472 sdot += (dofr[dj] * weff * or[dj]).f64();
1473 dz[ti * vd + h * dv + dj] += dofr[dj] * or[dj] * inv * w * silu_bwd(zv);
1474 }
1475 let coef = F::fromf(sdot / dv as f64) * inv * inv * inv;
1476 for dj in 0..dv {
1477 let zv = z[ti * vd + h * dv + dj];
1478 let weff = F::fromf(cfg.norm[dj] as f64) * silu(zv);
1479 do_o[dj] = inv * weff * dofr[dj] - or[dj] * coef;
1480 }
1481
1482 for x in dqh.iter_mut() {
1484 *x = F::ZERO;
1485 }
1486 for di in 0..dk {
1487 let qf = qrow[di] * invq;
1488 let row = &s_t[di * dv..(di + 1) * dv];
1489 let dsr = &mut ds[di * dv..(di + 1) * dv];
1490 let mut acc = F::ZERO;
1491 for dj in 0..dv {
1492 dsr[dj] += qf * do_o[dj];
1493 acc += row[dj] * do_o[dj];
1494 }
1495 dqh[di] = acc;
1496 }
1497
1498 for x in du.iter_mut() {
1501 *x = F::ZERO;
1502 }
1503 for x in dkh.iter_mut() {
1504 *x = F::ZERO;
1505 }
1506 for di in 0..dk {
1507 let kf = krow[di] * invk;
1508 let dsr = &ds[di * dv..(di + 1) * dv];
1509 let mut acc = F::ZERO;
1510 for dj in 0..dv {
1511 du[dj] += dsr[dj] * kf;
1512 acc += dsr[dj] * (vrow[dj] - kvr[dj]) * beta;
1513 }
1514 dkh[di] = acc;
1515 }
1516 let mut dbeta = F::ZERO;
1517 for dj in 0..dv {
1518 dbeta += du[dj] * (vrow[dj] - kvr[dj]);
1519 dcq[ti * c_dim + 2 * kd + h * dv + dj] += beta * du[dj];
1521 dkv[dj] = -(beta * du[dj]);
1522 }
1523
1524 let mut dg = F::ZERO;
1529 for di in 0..dk {
1530 let kf = krow[di] * invk;
1531 let spr = &s_prev[di * dv..(di + 1) * dv];
1532 let dsr = &mut ds[di * dv..(di + 1) * dv];
1533 let mut acc = F::ZERO;
1534 for dj in 0..dv {
1535 let dspre = dsr[dj] + kf * dkv[dj];
1536 acc += (spr[dj] * g) * dkv[dj];
1537 dg += dspre * spr[dj];
1538 dsr[dj] = g * dspre;
1539 }
1540 dkh[di] += acc;
1541 }
1542
1543 let sig = sigmoid_f(sp_arg[ti]);
1545 da[ti * nv + h] += dg * g * (-ea) * sig;
1546 db[ti * nv + h] += dbeta * beta * (F::ONE - beta);
1547
1548 let mut qdot = F::ZERO;
1552 let mut kdot = F::ZERO;
1553 for di in 0..dk {
1554 qdot += dqh[di] * qrow[di];
1555 kdot += dkh[di] * krow[di];
1556 }
1557 for di in 0..dk {
1558 dcq[ti * c_dim + ko * dk + di] +=
1559 invq * dqh[di] - qrow[di] * qdot * invq / nq2;
1560 dcq[ti * c_dim + kd + ko * dk + di] +=
1561 invk * dkh[di] - krow[di] * kdot * invk / nk2;
1562 }
1563 }
1564 }
1565}
1566
1567pub fn gdn_seq_fwd<F: Fp>(
1571 qkv: &[F],
1572 z: &[F],
1573 a: &[F],
1574 b: &[F],
1575 t: usize,
1576 cfg: &GdnSeqCfg,
1577 out: &mut [F],
1578) {
1579 let c_dim = cfg.c_dim();
1580 let mut pre = vec![F::ZERO; t * c_dim];
1581 let mut cq = vec![F::ZERO; t * c_dim];
1582 gdn_conv_fwd(qkv, t, c_dim, cfg.kk, cfg.conv, &mut pre, &mut cq);
1583 for ko in 0..cfg.nk {
1584 gdn_group_fwd(&cq, z, a, b, t, cfg, ko, out);
1585 }
1586}
1587
1588#[allow(clippy::too_many_arguments)]
1591pub fn gdn_seq_bwd<F: Fp>(
1592 qkv: &[F],
1593 z: &[F],
1594 a: &[F],
1595 b: &[F],
1596 t: usize,
1597 cfg: &GdnSeqCfg,
1598 dout: &[F],
1599 dqkv: &mut [F],
1600 dz: &mut [F],
1601 da: &mut [F],
1602 db: &mut [F],
1603) {
1604 let c_dim = cfg.c_dim();
1605 let mut pre = vec![F::ZERO; t * c_dim];
1606 let mut cq = vec![F::ZERO; t * c_dim];
1607 gdn_conv_fwd(qkv, t, c_dim, cfg.kk, cfg.conv, &mut pre, &mut cq);
1608 let mut dcq = vec![F::ZERO; t * c_dim];
1609 for ko in 0..cfg.nk {
1610 gdn_group_bwd(&cq, z, a, b, t, cfg, ko, dout, &mut dcq, dz, da, db);
1611 }
1612 gdn_conv_bwd(&pre, t, c_dim, cfg.kk, cfg.conv, &dcq, dqkv);
1613}