#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum HeadMap {
Tiled,
Grouped,
}
impl HeadMap {
pub fn k_head(self, v_head: usize, n_k_heads: usize, n_v_heads: usize) -> usize {
match self {
HeadMap::Tiled => v_head % n_k_heads,
HeadMap::Grouped => v_head / (n_v_heads / n_k_heads),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct DeltaDims {
pub n_k_heads: usize,
pub n_v_heads: usize,
pub head_dim: usize,
pub map: HeadMap,
}
impl DeltaDims {
pub fn state_len(self) -> usize {
self.n_v_heads * self.head_dim * self.head_dim
}
}
pub fn l2_normalize(x: &mut [f32], eps: f32) {
let sum: f64 = x.iter().map(|v| (*v as f64) * (*v as f64)).sum();
let scale = 1.0 / (sum as f32).sqrt().max(eps);
for v in x.iter_mut() {
*v *= scale;
}
}
#[allow(clippy::too_many_arguments)] pub fn delta_step(
dims: DeltaDims,
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(), n_k_heads * s);
assert_eq!(k.len(), n_k_heads * s);
assert_eq!(v.len(), n_v_heads * s);
assert_eq!(g.len(), n_v_heads);
assert_eq!(beta.len(), n_v_heads);
assert_eq!(out.len(), n_v_heads * s);
assert!(n_k_heads > 0 && n_v_heads.is_multiple_of(n_k_heads), ":308");
let scale = 1.0 / (s as f32).sqrt();
crate::par::chunks_mut2_by(state, out, s * s, s, 1, |h, st, oh| {
let kh = map.k_head(h, n_k_heads, n_v_heads);
let (qh, kk) = (&q[kh * s..(kh + 1) * s], &k[kh * s..(kh + 1) * s]);
let vh = &v[h * s..(h + 1) * s];
let decay = g[h].exp();
let mut d = vec![0.0f32; s];
for j in 0..s {
let row = &mut st[j * s..(j + 1) * s];
let pred = decay_and_dot(row, kk, decay);
d[j] = (vh[j] - pred) * beta[h];
}
for j in 0..s {
let row = &mut st[j * s..(j + 1) * s];
oh[j] = update_and_dot(row, kk, d[j], qh) * scale;
}
});
}
#[inline]
fn decay_and_dot(row: &mut [f32], k: &[f32], decay: f32) -> f32 {
let mut acc = [0.0f32; 4];
let (rb, rt) = row.as_chunks_mut::<4>();
let (kb, kt) = k.as_chunks::<4>();
for (r, kk) in rb.iter_mut().zip(kb) {
for l in 0..4 {
r[l] *= decay;
acc[l] += r[l] * kk[l];
}
}
let mut tail = 0.0f32;
for (r, kk) in rt.iter_mut().zip(kt) {
*r *= decay;
tail += *r * *kk;
}
acc[0] + acc[1] + acc[2] + acc[3] + tail
}
#[inline]
fn update_and_dot(row: &mut [f32], k: &[f32], d: f32, q: &[f32]) -> f32 {
let mut acc = [0.0f32; 4];
let (rb, rt) = row.as_chunks_mut::<4>();
let (kb, kt) = k.as_chunks::<4>();
let (qb, qt) = q.as_chunks::<4>();
for ((r, kk), qq) in rb.iter_mut().zip(kb).zip(qb) {
for l in 0..4 {
r[l] += kk[l] * d;
acc[l] += r[l] * qq[l];
}
}
let mut tail = 0.0f32;
for ((r, kk), qq) in rt.iter_mut().zip(kt).zip(qt) {
*r += *kk * d;
tail += *r * *qq;
}
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_delta_step_against_host_throughput() {
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,
};
let mut seed = 11u32;
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(n_k * s, 1.0);
let k = draw(n_k * s, 1.0);
let v = draw(n_v * s, 1.0);
let g: Vec<f32> = draw(n_v, 3.0).iter().map(|x| -x.abs()).collect();
let beta: Vec<f32> = draw(n_v, 1.0).iter().map(|x| 0.5 + 0.25 * x).collect();
let reps = 20;
let mut host_state = state0.clone();
let mut out = vec![0.0f32; n_v * s];
delta_step(dims, &mut host_state, &q, &k, &v, &g, &beta, &mut out);
let t = std::time::Instant::now();
for _ in 0..reps {
delta_step(dims, &mut host_state, &q, &k, &v, &g, &beta, &mut out);
}
let host_us = t.elapsed().as_secs_f64() * 1e6 / reps as f64;
let mut device_state = state0.clone();
ferrox_metal::gdn::launch_delta_step(shape, &mut device_state, &q, &k, &v, &g, &beta)
.expect("warm");
let t = std::time::Instant::now();
for _ in 0..reps {
ferrox_metal::gdn::launch_delta_step(shape, &mut device_state, &q, &k, &v, &g, &beta)
.expect("the kernel launches");
}
let device_us = t.elapsed().as_secs_f64() * 1e6 / reps as f64;
let gb = 2.0 * dims.state_len() as f64 * 4.0 / 1e9;
eprintln!(
"one row, {n_v}x{s}: host {host_us:.0} us ({:.1} GB/s), device {device_us:.0} us \
({:.1} GB/s, uploads and readback included)",
gb / (host_us / 1e6),
gb / (device_us / 1e6),
);
}
#[cfg(feature = "metal")]
#[test]
#[ignore = "needs a real Metal-capable GPU; run manually with --ignored on Apple Silicon"]
fn the_device_delta_step_matches_this_one() {
use ferrox_metal::gdn::{DeltaShape, HeadMapKind};
for (n_k, n_v, s, map, kind) in [
(
4usize,
48usize,
128usize,
HeadMap::Tiled,
HeadMapKind::Tiled,
),
(4, 48, 128, HeadMap::Grouped, HeadMapKind::Grouped),
(2, 2, 4, HeadMap::Tiled, HeadMapKind::Tiled),
] {
let dims = DeltaDims {
n_k_heads: n_k,
n_v_heads: n_v,
head_dim: s,
map,
};
let mut seed = 12345u32;
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(n_k * s, 1.0);
let k = draw(n_k * s, 1.0);
let v = draw(n_v * s, 1.0);
let g: Vec<f32> = draw(n_v, 1.0).iter().map(|x| -x.abs()).collect();
let beta: Vec<f32> = draw(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; n_v * s];
delta_step(dims, &mut host_state, &q, &k, &v, &g, &beta, &mut host_out);
let mut device_state = state0.clone();
let shape = DeltaShape {
n_k_heads: n_k,
n_v_heads: n_v,
head_dim: s,
map: kind,
};
let device_out = ferrox_metal::gdn::launch_delta_step(
shape,
&mut device_state,
&q,
&k,
&v,
&g,
&beta,
)
.expect("the kernel launches");
let tol = 2e-4;
for (i, (a, b)) in device_out.iter().zip(host_out.iter()).enumerate() {
assert!(
(a - b).abs() <= tol * b.abs().max(1.0),
"{n_v}x{s} {map:?} out[{i}]: device={a} host={b}"
);
}
for (i, (a, b)) in device_state.iter().zip(host_state.iter()).enumerate() {
assert!(
(a - b).abs() <= tol * b.abs().max(1.0),
"{n_v}x{s} {map:?} state[{i}]: device={a} host={b}"
);
}
}
}
#[test]
#[ignore = "a measurement, not an assertion; run with --nocapture"]
fn delta_step_throughput_probe() {
let dims = DeltaDims {
n_k_heads: 4,
n_v_heads: 48,
head_dim: 128,
map: HeadMap::Tiled,
};
let mut state = vec![0.01f32; dims.state_len()];
let q: Vec<f32> = (0..4 * 128).map(|i| (i as f32 * 0.01).sin()).collect();
let k = q.clone();
let v: Vec<f32> = (0..48 * 128).map(|i| (i as f32 * 0.02).cos()).collect();
let g = vec![-0.1f32; 48];
let beta = vec![0.5f32; 48];
let mut out = vec![0.0f32; 48 * 128];
for _ in 0..8 {
delta_step(dims, &mut state, &q, &k, &v, &g, &beta, &mut out);
}
let n = 200;
let t = std::time::Instant::now();
for _ in 0..n {
delta_step(dims, &mut state, &q, &k, &v, &g, &beta, &mut out);
}
let per = t.elapsed().as_secs_f64() / n as f64;
let flops = 48.0 * 2.0 * 2.0 * 128.0 * 128.0;
let bytes = (dims.state_len() * 4 * 2) as f64;
eprintln!(
"delta_step {:.1} us/call, {:.1} GFLOP/s, {:.1} GB/s of state",
per * 1e6,
flops / per / 1e9,
bytes / per / 1e9
);
}
#[test]
fn l2_normalize_clamps_the_divisor() {
let mut v = [3.0f32, 4.0];
l2_normalize(&mut v, 1e-6);
assert!((v[0] - 0.6).abs() < 1e-6 && (v[1] - 0.8).abs() < 1e-6);
let mut z = [0.0f32, 0.0];
l2_normalize(&mut z, 1e-6);
assert_eq!(z, [0.0, 0.0]);
}
#[test]
fn the_scalar_recurrence() {
let dims = DeltaDims {
n_k_heads: 1,
n_v_heads: 1,
head_dim: 1,
map: HeadMap::Tiled,
};
let mut st = vec![2.0f32];
let mut out = [0.0f32];
delta_step(
dims,
&mut st,
&[1.0],
&[1.0],
&[5.0],
&[0.0],
&[0.5],
&mut out,
);
assert!((st[0] - 3.5).abs() < 1e-6 && (out[0] - 3.5).abs() < 1e-6);
}
#[test]
fn the_two_head_maps_differ_and_are_the_documented_ones() {
assert_eq!(HeadMap::Tiled.k_head(3, 2, 4), 1);
assert_eq!(HeadMap::Grouped.k_head(3, 2, 4), 1);
assert_eq!(HeadMap::Tiled.k_head(1, 2, 4), 1);
assert_eq!(HeadMap::Grouped.k_head(1, 2, 4), 0);
let mk = |map| DeltaDims {
n_k_heads: 2,
n_v_heads: 4,
head_dim: 1,
map,
};
let (q, k) = ([1.0f32, 1.0], [1.0f32, 10.0]);
let v = [1.0f32; 4];
let mut out_t = [0.0f32; 4];
let mut out_g = [0.0f32; 4];
delta_step(
mk(HeadMap::Tiled),
&mut [0.0; 4],
&q,
&k,
&v,
&[0.0; 4],
&[1.0; 4],
&mut out_t,
);
delta_step(
mk(HeadMap::Grouped),
&mut [0.0; 4],
&q,
&k,
&v,
&[0.0; 4],
&[1.0; 4],
&mut out_g,
);
assert_eq!(out_t, [1.0, 10.0, 1.0, 10.0]);
assert_eq!(out_g, [1.0, 1.0, 10.0, 10.0]);
}
}