pub fn slopes(n_head: usize) -> Vec<f32> {
debug_assert!(n_head > 0);
let start = 2f32.powf(-(2f32.powf(-((n_head as f32).log2() - 3.0))));
(0..n_head).map(|h| start * start.powi(h as i32)).collect()
}
pub fn slope_scale(il: usize, n_layer: usize) -> f32 {
assert!(
n_layer > 1,
"minimax-01.cpp:288 divides by n_layer - 1; a one-layer model has no scale"
);
1.0 - il as f32 / (n_layer - 1) as f32 + 1e-5
}
pub fn lightning_step(q: &[f32], k: &[f32], v: &[f32], state: &mut [f32], decay: f32) -> Vec<f32> {
let d = q.len();
debug_assert_eq!(k.len(), d);
debug_assert_eq!(v.len(), d);
debug_assert_eq!(state.len(), d * d);
let mut out = vec![0.0f32; d];
for (a, &qa) in q.iter().enumerate() {
let qa = qa * decay;
if qa == 0.0 {
continue;
}
let row = &state[a * d..(a + 1) * d];
for (o, &s) in out.iter_mut().zip(row) {
*o += qa * s;
}
}
let qk: f32 = q.iter().zip(k).map(|(a, b)| a * b).sum();
for (o, &vb) in out.iter_mut().zip(v) {
*o += qk * vb;
}
for (a, &ka) in k.iter().enumerate() {
let row = &mut state[a * d..(a + 1) * d];
for (s, &vb) in row.iter_mut().zip(v) {
*s = *s * decay + ka * vb;
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_slopes_are_the_geometric_ladder() {
let s = slopes(8);
assert_eq!(s.len(), 8);
let start = 2f32.powf(-(2f32.powf(-((8f32).log2() - 3.0))));
assert!((s[0] - start).abs() < 1e-7);
assert!((s[1] - start * start).abs() < 1e-7);
for w in s.windows(2) {
assert!(w[1] < w[0], "the ladder decreases");
}
}
#[test]
fn the_layer_scale_walks_from_one_to_almost_zero() {
assert!((slope_scale(0, 4) - 1.000_01).abs() < 1e-6);
assert!((slope_scale(3, 4) - 1e-5).abs() < 1e-6);
}
#[test]
fn the_recurrence_equals_the_chunked_form_over_two_tokens() {
let d = 3;
let decay = 0.7f32;
let q0 = [0.3, -0.5, 0.9];
let q1 = [-0.2, 0.4, 0.1];
let k0 = [0.6, 0.1, -0.3];
let k1 = [0.2, -0.7, 0.5];
let v0 = [1.0, -2.0, 0.5];
let v1 = [-0.4, 0.8, 1.2];
let kv0: Vec<f32> = (0..d * d).map(|i| 0.1 * (i as f32) - 0.3).collect();
let mut state = kv0.clone();
let out0 = lightning_step(&q0, &k0, &v0, &mut state, decay);
let out1 = lightning_step(&q1, &k1, &v1, &mut state, decay);
let dot = |a: &[f32], b: &[f32]| -> f32 { a.iter().zip(b).map(|(x, y)| x * y).sum() };
let apply = |q: &[f32], scale: f32, kv: &[f32]| -> Vec<f32> {
(0..d)
.map(|b| (0..d).map(|a| q[a] * scale * kv[a * d + b]).sum())
.collect()
};
let want0: Vec<f32> = apply(&q0, decay, &kv0)
.iter()
.zip(&v0)
.map(|(x, v)| x + dot(&q0, &k0) * v)
.collect();
let want1: Vec<f32> = apply(&q1, decay * decay, &kv0)
.iter()
.enumerate()
.map(|(b, x)| x + decay * dot(&q1, &k0) * v0[b] + dot(&q1, &k1) * v1[b])
.collect();
for (got, want) in out0.iter().zip(&want0) {
assert!((got - want).abs() < 1e-6, "{got} vs {want}");
}
for (got, want) in out1.iter().zip(&want1) {
assert!((got - want).abs() < 1e-6, "{got} vs {want}");
}
for a in 0..d {
for b in 0..d {
let want = kv0[a * d + b] * decay * decay + decay * k0[a] * v0[b] + k1[a] * v1[b];
let got = state[a * d + b];
assert!(
(got - want).abs() < 1e-6,
"state[{a}][{b}]: {got} vs {want}"
);
}
}
}
}