use crate::cache::PagedKvStore;
fn rope_freqs(theta: f32, dim: usize) -> std::rc::Rc<Vec<f32>> {
type FreqTables = std::collections::HashMap<(u32, usize), std::rc::Rc<Vec<f32>>>;
thread_local! {
static CACHE: std::cell::RefCell<FreqTables> =
std::cell::RefCell::new(FreqTables::new());
}
let key = (theta.to_bits(), dim);
CACHE.with(|c| {
if let Some(hit) = c.borrow().get(&key) {
return std::rc::Rc::clone(hit);
}
let table: Vec<f32> = (0..dim / 2)
.map(|i| 1.0 / theta.powf((2 * i) as f32 / dim as f32))
.collect();
let table = std::rc::Rc::new(table);
c.borrow_mut().insert(key, std::rc::Rc::clone(&table));
table
})
}
pub fn apply_rope(vec: &mut [f32], pos: usize, theta: f32) {
let dim = vec.len();
let half = dim / 2;
let freqs = rope_freqs(theta, dim);
for i in 0..half {
let freq = freqs[i];
let angle = pos as f32 * freq;
let (sin, cos) = angle.sin_cos();
let a = vec[i];
let b = vec[i + half];
vec[i] = a * cos - b * sin;
vec[i + half] = a * sin + b * cos;
}
}
pub fn apply_rope_back(vec: &mut [f32], pos: usize, theta: f32) {
let dim = vec.len();
let half = dim / 2;
let freqs = rope_freqs(theta, dim);
for i in 0..half {
let freq = freqs[i];
let angle = pos as f32 * freq;
let (sin, cos) = angle.sin_cos();
let a = vec[i];
let b = vec[i + half];
vec[i] = a * cos + b * sin;
vec[i + half] = -a * sin + b * cos;
}
}
pub fn apply_rope_interleaved_back(vec: &mut [f32], pos: usize, theta: f32) {
let dim = vec.len();
let half = dim / 2;
let freqs = rope_freqs(theta, dim);
for i in 0..half {
let freq = freqs[i];
let angle = pos as f32 * freq;
let (sin, cos) = angle.sin_cos();
let a = vec[2 * i];
let b = vec[2 * i + 1];
vec[2 * i] = a * cos + b * sin;
vec[2 * i + 1] = -a * sin + b * cos;
}
}
#[inline]
fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len());
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("neon") {
return unsafe { dot_f32_neon(a, b) };
}
}
#[cfg(target_arch = "x86_64")]
{
if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
return unsafe { dot_f32_avx2(a, b) };
}
}
a.iter().zip(b).map(|(x, y)| x * y).sum()
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn dot_f32_neon(a: &[f32], b: &[f32]) -> f32 {
use std::arch::aarch64::*;
let n = a.len();
let mut acc = vdupq_n_f32(0.0);
let mut i = 0;
while i + 4 <= n {
let va = vld1q_f32(a.as_ptr().add(i));
let vb = vld1q_f32(b.as_ptr().add(i));
acc = vfmaq_f32(acc, va, vb);
i += 4;
}
let mut sum = vaddvq_f32(acc);
while i < n {
sum += a[i] * b[i];
i += 1;
}
sum
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn dot_f32_avx2(a: &[f32], b: &[f32]) -> f32 {
use std::arch::x86_64::*;
let n = a.len();
let mut acc = _mm256_setzero_ps();
let mut i = 0;
while i + 8 <= n {
let va = _mm256_loadu_ps(a.as_ptr().add(i));
let vb = _mm256_loadu_ps(b.as_ptr().add(i));
acc = _mm256_fmadd_ps(va, vb, acc);
i += 8;
}
let mut sum = hsum256_ps(acc);
while i < n {
sum += a[i] * b[i];
i += 1;
}
sum
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn hsum256_ps(acc: std::arch::x86_64::__m256) -> f32 {
use std::arch::x86_64::*;
let lo = _mm256_castps256_ps128(acc);
let hi = _mm256_extractf128_ps(acc, 1);
let mut s128 = _mm_add_ps(lo, hi);
s128 = _mm_add_ps(s128, _mm_movehl_ps(s128, s128));
s128 = _mm_add_ss(s128, _mm_shuffle_ps(s128, s128, 0x55));
_mm_cvtss_f32(s128)
}
fn online_attn_accumulate(
q_h: &[f32],
scale: f32,
head_dim: usize,
out_h: &mut [f32],
attn_softcap: Option<f32>,
sink: Option<f32>,
mut for_each_kv: impl FnMut(&mut dyn FnMut(&[f32], &[f32])),
) {
debug_assert_eq!(q_h.len(), head_dim);
debug_assert_eq!(out_h.len(), head_dim);
let mut m = f32::NEG_INFINITY;
let mut l = 0f32;
out_h.fill(0.0);
for_each_kv(&mut |k_t, v_t| {
let mut s = dot_f32(q_h, k_t) * scale;
if let Some(sc) = attn_softcap.filter(|&c| c > 0.0) {
s = sc * (s / sc).tanh();
}
let m_new = m.max(s);
let alpha = (m - m_new).exp();
let p = (s - m_new).exp();
l = l * alpha + p;
axpy_scale(out_h, alpha, v_t, p);
m = m_new;
});
if let Some(s) = sink {
let m_new = m.max(s);
let alpha = (m - m_new).exp();
l = l * alpha + (s - m_new).exp();
scale_inplace(out_h, alpha);
}
if l > 0.0 {
let inv = 1.0 / l;
scale_inplace(out_h, inv);
}
}
#[inline]
fn axpy_scale(out: &mut [f32], alpha: f32, v: &[f32], p: f32) {
debug_assert_eq!(out.len(), v.len());
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("neon") {
unsafe { axpy_scale_neon(out, alpha, v, p) };
return;
}
}
#[cfg(target_arch = "x86_64")]
{
if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
unsafe { axpy_scale_avx2(out, alpha, v, p) };
return;
}
}
for (o, &vv) in out.iter_mut().zip(v) {
*o = *o * alpha + p * vv;
}
}
#[inline]
fn axpy(out: &mut [f32], v: &[f32], p: f32) {
debug_assert_eq!(out.len(), v.len());
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("neon") {
unsafe { axpy_neon(out, v, p) };
return;
}
}
#[cfg(target_arch = "x86_64")]
{
if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
unsafe { axpy_avx2(out, v, p) };
return;
}
}
for (o, &vv) in out.iter_mut().zip(v) {
*o += p * vv;
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn axpy_neon(out: &mut [f32], v: &[f32], p: f32) {
use std::arch::aarch64::*;
let n = out.len();
let vp = vdupq_n_f32(p);
let mut i = 0;
while i + 4 <= n {
let o = vld1q_f32(out.as_ptr().add(i));
let vv = vld1q_f32(v.as_ptr().add(i));
vst1q_f32(out.as_mut_ptr().add(i), vfmaq_f32(o, vv, vp));
i += 4;
}
while i < n {
out[i] += p * v[i];
i += 1;
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn axpy_avx2(out: &mut [f32], v: &[f32], p: f32) {
use std::arch::x86_64::*;
let n = out.len();
let vp = _mm256_set1_ps(p);
let mut i = 0;
while i + 8 <= n {
let o = _mm256_loadu_ps(out.as_ptr().add(i));
let vv = _mm256_loadu_ps(v.as_ptr().add(i));
_mm256_storeu_ps(out.as_mut_ptr().add(i), _mm256_fmadd_ps(vv, vp, o));
i += 8;
}
while i < n {
out[i] += p * v[i];
i += 1;
}
}
#[inline]
fn scale_inplace(x: &mut [f32], s: f32) {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("neon") {
unsafe { scale_inplace_neon(x, s) };
return;
}
}
#[cfg(target_arch = "x86_64")]
{
if std::is_x86_feature_detected!("avx2") {
unsafe { scale_inplace_avx2(x, s) };
return;
}
}
for v in x.iter_mut() {
*v *= s;
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn axpy_scale_neon(out: &mut [f32], alpha: f32, v: &[f32], p: f32) {
use std::arch::aarch64::*;
let n = out.len();
let va = vdupq_n_f32(alpha);
let vp = vdupq_n_f32(p);
let mut i = 0;
while i + 4 <= n {
let o = vld1q_f32(out.as_ptr().add(i));
let vv = vld1q_f32(v.as_ptr().add(i));
let r = vfmaq_f32(vmulq_f32(o, va), vv, vp);
vst1q_f32(out.as_mut_ptr().add(i), r);
i += 4;
}
while i < n {
out[i] = out[i] * alpha + p * v[i];
i += 1;
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn scale_inplace_neon(x: &mut [f32], s: f32) {
use std::arch::aarch64::*;
let n = x.len();
let vs = vdupq_n_f32(s);
let mut i = 0;
while i + 4 <= n {
let v = vld1q_f32(x.as_ptr().add(i));
vst1q_f32(x.as_mut_ptr().add(i), vmulq_f32(v, vs));
i += 4;
}
while i < n {
x[i] *= s;
i += 1;
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn axpy_scale_avx2(out: &mut [f32], alpha: f32, v: &[f32], p: f32) {
use std::arch::x86_64::*;
let n = out.len();
let va = _mm256_set1_ps(alpha);
let vp = _mm256_set1_ps(p);
let mut i = 0;
while i + 8 <= n {
let o = _mm256_loadu_ps(out.as_ptr().add(i));
let vv = _mm256_loadu_ps(v.as_ptr().add(i));
let r = _mm256_fmadd_ps(vv, vp, _mm256_mul_ps(o, va));
_mm256_storeu_ps(out.as_mut_ptr().add(i), r);
i += 8;
}
while i < n {
out[i] = out[i] * alpha + p * v[i];
i += 1;
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn scale_inplace_avx2(x: &mut [f32], s: f32) {
use std::arch::x86_64::*;
let n = x.len();
let vs = _mm256_set1_ps(s);
let mut i = 0;
while i + 8 <= n {
let v = _mm256_loadu_ps(x.as_ptr().add(i));
_mm256_storeu_ps(x.as_mut_ptr().add(i), _mm256_mul_ps(v, vs));
i += 8;
}
while i < n {
x[i] *= s;
i += 1;
}
}
pub fn apply_rope_with_freq_factors(vec: &mut [f32], pos: usize, theta: f32, freq_factors: &[f32]) {
let dim = vec.len();
let half = dim / 2;
assert_eq!(
freq_factors.len(),
half,
"freq_factors must have one entry per rotation band (dim/2)"
);
let freqs = rope_freqs(theta, dim);
for i in 0..half {
let freq = freqs[i];
let angle = pos as f32 * freq / freq_factors[i];
let (sin, cos) = angle.sin_cos();
let a = vec[i];
let b = vec[i + half];
vec[i] = a * cos - b * sin;
vec[i + half] = a * sin + b * cos;
}
}
pub fn apply_rope_interleaved(vec: &mut [f32], pos: usize, theta: f32) {
let dim = vec.len();
let half = dim / 2;
let freqs = rope_freqs(theta, dim);
for i in 0..half {
let freq = freqs[i];
let angle = pos as f32 * freq;
let (sin, cos) = angle.sin_cos();
let a = vec[2 * i];
let b = vec[2 * i + 1];
vec[2 * i] = a * cos - b * sin;
vec[2 * i + 1] = a * sin + b * cos;
}
}
pub fn apply_rope_interleaved_with_freq_factors(
vec: &mut [f32],
pos: usize,
theta: f32,
freq_factors: &[f32],
) {
let dim = vec.len();
let half = dim / 2;
assert_eq!(
freq_factors.len(),
half,
"freq_factors must have one entry per rotation band (dim/2)"
);
let freqs = rope_freqs(theta, dim);
for i in 0..half {
let freq = freqs[i];
let angle = pos as f32 * freq / freq_factors[i];
let (sin, cos) = angle.sin_cos();
let a = vec[2 * i];
let b = vec[2 * i + 1];
vec[2 * i] = a * cos - b * sin;
vec[2 * i + 1] = a * sin + b * cos;
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct YarnScaling {
pub factor: f32,
pub beta_fast: f32,
pub beta_slow: f32,
pub orig_max_pos: usize,
pub truncate: bool,
}
impl YarnScaling {
pub fn new(factor: f32, orig_max_pos: usize) -> Self {
YarnScaling {
factor,
beta_fast: 32.0,
beta_slow: 1.0,
orig_max_pos,
truncate: true,
}
}
}
fn yarn_correction_dim(
num_rotations: f64,
rotary_dim: usize,
base: f64,
orig_max_pos: usize,
) -> f64 {
rotary_dim as f64 * (orig_max_pos as f64 / (num_rotations * 2.0 * std::f64::consts::PI)).ln()
/ (2.0 * base.ln())
}
pub fn yarn_correction_range(scaling: YarnScaling, rotary_dim: usize, base: f32) -> (f64, f64) {
let base = base as f64;
let mut low = yarn_correction_dim(
scaling.beta_fast as f64,
rotary_dim,
base,
scaling.orig_max_pos,
);
let mut high = yarn_correction_dim(
scaling.beta_slow as f64,
rotary_dim,
base,
scaling.orig_max_pos,
);
if scaling.truncate {
low = low.floor();
high = high.ceil();
}
low = low.max(0.0);
high = high.min(rotary_dim as f64 - 1.0);
if low == high {
high += 0.001;
}
(low, high)
}
pub fn yarn_freq_factors(scaling: YarnScaling, rotary_dim: usize, base: f32) -> Vec<f32> {
assert!(
rotary_dim > 0 && rotary_dim.is_multiple_of(2),
"rotary_dim must be a positive even number of channels, got {rotary_dim}"
);
assert!(
scaling.factor > 0.0,
"YaRN factor must be positive, got {}",
scaling.factor
);
let (low, high) = yarn_correction_range(scaling, rotary_dim, base);
let factor = scaling.factor as f64;
(0..rotary_dim / 2)
.map(|band| {
let ramp = ((band as f64 - low) / (high - low)).clamp(0.0, 1.0);
let freq_scale = ramp / factor + (1.0 - ramp);
(1.0 / freq_scale) as f32
})
.collect()
}
pub fn proportional_freq_factors(head_size: usize, rotary_dim: usize, base: f32) -> Vec<f32> {
assert!(
rotary_dim > 0 && rotary_dim.is_multiple_of(2) && rotary_dim <= head_size,
"rotary_dim {rotary_dim} must be positive, even, and no wider than head_size {head_size}"
);
let base = base as f64;
(0..rotary_dim / 2)
.map(|band| {
let exponent =
(2 * band) as f64 / head_size as f64 - (2 * band) as f64 / rotary_dim as f64;
base.powf(exponent) as f32
})
.collect()
}
pub fn causal_gqa_attention(
q: &[f32],
k_cache: &[f32],
v_cache: &[f32],
n_heads: usize,
n_kv_heads: usize,
head_dim: usize,
seq_len: usize,
) -> Vec<f32> {
causal_gqa_attention_softcap(
q, k_cache, v_cache, n_heads, n_kv_heads, head_dim, seq_len, None,
)
}
#[allow(clippy::too_many_arguments)]
pub fn causal_gqa_attention_softcap(
q: &[f32],
k_cache: &[f32],
v_cache: &[f32],
n_heads: usize,
n_kv_heads: usize,
head_dim: usize,
seq_len: usize,
attn_softcap: Option<f32>,
) -> Vec<f32> {
assert_eq!(q.len(), n_heads * head_dim);
assert_eq!(k_cache.len(), seq_len * n_kv_heads * head_dim);
assert_eq!(v_cache.len(), seq_len * n_kv_heads * head_dim);
let group_size = n_heads / n_kv_heads.max(1);
let scale = 1.0 / (head_dim as f32).sqrt();
let mut out = vec![0f32; n_heads * head_dim];
for h in 0..n_heads {
let kv_h = h / group_size.max(1);
let q_h = &q[h * head_dim..(h + 1) * head_dim];
let out_h = &mut out[h * head_dim..(h + 1) * head_dim];
online_attn_accumulate(q_h, scale, head_dim, out_h, attn_softcap, None, |visit| {
for t in 0..seq_len {
let k_t = &k_cache
[(t * n_kv_heads + kv_h) * head_dim..(t * n_kv_heads + kv_h + 1) * head_dim];
let v_t = &v_cache
[(t * n_kv_heads + kv_h) * head_dim..(t * n_kv_heads + kv_h + 1) * head_dim];
visit(k_t, v_t);
}
});
}
out
}
#[allow(clippy::too_many_arguments)]
pub fn causal_gqa_attention_windowed(
q: &[f32],
k_cache: &[f32],
v_cache: &[f32],
n_heads: usize,
n_kv_heads: usize,
head_dim: usize,
seq_len: usize,
window: usize,
) -> Vec<f32> {
causal_gqa_attention_windowed_softcap(
q, k_cache, v_cache, n_heads, n_kv_heads, head_dim, seq_len, window, None,
)
}
#[allow(clippy::too_many_arguments)]
pub fn causal_gqa_attention_windowed_softcap(
q: &[f32],
k_cache: &[f32],
v_cache: &[f32],
n_heads: usize,
n_kv_heads: usize,
head_dim: usize,
seq_len: usize,
window: usize,
attn_softcap: Option<f32>,
) -> Vec<f32> {
assert_eq!(q.len(), n_heads * head_dim);
assert_eq!(k_cache.len(), seq_len * n_kv_heads * head_dim);
assert_eq!(v_cache.len(), seq_len * n_kv_heads * head_dim);
assert!(window > 0, "window must be positive");
let group_size = n_heads / n_kv_heads.max(1);
let scale = 1.0 / (head_dim as f32).sqrt();
let mut out = vec![0f32; n_heads * head_dim];
let window_start = seq_len.saturating_sub(window);
for h in 0..n_heads {
let kv_h = h / group_size.max(1);
let q_h = &q[h * head_dim..(h + 1) * head_dim];
let out_h = &mut out[h * head_dim..(h + 1) * head_dim];
online_attn_accumulate(q_h, scale, head_dim, out_h, attn_softcap, None, |visit| {
for t in window_start..seq_len {
let k_t = &k_cache
[(t * n_kv_heads + kv_h) * head_dim..(t * n_kv_heads + kv_h + 1) * head_dim];
let v_t = &v_cache
[(t * n_kv_heads + kv_h) * head_dim..(t * n_kv_heads + kv_h + 1) * head_dim];
visit(k_t, v_t);
}
});
}
out
}
#[allow(clippy::too_many_arguments)]
pub fn causal_gqa_attention_sinks(
q: &[f32],
k_cache: &[f32],
v_cache: &[f32],
n_heads: usize,
n_kv_heads: usize,
head_dim: usize,
seq_len: usize,
window: Option<usize>,
sinks: &[f32],
) -> Vec<f32> {
assert_eq!(q.len(), n_heads * head_dim);
assert_eq!(k_cache.len(), seq_len * n_kv_heads * head_dim);
assert_eq!(v_cache.len(), seq_len * n_kv_heads * head_dim);
assert_eq!(
sinks.len(),
n_heads,
"attention sinks are per query head (llama.cpp `attn_sinks` is {{n_head}})"
);
let group_size = n_heads / n_kv_heads.max(1);
let scale = 1.0 / (head_dim as f32).sqrt();
let mut out = vec![0f32; n_heads * head_dim];
let start = match window {
Some(w) => {
assert!(w > 0, "window must be positive");
seq_len.saturating_sub(w)
}
None => 0,
};
for h in 0..n_heads {
let kv_h = h / group_size.max(1);
let q_h = &q[h * head_dim..(h + 1) * head_dim];
let sink = sinks[h];
let out_h = &mut out[h * head_dim..(h + 1) * head_dim];
online_attn_accumulate(q_h, scale, head_dim, out_h, None, Some(sink), |visit| {
for t in start..seq_len {
let k_t = &k_cache
[(t * n_kv_heads + kv_h) * head_dim..(t * n_kv_heads + kv_h + 1) * head_dim];
let v_t = &v_cache
[(t * n_kv_heads + kv_h) * head_dim..(t * n_kv_heads + kv_h + 1) * head_dim];
visit(k_t, v_t);
}
});
}
out
}
pub fn causal_gqa_attention_prefill(
q: &[f32],
k_cache: &[f32],
v_cache: &[f32],
n_heads: usize,
n_kv_heads: usize,
head_dim: usize,
seq_len: usize,
) -> Vec<f32> {
assert_eq!(q.len(), seq_len * n_heads * head_dim);
assert_eq!(k_cache.len(), seq_len * n_kv_heads * head_dim);
assert_eq!(v_cache.len(), seq_len * n_kv_heads * head_dim);
let q_stride = n_heads * head_dim;
let kv_stride = n_kv_heads * head_dim;
let mut out = vec![0f32; seq_len * q_stride];
for t in 0..seq_len {
let q_t = &q[t * q_stride..(t + 1) * q_stride];
let k_prefix = &k_cache[..(t + 1) * kv_stride];
let v_prefix = &v_cache[..(t + 1) * kv_stride];
let attn = causal_gqa_attention(
q_t,
k_prefix,
v_prefix,
n_heads,
n_kv_heads,
head_dim,
t + 1,
);
out[t * q_stride..(t + 1) * q_stride].copy_from_slice(&attn);
}
out
}
#[allow(clippy::too_many_arguments)]
pub fn causal_gqa_attention_prefill_shared_kv(
q: &[f32],
k_cache: &[f32],
v_cache: &[f32],
n_heads: usize,
n_kv_heads: usize,
head_dim: usize,
n_q: usize,
kv_prefix: usize,
attn_softcap: Option<f32>,
) -> Vec<f32> {
causal_gqa_attention_prefill_shared_kv_windowed(
q,
k_cache,
v_cache,
n_heads,
n_kv_heads,
head_dim,
n_q,
kv_prefix,
attn_softcap,
None,
)
}
#[allow(clippy::too_many_arguments)]
pub fn causal_gqa_attention_prefill_shared_kv_windowed(
q: &[f32],
k_cache: &[f32],
v_cache: &[f32],
n_heads: usize,
n_kv_heads: usize,
head_dim: usize,
n_q: usize,
kv_prefix: usize,
attn_softcap: Option<f32>,
window: Option<usize>,
) -> Vec<f32> {
use rayon::prelude::*;
let q_stride = n_heads * head_dim;
let kv_stride = n_kv_heads * head_dim;
assert_eq!(q.len(), n_q * q_stride);
let kv_len = kv_prefix + n_q;
assert!(k_cache.len() >= kv_len * kv_stride);
assert!(v_cache.len() >= kv_len * kv_stride);
let group_size = n_heads / n_kv_heads.max(1);
let scale = 1.0 / (head_dim as f32).sqrt();
let mut out = vec![0f32; n_q * q_stride];
struct OutPtr(*mut f32);
unsafe impl Send for OutPtr {}
unsafe impl Sync for OutPtr {}
impl OutPtr {
#[inline]
unsafe fn write(&self, off: usize, src: &[f32]) {
std::ptr::copy_nonoverlapping(src.as_ptr(), self.0.add(off), src.len());
}
}
#[derive(Default)]
struct Scratch {
q_tile: Vec<f32>,
scores: Vec<f32>,
acc: Vec<f32>,
}
const Q_BLOCK: usize = 8;
let n_blocks = n_q.div_ceil(Q_BLOCK);
let out_w = OutPtr(out.as_mut_ptr());
let softcap = attn_softcap.filter(|&c| c > 0.0);
(0..n_blocks * n_heads)
.into_par_iter()
.with_min_len(1)
.for_each_init(Scratch::default, |scratch, task| {
let Scratch {
q_tile,
scores,
acc,
} = scratch;
let blk = task / n_heads;
let h = task % n_heads;
let kv_h = h / group_size.max(1);
let b_start = blk * Q_BLOCK;
let b_end = (b_start + Q_BLOCK).min(n_q);
let n_b = b_end - b_start;
let t_hi = kv_prefix + b_end;
let t_lo = match window {
Some(w) => (kv_prefix + b_start + 1).saturating_sub(w),
None => 0,
};
let span = t_hi - t_lo;
let kv_off = t_lo * kv_stride + kv_h * head_dim;
q_tile.clear();
for b in b_start..b_end {
q_tile.extend_from_slice(&q[b * q_stride + h * head_dim..][..head_dim]);
}
scores.resize(n_b * span, 0.0);
qk_tile(
q_tile, n_b, head_dim, k_cache, kv_off, kv_stride, span, scale, scores,
);
let mut norms = [0f32; Q_BLOCK];
for b in b_start..b_end {
let causal_len = kv_prefix + b + 1;
let t_start = match window {
Some(w) => causal_len.saturating_sub(w),
None => 0,
};
let row = &mut scores[(b - b_start) * span..][..span];
let lo = t_start - t_lo;
let hi = causal_len - t_lo;
row[..lo].fill(0.0);
row[hi..].fill(0.0);
let live = &mut row[lo..hi];
if let Some(sc) = softcap {
for s in live.iter_mut() {
*s = sc * (*s / sc).tanh();
}
}
norms[b - b_start] = softmax_row_exp_sum(live);
}
acc.resize(n_b * head_dim, 0.0);
acc.fill(0.0);
pv_tile(scores, n_b, span, v_cache, kv_off, kv_stride, head_dim, acc);
for b in b_start..b_end {
let out_h = &mut acc[(b - b_start) * head_dim..][..head_dim];
let l = norms[b - b_start];
if l > 0.0 {
scale_inplace(out_h, 1.0 / l);
}
unsafe {
out_w.write(b * q_stride + h * head_dim, out_h);
}
}
});
out
}
#[inline]
fn softmax_row_exp_sum(x: &mut [f32]) -> f32 {
if x.is_empty() {
return 0.0;
}
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("neon") {
return unsafe { softmax_row_exp_sum_neon(x) };
}
}
#[cfg(target_arch = "x86_64")]
{
if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
return unsafe { softmax_row_exp_sum_avx2(x) };
}
}
softmax_row_exp_sum_scalar(x)
}
fn softmax_row_exp_sum_scalar(x: &mut [f32]) -> f32 {
let m = x.iter().fold(f32::NEG_INFINITY, |a, &s| a.max(s));
let mut l = 0f32;
for s in x.iter_mut() {
*s = (*s - m).exp();
l += *s;
}
l
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
#[inline]
unsafe fn expf_neon(x: std::arch::aarch64::float32x4_t) -> std::arch::aarch64::float32x4_t {
use std::arch::aarch64::*;
let x = vmaxq_f32(x, vdupq_n_f32(EXP_MIN_ARG));
let r = vdupq_n_f32(EXP_SHIFT);
let z = vfmaq_f32(r, x, vdupq_n_f32(EXP_LOG2E));
let n = vsubq_f32(z, r);
let b = vfmsq_f32(
vfmsq_f32(x, n, vdupq_n_f32(EXP_LN2_HI)),
n,
vdupq_n_f32(EXP_LN2_LO),
);
let e = vshlq_n_u32::<23>(vreinterpretq_u32_f32(z));
let k = vreinterpretq_f32_u32(vaddq_u32(e, vreinterpretq_u32_f32(vdupq_n_f32(1.0))));
let u = vmulq_f32(b, b);
let j = vfmaq_f32(
vmulq_f32(vdupq_n_f32(EXP_C0), b),
vfmaq_f32(
vfmaq_f32(vdupq_n_f32(EXP_C1), vdupq_n_f32(EXP_C2), b),
vfmaq_f32(vdupq_n_f32(EXP_C3), vdupq_n_f32(EXP_C4), b),
u,
),
u,
);
vfmaq_f32(k, j, k)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
#[inline]
unsafe fn expf_avx2(x: std::arch::x86_64::__m256) -> std::arch::x86_64::__m256 {
use std::arch::x86_64::*;
let x = _mm256_max_ps(x, _mm256_set1_ps(EXP_MIN_ARG));
let r = _mm256_set1_ps(EXP_SHIFT);
let z = _mm256_fmadd_ps(x, _mm256_set1_ps(EXP_LOG2E), r);
let n = _mm256_sub_ps(z, r);
let b = _mm256_fnmadd_ps(
n,
_mm256_set1_ps(EXP_LN2_LO),
_mm256_fnmadd_ps(n, _mm256_set1_ps(EXP_LN2_HI), x),
);
let e = _mm256_slli_epi32::<23>(_mm256_castps_si256(z));
let k = _mm256_castsi256_ps(_mm256_add_epi32(
e,
_mm256_castps_si256(_mm256_set1_ps(1.0)),
));
let u = _mm256_mul_ps(b, b);
let j = _mm256_fmadd_ps(
_mm256_fmadd_ps(
_mm256_fmadd_ps(_mm256_set1_ps(EXP_C4), b, _mm256_set1_ps(EXP_C3)),
u,
_mm256_fmadd_ps(_mm256_set1_ps(EXP_C2), b, _mm256_set1_ps(EXP_C1)),
),
u,
_mm256_mul_ps(_mm256_set1_ps(EXP_C0), b),
);
_mm256_fmadd_ps(j, k, k)
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
mod exp_consts {
pub use crate::vexp::*;
pub const EXP_MIN_ARG: f32 = -87.0;
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
use exp_consts::*;
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn softmax_row_exp_sum_neon(x: &mut [f32]) -> f32 {
use std::arch::aarch64::*;
let n = x.len();
let p = x.as_mut_ptr();
let nv = n & !3;
let mut mv = vdupq_n_f32(f32::NEG_INFINITY);
let mut i = 0;
while i < nv {
mv = vmaxq_f32(mv, vld1q_f32(p.add(i)));
i += 4;
}
let mut m = if nv == 0 {
f32::NEG_INFINITY
} else {
vmaxvq_f32(mv)
};
for j in nv..n {
m = m.max(*p.add(j));
}
let mvec = vdupq_n_f32(m);
let mut sv = vdupq_n_f32(0.0);
let mut i = 0;
while i < nv {
let e = expf_neon(vsubq_f32(vld1q_f32(p.add(i)), mvec));
vst1q_f32(p.add(i), e);
sv = vaddq_f32(sv, e);
i += 4;
}
let mut l = vaddvq_f32(sv);
if nv < n {
let mut buf = [0f32; 4];
for (j, slot) in (nv..n).zip(buf.iter_mut()) {
*slot = *p.add(j) - m;
}
vst1q_f32(buf.as_mut_ptr(), expf_neon(vld1q_f32(buf.as_ptr())));
for (j, &e) in (nv..n).zip(buf.iter()) {
*p.add(j) = e;
l += e;
}
}
l
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn softmax_row_exp_sum_avx2(x: &mut [f32]) -> f32 {
use std::arch::x86_64::*;
let n = x.len();
let p = x.as_mut_ptr();
let nv = n & !7;
let mut mv = _mm256_set1_ps(f32::NEG_INFINITY);
let mut i = 0;
while i < nv {
mv = _mm256_max_ps(mv, _mm256_loadu_ps(p.add(i)));
i += 8;
}
let mut m = if nv == 0 {
f32::NEG_INFINITY
} else {
let mut lanes = [0f32; 8];
_mm256_storeu_ps(lanes.as_mut_ptr(), mv);
lanes.iter().fold(f32::NEG_INFINITY, |a, &s| a.max(s))
};
for j in nv..n {
m = m.max(*p.add(j));
}
let mvec = _mm256_set1_ps(m);
let mut sv = _mm256_setzero_ps();
let mut i = 0;
while i < nv {
let e = expf_avx2(_mm256_sub_ps(_mm256_loadu_ps(p.add(i)), mvec));
_mm256_storeu_ps(p.add(i), e);
sv = _mm256_add_ps(sv, e);
i += 8;
}
let mut l = hsum256_ps(sv);
if nv < n {
let mut buf = [0f32; 8];
for (j, slot) in (nv..n).zip(buf.iter_mut()) {
*slot = *p.add(j) - m;
}
_mm256_storeu_ps(buf.as_mut_ptr(), expf_avx2(_mm256_loadu_ps(buf.as_ptr())));
for (j, &e) in (nv..n).zip(buf.iter()) {
*p.add(j) = e;
l += e;
}
}
l
}
#[inline(never)]
#[allow(clippy::too_many_arguments)]
fn qk_tile(
q_tile: &[f32],
n_b: usize,
head_dim: usize,
k: &[f32],
k_off: usize,
k_stride: usize,
span: usize,
scale: f32,
scores: &mut [f32],
) {
debug_assert_eq!(q_tile.len(), n_b * head_dim);
debug_assert_eq!(scores.len(), n_b * span);
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("neon") {
unsafe {
qk_tile_neon(
q_tile, n_b, head_dim, k, k_off, k_stride, span, scale, scores,
)
};
return;
}
}
#[cfg(target_arch = "x86_64")]
{
if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
unsafe {
qk_tile_avx2(
q_tile, n_b, head_dim, k, k_off, k_stride, span, scale, scores,
)
};
return;
}
}
qk_rows(
q_tile,
head_dim,
k,
k_off,
k_stride,
span,
scale,
scores,
0..n_b,
0..span,
);
}
#[allow(clippy::too_many_arguments)]
fn qk_rows(
q_tile: &[f32],
head_dim: usize,
k: &[f32],
k_off: usize,
k_stride: usize,
span: usize,
scale: f32,
scores: &mut [f32],
rows: std::ops::Range<usize>,
cols: std::ops::Range<usize>,
) {
for b in rows {
let q_b = &q_tile[b * head_dim..][..head_dim];
for t in cols.clone() {
let k_t = &k[k_off + t * k_stride..][..head_dim];
scores[b * span + t] = dot_f32(q_b, k_t) * scale;
}
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
#[allow(clippy::too_many_arguments)]
unsafe fn qk_tile_neon(
q_tile: &[f32],
n_b: usize,
head_dim: usize,
k: &[f32],
k_off: usize,
k_stride: usize,
span: usize,
scale: f32,
scores: &mut [f32],
) {
use std::arch::aarch64::*;
let qp = q_tile.as_ptr();
let kp = k.as_ptr().add(k_off);
let sp = scores.as_mut_ptr();
let bt = n_b & !3;
let tt = span & !3;
let dv = head_dim & !3;
let mut t0 = 0;
while t0 < tt {
let k0 = kp.add(t0 * k_stride);
let k1 = k0.add(k_stride);
let k2 = k1.add(k_stride);
let k3 = k2.add(k_stride);
let mut b0 = 0;
while b0 < bt {
let a0 = qp.add(b0 * head_dim);
let a1 = a0.add(head_dim);
let a2 = a1.add(head_dim);
let a3 = a2.add(head_dim);
let z = vdupq_n_f32(0.0);
let (mut c00, mut c01, mut c02, mut c03) = (z, z, z, z);
let (mut c10, mut c11, mut c12, mut c13) = (z, z, z, z);
let (mut c20, mut c21, mut c22, mut c23) = (z, z, z, z);
let (mut c30, mut c31, mut c32, mut c33) = (z, z, z, z);
let mut d = 0;
while d < dv {
let av0 = vld1q_f32(a0.add(d));
let av1 = vld1q_f32(a1.add(d));
let av2 = vld1q_f32(a2.add(d));
let av3 = vld1q_f32(a3.add(d));
let kv0 = vld1q_f32(k0.add(d));
c00 = vfmaq_f32(c00, av0, kv0);
c10 = vfmaq_f32(c10, av1, kv0);
c20 = vfmaq_f32(c20, av2, kv0);
c30 = vfmaq_f32(c30, av3, kv0);
let kv1 = vld1q_f32(k1.add(d));
c01 = vfmaq_f32(c01, av0, kv1);
c11 = vfmaq_f32(c11, av1, kv1);
c21 = vfmaq_f32(c21, av2, kv1);
c31 = vfmaq_f32(c31, av3, kv1);
let kv2 = vld1q_f32(k2.add(d));
c02 = vfmaq_f32(c02, av0, kv2);
c12 = vfmaq_f32(c12, av1, kv2);
c22 = vfmaq_f32(c22, av2, kv2);
c32 = vfmaq_f32(c32, av3, kv2);
let kv3 = vld1q_f32(k3.add(d));
c03 = vfmaq_f32(c03, av0, kv3);
c13 = vfmaq_f32(c13, av1, kv3);
c23 = vfmaq_f32(c23, av2, kv3);
c33 = vfmaq_f32(c33, av3, kv3);
d += 4;
}
let mut r = [
[
vaddvq_f32(c00),
vaddvq_f32(c01),
vaddvq_f32(c02),
vaddvq_f32(c03),
],
[
vaddvq_f32(c10),
vaddvq_f32(c11),
vaddvq_f32(c12),
vaddvq_f32(c13),
],
[
vaddvq_f32(c20),
vaddvq_f32(c21),
vaddvq_f32(c22),
vaddvq_f32(c23),
],
[
vaddvq_f32(c30),
vaddvq_f32(c31),
vaddvq_f32(c32),
vaddvq_f32(c33),
],
];
let arow = [a0, a1, a2, a3];
let krow = [k0, k1, k2, k3];
for d in dv..head_dim {
for (i, ai) in arow.iter().enumerate() {
let av = *ai.add(d);
for (j, kj) in krow.iter().enumerate() {
r[i][j] += av * *kj.add(d);
}
}
}
for (i, ri) in r.iter().enumerate() {
for (j, v) in ri.iter().enumerate() {
*sp.add((b0 + i) * span + t0 + j) = v * scale;
}
}
b0 += 4;
}
t0 += 4;
}
qk_rows(
q_tile,
head_dim,
k,
k_off,
k_stride,
span,
scale,
scores,
0..bt,
tt..span,
);
qk_rows(
q_tile,
head_dim,
k,
k_off,
k_stride,
span,
scale,
scores,
bt..n_b,
0..span,
);
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
#[allow(clippy::too_many_arguments)]
unsafe fn qk_tile_avx2(
q_tile: &[f32],
n_b: usize,
head_dim: usize,
k: &[f32],
k_off: usize,
k_stride: usize,
span: usize,
scale: f32,
scores: &mut [f32],
) {
use std::arch::x86_64::*;
let qp = q_tile.as_ptr();
let kp = k.as_ptr().add(k_off);
let sp = scores.as_mut_ptr();
let bt = n_b & !3;
let tt = span & !1;
let dv = head_dim & !7;
let mut t0 = 0;
while t0 < tt {
let k0 = kp.add(t0 * k_stride);
let k1 = k0.add(k_stride);
let mut b0 = 0;
while b0 < bt {
let a0 = qp.add(b0 * head_dim);
let a1 = a0.add(head_dim);
let a2 = a1.add(head_dim);
let a3 = a2.add(head_dim);
let z = _mm256_setzero_ps();
let (mut c00, mut c01) = (z, z);
let (mut c10, mut c11) = (z, z);
let (mut c20, mut c21) = (z, z);
let (mut c30, mut c31) = (z, z);
let mut d = 0;
while d < dv {
let av0 = _mm256_loadu_ps(a0.add(d));
let av1 = _mm256_loadu_ps(a1.add(d));
let av2 = _mm256_loadu_ps(a2.add(d));
let av3 = _mm256_loadu_ps(a3.add(d));
let kv0 = _mm256_loadu_ps(k0.add(d));
c00 = _mm256_fmadd_ps(av0, kv0, c00);
c10 = _mm256_fmadd_ps(av1, kv0, c10);
c20 = _mm256_fmadd_ps(av2, kv0, c20);
c30 = _mm256_fmadd_ps(av3, kv0, c30);
let kv1 = _mm256_loadu_ps(k1.add(d));
c01 = _mm256_fmadd_ps(av0, kv1, c01);
c11 = _mm256_fmadd_ps(av1, kv1, c11);
c21 = _mm256_fmadd_ps(av2, kv1, c21);
c31 = _mm256_fmadd_ps(av3, kv1, c31);
d += 8;
}
let mut r = [
[hsum256_ps(c00), hsum256_ps(c01)],
[hsum256_ps(c10), hsum256_ps(c11)],
[hsum256_ps(c20), hsum256_ps(c21)],
[hsum256_ps(c30), hsum256_ps(c31)],
];
let arow = [a0, a1, a2, a3];
let krow = [k0, k1];
for d in dv..head_dim {
for (i, ai) in arow.iter().enumerate() {
let av = *ai.add(d);
for (j, kj) in krow.iter().enumerate() {
r[i][j] += av * *kj.add(d);
}
}
}
for (i, ri) in r.iter().enumerate() {
for (j, v) in ri.iter().enumerate() {
*sp.add((b0 + i) * span + t0 + j) = v * scale;
}
}
b0 += 4;
}
t0 += 2;
}
qk_rows(
q_tile,
head_dim,
k,
k_off,
k_stride,
span,
scale,
scores,
0..bt,
tt..span,
);
qk_rows(
q_tile,
head_dim,
k,
k_off,
k_stride,
span,
scale,
scores,
bt..n_b,
0..span,
);
}
#[inline(never)]
#[allow(clippy::too_many_arguments)]
fn pv_tile(
p: &[f32],
n_b: usize,
span: usize,
v: &[f32],
v_off: usize,
v_stride: usize,
head_dim: usize,
acc: &mut [f32],
) {
debug_assert_eq!(p.len(), n_b * span);
debug_assert_eq!(acc.len(), n_b * head_dim);
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("neon") {
unsafe { pv_tile_neon(p, n_b, span, v, v_off, v_stride, head_dim, acc) };
return;
}
}
#[cfg(target_arch = "x86_64")]
{
if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
unsafe { pv_tile_avx2(p, n_b, span, v, v_off, v_stride, head_dim, acc) };
return;
}
}
pv_rows(p, span, v, v_off, v_stride, head_dim, acc, 0..n_b);
}
#[allow(clippy::too_many_arguments)]
fn pv_rows(
p: &[f32],
span: usize,
v: &[f32],
v_off: usize,
v_stride: usize,
head_dim: usize,
acc: &mut [f32],
rows: std::ops::Range<usize>,
) {
for b in rows {
let out_b = &mut acc[b * head_dim..][..head_dim];
for t in 0..span {
let w = p[b * span + t];
if w == 0.0 {
continue;
}
axpy(out_b, &v[v_off + t * v_stride..][..head_dim], w);
}
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
#[allow(clippy::too_many_arguments)]
unsafe fn pv_tile_neon(
p: &[f32],
n_b: usize,
span: usize,
v: &[f32],
v_off: usize,
v_stride: usize,
head_dim: usize,
acc: &mut [f32],
) {
use std::arch::aarch64::*;
let vp = v.as_ptr().add(v_off);
let pp = p.as_ptr();
let ap = acc.as_mut_ptr();
let bt = n_b & !7;
let dv = head_dim & !7;
let dv4 = head_dim & !3;
let mut b0 = 0;
while b0 < bt {
let mut d0 = 0;
while d0 < dv {
let mut c0l = vld1q_f32(ap.add(b0 * head_dim + d0));
let mut c0h = vld1q_f32(ap.add(b0 * head_dim + d0 + 4));
let mut c1l = vld1q_f32(ap.add((b0 + 1) * head_dim + d0));
let mut c1h = vld1q_f32(ap.add((b0 + 1) * head_dim + d0 + 4));
let mut c2l = vld1q_f32(ap.add((b0 + 2) * head_dim + d0));
let mut c2h = vld1q_f32(ap.add((b0 + 2) * head_dim + d0 + 4));
let mut c3l = vld1q_f32(ap.add((b0 + 3) * head_dim + d0));
let mut c3h = vld1q_f32(ap.add((b0 + 3) * head_dim + d0 + 4));
let mut c4l = vld1q_f32(ap.add((b0 + 4) * head_dim + d0));
let mut c4h = vld1q_f32(ap.add((b0 + 4) * head_dim + d0 + 4));
let mut c5l = vld1q_f32(ap.add((b0 + 5) * head_dim + d0));
let mut c5h = vld1q_f32(ap.add((b0 + 5) * head_dim + d0 + 4));
let mut c6l = vld1q_f32(ap.add((b0 + 6) * head_dim + d0));
let mut c6h = vld1q_f32(ap.add((b0 + 6) * head_dim + d0 + 4));
let mut c7l = vld1q_f32(ap.add((b0 + 7) * head_dim + d0));
let mut c7h = vld1q_f32(ap.add((b0 + 7) * head_dim + d0 + 4));
for t in 0..span {
let vr = vp.add(t * v_stride + d0);
let v0 = vld1q_f32(vr);
let v1 = vld1q_f32(vr.add(4));
let s0 = vdupq_n_f32(*pp.add(b0 * span + t));
c0l = vfmaq_f32(c0l, v0, s0);
c0h = vfmaq_f32(c0h, v1, s0);
let s1 = vdupq_n_f32(*pp.add((b0 + 1) * span + t));
c1l = vfmaq_f32(c1l, v0, s1);
c1h = vfmaq_f32(c1h, v1, s1);
let s2 = vdupq_n_f32(*pp.add((b0 + 2) * span + t));
c2l = vfmaq_f32(c2l, v0, s2);
c2h = vfmaq_f32(c2h, v1, s2);
let s3 = vdupq_n_f32(*pp.add((b0 + 3) * span + t));
c3l = vfmaq_f32(c3l, v0, s3);
c3h = vfmaq_f32(c3h, v1, s3);
let s4 = vdupq_n_f32(*pp.add((b0 + 4) * span + t));
c4l = vfmaq_f32(c4l, v0, s4);
c4h = vfmaq_f32(c4h, v1, s4);
let s5 = vdupq_n_f32(*pp.add((b0 + 5) * span + t));
c5l = vfmaq_f32(c5l, v0, s5);
c5h = vfmaq_f32(c5h, v1, s5);
let s6 = vdupq_n_f32(*pp.add((b0 + 6) * span + t));
c6l = vfmaq_f32(c6l, v0, s6);
c6h = vfmaq_f32(c6h, v1, s6);
let s7 = vdupq_n_f32(*pp.add((b0 + 7) * span + t));
c7l = vfmaq_f32(c7l, v0, s7);
c7h = vfmaq_f32(c7h, v1, s7);
}
vst1q_f32(ap.add(b0 * head_dim + d0), c0l);
vst1q_f32(ap.add(b0 * head_dim + d0 + 4), c0h);
vst1q_f32(ap.add((b0 + 1) * head_dim + d0), c1l);
vst1q_f32(ap.add((b0 + 1) * head_dim + d0 + 4), c1h);
vst1q_f32(ap.add((b0 + 2) * head_dim + d0), c2l);
vst1q_f32(ap.add((b0 + 2) * head_dim + d0 + 4), c2h);
vst1q_f32(ap.add((b0 + 3) * head_dim + d0), c3l);
vst1q_f32(ap.add((b0 + 3) * head_dim + d0 + 4), c3h);
vst1q_f32(ap.add((b0 + 4) * head_dim + d0), c4l);
vst1q_f32(ap.add((b0 + 4) * head_dim + d0 + 4), c4h);
vst1q_f32(ap.add((b0 + 5) * head_dim + d0), c5l);
vst1q_f32(ap.add((b0 + 5) * head_dim + d0 + 4), c5h);
vst1q_f32(ap.add((b0 + 6) * head_dim + d0), c6l);
vst1q_f32(ap.add((b0 + 6) * head_dim + d0 + 4), c6h);
vst1q_f32(ap.add((b0 + 7) * head_dim + d0), c7l);
vst1q_f32(ap.add((b0 + 7) * head_dim + d0 + 4), c7h);
d0 += 8;
}
if dv < head_dim {
for t in 0..span {
for i in 0..8 {
let w = *pp.add((b0 + i) * span + t);
if w == 0.0 {
continue;
}
let row = ap.add((b0 + i) * head_dim);
for d in dv..dv4 {
*row.add(d) = f32::mul_add(w, *vp.add(t * v_stride + d), *row.add(d));
}
for d in dv4..head_dim {
*row.add(d) += w * *vp.add(t * v_stride + d);
}
}
}
}
b0 += 8;
}
pv_rows(p, span, v, v_off, v_stride, head_dim, acc, bt..n_b);
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
#[allow(clippy::too_many_arguments)]
unsafe fn pv_tile_avx2(
p: &[f32],
n_b: usize,
span: usize,
v: &[f32],
v_off: usize,
v_stride: usize,
head_dim: usize,
acc: &mut [f32],
) {
use std::arch::x86_64::*;
let vp = v.as_ptr().add(v_off);
let pp = p.as_ptr();
let ap = acc.as_mut_ptr();
let bt = n_b & !7;
let dv = head_dim & !7;
let mut b0 = 0;
while b0 < bt {
let mut d0 = 0;
while d0 < dv {
let mut c0 = _mm256_loadu_ps(ap.add(b0 * head_dim + d0));
let mut c1 = _mm256_loadu_ps(ap.add((b0 + 1) * head_dim + d0));
let mut c2 = _mm256_loadu_ps(ap.add((b0 + 2) * head_dim + d0));
let mut c3 = _mm256_loadu_ps(ap.add((b0 + 3) * head_dim + d0));
let mut c4 = _mm256_loadu_ps(ap.add((b0 + 4) * head_dim + d0));
let mut c5 = _mm256_loadu_ps(ap.add((b0 + 5) * head_dim + d0));
let mut c6 = _mm256_loadu_ps(ap.add((b0 + 6) * head_dim + d0));
let mut c7 = _mm256_loadu_ps(ap.add((b0 + 7) * head_dim + d0));
for t in 0..span {
let vv = _mm256_loadu_ps(vp.add(t * v_stride + d0));
c0 = _mm256_fmadd_ps(vv, _mm256_set1_ps(*pp.add(b0 * span + t)), c0);
c1 = _mm256_fmadd_ps(vv, _mm256_set1_ps(*pp.add((b0 + 1) * span + t)), c1);
c2 = _mm256_fmadd_ps(vv, _mm256_set1_ps(*pp.add((b0 + 2) * span + t)), c2);
c3 = _mm256_fmadd_ps(vv, _mm256_set1_ps(*pp.add((b0 + 3) * span + t)), c3);
c4 = _mm256_fmadd_ps(vv, _mm256_set1_ps(*pp.add((b0 + 4) * span + t)), c4);
c5 = _mm256_fmadd_ps(vv, _mm256_set1_ps(*pp.add((b0 + 5) * span + t)), c5);
c6 = _mm256_fmadd_ps(vv, _mm256_set1_ps(*pp.add((b0 + 6) * span + t)), c6);
c7 = _mm256_fmadd_ps(vv, _mm256_set1_ps(*pp.add((b0 + 7) * span + t)), c7);
}
_mm256_storeu_ps(ap.add(b0 * head_dim + d0), c0);
_mm256_storeu_ps(ap.add((b0 + 1) * head_dim + d0), c1);
_mm256_storeu_ps(ap.add((b0 + 2) * head_dim + d0), c2);
_mm256_storeu_ps(ap.add((b0 + 3) * head_dim + d0), c3);
_mm256_storeu_ps(ap.add((b0 + 4) * head_dim + d0), c4);
_mm256_storeu_ps(ap.add((b0 + 5) * head_dim + d0), c5);
_mm256_storeu_ps(ap.add((b0 + 6) * head_dim + d0), c6);
_mm256_storeu_ps(ap.add((b0 + 7) * head_dim + d0), c7);
d0 += 8;
}
if dv < head_dim {
for t in 0..span {
for i in 0..8 {
let w = *pp.add((b0 + i) * span + t);
if w == 0.0 {
continue;
}
let row = ap.add((b0 + i) * head_dim);
for d in dv..head_dim {
*row.add(d) += w * *vp.add(t * v_stride + d);
}
}
}
}
b0 += 8;
}
pv_rows(p, span, v, v_off, v_stride, head_dim, acc, bt..n_b);
}
pub fn causal_gqa_attention_paged(
q: &[f32],
store: &PagedKvStore,
block_table: &[usize],
n_heads: usize,
n_kv_heads: usize,
head_dim: usize,
seq_len: usize,
) -> Vec<f32> {
assert_eq!(q.len(), n_heads * head_dim);
let block_size = store.block_size();
assert!(
block_table.len() * block_size >= seq_len,
"block table too short for seq_len"
);
let group_size = n_heads / n_kv_heads.max(1);
let scale = 1.0 / (head_dim as f32).sqrt();
let mut out = vec![0f32; n_heads * head_dim];
for h in 0..n_heads {
let kv_h = h / group_size.max(1);
let q_h = &q[h * head_dim..(h + 1) * head_dim];
let out_h = &mut out[h * head_dim..(h + 1) * head_dim];
online_attn_accumulate(q_h, scale, head_dim, out_h, None, None, |visit| {
for t in 0..seq_len {
let block_id = block_table[t / block_size];
let offset = t % block_size;
let k_row = store.k_row(block_id, offset);
let v_row = store.v_row(block_id, offset);
let k_t = &k_row[kv_h * head_dim..(kv_h + 1) * head_dim];
let v_t = &v_row[kv_h * head_dim..(kv_h + 1) * head_dim];
visit(k_t, v_t);
}
});
}
out
}
#[allow(clippy::too_many_arguments)]
pub fn causal_gqa_attention_paged_sinks(
q: &[f32],
store: &PagedKvStore,
block_table: &[usize],
n_heads: usize,
n_kv_heads: usize,
head_dim: usize,
seq_len: usize,
window: Option<usize>,
sinks: Option<&[f32]>,
attn_softcap: Option<f32>,
) -> Vec<f32> {
assert_eq!(q.len(), n_heads * head_dim);
let block_size = store.block_size();
assert!(
block_table.len() * block_size >= seq_len,
"block table too short for seq_len"
);
if let Some(sinks) = sinks {
assert_eq!(
sinks.len(),
n_heads,
"attention sinks are per query head (llama.cpp `attn_sinks` is {{n_head}})"
);
}
let group_size = n_heads / n_kv_heads.max(1);
let scale = 1.0 / (head_dim as f32).sqrt();
let mut out = vec![0f32; n_heads * head_dim];
let start = match window {
Some(w) => {
assert!(w > 0, "window must be positive");
seq_len.saturating_sub(w)
}
None => 0,
};
for h in 0..n_heads {
let kv_h = h / group_size.max(1);
let q_h = &q[h * head_dim..(h + 1) * head_dim];
let sink = sinks.map(|s| s[h]);
let out_h = &mut out[h * head_dim..(h + 1) * head_dim];
online_attn_accumulate(q_h, scale, head_dim, out_h, attn_softcap, sink, |visit| {
for t in start..seq_len {
let block_id = block_table[t / block_size];
let offset = t % block_size;
let k_row = store.k_row(block_id, offset);
let v_row = store.v_row(block_id, offset);
let k_t = &k_row[kv_h * head_dim..(kv_h + 1) * head_dim];
let v_t = &v_row[kv_h * head_dim..(kv_h + 1) * head_dim];
visit(k_t, v_t);
}
});
}
out
}
pub fn causal_mla_attention(
q: &[f32],
k_cache: &[f32],
v_cache: &[f32],
n_heads: usize,
qk_head_dim: usize,
v_head_dim: usize,
seq_len: usize,
) -> Vec<f32> {
mla_attention_inner(
q,
k_cache,
v_cache,
n_heads,
qk_head_dim,
v_head_dim,
seq_len,
None,
None,
)
}
#[allow(clippy::too_many_arguments)]
fn mla_attention_inner(
q: &[f32],
k_cache: &[f32],
v_cache: &[f32],
n_heads: usize,
qk_head_dim: usize,
v_head_dim: usize,
seq_len: usize,
visible: Option<&[usize]>,
sinks: Option<&[f32]>,
) -> Vec<f32> {
assert_eq!(q.len(), n_heads * qk_head_dim);
assert_eq!(k_cache.len(), seq_len * n_heads * qk_head_dim);
assert_eq!(v_cache.len(), seq_len * n_heads * v_head_dim);
if let Some(visible) = visible {
assert!(
visible.iter().all(|&t| t < seq_len),
"visible positions must be within seq_len"
);
}
if let Some(sinks) = sinks {
assert_eq!(
sinks.len(),
n_heads,
"one sink logit per query head, or none at all"
);
}
let n_positions = visible.map_or(seq_len, |v| v.len());
let position_at = |i: usize| visible.map_or(i, |v| v[i]);
let scale = 1.0 / (qk_head_dim as f32).sqrt();
let mut out = vec![0f32; n_heads * v_head_dim];
for h in 0..n_heads {
let q_h = &q[h * qk_head_dim..(h + 1) * qk_head_dim];
let mut scores = vec![0f32; n_positions];
for (i, score) in scores.iter_mut().enumerate() {
let t = position_at(i);
let k_t =
&k_cache[(t * n_heads + h) * qk_head_dim..(t * n_heads + h + 1) * qk_head_dim];
let mut dot = 0f32;
for d in 0..qk_head_dim {
dot += q_h[d] * k_t[d];
}
*score = dot * scale;
}
let sink = sinks.map(|s| s[h]);
let mut max = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
if let Some(s) = sink {
max = max.max(s);
}
let mut sum = 0f32;
for s in scores.iter_mut() {
*s = (*s - max).exp();
sum += *s;
}
if let Some(s) = sink {
sum += (s - max).exp();
}
if sum > 0.0 {
for s in scores.iter_mut() {
*s /= sum;
}
}
let out_h = &mut out[h * v_head_dim..(h + 1) * v_head_dim];
for (i, &w) in scores.iter().enumerate() {
let t = position_at(i);
let v_t = &v_cache[(t * n_heads + h) * v_head_dim..(t * n_heads + h + 1) * v_head_dim];
for d in 0..v_head_dim {
out_h[d] += w * v_t[d];
}
}
}
out
}
#[allow(clippy::too_many_arguments)]
pub fn causal_mla_attention_sinks(
q: &[f32],
k_cache: &[f32],
v_cache: &[f32],
n_heads: usize,
qk_head_dim: usize,
v_head_dim: usize,
seq_len: usize,
sinks: Option<&[f32]>,
) -> Vec<f32> {
mla_attention_inner(
q,
k_cache,
v_cache,
n_heads,
qk_head_dim,
v_head_dim,
seq_len,
None,
sinks,
)
}
#[allow(clippy::too_many_arguments)]
pub fn causal_mla_attention_sparse_sinks(
q: &[f32],
k_cache: &[f32],
v_cache: &[f32],
n_heads: usize,
qk_head_dim: usize,
v_head_dim: usize,
seq_len: usize,
visible: &[usize],
sinks: Option<&[f32]>,
) -> Vec<f32> {
mla_attention_inner(
q,
k_cache,
v_cache,
n_heads,
qk_head_dim,
v_head_dim,
seq_len,
Some(visible),
sinks,
)
}
pub fn lightning_indexer_topk(
indexer_q: &[Vec<f32>],
indexer_keys: &[Vec<f32>],
indexer_weights: &[f32],
top_k: usize,
) -> Vec<usize> {
let n_heads = indexer_q.len();
assert_eq!(indexer_weights.len(), n_heads);
let index_head_dim = indexer_q.first().map_or(0, |q| q.len());
let scale = 1.0 / ((index_head_dim * n_heads) as f32).sqrt();
let mut scored: Vec<(usize, f32)> = indexer_keys
.iter()
.enumerate()
.map(|(j, k)| {
let score: f32 = indexer_q
.iter()
.zip(indexer_weights.iter())
.map(|(q, w)| {
let dot: f32 = q.iter().zip(k.iter()).map(|(a, b)| a * b).sum();
dot.max(0.0) * w * scale
})
.sum();
(j, score)
})
.collect();
let keep = top_k.min(scored.len());
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
let mut kept: Vec<usize> = scored.into_iter().take(keep).map(|(j, _)| j).collect();
kept.sort_unstable();
kept
}
#[allow(clippy::too_many_arguments)]
pub fn causal_mla_attention_sparse(
q: &[f32],
k_cache: &[f32],
v_cache: &[f32],
n_heads: usize,
qk_head_dim: usize,
v_head_dim: usize,
seq_len: usize,
visible: &[usize],
) -> Vec<f32> {
mla_attention_inner(
q,
k_cache,
v_cache,
n_heads,
qk_head_dim,
v_head_dim,
seq_len,
Some(visible),
None,
)
}
#[cfg(test)]
mod tests {
#[test]
fn an_mla_sink_removes_weight_from_the_real_keys_instead_of_moving_it() {
let (n_heads, qk, vd, seq) = (2, 2, 2, 2);
let q = vec![1.0, 0.0, 0.0, 1.0];
let k = vec![1.0, 0.0, 0.0, 1.0, 1.0, 0.0, 0.0, 1.0];
let v = vec![4.0, 8.0, 1.0, 2.0, 4.0, 8.0, 1.0, 2.0];
let plain = super::causal_mla_attention(&q, &k, &v, n_heads, qk, vd, seq);
let none = super::causal_mla_attention_sinks(&q, &k, &v, n_heads, qk, vd, seq, None);
assert_eq!(plain, none, "no sink must be exactly the old path");
let big = super::causal_mla_attention_sinks(
&q,
&k,
&v,
n_heads,
qk,
vd,
seq,
Some(&[40.0, 40.0]),
);
for (b, p) in big.iter().zip(plain.iter()) {
assert!(b.abs() < 1e-6, "a dominant sink leaves ~0, got {b} vs {p}");
}
let tiny = super::causal_mla_attention_sinks(
&q,
&k,
&v,
n_heads,
qk,
vd,
seq,
Some(&[-40.0, -40.0]),
);
for (s, p) in tiny.iter().zip(plain.iter()) {
assert!((s - p).abs() < 1e-5, "negligible sink: {s} vs {p}");
}
}
#[test]
fn each_head_gets_its_own_mla_sink() {
let (n_heads, qk, vd, seq) = (2, 1, 1, 1);
let q = vec![1.0, 1.0];
let k = vec![1.0, 1.0];
let v = vec![5.0, 5.0];
let out = super::causal_mla_attention_sinks(
&q,
&k,
&v,
n_heads,
qk,
vd,
seq,
Some(&[40.0, -40.0]),
);
assert!(out[0].abs() < 1e-6, "head 0 declined: {}", out[0]);
assert!(
(out[1] - 5.0).abs() < 1e-5,
"head 1 attended normally: {}",
out[1]
);
}
#[test]
fn the_mla_sink_logit_is_not_scaled_by_the_head_width() {
let sink = 0.0f32;
let mut outs = Vec::new();
for qk in [1usize, 4, 16] {
let q = vec![0.0; qk];
let k = vec![0.0; qk];
let v = vec![10.0];
outs.push(super::causal_mla_attention_sinks(&q, &k, &v, 1, qk, 1, 1, Some(&[sink]))[0]);
}
for o in &outs {
assert!((o - 5.0).abs() < 1e-5, "expected 5.0, got {o}");
}
}
#[test]
fn the_sparse_mla_path_honours_a_sink_over_the_selected_positions() {
let (n_heads, qk, vd, seq) = (1, 1, 1, 3);
let q = vec![1.0];
let k = vec![1.0, 1.0, 1.0];
let v = vec![2.0, 4.0, 6.0];
let visible = [0usize, 2];
let plain = super::causal_mla_attention_sparse(&q, &k, &v, n_heads, qk, vd, seq, &visible);
let none = super::causal_mla_attention_sparse_sinks(
&q, &k, &v, n_heads, qk, vd, seq, &visible, None,
);
assert_eq!(plain, none);
assert!((plain[0] - 4.0).abs() < 1e-5, "mean of 2 and 6");
let sunk = super::causal_mla_attention_sparse_sinks(
&q,
&k,
&v,
n_heads,
qk,
vd,
seq,
&visible,
Some(&[40.0]),
);
assert!(sunk[0].abs() < 1e-6, "a dominant sink leaves ~0");
}
#[test]
fn prefill_shared_kv_matches_per_query_reference() {
let n_heads = 6;
let n_kv_heads = 2;
let head_dim = 16;
let n_q = 19;
let kv_prefix = 5;
let kv_len = kv_prefix + n_q;
let q_stride = n_heads * head_dim;
let kv_stride = n_kv_heads * head_dim;
let q: Vec<f32> = (0..n_q * q_stride)
.map(|i| ((i as f32) * 0.013 - 0.7).sin() * 1.3)
.collect();
let k_cache: Vec<f32> = (0..kv_len * kv_stride)
.map(|i| ((i as f32) * 0.017 - 0.3).cos() * 1.1)
.collect();
let v_cache: Vec<f32> = (0..kv_len * kv_stride)
.map(|i| ((i as f32) * 0.011 + 0.2).sin() * 0.9)
.collect();
for softcap in [None, Some(30.0)] {
let got = super::causal_gqa_attention_prefill_shared_kv(
&q, &k_cache, &v_cache, n_heads, n_kv_heads, head_dim, n_q, kv_prefix, softcap,
);
assert_eq!(got.len(), n_q * q_stride);
for b in 0..n_q {
let causal_len = kv_prefix + b + 1;
let want = super::causal_gqa_attention_softcap(
&q[b * q_stride..(b + 1) * q_stride],
&k_cache[..causal_len * kv_stride],
&v_cache[..causal_len * kv_stride],
n_heads,
n_kv_heads,
head_dim,
causal_len,
softcap,
);
for (i, (g, w)) in got[b * q_stride..(b + 1) * q_stride]
.iter()
.zip(want.iter())
.enumerate()
{
assert!(
(g - w).abs() < 1e-5,
"softcap {softcap:?} query {b} slot {i}: blocked {g} vs online {w}"
);
}
}
}
}
#[test]
fn windowed_prefill_shared_kv_matches_the_per_query_windowed_reference() {
let n_heads = 6;
let n_kv_heads = 2;
let head_dim = 16;
let n_q = 19;
let kv_prefix = 5;
let kv_len = kv_prefix + n_q;
let q_stride = n_heads * head_dim;
let kv_stride = n_kv_heads * head_dim;
let q: Vec<f32> = (0..n_q * q_stride)
.map(|i| ((i as f32) * 0.013 - 0.7).sin() * 1.3)
.collect();
let k_cache: Vec<f32> = (0..kv_len * kv_stride)
.map(|i| ((i as f32) * 0.017 - 0.3).cos() * 1.1)
.collect();
let v_cache: Vec<f32> = (0..kv_len * kv_stride)
.map(|i| ((i as f32) * 0.011 + 0.2).sin() * 0.9)
.collect();
for window in [1usize, 3, 7, kv_prefix, kv_len, kv_len + 8] {
for softcap in [None, Some(30.0)] {
let got = super::causal_gqa_attention_prefill_shared_kv_windowed(
&q,
&k_cache,
&v_cache,
n_heads,
n_kv_heads,
head_dim,
n_q,
kv_prefix,
softcap,
Some(window),
);
assert_eq!(got.len(), n_q * q_stride);
for b in 0..n_q {
let causal_len = kv_prefix + b + 1;
let want = super::causal_gqa_attention_windowed_softcap(
&q[b * q_stride..(b + 1) * q_stride],
&k_cache[..causal_len * kv_stride],
&v_cache[..causal_len * kv_stride],
n_heads,
n_kv_heads,
head_dim,
causal_len,
window,
softcap,
);
for (i, (g, w)) in got[b * q_stride..(b + 1) * q_stride]
.iter()
.zip(want.iter())
.enumerate()
{
assert!(
(g - w).abs() < 1e-5,
"window {window} softcap {softcap:?} query {b} slot {i}: \
blocked {g} vs per-query {w}"
);
}
}
}
}
let windowed = super::causal_gqa_attention_prefill_shared_kv_windowed(
&q,
&k_cache,
&v_cache,
n_heads,
n_kv_heads,
head_dim,
n_q,
kv_prefix,
None,
Some(kv_len + 8),
);
let full = super::causal_gqa_attention_prefill_shared_kv(
&q, &k_cache, &v_cache, n_heads, n_kv_heads, head_dim, n_q, kv_prefix, None,
);
assert_eq!(windowed, full);
}
#[allow(clippy::too_many_arguments)]
fn prefill_query_outer_reference(
q: &[f32],
k_cache: &[f32],
v_cache: &[f32],
n_heads: usize,
n_kv_heads: usize,
head_dim: usize,
n_q: usize,
kv_prefix: usize,
attn_softcap: Option<f32>,
window: Option<usize>,
) -> Vec<f32> {
let q_stride = n_heads * head_dim;
let group_size = n_heads / n_kv_heads.max(1);
let scale = 1.0 / (head_dim as f32).sqrt();
let softcap = attn_softcap.filter(|&c| c > 0.0);
let mut out = vec![0f32; n_q * q_stride];
let mut acc = vec![0f32; head_dim];
for h in 0..n_heads {
let kv_h = h / group_size.max(1);
for b in 0..n_q {
let causal_len = kv_prefix + b + 1;
let t_start = match window {
Some(w) => causal_len.saturating_sub(w),
None => 0,
};
let q_h = &q[b * q_stride + h * head_dim..][..head_dim];
let mut scores = vec![0f32; causal_len - t_start];
for (i, s) in scores.iter_mut().enumerate() {
let base = ((t_start + i) * n_kv_heads + kv_h) * head_dim;
let mut v = super::dot_f32(q_h, &k_cache[base..base + head_dim]) * scale;
if let Some(sc) = softcap {
v = sc * (v / sc).tanh();
}
*s = v;
}
let l = super::softmax_row_exp_sum(&mut scores);
acc.fill(0.0);
for (i, &p) in scores.iter().enumerate() {
let base = ((t_start + i) * n_kv_heads + kv_h) * head_dim;
super::axpy(&mut acc, &v_cache[base..base + head_dim], p);
}
if l > 0.0 {
super::scale_inplace(&mut acc, 1.0 / l);
}
out[b * q_stride + h * head_dim..][..head_dim].copy_from_slice(&acc);
}
}
out
}
#[test]
fn position_outer_prefill_is_bit_identical_to_the_query_outer_form() {
let n_heads = 6;
let n_kv_heads = 2;
let head_dim = 16;
let n_q = 19;
let kv_prefix = 5;
let kv_len = kv_prefix + n_q;
let q_stride = n_heads * head_dim;
let kv_stride = n_kv_heads * head_dim;
let q: Vec<f32> = (0..n_q * q_stride)
.map(|i| ((i as f32) * 0.013 - 0.7).sin() * 1.3)
.collect();
let k_cache: Vec<f32> = (0..kv_len * kv_stride)
.map(|i| ((i as f32) * 0.017 - 0.3).cos() * 1.1)
.collect();
let v_cache: Vec<f32> = (0..kv_len * kv_stride)
.map(|i| ((i as f32) * 0.011 + 0.2).sin() * 0.9)
.collect();
for window in [None, Some(1), Some(3), Some(8), Some(9), Some(kv_len + 4)] {
for softcap in [None, Some(30.0)] {
let got = super::causal_gqa_attention_prefill_shared_kv_windowed(
&q, &k_cache, &v_cache, n_heads, n_kv_heads, head_dim, n_q, kv_prefix, softcap,
window,
);
let want = prefill_query_outer_reference(
&q, &k_cache, &v_cache, n_heads, n_kv_heads, head_dim, n_q, kv_prefix, softcap,
window,
);
assert_eq!(got, want, "window {window:?} softcap {softcap:?}");
}
}
}
#[test]
fn tiled_prefill_gemm_is_bit_identical_across_awkward_shapes() {
let shapes = [
(4usize, 4usize, 64usize),
(6, 2, 64),
(5, 1, 80),
(3, 3, 128),
(2, 1, 256),
];
let batches = [(19usize, 5usize), (8, 0), (3, 7), (16, 1), (7, 0)];
for &(n_heads, n_kv_heads, head_dim) in &shapes {
let q_stride = n_heads * head_dim;
let kv_stride = n_kv_heads * head_dim;
for &(n_q, kv_prefix) in &batches {
let kv_len = kv_prefix + n_q;
let q: Vec<f32> = (0..n_q * q_stride)
.map(|i| ((i as f32) * 0.013 - 0.7).sin() * 1.3)
.collect();
let k_cache: Vec<f32> = (0..kv_len * kv_stride)
.map(|i| ((i as f32) * 0.017 - 0.3).cos() * 1.1)
.collect();
let v_cache: Vec<f32> = (0..kv_len * kv_stride)
.map(|i| ((i as f32) * 0.011 + 0.2).sin() * 0.9)
.collect();
for window in [None, Some(2), Some(5), Some(9), Some(kv_len + 3)] {
for softcap in [None, Some(30.0)] {
let got = super::causal_gqa_attention_prefill_shared_kv_windowed(
&q, &k_cache, &v_cache, n_heads, n_kv_heads, head_dim, n_q, kv_prefix,
softcap, window,
);
let want = prefill_query_outer_reference(
&q, &k_cache, &v_cache, n_heads, n_kv_heads, head_dim, n_q, kv_prefix,
softcap, window,
);
assert_eq!(
got, want,
"heads {n_heads}/{n_kv_heads} head_dim {head_dim} n_q {n_q} \
kv_prefix {kv_prefix} window {window:?} softcap {softcap:?}"
);
}
}
}
}
}
#[test]
fn vectorised_softmax_row_matches_the_scalar_libm_form() {
fn libm_reference(x: &[f32]) -> (Vec<f32>, f32) {
let m = x.iter().fold(f32::NEG_INFINITY, |a, &s| a.max(s));
let out: Vec<f32> = x.iter().map(|s| (s - m).exp()).collect();
let mut l = 0f32;
for &e in out.iter() {
l += e;
}
(out, l)
}
for spread in [1.0f32, 8.0, 200.0, 0.0] {
for n in (0..=17).chain([31, 32, 33, 64, 127, 512]) {
let row: Vec<f32> = (0..n)
.map(|i| ((i as f32) * 0.37 - 1.1).sin() * spread)
.collect();
let (want, want_l) = libm_reference(&row);
let mut got = row.clone();
let got_l = super::softmax_row_exp_sum(&mut got);
for (i, (g, w)) in got.iter().zip(want.iter()).enumerate() {
assert!(
(g - w).abs() <= 1e-6 * w + 1e-30,
"spread {spread} n {n} slot {i}: vector {g} vs libm {w}"
);
}
assert!(
(got_l - want_l).abs() <= 1e-5 * want_l.max(1.0),
"spread {spread} n {n}: sum {got_l} vs libm {want_l}"
);
if n == 0 {
assert_eq!(got_l, 0.0, "an empty visible range normalises to nothing");
}
if spread == 0.0 && n > 0 {
for (i, g) in got.iter().enumerate() {
assert_eq!(*g, 1.0, "n {n} slot {i}: exp(0) must be exact");
}
}
}
}
}
#[test]
fn the_vectorised_softmax_is_no_less_accurate_than_the_scalar_one() {
for spread in [1.0f64, 6.0, 20.0] {
for n in [64usize, 253, 512] {
let row: Vec<f64> = (0..n)
.map(|i| ((i as f64) * 0.37 - 1.1).sin() * spread)
.collect();
let m64 = row.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
let truth: Vec<f64> = row.iter().map(|s| (s - m64).exp()).collect();
let truth_l: f64 = truth.iter().sum();
let f32_row: Vec<f32> = row.iter().map(|&s| s as f32).collect();
let mut vector = f32_row.clone();
let vector_l = super::softmax_row_exp_sum(&mut vector);
let mut scalar = f32_row.clone();
let scalar_l = super::softmax_row_exp_sum_scalar(&mut scalar);
let worst = |got: &[f32]| -> f64 {
got.iter()
.zip(truth.iter())
.map(|(&g, &t)| ((g as f64) - t).abs() / t)
.fold(0.0, f64::max)
};
let (ev, es) = (worst(&vector), worst(&scalar));
let lv = ((vector_l as f64) - truth_l).abs() / truth_l;
let ls = ((scalar_l as f64) - truth_l).abs() / truth_l;
let eps = f64::from(f32::EPSILON);
assert!(
ev <= 2.0 * es.max(eps) && ev <= 64.0 * eps,
"spread {spread} n {n}: vector probabilities err {ev:e} \
against scalar {es:e}"
);
assert!(
lv <= ls.max(eps),
"spread {spread} n {n}: vector normaliser err {lv:e} \
against scalar {ls:e}"
);
}
}
}
#[test]
fn softmax_row_finds_its_maximum_in_every_lane_position() {
for n in 1usize..=20 {
for peak in 0..n {
let mut row: Vec<f32> = (0..n).map(|i| -(i as f32) - 3.0).collect();
row[peak] = 12.5;
let l = super::softmax_row_exp_sum(&mut row);
assert_eq!(row[peak], 1.0, "n {n} peak {peak}: the max term is exp(0)");
for (i, &p) in row.iter().enumerate() {
assert!(p <= 1.0, "n {n} peak {peak} slot {i}: {p} exceeds the max");
}
assert!(l >= 1.0, "n {n} peak {peak}: sum {l} must include the max");
}
}
}
use super::*;
#[test]
fn simd_dot_f32_matches_scalar_across_lengths() {
for n in [1usize, 3, 4, 7, 8, 15, 16, 63, 128, 129] {
let a: Vec<f32> = (0..n).map(|i| ((i as f32) * 0.31 - 2.0).sin()).collect();
let b: Vec<f32> = (0..n).map(|i| ((i as f32) * 0.17 + 1.0).cos()).collect();
let simd = dot_f32(&a, &b);
let scalar: f32 = a.iter().zip(&b).map(|(x, y)| x * y).sum();
assert!(
(simd - scalar).abs() <= 1e-4 * scalar.abs().max(1.0),
"n={n} simd={simd} scalar={scalar}"
);
}
}
#[test]
fn rope_preserves_vector_norm() {
let mut v = vec![1.0, 2.0, 3.0, 4.0];
let norm_before: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
apply_rope(&mut v, 5, 10000.0);
let norm_after: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!(
(norm_before - norm_after).abs() < 1e-4,
"RoPE is a rotation and must preserve norm"
);
}
#[test]
fn rope_back_inverts_rope() {
let original = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let mut v = original.clone();
apply_rope(&mut v, 11, 10000.0);
apply_rope_back(&mut v, 11, 10000.0);
for (a, b) in v.iter().zip(original.iter()) {
assert!((a - b).abs() < 1e-5, "{a} vs {b}");
}
}
#[test]
fn rope_interleaved_back_inverts_interleaved() {
let original = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let mut v = original.clone();
apply_rope_interleaved(&mut v, 11, 10000.0);
apply_rope_interleaved_back(&mut v, 11, 10000.0);
for (a, b) in v.iter().zip(original.iter()) {
assert!((a - b).abs() < 1e-5, "{a} vs {b}");
}
}
#[test]
fn rope_at_position_zero_is_identity() {
let mut v = vec![1.0, 2.0, 3.0, 4.0];
let original = v.clone();
apply_rope(&mut v, 0, 10000.0);
for (a, b) in v.iter().zip(original.iter()) {
assert!((a - b).abs() < 1e-5);
}
}
#[test]
fn rope_with_all_ones_freq_factors_matches_plain_rope() {
let mut with_factors = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let mut plain = with_factors.clone();
let ones = vec![1.0; 3];
apply_rope_with_freq_factors(&mut with_factors, 7, 10000.0, &ones);
apply_rope(&mut plain, 7, 10000.0);
for (a, b) in with_factors.iter().zip(plain.iter()) {
assert!((a - b).abs() < 1e-5, "{a} vs {b}");
}
}
#[test]
fn rope_with_freq_factors_diverges_from_plain_rope_when_factors_are_not_one() {
let mut with_factors = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let mut plain = with_factors.clone();
let factors = vec![0.5, 2.0, 1.0];
apply_rope_with_freq_factors(&mut with_factors, 7, 10000.0, &factors);
apply_rope(&mut plain, 7, 10000.0);
let differs = with_factors
.iter()
.zip(plain.iter())
.any(|(a, b)| (a - b).abs() > 1e-4);
assert!(differs, "non-1.0 freq_factors must change the rotation");
}
#[test]
fn rope_with_freq_factors_preserves_vector_norm() {
let mut v = vec![1.0, 2.0, 3.0, 4.0];
let norm_before: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
apply_rope_with_freq_factors(&mut v, 5, 10000.0, &[0.8, 1.3]);
let norm_after: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm_before - norm_after).abs() < 1e-4);
}
#[test]
fn yarn_high_is_clamped_to_rotary_dim_minus_one_not_half_minus_one() {
let scaling = YarnScaling::new(8.0, 131_072);
let (low, high) = yarn_correction_range(scaling, 64, 10_000.0);
assert!((low - 22.0).abs() < 1e-9, "low was {low}");
assert!((high - 35.0).abs() < 1e-9, "high was {high}");
let factors = yarn_freq_factors(scaling, 64, 10_000.0);
let last = factors[31];
let ramp = (31.0 - 22.0) / (35.0 - 22.0);
let want = 1.0 / (ramp / 8.0 + (1.0 - ramp));
assert!(
(last - want).abs() < 1e-4,
"last band divisor {last} must be the reference's {want}"
);
assert!(
(last - 8.0).abs() > 1.0,
"clamping high to rotary_dim/2 - 1 would fully interpolate this band \
(divisor 8.0, the whole factor); got {last}"
);
}
#[test]
fn yarn_freq_factors_match_the_reference_ramp_formula_band_by_band() {
let scaling = YarnScaling::new(8.0, 131_072);
let factors = yarn_freq_factors(scaling, 64, 10_000.0);
assert_eq!(factors.len(), 32, "one divisor per rotation band");
for band in [0usize, 10, 22] {
assert!(
(factors[band] - 1.0).abs() < 1e-6,
"band {band} is at or below low=22 and must be left extrapolated, \
got {}",
factors[band]
);
}
for band in [23usize, 27, 31] {
let ramp = (band as f32 - 22.0) / (35.0 - 22.0);
let want = 1.0 / (ramp / 8.0 + (1.0 - ramp));
assert!(
(factors[band] - want).abs() < 1e-4,
"band {band}: got {}, reference {want}",
factors[band]
);
}
}
#[test]
fn yarn_nudges_a_collapsed_correction_range_instead_of_flooring_the_gap_at_one() {
let scaling = YarnScaling {
beta_slow: 32.0,
truncate: false,
..YarnScaling::new(8.0, 131_072)
};
let (low, high) = yarn_correction_range(scaling, 64, 10_000.0);
assert!((low - 22.513_44).abs() < 1e-4, "low was {low}");
assert!(
(high - low - 0.001).abs() < 1e-9,
"high must be low + 0.001, got {high}"
);
let factors = yarn_freq_factors(scaling, 64, 10_000.0);
assert!(
(factors[22] - 1.0).abs() < 1e-6,
"band below the step must be untouched, got {}",
factors[22]
);
assert!(
(factors[23] - 8.0).abs() < 1e-4,
"band above the step must take the whole factor (a gap of 1 would \
give 1.7414); got {}",
factors[23]
);
}
#[test]
fn yarn_with_a_factor_of_one_leaves_every_band_untouched() {
let factors = yarn_freq_factors(YarnScaling::new(1.0, 4096), 32, 10_000.0);
for (band, f) in factors.iter().enumerate() {
assert!((f - 1.0).abs() < 1e-6, "band {band} moved to {f}");
}
}
#[test]
fn yarn_divisors_reproduce_the_references_rewritten_frequencies() {
let scaling = YarnScaling::new(8.0, 131_072);
let factors = yarn_freq_factors(scaling, 64, 10_000.0);
let band = 31usize;
let pos = 1024usize;
let mut v = vec![0.0f32; 64];
v[band] = 1.0;
apply_rope_with_freq_factors(&mut v, pos, 10_000.0, &factors);
let ramp = (band as f64 - 22.0) / (35.0 - 22.0);
let inv_freq = 1.0 / 10_000f64.powf((2 * band) as f64 / 64.0);
let inv_freq_new = inv_freq * (ramp / 8.0 + (1.0 - ramp));
let angle = pos as f64 * inv_freq_new;
assert!(
(v[band] as f64 - angle.cos()).abs() < 1e-5,
"cos: {} vs {}",
v[band],
angle.cos()
);
assert!(
(v[band + 32] as f64 - angle.sin()).abs() < 1e-5,
"sin: {} vs {}",
v[band + 32],
angle.sin()
);
}
#[test]
fn proportional_freq_factors_respace_frequencies_over_the_full_head() {
let factors = proportional_freq_factors(128, 96, 10_000.0);
assert_eq!(factors.len(), 48, "one divisor per rotated band");
assert!((factors[0] - 1.0).abs() < 1e-6, "band 0 is 1/1");
assert!(
(factors[1] - 0.953_161_9).abs() < 1e-5,
"band 1 was {}",
factors[1]
);
assert!(
(factors[47] - 0.104_913_97).abs() < 1e-5,
"last band was {}",
factors[47]
);
}
#[test]
fn proportional_freq_factors_are_all_ones_when_the_whole_head_rotates() {
for f in proportional_freq_factors(128, 128, 500_000.0) {
assert!((f - 1.0).abs() < 1e-6, "full-width band moved to {f}");
}
}
#[test]
fn rope_interleaved_preserves_vector_norm() {
let mut v = vec![1.0, 2.0, 3.0, 4.0];
let norm_before: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
apply_rope_interleaved(&mut v, 5, 10000.0);
let norm_after: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!(
(norm_before - norm_after).abs() < 1e-4,
"RoPE is a rotation and must preserve norm"
);
}
#[test]
fn rope_interleaved_at_position_zero_is_identity() {
let mut v = vec![1.0, 2.0, 3.0, 4.0];
let original = v.clone();
apply_rope_interleaved(&mut v, 0, 10000.0);
for (a, b) in v.iter().zip(original.iter()) {
assert!((a - b).abs() < 1e-5);
}
}
#[test]
fn rope_interleaved_rotates_adjacent_pairs_not_split_halves() {
let mut interleaved = vec![1.0, 0.0, 0.0, 1.0];
let mut split_half = interleaved.clone();
apply_rope_interleaved(&mut interleaved, 3, 10000.0);
apply_rope(&mut split_half, 3, 10000.0);
let differs = interleaved
.iter()
.zip(split_half.iter())
.any(|(a, b)| (a - b).abs() > 1e-4);
assert!(differs, "the two RoPE conventions must not coincide here");
}
#[test]
fn rope_interleaved_with_all_ones_freq_factors_matches_plain_interleaved() {
let mut with_factors = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let mut plain = with_factors.clone();
let ones = vec![1.0; 3];
apply_rope_interleaved_with_freq_factors(&mut with_factors, 7, 10000.0, &ones);
apply_rope_interleaved(&mut plain, 7, 10000.0);
for (a, b) in with_factors.iter().zip(plain.iter()) {
assert!((a - b).abs() < 1e-5, "{a} vs {b}");
}
}
#[test]
fn rope_interleaved_with_freq_factors_diverges_when_factors_are_not_one() {
let mut with_factors = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let mut plain = with_factors.clone();
let factors = vec![0.5, 2.0, 1.0];
apply_rope_interleaved_with_freq_factors(&mut with_factors, 7, 10000.0, &factors);
apply_rope_interleaved(&mut plain, 7, 10000.0);
let differs = with_factors
.iter()
.zip(plain.iter())
.any(|(a, b)| (a - b).abs() > 1e-4);
assert!(differs, "non-1.0 freq_factors must change the rotation");
}
#[test]
fn rope_interleaved_with_freq_factors_preserves_vector_norm() {
let mut v = vec![1.0, 2.0, 3.0, 4.0];
let norm_before: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
apply_rope_interleaved_with_freq_factors(&mut v, 5, 10000.0, &[0.8, 1.3]);
let norm_after: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm_before - norm_after).abs() < 1e-4);
}
#[test]
fn attention_with_single_position_returns_that_value() {
let q = vec![1.0, 0.0]; let k_cache = vec![0.5, 0.5]; let v_cache = vec![9.0, -3.0];
let out = causal_gqa_attention(&q, &k_cache, &v_cache, 1, 1, 2, 1);
assert!((out[0] - 9.0).abs() < 1e-4);
assert!((out[1] - (-3.0)).abs() < 1e-4);
}
#[test]
fn gqa_group_mapping_shares_kv_heads_correctly() {
let head_dim = 2;
let q = vec![
1.0, 0.0, 1.0, 0.0, 1.0, 0.0, 1.0, 0.0, ];
let k_cache = vec![1.0, 0.0, 1.0, 0.0];
let v_cache = vec![100.0, 100.0, 200.0, 200.0];
let out = causal_gqa_attention(&q, &k_cache, &v_cache, 4, 2, head_dim, 1);
assert_eq!(&out[0..2], &[100.0, 100.0][..]);
assert_eq!(&out[2..4], &[100.0, 100.0][..]);
assert_eq!(&out[4..6], &[200.0, 200.0][..]);
assert_eq!(&out[6..8], &[200.0, 200.0][..]);
}
#[test]
fn prefill_gqa_matches_per_token_causal() {
let n_heads = 4;
let n_kv_heads = 2;
let head_dim = 4;
let seq_len = 5;
let q: Vec<f32> = (0..seq_len * n_heads * head_dim)
.map(|i| (i as f32 * 0.13).sin())
.collect();
let k: Vec<f32> = (0..seq_len * n_kv_heads * head_dim)
.map(|i| (i as f32 * 0.19).cos())
.collect();
let v: Vec<f32> = (0..seq_len * n_kv_heads * head_dim)
.map(|i| (i as f32 * 0.07).sin())
.collect();
let batched =
causal_gqa_attention_prefill(&q, &k, &v, n_heads, n_kv_heads, head_dim, seq_len);
let q_stride = n_heads * head_dim;
let kv_stride = n_kv_heads * head_dim;
for t in 0..seq_len {
let expect = causal_gqa_attention(
&q[t * q_stride..(t + 1) * q_stride],
&k[..(t + 1) * kv_stride],
&v[..(t + 1) * kv_stride],
n_heads,
n_kv_heads,
head_dim,
t + 1,
);
let got = &batched[t * q_stride..(t + 1) * q_stride];
for (a, b) in got.iter().zip(expect.iter()) {
assert!((a - b).abs() < 1e-5, "t={t}: {a} vs {b}");
}
}
}
#[test]
fn windowed_attention_with_window_covering_full_history_matches_full_causal() {
let n_heads = 2;
let n_kv_heads = 1;
let head_dim = 3;
let seq_len = 4;
let q: Vec<f32> = (0..n_heads * head_dim)
.map(|i| (i as f32 * 0.3).sin())
.collect();
let k_cache: Vec<f32> = (0..seq_len * n_kv_heads * head_dim)
.map(|i| (i as f32 * 0.17).cos())
.collect();
let v_cache: Vec<f32> = (0..seq_len * n_kv_heads * head_dim)
.map(|i| (i as f32 * 0.11).sin())
.collect();
let full = causal_gqa_attention(
&q, &k_cache, &v_cache, n_heads, n_kv_heads, head_dim, seq_len,
);
let windowed = causal_gqa_attention_windowed(
&q, &k_cache, &v_cache, n_heads, n_kv_heads, head_dim, seq_len, seq_len,
);
assert_eq!(full.len(), windowed.len());
for (a, b) in full.iter().zip(windowed.iter()) {
assert_eq!(
a.to_bits(),
b.to_bits(),
"window >= seq_len must be bit-identical to full causal"
);
}
}
#[test]
fn windowed_attention_ignores_positions_outside_the_window() {
let head_dim = 2;
let q = vec![1.0, 0.0];
let k_cache = vec![9.0, -9.0, 0.5, 0.5, -3.0, 7.0]; let v_cache = vec![10.0, 20.0, 30.0, 40.0, 50.0, 60.0];
let out = causal_gqa_attention_windowed(&q, &k_cache, &v_cache, 1, 1, head_dim, 3, 1);
assert!((out[0] - 50.0).abs() < 1e-4);
assert!((out[1] - 60.0).abs() < 1e-4);
}
#[test]
fn paged_attention_matches_contiguous_attention_bit_identical() {
use crate::cache::{PagedKvCache, PagedKvStore};
let n_heads = 4;
let n_kv_heads = 2;
let head_dim = 3;
let block_size = 2;
let seq_len = 5;
let k_flat: Vec<f32> = (0..seq_len * n_kv_heads * head_dim)
.map(|i| ((i * 7 + 1) % 13) as f32 * 0.1)
.collect();
let v_flat: Vec<f32> = (0..seq_len * n_kv_heads * head_dim)
.map(|i| ((i * 5 + 3) % 11) as f32 * 0.1)
.collect();
let q: Vec<f32> = (0..n_heads * head_dim)
.map(|i| ((i * 3 + 2) % 9) as f32 * 0.1)
.collect();
let contiguous =
causal_gqa_attention(&q, &k_flat, &v_flat, n_heads, n_kv_heads, head_dim, seq_len);
let mut store = PagedKvStore::new(block_size, seq_len, n_kv_heads, head_dim);
let mut cache = PagedKvCache::new();
for t in 0..seq_len {
let start = t * n_kv_heads * head_dim;
let end = start + n_kv_heads * head_dim;
cache
.push(&mut store, &k_flat[start..end], &v_flat[start..end])
.expect("store sized for seq_len blocks, must not exhaust");
}
let paged = causal_gqa_attention_paged(
&q,
&store,
cache.block_table(),
n_heads,
n_kv_heads,
head_dim,
seq_len,
);
assert_eq!(contiguous.len(), paged.len());
for (a, b) in contiguous.iter().zip(paged.iter()) {
assert_eq!(a.to_bits(), b.to_bits(), "paged path must be bit-identical");
}
}
fn paged_fixture(
seq_len: usize,
n_kv_heads: usize,
head_dim: usize,
block_size: usize,
) -> (Vec<f32>, Vec<f32>, crate::cache::PagedKvStore, Vec<usize>) {
use crate::cache::{PagedKvCache, PagedKvStore};
let k_flat: Vec<f32> = (0..seq_len * n_kv_heads * head_dim)
.map(|i| ((i * 7 + 1) % 13) as f32 * 0.1)
.collect();
let v_flat: Vec<f32> = (0..seq_len * n_kv_heads * head_dim)
.map(|i| ((i * 5 + 3) % 11) as f32 * 0.1)
.collect();
let mut store = PagedKvStore::new(block_size, seq_len, n_kv_heads, head_dim);
let mut cache = PagedKvCache::new();
for t in 0..seq_len {
let start = t * n_kv_heads * head_dim;
let end = start + n_kv_heads * head_dim;
cache
.push(&mut store, &k_flat[start..end], &v_flat[start..end])
.expect("store sized for seq_len blocks, must not exhaust");
}
let table = cache.block_table().to_vec();
(k_flat, v_flat, store, table)
}
#[test]
fn the_paged_sink_term_is_bit_identical_to_the_contiguous_one() {
let (n_heads, n_kv_heads, head_dim, seq_len, block_size) = (4, 2, 3, 5, 2);
let (k_flat, v_flat, store, table) =
paged_fixture(seq_len, n_kv_heads, head_dim, block_size);
let q: Vec<f32> = (0..n_heads * head_dim)
.map(|i| ((i * 3 + 2) % 9) as f32 * 0.1)
.collect();
let sinks = vec![-30.0f32, 0.0, 1.5, 30.0];
let contiguous = causal_gqa_attention_sinks(
&q, &k_flat, &v_flat, n_heads, n_kv_heads, head_dim, seq_len, None, &sinks,
);
let paged = causal_gqa_attention_paged_sinks(
&q,
&store,
&table,
n_heads,
n_kv_heads,
head_dim,
seq_len,
None,
Some(&sinks),
None,
);
assert_eq!(contiguous.len(), paged.len());
for (i, (a, b)) in contiguous.iter().zip(paged.iter()).enumerate() {
assert_eq!(a.to_bits(), b.to_bits(), "element {i}: {a} vs {b}");
}
}
#[test]
fn the_paged_window_arm_is_bit_identical_to_the_contiguous_one() {
let (n_heads, n_kv_heads, head_dim, seq_len, block_size) = (4, 2, 3, 7, 2);
let (k_flat, v_flat, store, table) =
paged_fixture(seq_len, n_kv_heads, head_dim, block_size);
let q: Vec<f32> = (0..n_heads * head_dim)
.map(|i| ((i * 3 + 2) % 9) as f32 * 0.1)
.collect();
let sinks = vec![0.5f32; n_heads];
for window in [1usize, 2, 3, 6, 7, 99] {
let contiguous = causal_gqa_attention_sinks(
&q,
&k_flat,
&v_flat,
n_heads,
n_kv_heads,
head_dim,
seq_len,
Some(window),
&sinks,
);
let paged = causal_gqa_attention_paged_sinks(
&q,
&store,
&table,
n_heads,
n_kv_heads,
head_dim,
seq_len,
Some(window),
Some(&sinks),
None,
);
for (i, (a, b)) in contiguous.iter().zip(paged.iter()).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"window {window} element {i}: {a} vs {b}"
);
}
}
}
#[test]
fn the_paged_sink_kernel_without_sinks_or_window_is_the_plain_paged_kernel() {
let (n_heads, n_kv_heads, head_dim, seq_len, block_size) = (4, 2, 3, 5, 2);
let (_k, _v, store, table) = paged_fixture(seq_len, n_kv_heads, head_dim, block_size);
let q: Vec<f32> = (0..n_heads * head_dim)
.map(|i| ((i * 3 + 2) % 9) as f32 * 0.1)
.collect();
let plain =
causal_gqa_attention_paged(&q, &store, &table, n_heads, n_kv_heads, head_dim, seq_len);
let via_sinks = causal_gqa_attention_paged_sinks(
&q, &store, &table, n_heads, n_kv_heads, head_dim, seq_len, None, None, None,
);
for (i, (a, b)) in plain.iter().zip(via_sinks.iter()).enumerate() {
assert_eq!(a.to_bits(), b.to_bits(), "element {i}");
}
}
#[test]
fn mla_attention_with_single_position_returns_that_value() {
let q = vec![1.0, 0.0, 0.0, 0.0, 0.0]; let k_cache = vec![0.2, 0.2, 0.2, 0.2, 0.2]; let v_cache = vec![9.0, -3.0, 1.0]; let out = causal_mla_attention(&q, &k_cache, &v_cache, 1, 5, 3, 1);
assert_eq!(out.len(), 3);
assert!((out[0] - 9.0).abs() < 1e-4);
assert!((out[1] - (-3.0)).abs() < 1e-4);
assert!((out[2] - 1.0).abs() < 1e-4);
}
#[test]
fn mla_attention_every_head_gets_its_own_kv_no_grouping() {
let qk_head_dim = 2;
let v_head_dim = 2;
let q = vec![1.0, 0.0, 1.0, 0.0]; let k_cache = vec![1.0, 0.0, 1.0, 0.0]; let v_cache = vec![100.0, 100.0, 200.0, 200.0];
let out = causal_mla_attention(&q, &k_cache, &v_cache, 2, qk_head_dim, v_head_dim, 1);
assert_eq!(&out[0..2], &[100.0, 100.0][..]);
assert_eq!(&out[2..4], &[200.0, 200.0][..]);
}
#[test]
fn lightning_indexer_topk_keeps_all_positions_when_top_k_covers_them() {
let indexer_q = vec![vec![1.0, 0.0]];
let indexer_keys = vec![vec![1.0, 0.0], vec![0.5, 0.5], vec![0.1, 0.9]];
let indexer_weights = vec![1.0];
let kept = lightning_indexer_topk(&indexer_q, &indexer_keys, &indexer_weights, 10);
assert_eq!(kept, vec![0, 1, 2]);
}
#[test]
fn lightning_indexer_topk_selects_highest_scoring_positions() {
let indexer_q = vec![vec![1.0, 0.0]];
let indexer_keys = vec![vec![1.0, 0.0], vec![0.0, 1.0], vec![0.9, 0.1]];
let indexer_weights = vec![1.0];
let kept = lightning_indexer_topk(&indexer_q, &indexer_keys, &indexer_weights, 2);
assert_eq!(kept, vec![0, 2]);
}
#[test]
fn lightning_indexer_topk_relu_zeroes_negative_dot_products() {
let indexer_q = vec![vec![1.0, 0.0]];
let indexer_keys = vec![vec![0.3, 0.0], vec![-1.0, 0.0]];
let indexer_weights = vec![1.0];
let kept = lightning_indexer_topk(&indexer_q, &indexer_keys, &indexer_weights, 1);
assert_eq!(kept, vec![0]);
}
#[test]
fn mla_attention_sparse_with_all_positions_visible_matches_full_causal() {
let qk_head_dim = 3;
let v_head_dim = 2;
let seq_len = 4;
let n_heads = 2;
let q: Vec<f32> = (0..n_heads * qk_head_dim).map(|i| i as f32 * 0.1).collect();
let k_cache: Vec<f32> = (0..seq_len * n_heads * qk_head_dim)
.map(|i| (i as f32 * 0.05).sin())
.collect();
let v_cache: Vec<f32> = (0..seq_len * n_heads * v_head_dim)
.map(|i| (i as f32 * 0.05).cos())
.collect();
let full = causal_mla_attention(
&q,
&k_cache,
&v_cache,
n_heads,
qk_head_dim,
v_head_dim,
seq_len,
);
let visible: Vec<usize> = (0..seq_len).collect();
let sparse = causal_mla_attention_sparse(
&q,
&k_cache,
&v_cache,
n_heads,
qk_head_dim,
v_head_dim,
seq_len,
&visible,
);
assert_eq!(full.len(), sparse.len());
for (a, b) in full.iter().zip(sparse.iter()) {
assert!((a - b).abs() < 1e-6, "full={a} sparse={b}");
}
}
#[test]
fn mla_attention_sparse_ignores_positions_outside_visible_set() {
let qk_head_dim = 2;
let v_head_dim = 1;
let q = vec![1.0, 0.0];
let k_cache = vec![1.0, 0.0, 1.0, 0.0]; let v_cache = vec![5.0, 999.0]; let out = causal_mla_attention_sparse(
&q,
&k_cache,
&v_cache,
1,
qk_head_dim,
v_head_dim,
2,
&[0],
);
assert_eq!(out.len(), 1);
assert!((out[0] - 5.0).abs() < 1e-6);
}
#[test]
fn attn_logit_softcap_changes_output_vs_uncapped() {
let n_heads = 2;
let n_kv_heads = 1;
let head_dim = 4;
let seq_len = 3;
let q: Vec<f32> = (0..n_heads * head_dim)
.map(|i| (i as f32 + 1.0) * 2.5)
.collect();
let k: Vec<f32> = (0..seq_len * n_kv_heads * head_dim)
.map(|i| (i as f32 * 0.7).sin() * 3.0)
.collect();
let v: Vec<f32> = (0..seq_len * n_kv_heads * head_dim)
.map(|i| (i as f32 * 0.3).cos())
.collect();
let plain = causal_gqa_attention(&q, &k, &v, n_heads, n_kv_heads, head_dim, seq_len);
let capped = causal_gqa_attention_softcap(
&q,
&k,
&v,
n_heads,
n_kv_heads,
head_dim,
seq_len,
Some(30.0),
);
assert_eq!(plain.len(), capped.len());
let differs = plain
.iter()
.zip(capped.iter())
.any(|(a, b)| (a - b).abs() > 1e-5);
assert!(differs, "softcap must change attention output");
let none =
causal_gqa_attention_softcap(&q, &k, &v, n_heads, n_kv_heads, head_dim, seq_len, None);
for (a, b) in plain.iter().zip(none.iter()) {
assert!((a - b).abs() < 1e-6);
}
}
}