#[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::*;
#[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]);
}
}