use crate::pool::Pool;
pub trait Fp:
Copy
+ PartialOrd
+ core::ops::Add<Output = Self>
+ core::ops::Sub<Output = Self>
+ core::ops::Mul<Output = Self>
+ core::ops::Div<Output = Self>
+ core::ops::Neg<Output = Self>
+ core::ops::AddAssign
+ core::ops::MulAssign
+ Send
+ Sync
+ 'static
{
const ZERO: Self;
const ONE: Self;
fn exp(self) -> Self;
fn sqrt(self) -> Self;
fn maxf(self, o: Self) -> Self;
fn fromf(x: f64) -> Self;
fn f64(self) -> f64;
}
impl Fp for f32 {
const ZERO: Self = 0.0;
const ONE: Self = 1.0;
#[inline]
fn exp(self) -> Self {
f32::exp(self)
}
#[inline]
fn sqrt(self) -> Self {
f32::sqrt(self)
}
#[inline]
fn maxf(self, o: Self) -> Self {
f32::max(self, o)
}
#[inline]
fn fromf(x: f64) -> Self {
x as f32
}
#[inline]
fn f64(self) -> f64 {
self as f64
}
}
impl Fp for f64 {
const ZERO: Self = 0.0;
const ONE: Self = 1.0;
#[inline]
fn exp(self) -> Self {
f64::exp(self)
}
#[inline]
fn sqrt(self) -> Self {
f64::sqrt(self)
}
#[inline]
fn maxf(self, o: Self) -> Self {
f64::max(self, o)
}
#[inline]
fn fromf(x: f64) -> Self {
x
}
#[inline]
fn f64(self) -> f64 {
self
}
}
#[inline]
fn dot<F: Fp>(a: &[F], b: &[F]) -> F {
let mut s = F::ZERO;
for (x, y) in a.iter().zip(b) {
s += *x * *y;
}
s
}
pub fn matmul_nt<F: Fp>(x: &[F], w: &[F], y: &mut [F], n: usize, k: usize, m: usize) {
for i in 0..n {
let xr = &x[i * k..(i + 1) * k];
for o in 0..m {
y[i * m + o] = dot(xr, &w[o * k..(o + 1) * k]);
}
}
}
pub fn matmul_nt_dx<F: Fp>(dy: &[F], w: &[F], dx: &mut [F], n: usize, k: usize, m: usize) {
for i in 0..n {
let dxr = &mut dx[i * k..(i + 1) * k];
for o in 0..m {
let g = dy[i * m + o];
for (d, wv) in dxr.iter_mut().zip(&w[o * k..(o + 1) * k]) {
*d += g * *wv;
}
}
}
}
pub fn matmul_nt_dw<F: Fp>(dy: &[F], x: &[F], dw: &mut [F], n: usize, k: usize, m: usize) {
for i in 0..n {
let xr = &x[i * k..(i + 1) * k];
for o in 0..m {
let g = dy[i * m + o];
for (d, xv) in dw[o * k..(o + 1) * k].iter_mut().zip(xr) {
*d += g * *xv;
}
}
}
}
const GEMM_BLOCK: usize = 128;
struct SendMut<T>(*mut T);
unsafe impl<T> Send for SendMut<T> {}
unsafe impl<T> Sync for SendMut<T> {}
impl<T> SendMut<T> {
#[inline]
#[allow(clippy::mut_from_ref)]
unsafe fn slice(&self, off: usize, len: usize) -> &mut [T] {
unsafe { std::slice::from_raw_parts_mut(self.0.add(off), len) }
}
}
pub fn gemm_nt(x: &[f32], w: &[f32], y: &mut [f32], n: usize, k: usize, m: usize, pool: Option<&Pool>) {
debug_assert_eq!(x.len(), n * k);
debug_assert_eq!(w.len(), m * k);
debug_assert_eq!(y.len(), n * m);
let nb = n.div_ceil(GEMM_BLOCK);
let block = |r0: usize, r1: usize, y: &mut [f32]| {
for o in 0..m {
let wr = &w[o * k..(o + 1) * k];
for i in r0..r1 {
y[(i - r0) * m + o] = crate::attention::dot_f32(&x[i * k..(i + 1) * k], wr);
}
}
};
match pool {
Some(p) if nb > 1 => {
let yp = SendMut(y.as_mut_ptr());
p.run(&|widx, nw| {
for bi in (widx..nb).step_by(nw) {
let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
let ys = unsafe { yp.slice(r0 * m, (r1 - r0) * m) };
block(r0, r1, ys);
}
});
}
_ => {
for bi in 0..nb {
let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
block(r0, r1, &mut y[r0 * m..r1 * m]);
}
}
}
}
pub fn gemm_dx(dy: &[f32], w: &[f32], dx: &mut [f32], n: usize, k: usize, m: usize, pool: Option<&Pool>) {
debug_assert_eq!(dy.len(), n * m);
debug_assert_eq!(w.len(), m * k);
debug_assert_eq!(dx.len(), n * k);
let nb = n.div_ceil(GEMM_BLOCK);
let block = |r0: usize, r1: usize, dxs: &mut [f32]| {
for o in 0..m {
let wr = &w[o * k..(o + 1) * k];
for i in r0..r1 {
let g = dy[i * m + o];
if g != 0.0 {
crate::attention::axpy_f32(&mut dxs[(i - r0) * k..(i - r0 + 1) * k], wr, g);
}
}
}
};
match pool {
Some(p) if nb > 1 => {
let dxp = SendMut(dx.as_mut_ptr());
p.run(&|widx, nw| {
for bi in (widx..nb).step_by(nw) {
let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
let dxs = unsafe { dxp.slice(r0 * k, (r1 - r0) * k) };
block(r0, r1, dxs);
}
});
}
_ => {
for bi in 0..nb {
let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
block(r0, r1, &mut dx[r0 * k..r1 * k]);
}
}
}
}
pub fn gemm_dw(dy: &[f32], x: &[f32], dw: &mut [f32], n: usize, k: usize, m: usize, pool: Option<&Pool>) {
debug_assert_eq!(dy.len(), n * m);
debug_assert_eq!(x.len(), n * k);
debug_assert_eq!(dw.len(), m * k);
let range = |o0: usize, o1: usize, dws: &mut [f32]| {
let nb = n.div_ceil(GEMM_BLOCK);
for bi in 0..nb {
let (r0, r1) = (bi * GEMM_BLOCK, ((bi + 1) * GEMM_BLOCK).min(n));
for o in o0..o1 {
let dwr = &mut dws[(o - o0) * k..(o - o0 + 1) * k];
for i in r0..r1 {
let g = dy[i * m + o];
if g != 0.0 {
crate::attention::axpy_f32(dwr, &x[i * k..(i + 1) * k], g);
}
}
}
}
};
match pool {
Some(p) if m >= 8 => {
let dwp = SendMut(dw.as_mut_ptr());
p.run(&|widx, nw| {
let (o0, o1) = (widx * m / nw, (widx + 1) * m / nw);
if o0 < o1 {
let dws = unsafe { dwp.slice(o0 * k, (o1 - o0) * k) };
range(o0, o1, dws);
}
});
}
_ => range(0, m, dw),
}
}
#[inline]
pub fn silu<F: Fp>(x: F) -> F {
x / (F::ONE + (-x).exp())
}
#[inline]
pub fn silu_bwd<F: Fp>(x: F) -> F {
let s = F::ONE / (F::ONE + (-x).exp());
s * (F::ONE + x * (F::ONE - s))
}
pub fn rmsnorm_fwd<F: Fp>(x: &[F], w: &[F], eps: f64, gemma: bool, y: &mut [F], inv_out: &mut [F]) {
let d = w.len();
let n = x.len() / d;
for r in 0..n {
let xr = &x[r * d..(r + 1) * d];
let mut ss = 0f64;
for v in xr {
ss += v.f64() * v.f64();
}
let inv = F::fromf(1.0 / (ss / d as f64 + eps).sqrt());
inv_out[r] = inv;
let yr = &mut y[r * d..(r + 1) * d];
for j in 0..d {
let weff = if gemma { F::ONE + w[j] } else { w[j] };
yr[j] = xr[j] * inv * weff;
}
}
}
pub fn rmsnorm_bwd<F: Fp>(
x: &[F],
w: &[F],
inv: &[F],
dy: &[F],
gemma: bool,
dx: &mut [F],
mut dw: Option<&mut [F]>,
) {
let d = w.len();
let n = x.len() / d;
for r in 0..n {
let xr = &x[r * d..(r + 1) * d];
let dyr = &dy[r * d..(r + 1) * d];
let iv = inv[r];
let mut s = 0f64;
for j in 0..d {
let weff = if gemma { F::ONE + w[j] } else { w[j] };
s += (dyr[j] * weff * xr[j]).f64();
}
let coef = F::fromf(s / d as f64) * iv * iv * iv;
let dxr = &mut dx[r * d..(r + 1) * d];
for j in 0..d {
let weff = if gemma { F::ONE + w[j] } else { w[j] };
dxr[j] += iv * weff * dyr[j] - xr[j] * coef;
}
if let Some(dwv) = dw.as_deref_mut() {
for j in 0..d {
dwv[j] += dyr[j] * xr[j] * iv;
}
}
}
}
pub fn rope_fwd<F: Fp>(x: &mut [F], position: usize, inv_freq: &[f64]) {
let half = inv_freq.len();
for (i, &freq) in inv_freq.iter().enumerate() {
let angle = position as f64 * freq;
let (sin, cos) = (F::fromf(angle.sin()), F::fromf(angle.cos()));
let x0 = x[i];
let x1 = x[i + half];
x[i] = x0 * cos - x1 * sin;
x[i + half] = x0 * sin + x1 * cos;
}
}
pub fn rope_bwd<F: Fp>(dy: &mut [F], position: usize, inv_freq: &[f64]) {
let half = inv_freq.len();
for (i, &freq) in inv_freq.iter().enumerate() {
let angle = position as f64 * freq;
let (sin, cos) = (F::fromf(angle.sin()), F::fromf(angle.cos()));
let g0 = dy[i];
let g1 = dy[i + half];
dy[i] = g0 * cos + g1 * sin;
dy[i + half] = -g0 * sin + g1 * cos;
}
}
pub fn seg_means<F: Fp>(x: &[F], t: usize, d: usize, m: usize, out: &mut [F]) {
for i in 0..m {
let (lo, hi) = (i * t / m, (i + 1) * t / m);
let or = &mut out[i * d..(i + 1) * d];
for v in or.iter_mut() {
*v = F::ZERO;
}
for j in lo..hi {
for c in 0..d {
or[c] += x[j * d + c];
}
}
let inv = F::fromf(1.0 / (hi - lo) as f64);
for v in or.iter_mut() {
*v *= inv;
}
}
}
pub fn seg_means_bwd<F: Fp>(dl: &[F], t: usize, d: usize, m: usize, dx: &mut [F]) {
for i in 0..m {
let (lo, hi) = (i * t / m, (i + 1) * t / m);
let inv = F::fromf(1.0 / (hi - lo) as f64);
let dlr = &dl[i * d..(i + 1) * d];
for j in lo..hi {
for c in 0..d {
dx[j * d + c] += dlr[c] * inv;
}
}
}
}
#[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]) {
let scale = F::fromf(1.0 / (d as f64).sqrt());
let mut row = vec![F::ZERO; t];
for ti in 0..t {
let qr = &q[ti * d..(ti + 1) * d];
let mut mx = F::fromf(f64::NEG_INFINITY);
for j in 0..=ti {
let s = dot(qr, &k[j * d..(j + 1) * d]) * scale;
row[j] = s;
mx = mx.maxf(s);
}
let mut den = F::ZERO;
for j in 0..=ti {
row[j] = (row[j] - mx).exp();
den += row[j];
}
let or = &mut out[ti * dv..(ti + 1) * dv];
for o in or.iter_mut() {
*o = F::ZERO;
}
for j in 0..=ti {
let p = row[j] / den;
for (o, vv) in or.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
*o += p * *vv;
}
}
}
}
#[allow(clippy::too_many_arguments, clippy::needless_range_loop)]
pub fn attn_head_bwd<F: Fp>(
q: &[F],
k: &[F],
v: &[F],
dout: &[F],
t: usize,
d: usize,
dv: usize,
dq: &mut [F],
dk: &mut [F],
dvv: &mut [F],
) {
let scale = F::fromf(1.0 / (d as f64).sqrt());
let mut row = vec![F::ZERO; t];
for ti in 0..t {
let qr = &q[ti * d..(ti + 1) * d];
let mut mx = F::fromf(f64::NEG_INFINITY);
for j in 0..=ti {
let s = dot(qr, &k[j * d..(j + 1) * d]) * scale;
row[j] = s;
mx = mx.maxf(s);
}
let mut den = F::ZERO;
for j in 0..=ti {
row[j] = (row[j] - mx).exp();
den += row[j];
}
let dor = &dout[ti * dv..(ti + 1) * dv];
let mut pdp = F::ZERO;
let mut dp = vec![F::ZERO; ti + 1];
for j in 0..=ti {
let p = row[j] / den;
row[j] = p; dp[j] = dot(dor, &v[j * dv..(j + 1) * dv]);
pdp += p * dp[j];
}
let dqr = &mut dq[ti * d..(ti + 1) * d];
for j in 0..=ti {
let p = row[j];
for (dvo, o) in dvv[j * dv..(j + 1) * dv].iter_mut().zip(dor) {
*dvo += p * *o;
}
let ds = p * (dp[j] - pdp) * scale;
let kr = &k[j * d..(j + 1) * d];
for c in 0..d {
dqr[c] += ds * kr[c];
}
let dkr = &mut dk[j * d..(j + 1) * d];
for c in 0..d {
dkr[c] += ds * qr[c];
}
}
}
}
#[derive(Clone, Copy, Debug)]
pub struct NysCfg {
pub m: usize,
pub w: usize,
pub sink: usize,
pub prefill: Option<usize>,
}
impl NysCfg {
#[inline]
pub fn prefill_len(&self, t: usize) -> usize {
self.prefill.unwrap_or(t / 2).clamp(1, t)
}
}
const NYS_DEN_EPS: f64 = 1e-30;
struct NysGraph<F: Fp> {
m_eff: usize,
tp: usize,
q_l: Vec<F>,
k_l: Vec<F>,
mu: Vec<F>,
fu: Vec<F>,
e: Vec<F>,
fumu: Vec<F>,
far_keep: Vec<bool>,
wmat: Vec<F>,
c_row: Vec<F>,
den: Vec<F>,
}
#[inline]
fn nys_exact_row(ti: usize, tp: usize) -> bool {
ti < tp
}
#[inline]
fn nys_near(ti: usize, j: usize, w: usize, sink: usize) -> bool {
ti - j < w || j < sink
}
fn nys_graph<F: Fp>(
q: &[F],
k: &[F],
t: usize,
d: usize,
cfg: &NysCfg,
mu_override: Option<&[F]>,
) -> NysGraph<F> {
let scale = 1.0 / (d as f64).sqrt();
let fscale = F::fromf(scale);
let tp = cfg.prefill_len(t);
let m_eff = (tp / 8).clamp(4, cfg.m);
let mut q_l = vec![F::ZERO; m_eff * d];
let mut k_l = vec![F::ZERO; m_eff * d];
seg_means(&q[..tp * d], tp, d, m_eff, &mut q_l);
seg_means(&k[..tp * d], tp, d, m_eff, &mut k_l);
let mut au = vec![0f64; m_eff * m_eff];
for i in 0..m_eff {
for j in 0..m_eff {
let mut s = 0f64;
for c in 0..d {
s += q_l[i * d + c].f64() * k_l[j * d + c].f64();
}
au[i * m_eff + j] = (s * scale).exp();
}
}
let mu: Vec<F> = match mu_override {
Some(m) => m.to_vec(),
None => crate::nystrom::ridge_pinv(&au, m_eff)
.iter()
.map(|&x| F::fromf(x))
.collect(),
};
let mut fu = vec![F::ZERO; t * m_eff];
for ti in 0..t {
for i in 0..m_eff {
fu[ti * m_eff + i] =
(dot(&q[ti * d..(ti + 1) * d], &k_l[i * d..(i + 1) * d]) * fscale).exp();
}
}
let mut e = vec![F::ZERO; m_eff * t];
for i in 0..m_eff {
for j in 0..t {
e[i * t + j] =
(dot(&q_l[i * d..(i + 1) * d], &k[j * d..(j + 1) * d]) * fscale).exp();
}
}
let mut fumu = vec![F::ZERO; t * m_eff];
matmul_nt(
&fu,
&transpose(&mu, m_eff, m_eff),
&mut fumu,
t,
m_eff,
m_eff,
);
let mut a = vec![F::ZERO; t * t];
for ti in tp..t {
let fr = &fumu[ti * m_eff..(ti + 1) * m_eff];
let ar = &mut a[ti * t..ti * t + ti + 1]; for (i, &f) in fr.iter().enumerate() {
let er = &e[i * t..i * t + ti + 1];
for (av, ev) in ar.iter_mut().zip(er) {
*av += f * *ev;
}
}
}
let mut wmat = vec![F::ZERO; t * t];
let mut c_row = vec![F::ZERO; t];
let mut den = vec![F::ZERO; t];
let mut far_keep = vec![false; t];
let mut lg_row = vec![F::ZERO; t];
for ti in 0..t {
let qr = &q[ti * d..(ti + 1) * d];
let mut c = F::fromf(f64::NEG_INFINITY);
for j in 0..=ti {
let s = dot(qr, &k[j * d..(j + 1) * d]) * fscale;
lg_row[j] = s;
c = c.maxf(s);
}
c_row[ti] = c;
let emc = (-c).exp();
let exact_row = nys_exact_row(ti, tp);
let mut far_sum = F::ZERO;
if !exact_row {
for j in 0..=ti {
if !nys_near(ti, j, cfg.w, cfg.sink) {
far_sum += a[ti * t + j];
}
}
}
let keep = !exact_row && far_sum.f64() >= 0.0;
far_keep[ti] = keep;
let wr = &mut wmat[ti * t..(ti + 1) * t];
let mut dsum = F::ZERO;
for j in 0..=ti {
let wv = if exact_row || nys_near(ti, j, cfg.w, cfg.sink) {
(lg_row[j] - c).exp()
} else if keep {
a[ti * t + j] * emc
} else {
F::ZERO
};
wr[j] = wv;
dsum += wv;
}
den[ti] = dsum.maxf(F::fromf(NYS_DEN_EPS));
}
NysGraph { m_eff, tp, q_l, k_l, mu, fu, e, fumu, far_keep, wmat, c_row, den }
}
fn transpose<F: Fp>(x: &[F], rows: usize, cols: usize) -> Vec<F> {
let mut out = vec![F::ZERO; rows * cols];
for r in 0..rows {
for c in 0..cols {
out[c * rows + r] = x[r * cols + c];
}
}
out
}
#[allow(clippy::too_many_arguments)]
pub fn nystrom_head_fwd<F: Fp>(
q: &[F],
k: &[F],
v: &[F],
t: usize,
d: usize,
dv: usize,
cfg: &NysCfg,
out: &mut [F],
) {
if nys_degenerate(t, cfg) {
attn_head_fwd(q, k, v, t, d, dv, out);
return;
}
nystrom_head_fwd_mu(q, k, v, t, d, dv, cfg, None, out);
}
#[inline]
fn nys_degenerate(t: usize, cfg: &NysCfg) -> bool {
cfg.prefill_len(t) <= cfg.w + cfg.sink + 8
}
#[doc(hidden)]
pub fn nystrom_mu_for_test<F: Fp>(q: &[F], k: &[F], t: usize, d: usize, cfg: &NysCfg) -> Vec<F> {
nys_graph(q, k, t, d, cfg, None).mu
}
#[doc(hidden)]
#[allow(clippy::too_many_arguments)]
pub fn nystrom_head_fwd_mu<F: Fp>(
q: &[F],
k: &[F],
v: &[F],
t: usize,
d: usize,
dv: usize,
cfg: &NysCfg,
mu_override: Option<&[F]>,
out: &mut [F],
) {
let g = nys_graph(q, k, t, d, cfg, mu_override);
for ti in 0..t {
let wr = &g.wmat[ti * t..(ti + 1) * t];
let den = g.den[ti];
let or = &mut out[ti * dv..(ti + 1) * dv];
for o in or.iter_mut() {
*o = F::ZERO;
}
for j in 0..=ti {
let p = wr[j] / den;
if p.f64() != 0.0 {
for (o, vv) in or.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
*o += p * *vv;
}
}
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn nystrom_head_bwd<F: Fp>(
q: &[F],
k: &[F],
v: &[F],
dout: &[F],
t: usize,
d: usize,
dv: usize,
cfg: &NysCfg,
dq: &mut [F],
dk: &mut [F],
dvv: &mut [F],
) {
if nys_degenerate(t, cfg) {
attn_head_bwd(q, k, v, dout, t, d, dv, dq, dk, dvv);
return;
}
nystrom_head_bwd_mu(q, k, v, dout, t, d, dv, cfg, None, dq, dk, dvv);
}
#[doc(hidden)]
#[allow(clippy::too_many_arguments)]
pub fn nystrom_head_bwd_mu<F: Fp>(
q: &[F],
k: &[F],
v: &[F],
dout: &[F],
t: usize,
d: usize,
dv: usize,
cfg: &NysCfg,
mu_override: Option<&[F]>,
dq: &mut [F],
dk: &mut [F],
dvv: &mut [F],
) {
let scale = F::fromf(1.0 / (d as f64).sqrt());
let g = nys_graph(q, k, t, d, cfg, mu_override);
let m_eff = g.m_eff;
let mut dwmat = vec![F::ZERO; t * t];
let mut out_row = vec![F::ZERO; dv];
for ti in 0..t {
let wr = &g.wmat[ti * t..(ti + 1) * t];
let den = g.den[ti];
for o in out_row.iter_mut() {
*o = F::ZERO;
}
for j in 0..=ti {
let p = wr[j] / den;
if p.f64() != 0.0 {
for (o, vv) in out_row.iter_mut().zip(&v[j * dv..(j + 1) * dv]) {
*o += p * *vv;
}
}
}
let dor = &dout[ti * dv..(ti + 1) * dv];
let dwr = &mut dwmat[ti * t..(ti + 1) * t];
for j in 0..=ti {
let p = wr[j] / den;
if p.f64() != 0.0 {
for (dvo, o) in dvv[j * dv..(j + 1) * dv].iter_mut().zip(dor) {
*dvo += p * *o;
}
}
let mut s = F::ZERO;
for c in 0..dv {
s += dor[c] * (v[j * dv + c] - out_row[c]);
}
dwr[j] = s / den;
}
}
for ti in 0..t {
let qr = &q[ti * d..(ti + 1) * d];
let dqr_base = ti * d;
for j in 0..=ti {
if !(nys_exact_row(ti, g.tp) || nys_near(ti, j, cfg.w, cfg.sink)) {
continue;
}
let dlg = dwmat[ti * t + j] * g.wmat[ti * t + j] * scale;
if dlg.f64() == 0.0 {
continue;
}
let kr = &k[j * d..(j + 1) * d];
for c in 0..d {
dq[dqr_base + c] += dlg * kr[c];
}
let dkr = &mut dk[j * d..(j + 1) * d];
for c in 0..d {
dkr[c] += dlg * qr[c];
}
}
}
let mut da = vec![F::ZERO; t * t];
for ti in g.tp..t {
if !g.far_keep[ti] {
continue; }
let emc = (-g.c_row[ti]).exp();
for j in 0..=ti {
if nys_near(ti, j, cfg.w, cfg.sink) {
continue;
}
da[ti * t + j] = dwmat[ti * t + j] * emc;
}
}
let mut dfumu = vec![F::ZERO; t * m_eff];
for ti in 0..t {
let dar = &da[ti * t..ti * t + ti + 1];
let dfr = &mut dfumu[ti * m_eff..(ti + 1) * m_eff];
for (i, df) in dfr.iter_mut().enumerate() {
let er = &g.e[i * t..i * t + ti + 1];
let mut s = F::ZERO;
for (av, ev) in dar.iter().zip(er) {
s += *av * *ev;
}
*df = s;
}
}
let mut dfu = vec![F::ZERO; t * m_eff];
matmul_nt(&dfumu, &g.mu, &mut dfu, t, m_eff, m_eff);
let mut de = vec![F::ZERO; m_eff * t];
for ti in 0..t {
let dar = &da[ti * t..ti * t + ti + 1];
let fr = &g.fumu[ti * m_eff..(ti + 1) * m_eff];
for (i, &f) in fr.iter().enumerate() {
if f.f64() == 0.0 {
continue;
}
let der = &mut de[i * t..i * t + ti + 1];
for (dev, av) in der.iter_mut().zip(dar) {
*dev += f * *av;
}
}
}
let mut dq_l = vec![F::ZERO; m_eff * d];
let mut dk_l = vec![F::ZERO; m_eff * d];
for ti in 0..t {
let qr = &q[ti * d..(ti + 1) * d];
for i in 0..m_eff {
let dlg = dfu[ti * m_eff + i] * g.fu[ti * m_eff + i] * scale;
if dlg.f64() == 0.0 {
continue;
}
let klr = &g.k_l[i * d..(i + 1) * d];
for c in 0..d {
dq[ti * d + c] += dlg * klr[c];
}
let dklr = &mut dk_l[i * d..(i + 1) * d];
for c in 0..d {
dklr[c] += dlg * qr[c];
}
}
}
for i in 0..m_eff {
let qlr = &g.q_l[i * d..(i + 1) * d];
for j in 0..t {
let dlg = de[i * t + j] * g.e[i * t + j] * scale;
if dlg.f64() == 0.0 {
continue;
}
let kr = &k[j * d..(j + 1) * d];
let dqlr = &mut dq_l[i * d..(i + 1) * d];
for c in 0..d {
dqlr[c] += dlg * kr[c];
}
for c in 0..d {
dk[j * d + c] += dlg * qlr[c];
}
}
}
let tp = g.tp;
seg_means_bwd(&dq_l, tp, d, m_eff, &mut dq[..tp * d]);
seg_means_bwd(&dk_l, tp, d, m_eff, &mut dk[..tp * d]);
}
pub fn ce_kl_position<F: Fp>(
s_logits: &[F],
t_logits: &[F],
target: usize,
kl_w: f64,
inv_n: f64,
dlogits: &mut [F],
) -> (f64, f64) {
let vsz = s_logits.len();
debug_assert_eq!(t_logits.len(), vsz);
let mut smax = f64::NEG_INFINITY;
let mut tmax = f64::NEG_INFINITY;
for i in 0..vsz {
smax = smax.max(s_logits[i].f64());
tmax = tmax.max(t_logits[i].f64());
}
let mut ssum = 0f64;
let mut tsum = 0f64;
for i in 0..vsz {
ssum += (s_logits[i].f64() - smax).exp();
tsum += (t_logits[i].f64() - tmax).exp();
}
let slz = smax + ssum.ln();
let tlz = tmax + tsum.ln();
let ce = slz - s_logits[target].f64();
let mut kl = 0f64;
for i in 0..vsz {
let ls = s_logits[i].f64() - slz;
let lt = t_logits[i].f64() - tlz;
let pt = lt.exp();
let ps = ls.exp();
if pt > 0.0 {
kl += pt * (lt - ls);
}
let mut gd = (1.0 - kl_w) * ps + kl_w * (ps - pt);
if i == target {
gd -= 1.0 - kl_w;
}
dlogits[i] = F::fromf(gd * inv_n);
}
(ce, kl)
}
pub struct GdnSeqCfg<'a> {
pub nv: usize,
pub nk: usize,
pub dk: usize,
pub dv: usize,
pub kk: usize,
pub rms_eps: f64,
pub conv: &'a [f32],
pub a_log: &'a [f32],
pub dt_bias: &'a [f32],
pub norm: &'a [f32],
}
impl GdnSeqCfg<'_> {
pub fn c_dim(&self) -> usize {
2 * self.nk * self.dk + self.nv * self.dv
}
}
#[inline]
fn softplus_f<F: Fp>(x: F) -> F {
if x.f64() > 20.0 {
x
} else {
F::fromf(x.f64().exp().ln_1p())
}
}
#[inline]
fn sigmoid_f<F: Fp>(x: F) -> F {
F::ONE / (F::ONE + (-x).exp())
}
pub fn gdn_conv_fwd<F: Fp>(
raw: &[F],
t: usize,
c_dim: usize,
kk: usize,
conv: &[f32],
pre: &mut [F],
cq: &mut [F],
) {
for ti in 0..t {
for c in 0..c_dim {
let taps = &conv[c * kk..(c + 1) * kk];
let mut acc = F::ZERO;
for (j, &tap) in taps.iter().enumerate() {
let p = ti as isize - (kk as isize - 1) + j as isize;
if p >= 0 {
acc += raw[p as usize * c_dim + c] * F::fromf(tap as f64);
}
}
pre[ti * c_dim + c] = acc;
cq[ti * c_dim + c] = silu(acc);
}
}
}
pub fn gdn_conv_bwd<F: Fp>(
pre: &[F],
t: usize,
c_dim: usize,
kk: usize,
conv: &[f32],
dcq: &[F],
draw: &mut [F],
) {
for ti in 0..t {
for c in 0..c_dim {
let g = dcq[ti * c_dim + c];
if g.f64() == 0.0 {
continue;
}
let dp = g * silu_bwd(pre[ti * c_dim + c]);
let taps = &conv[c * kk..(c + 1) * kk];
for (j, &tap) in taps.iter().enumerate() {
let p = ti as isize - (kk as isize - 1) + j as isize;
if p >= 0 {
draw[p as usize * c_dim + c] += dp * F::fromf(tap as f64);
}
}
}
}
}
#[inline]
fn gdn_inv<F: Fp>(x: &[F], extra_scale: f64) -> (F, F) {
let mut n2 = F::ZERO;
for v in x {
n2 += *v * *v;
}
let n2e = n2 + F::fromf(1e-6);
let inv = F::ONE / (n2e.sqrt() * F::fromf(extra_scale));
(inv, n2e)
}
#[allow(clippy::too_many_arguments)]
pub fn gdn_group_fwd<F: Fp>(
cq: &[F],
z: &[F],
a: &[F],
b: &[F],
t: usize,
cfg: &GdnSeqCfg,
ko: usize,
out: &mut [F],
) {
let (nv, nk, dk, dv) = (cfg.nv, cfg.nk, cfg.dk, cfg.dv);
let c_dim = cfg.c_dim();
let kd = nk * dk;
let rep = nv / nk;
let vd = nv * dv;
let sqdk = (dk as f64).sqrt();
for hh in 0..rep {
let h = ko * rep + hh;
let ea = F::fromf((cfg.a_log[h] as f64).exp());
let mut s = vec![F::ZERO; dk * dv];
let mut kv = vec![F::ZERO; dv];
let mut o = vec![F::ZERO; dv];
for ti in 0..t {
let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
let (invq, _) = gdn_inv(qrow, sqdk);
let (invk, _) = gdn_inv(krow, 1.0);
let g = (-ea * softplus_f(a[ti * nv + h] + F::fromf(cfg.dt_bias[h] as f64))).exp();
let beta = sigmoid_f(b[ti * nv + h]);
for x in kv.iter_mut() {
*x = F::ZERO;
}
for di in 0..dk {
let kf = krow[di] * invk;
let row = &mut s[di * dv..(di + 1) * dv];
for dj in 0..dv {
row[dj] *= g;
kv[dj] += row[dj] * kf;
}
}
for x in o.iter_mut() {
*x = F::ZERO;
}
for di in 0..dk {
let kf = krow[di] * invk;
let qf = qrow[di] * invq;
let row = &mut s[di * dv..(di + 1) * dv];
for dj in 0..dv {
row[dj] += kf * (vrow[dj] - kv[dj]) * beta;
o[dj] += qf * row[dj];
}
}
let mut ss = 0f64;
for v in &o {
ss += v.f64() * v.f64();
}
let inv = F::fromf(1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt());
for dj in 0..dv {
let zv = z[ti * vd + h * dv + dj];
out[ti * vd + h * dv + dj] =
o[dj] * inv * F::fromf(cfg.norm[dj] as f64) * silu(zv);
}
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn gdn_group_bwd<F: Fp>(
cq: &[F],
z: &[F],
a: &[F],
b: &[F],
t: usize,
cfg: &GdnSeqCfg,
ko: usize,
dout: &[F],
dcq: &mut [F],
dz: &mut [F],
da: &mut [F],
db: &mut [F],
) {
let (nv, nk, dk, dv) = (cfg.nv, cfg.nk, cfg.dk, cfg.dv);
let c_dim = cfg.c_dim();
let kd = nk * dk;
let rep = nv / nk;
let vd = nv * dv;
let sqdk = (dk as f64).sqrt();
for hh in 0..rep {
let h = ko * rep + hh;
let ea = F::fromf((cfg.a_log[h] as f64).exp());
let mut s_hist = vec![F::ZERO; (t + 1) * dk * dv]; let mut kv_hist = vec![F::ZERO; t * dv];
let mut o_hist = vec![F::ZERO; t * dv];
let mut g_v = vec![F::ZERO; t];
let mut beta_v = vec![F::ZERO; t];
let mut sp_arg = vec![F::ZERO; t]; for ti in 0..t {
let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
let (invq, _) = gdn_inv(qrow, sqdk);
let (invk, _) = gdn_inv(krow, 1.0);
let arg = a[ti * nv + h] + F::fromf(cfg.dt_bias[h] as f64);
let g = (-ea * softplus_f(arg)).exp();
let beta = sigmoid_f(b[ti * nv + h]);
sp_arg[ti] = arg;
g_v[ti] = g;
beta_v[ti] = beta;
let (prev, cur) = s_hist.split_at_mut((ti + 1) * dk * dv);
let sp = &prev[ti * dk * dv..];
let sn = &mut cur[..dk * dv];
let kvr = &mut kv_hist[ti * dv..(ti + 1) * dv];
for di in 0..dk {
let kf = krow[di] * invk;
for dj in 0..dv {
let dec = sp[di * dv + dj] * g;
sn[di * dv + dj] = dec;
kvr[dj] += dec * kf;
}
}
let or = &mut o_hist[ti * dv..(ti + 1) * dv];
for di in 0..dk {
let kf = krow[di] * invk;
let qf = qrow[di] * invq;
for dj in 0..dv {
let sv = sn[di * dv + dj] + kf * (vrow[dj] - kvr[dj]) * beta_v[ti];
sn[di * dv + dj] = sv;
or[dj] += qf * sv;
}
}
}
let mut ds = vec![F::ZERO; dk * dv];
let mut do_o = vec![F::ZERO; dv];
let mut du = vec![F::ZERO; dv];
let mut dkv = vec![F::ZERO; dv];
let mut dqh = vec![F::ZERO; dk];
let mut dkh = vec![F::ZERO; dk];
for ti in (0..t).rev() {
let qrow = &cq[ti * c_dim + ko * dk..ti * c_dim + (ko + 1) * dk];
let krow = &cq[ti * c_dim + kd + ko * dk..ti * c_dim + kd + (ko + 1) * dk];
let vrow = &cq[ti * c_dim + 2 * kd + h * dv..ti * c_dim + 2 * kd + (h + 1) * dv];
let (invq, nq2) = gdn_inv(qrow, sqdk);
let (invk, nk2) = gdn_inv(krow, 1.0);
let g = g_v[ti];
let beta = beta_v[ti];
let s_t = &s_hist[(ti + 1) * dk * dv..(ti + 2) * dk * dv];
let s_prev = &s_hist[ti * dk * dv..(ti + 1) * dk * dv];
let kvr = &kv_hist[ti * dv..(ti + 1) * dv];
let or = &o_hist[ti * dv..(ti + 1) * dv];
let mut ss = 0f64;
for v in or {
ss += v.f64() * v.f64();
}
let inv = F::fromf(1.0 / (ss / dv as f64 + cfg.rms_eps).sqrt());
let dofr = &dout[ti * vd + h * dv..ti * vd + (h + 1) * dv];
let mut sdot = 0f64; for dj in 0..dv {
let zv = z[ti * vd + h * dv + dj];
let w = F::fromf(cfg.norm[dj] as f64);
let weff = w * silu(zv);
sdot += (dofr[dj] * weff * or[dj]).f64();
dz[ti * vd + h * dv + dj] += dofr[dj] * or[dj] * inv * w * silu_bwd(zv);
}
let coef = F::fromf(sdot / dv as f64) * inv * inv * inv;
for dj in 0..dv {
let zv = z[ti * vd + h * dv + dj];
let weff = F::fromf(cfg.norm[dj] as f64) * silu(zv);
do_o[dj] = inv * weff * dofr[dj] - or[dj] * coef;
}
for x in dqh.iter_mut() {
*x = F::ZERO;
}
for di in 0..dk {
let qf = qrow[di] * invq;
let row = &s_t[di * dv..(di + 1) * dv];
let dsr = &mut ds[di * dv..(di + 1) * dv];
let mut acc = F::ZERO;
for dj in 0..dv {
dsr[dj] += qf * do_o[dj];
acc += row[dj] * do_o[dj];
}
dqh[di] = acc;
}
for x in du.iter_mut() {
*x = F::ZERO;
}
for x in dkh.iter_mut() {
*x = F::ZERO;
}
for di in 0..dk {
let kf = krow[di] * invk;
let dsr = &ds[di * dv..(di + 1) * dv];
let mut acc = F::ZERO;
for dj in 0..dv {
du[dj] += dsr[dj] * kf;
acc += dsr[dj] * (vrow[dj] - kvr[dj]) * beta;
}
dkh[di] = acc;
}
let mut dbeta = F::ZERO;
for dj in 0..dv {
dbeta += du[dj] * (vrow[dj] - kvr[dj]);
dcq[ti * c_dim + 2 * kd + h * dv + dj] += beta * du[dj];
dkv[dj] = -(beta * du[dj]);
}
let mut dg = F::ZERO;
for di in 0..dk {
let kf = krow[di] * invk;
let spr = &s_prev[di * dv..(di + 1) * dv];
let dsr = &mut ds[di * dv..(di + 1) * dv];
let mut acc = F::ZERO;
for dj in 0..dv {
let dspre = dsr[dj] + kf * dkv[dj];
acc += (spr[dj] * g) * dkv[dj];
dg += dspre * spr[dj];
dsr[dj] = g * dspre;
}
dkh[di] += acc;
}
let sig = sigmoid_f(sp_arg[ti]);
da[ti * nv + h] += dg * g * (-ea) * sig;
db[ti * nv + h] += dbeta * beta * (F::ONE - beta);
let mut qdot = F::ZERO;
let mut kdot = F::ZERO;
for di in 0..dk {
qdot += dqh[di] * qrow[di];
kdot += dkh[di] * krow[di];
}
for di in 0..dk {
dcq[ti * c_dim + ko * dk + di] +=
invq * dqh[di] - qrow[di] * qdot * invq / nq2;
dcq[ti * c_dim + kd + ko * dk + di] +=
invk * dkh[di] - krow[di] * kdot * invk / nk2;
}
}
}
}
pub fn gdn_seq_fwd<F: Fp>(
qkv: &[F],
z: &[F],
a: &[F],
b: &[F],
t: usize,
cfg: &GdnSeqCfg,
out: &mut [F],
) {
let c_dim = cfg.c_dim();
let mut pre = vec![F::ZERO; t * c_dim];
let mut cq = vec![F::ZERO; t * c_dim];
gdn_conv_fwd(qkv, t, c_dim, cfg.kk, cfg.conv, &mut pre, &mut cq);
for ko in 0..cfg.nk {
gdn_group_fwd(&cq, z, a, b, t, cfg, ko, out);
}
}
#[allow(clippy::too_many_arguments)]
pub fn gdn_seq_bwd<F: Fp>(
qkv: &[F],
z: &[F],
a: &[F],
b: &[F],
t: usize,
cfg: &GdnSeqCfg,
dout: &[F],
dqkv: &mut [F],
dz: &mut [F],
da: &mut [F],
db: &mut [F],
) {
let c_dim = cfg.c_dim();
let mut pre = vec![F::ZERO; t * c_dim];
let mut cq = vec![F::ZERO; t * c_dim];
gdn_conv_fwd(qkv, t, c_dim, cfg.kk, cfg.conv, &mut pre, &mut cq);
let mut dcq = vec![F::ZERO; t * c_dim];
for ko in 0..cfg.nk {
gdn_group_bwd(&cq, z, a, b, t, cfg, ko, dout, &mut dcq, dz, da, db);
}
gdn_conv_bwd(&pre, t, c_dim, cfg.kk, cfg.conv, &dcq, dqkv);
}