use crate::gdn::DeltaDims;
pub const CHUNK: usize = 32;
pub const DEVICE_ROWS: usize = 32;
#[allow(clippy::too_many_arguments)]
pub fn delta_chunk_rows(
dims: DeltaDims,
rows: usize,
state: &mut crate::recurrent_state::AlignedF32,
q: &[f32],
k: &[f32],
v: &[f32],
g: &[f32],
beta: &[f32],
out: &mut [f32],
) {
#[cfg(feature = "metal")]
if rows >= DEVICE_ROWS && dims.head_dim <= ferrox_metal::gdn_chunk::MAX_HEAD_DIM {
let shape = ferrox_metal::gdn::DeltaShape {
n_k_heads: dims.n_k_heads,
n_v_heads: dims.n_v_heads,
head_dim: dims.head_dim,
map: match dims.map {
crate::gdn::HeadMap::Tiled => ferrox_metal::gdn::HeadMapKind::Tiled,
crate::gdn::HeadMap::Grouped => ferrox_metal::gdn::HeadMapKind::Grouped,
},
};
let (bytes, ptr) = (state.alloc_bytes(), state.as_ptr());
let done = unsafe {
ferrox_metal::gdn_chunk::launch_delta_chunk(shape, rows, ptr, bytes, q, k, v, g, beta)
};
if let Ok(o) = done {
out.copy_from_slice(&o);
return;
}
}
delta_chunk(dims, rows, state, q, k, v, g, beta, out);
}
#[allow(clippy::too_many_arguments)]
pub fn delta_chunk(
dims: DeltaDims,
rows: usize,
state: &mut [f32],
q: &[f32],
k: &[f32],
v: &[f32],
g: &[f32],
beta: &[f32],
out: &mut [f32],
) {
let DeltaDims {
n_k_heads,
n_v_heads,
head_dim: s,
map,
} = dims;
assert_eq!(state.len(), dims.state_len());
assert_eq!(q.len(), rows * n_k_heads * s);
assert_eq!(k.len(), rows * n_k_heads * s);
assert_eq!(v.len(), rows * n_v_heads * s);
assert_eq!(g.len(), rows * n_v_heads);
assert_eq!(beta.len(), rows * n_v_heads);
assert_eq!(out.len(), rows * n_v_heads * s);
let scale = 1.0 / (s as f32).sqrt();
let key_row = n_k_heads * s;
let val_row = n_v_heads * s;
#[derive(Clone, Copy)]
struct HeadOut(*mut f32);
unsafe impl Send for HeadOut {}
unsafe impl Sync for HeadOut {}
let out_ptr = HeadOut(out.as_mut_ptr());
crate::par::chunks_mut(state, s * s, 1, |h, st| {
let out_ptr = { out_ptr };
let kh = map.k_head(h, n_k_heads, n_v_heads);
let mut m = vec![0.0f32; s * CHUNK]; let mut n = vec![0.0f32; s * CHUNK]; let mut d = vec![0.0f32; CHUNK * s]; let mut gram = vec![0.0f32; CHUNK * CHUNK]; let mut qk = vec![0.0f32; CHUNK * CHUNK]; let mut ratio = vec![0.0f32; CHUNK * CHUNK]; let mut a = [0.0f32; CHUNK];
let mut start = 0;
while start < rows {
let c = CHUNK.min(rows - start);
let krow = |t: usize| &k[(start + t) * key_row + kh * s..][..s];
let qrow = |t: usize| &q[(start + t) * key_row + kh * s..][..s];
for t in 0..c {
let decay = g[(start + t) * n_v_heads + h].exp();
ratio[t * CHUNK + t] = 1.0;
for u in (0..t).rev() {
ratio[t * CHUNK + u] =
ratio[t * CHUNK + u + 1] * g[(start + u + 1) * n_v_heads + h].exp();
}
a[t] = if t == 0 { decay } else { a[t - 1] * decay };
}
for j in 0..s {
let row = &st[j * s..(j + 1) * s];
let mut t = 0;
while t + 4 <= c {
let (m0, m1, m2, m3) =
dot4(row, krow(t), krow(t + 1), krow(t + 2), krow(t + 3));
m[j * CHUNK + t] = m0;
m[j * CHUNK + t + 1] = m1;
m[j * CHUNK + t + 2] = m2;
m[j * CHUNK + t + 3] = m3;
let (n0, n1, n2, n3) =
dot4(row, qrow(t), qrow(t + 1), qrow(t + 2), qrow(t + 3));
n[j * CHUNK + t] = n0;
n[j * CHUNK + t + 1] = n1;
n[j * CHUNK + t + 2] = n2;
n[j * CHUNK + t + 3] = n3;
t += 4;
}
while t < c {
m[j * CHUNK + t] = dot(row, krow(t));
n[j * CHUNK + t] = dot(row, qrow(t));
t += 1;
}
}
for t in 0..c {
for u in 0..=t {
gram[t * CHUNK + u] = dot(krow(t), krow(u));
qk[t * CHUNK + u] = dot(qrow(t), krow(u));
}
}
for t in 0..c {
let b = beta[(start + t) * n_v_heads + h];
let (d_head, d_tail) = d.split_at_mut(t * s);
let d_t = &mut d_tail[..s];
let o_t = unsafe {
std::slice::from_raw_parts_mut(out_ptr.0.add((start + t) * val_row + h * s), s)
};
let v_t = &v[(start + t) * val_row + h * s..][..s];
for j in 0..s {
let mut pred = a[t] * m[j * CHUNK + t];
let mut read = a[t] * n[j * CHUNK + t];
for u in 0..t {
let w = ratio[t * CHUNK + u];
let du = d_head[u * s + j];
pred += w * gram[t * CHUNK + u] * du;
read += w * qk[t * CHUNK + u] * du;
}
let dj = b * (v_t[j] - pred);
d_t[j] = dj;
o_t[j] = (read + qk[t * CHUNK + t] * dj) * scale;
}
}
let last = c - 1;
for j in 0..s {
let row = &mut st[j * s..(j + 1) * s];
for x in row.iter_mut() {
*x *= a[last];
}
for t in 0..c {
let w = ratio[last * CHUNK + t] * d[t * s + j];
if w == 0.0 {
continue;
}
let kt = krow(t);
for (x, kx) in row.iter_mut().zip(kt) {
*x += w * kx;
}
}
}
start += c;
}
});
}
#[inline]
fn dot4(row: &[f32], a: &[f32], b: &[f32], c: &[f32], d: &[f32]) -> (f32, f32, f32, f32) {
let mut acc = [[0.0f32; 4]; 4];
let (rb, rt) = row.as_chunks::<4>();
let (ab, at) = a.as_chunks::<4>();
let (bb, bt) = b.as_chunks::<4>();
let (cb, ct) = c.as_chunks::<4>();
let (db, dt) = d.as_chunks::<4>();
for ((((r, x), y), z), w) in rb.iter().zip(ab).zip(bb).zip(cb).zip(db) {
for l in 0..4 {
acc[0][l] += r[l] * x[l];
acc[1][l] += r[l] * y[l];
acc[2][l] += r[l] * z[l];
acc[3][l] += r[l] * w[l];
}
}
let mut tail = [0.0f32; 4];
for ((((r, x), y), z), w) in rt.iter().zip(at).zip(bt).zip(ct).zip(dt) {
tail[0] += *r * *x;
tail[1] += *r * *y;
tail[2] += *r * *z;
tail[3] += *r * *w;
}
let sum = |v: [f32; 4], t: f32| v[0] + v[1] + v[2] + v[3] + t;
(
sum(acc[0], tail[0]),
sum(acc[1], tail[1]),
sum(acc[2], tail[2]),
sum(acc[3], tail[3]),
)
}
#[inline]
fn dot(a: &[f32], b: &[f32]) -> f32 {
let mut acc = [0.0f32; 4];
let (ab, at) = a.as_chunks::<4>();
let (bb, bt) = b.as_chunks::<4>();
for (x, y) in ab.iter().zip(bb) {
for l in 0..4 {
acc[l] += x[l] * y[l];
}
}
let mut tail = 0.0f32;
for (x, y) in at.iter().zip(bt) {
tail += *x * *y;
}
acc[0] + acc[1] + acc[2] + acc[3] + tail
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(feature = "metal")]
#[test]
#[ignore = "needs a real Metal-capable GPU; run manually with --ignored on Apple Silicon"]
fn device_chunk_against_host_chunk_throughput() {
use crate::gdn::HeadMap;
use crate::recurrent_state::AlignedF32;
use ferrox_metal::gdn::{DeltaShape, HeadMapKind};
let (n_k, n_v, s) = (4usize, 48usize, 128usize);
let dims = DeltaDims {
n_k_heads: n_k,
n_v_heads: n_v,
head_dim: s,
map: HeadMap::Tiled,
};
let shape = DeltaShape {
n_k_heads: n_k,
n_v_heads: n_v,
head_dim: s,
map: HeadMapKind::Tiled,
};
for rows in [1usize, 2, 8, 32, 64, 128, 512, 1201] {
let mut seed = 5u32;
let mut draw = |n: usize, scale: f32| -> Vec<f32> {
(0..n)
.map(|_| {
seed = seed.wrapping_mul(1664525).wrapping_add(1013904223);
((seed >> 8) as f32 / 8388608.0 - 1.0) * scale
})
.collect()
};
let state0 = draw(dims.state_len(), 0.5);
let q = draw(rows * n_k * s, 1.0);
let k = draw(rows * n_k * s, 1.0);
let v = draw(rows * n_v * s, 1.0);
let g: Vec<f32> = draw(rows * n_v, 3.0).iter().map(|x| -x.abs()).collect();
let beta: Vec<f32> = draw(rows * n_v, 1.0)
.iter()
.map(|x| 0.5 + 0.25 * x)
.collect();
let mut out = vec![0.0f32; rows * n_v * s];
let mut host_state = state0.clone();
let mut host = || {
host_state.copy_from_slice(&state0);
let t = std::time::Instant::now();
delta_chunk(dims, rows, &mut host_state, &q, &k, &v, &g, &beta, &mut out);
t.elapsed().as_secs_f64() * 1e3
};
host();
let host_ms = host();
let mut device_state = AlignedF32::zeros(state0.len());
device_state.copy_from_slice(&state0);
let (bytes, ptr) = (device_state.alloc_bytes(), device_state.as_ptr());
let warm = unsafe {
ferrox_metal::gdn_chunk::launch_delta_chunk(
shape, rows, ptr, bytes, &q, &k, &v, &g, &beta,
)
};
warm.expect("the kernel launches");
device_state.copy_from_slice(&state0);
let t = std::time::Instant::now();
let _ = unsafe {
ferrox_metal::gdn_chunk::launch_delta_chunk(
shape, rows, ptr, bytes, &q, &k, &v, &g, &beta,
)
}
.expect("the kernel launches");
let device_ms = t.elapsed().as_secs_f64() * 1e3;
eprintln!(
"{rows} rows: host chunk {host_ms:.2} ms, device chunk {device_ms:.2} ms, \
{:.2}x",
host_ms / device_ms
);
}
}
#[cfg(feature = "metal")]
#[test]
#[ignore = "needs a real Metal-capable GPU; run manually with --ignored on Apple Silicon"]
fn the_device_chunk_matches_this_one() {
use crate::gdn::HeadMap;
use crate::recurrent_state::AlignedF32;
use ferrox_metal::gdn::{DeltaShape, HeadMapKind};
for (n_k, n_v, s, rows, map, kind) in [
(
4usize,
48usize,
128usize,
37usize,
HeadMap::Tiled,
HeadMapKind::Tiled,
),
(4, 48, 128, 37, HeadMap::Grouped, HeadMapKind::Grouped),
(2, 2, 4, 3, HeadMap::Tiled, HeadMapKind::Tiled),
(2, 4, 16, 70, HeadMap::Grouped, HeadMapKind::Grouped),
] {
let dims = DeltaDims {
n_k_heads: n_k,
n_v_heads: n_v,
head_dim: s,
map,
};
let mut seed = 987_654_321u32;
let mut draw = |n: usize, scale: f32| -> Vec<f32> {
(0..n)
.map(|_| {
seed = seed.wrapping_mul(1664525).wrapping_add(1013904223);
((seed >> 8) as f32 / 8388608.0 - 1.0) * scale
})
.collect()
};
let state0 = draw(dims.state_len(), 0.5);
let q = draw(rows * n_k * s, 1.0);
let k = draw(rows * n_k * s, 1.0);
let v = draw(rows * n_v * s, 1.0);
let g: Vec<f32> = draw(rows * n_v, 1.0).iter().map(|x| -x.abs()).collect();
let beta: Vec<f32> = draw(rows * n_v, 1.0)
.iter()
.map(|x| 0.5 + 0.25 * x)
.collect();
let mut host_state = state0.clone();
let mut host_out = vec![0.0f32; rows * n_v * s];
delta_chunk(
dims,
rows,
&mut host_state,
&q,
&k,
&v,
&g,
&beta,
&mut host_out,
);
let mut device_state = AlignedF32::zeros(state0.len());
device_state.copy_from_slice(&state0);
let shape = DeltaShape {
n_k_heads: n_k,
n_v_heads: n_v,
head_dim: s,
map: kind,
};
let bytes = device_state.alloc_bytes();
let ptr = device_state.as_ptr();
let device_out = unsafe {
ferrox_metal::gdn_chunk::launch_delta_chunk(
shape, rows, ptr, bytes, &q, &k, &v, &g, &beta,
)
}
.expect("the kernel launches");
let mut seq_state = state0.clone();
let mut seq_out = vec![0.0f32; rows * n_v * s];
for t in 0..rows {
crate::gdn::delta_step(
dims,
&mut seq_state,
&q[t * n_k * s..][..n_k * s],
&k[t * n_k * s..][..n_k * s],
&v[t * n_v * s..][..n_v * s],
&g[t * n_v..][..n_v],
&beta[t * n_v..][..n_v],
&mut seq_out[t * n_v * s..][..n_v * s],
);
}
let worst = |xs: &[f32], ys: &[f32]| -> f32 {
xs.iter()
.zip(ys)
.map(|(a, b)| (a - b).abs() / b.abs().max(1.0))
.fold(0.0f32, f32::max)
};
let (dev_out, dev_state) = (
worst(&device_out, &seq_out),
worst(&device_state, &seq_state),
);
let (hst_out, hst_state) = (worst(&host_out, &seq_out), worst(&host_state, &seq_state));
eprintln!(
"{n_v}x{s} rows {rows} {map:?}: device {dev_out:.2e}/{dev_state:.2e} \
host {hst_out:.2e}/{hst_state:.2e} against the sequential rule"
);
let tol = 5e-3;
for (what, got) in [
("device out", dev_out),
("device state", dev_state),
("host out", hst_out),
("host state", hst_state),
] {
assert!(
got <= tol,
"{n_v}x{s} rows {rows} {map:?} {what}: {got:.3e} over {tol:.0e}"
);
}
}
}
use crate::gdn::{delta_step, HeadMap};
#[test]
fn a_chunk_is_the_sequential_recurrence() {
for (n_k, n_v, s, rows, map) in [
(2usize, 4usize, 8usize, 3usize, HeadMap::Tiled),
(2, 4, 8, CHUNK, HeadMap::Tiled),
(2, 4, 8, CHUNK + 5, HeadMap::Grouped),
(4, 8, 16, 2 * CHUNK + 1, HeadMap::Tiled),
(1, 3, 32, 7, HeadMap::Grouped),
] {
let dims = DeltaDims {
n_k_heads: n_k,
n_v_heads: n_v,
head_dim: s,
map,
};
let mut seed = 7u32 + rows as u32;
let mut draw = |n: usize, scale: f32| -> Vec<f32> {
(0..n)
.map(|_| {
seed = seed.wrapping_mul(1664525).wrapping_add(1013904223);
((seed >> 8) as f32 / 8388608.0 - 1.0) * scale
})
.collect()
};
let state0 = draw(dims.state_len(), 0.4);
let q = draw(rows * n_k * s, 1.0);
let k = draw(rows * n_k * s, 1.0);
let v = draw(rows * n_v * s, 1.0);
let g: Vec<f32> = draw(rows * n_v, 3.0).iter().map(|x| -x.abs()).collect();
let beta: Vec<f32> = draw(rows * n_v, 1.0)
.iter()
.map(|x| 0.5 + 0.3 * x)
.collect();
let mut seq_state = state0.clone();
let mut seq_out = vec![0.0f32; rows * n_v * s];
for r in 0..rows {
let mut o = vec![0.0f32; n_v * s];
delta_step(
dims,
&mut seq_state,
&q[r * n_k * s..(r + 1) * n_k * s],
&k[r * n_k * s..(r + 1) * n_k * s],
&v[r * n_v * s..(r + 1) * n_v * s],
&g[r * n_v..(r + 1) * n_v],
&beta[r * n_v..(r + 1) * n_v],
&mut o,
);
seq_out[r * n_v * s..(r + 1) * n_v * s].copy_from_slice(&o);
}
let mut chunk_state = state0.clone();
let mut chunk_out = vec![0.0f32; rows * n_v * s];
delta_chunk(
dims,
rows,
&mut chunk_state,
&q,
&k,
&v,
&g,
&beta,
&mut chunk_out,
);
let tol = 2e-4;
for (i, (a, b)) in chunk_out.iter().zip(seq_out.iter()).enumerate() {
assert!(
(a - b).abs() <= tol * b.abs().max(1.0),
"{n_v}x{s} rows={rows} out[{i}]: chunk={a} seq={b}"
);
}
for (i, (a, b)) in chunk_state.iter().zip(seq_state.iter()).enumerate() {
assert!(
(a - b).abs() <= tol * b.abs().max(1.0),
"{n_v}x{s} rows={rows} state[{i}]: chunk={a} seq={b}"
);
}
}
}
#[test]
#[ignore = "a measurement, not an assertion; run with --nocapture"]
fn chunked_against_sequential_throughput() {
let dims = DeltaDims {
n_k_heads: 4,
n_v_heads: 48,
head_dim: 128,
map: HeadMap::Tiled,
};
let rows = 128;
let s = 128;
let state0 = vec![0.01f32; dims.state_len()];
let q: Vec<f32> = (0..rows * 4 * s).map(|i| (i as f32 * 0.01).sin()).collect();
let k: Vec<f32> = (0..rows * 4 * s).map(|i| (i as f32 * 0.02).cos()).collect();
let v: Vec<f32> = (0..rows * 48 * s)
.map(|i| (i as f32 * 0.03).sin())
.collect();
let g = vec![-0.05f32; rows * 48];
let beta = vec![0.5f32; rows * 48];
let mut st = state0.clone();
let mut out = vec![0.0f32; rows * 48 * s];
let t = std::time::Instant::now();
for r in 0..rows {
let mut o = vec![0.0f32; 48 * s];
delta_step(
dims,
&mut st,
&q[r * 4 * s..(r + 1) * 4 * s],
&k[r * 4 * s..(r + 1) * 4 * s],
&v[r * 48 * s..(r + 1) * 48 * s],
&g[r * 48..(r + 1) * 48],
&beta[r * 48..(r + 1) * 48],
&mut o,
);
out[r * 48 * s..(r + 1) * 48 * s].copy_from_slice(&o);
}
let seq = t.elapsed().as_secs_f64();
let mut st2 = state0.clone();
let mut out2 = vec![0.0f32; rows * 48 * s];
let t = std::time::Instant::now();
delta_chunk(dims, rows, &mut st2, &q, &k, &v, &g, &beta, &mut out2);
let chunked = t.elapsed().as_secs_f64();
eprintln!(
"{rows} rows: sequential {:.1} ms, chunked(C={CHUNK}) {:.1} ms, {:.2}x",
seq * 1e3,
chunked * 1e3,
seq / chunked
);
}
#[test]
fn a_vanishing_decay_is_forgetting_and_not_a_division() {
let dims = DeltaDims {
n_k_heads: 1,
n_v_heads: 2,
head_dim: 8,
map: HeadMap::Tiled,
};
let rows = CHUNK + 3;
let state0: Vec<f32> = (0..dims.state_len())
.map(|i| 0.1 + i as f32 * 0.01)
.collect();
let q = vec![0.3f32; rows * 8];
let k = vec![0.2f32; rows * 8];
let v = vec![0.7f32; rows * 2 * 8];
let g = vec![-90.0f32; rows * 2];
let beta = vec![0.5f32; rows * 2];
let mut seq_state = state0.clone();
let mut seq_out = vec![0.0f32; rows * 2 * 8];
for r in 0..rows {
let mut o = vec![0.0f32; 2 * 8];
delta_step(
dims,
&mut seq_state,
&q[r * 8..(r + 1) * 8],
&k[r * 8..(r + 1) * 8],
&v[r * 16..(r + 1) * 16],
&g[r * 2..(r + 1) * 2],
&beta[r * 2..(r + 1) * 2],
&mut o,
);
seq_out[r * 16..(r + 1) * 16].copy_from_slice(&o);
}
let mut chunk_state = state0.clone();
let mut chunk_out = vec![0.0f32; rows * 2 * 8];
delta_chunk(
dims,
rows,
&mut chunk_state,
&q,
&k,
&v,
&g,
&beta,
&mut chunk_out,
);
for (a, b) in chunk_out.iter().zip(seq_out.iter()) {
assert!((a - b).abs() <= 1e-5, "chunk={a} seq={b}");
assert!(a.is_finite(), "no division by a vanished decay");
}
for (a, b) in chunk_state.iter().zip(seq_state.iter()) {
assert!((a - b).abs() <= 1e-5 && a.is_finite());
}
}
}