use crate::forward::cpu::matmul_bt;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct NsaConfig {
pub num_heads: usize,
pub num_kv_heads: usize,
pub head_dim: usize,
pub compress_block: usize,
pub compress_stride: usize,
pub select_block: usize,
pub num_selected: usize,
pub window: usize,
}
impl NsaConfig {
pub fn validate(&self) {
assert!(self.num_heads > 0, "num_heads must be > 0");
assert!(self.num_kv_heads > 0, "num_kv_heads must be > 0");
assert!(self.head_dim > 0, "head_dim must be > 0");
assert!(self.compress_block > 0, "compress_block (l) must be > 0");
assert!(self.compress_stride > 0, "compress_stride (d) must be > 0");
assert!(self.select_block > 0, "select_block (l') must be > 0");
assert!(
self.num_selected >= 3,
"num_selected (n={}) must be >= 3 — the forced-block scheme needs \
1 initial + 2 local blocks",
self.num_selected
);
assert!(self.window > 0, "window (w) must be > 0");
assert_eq!(
self.num_heads % self.num_kv_heads,
0,
"num_heads ({}) must be divisible by num_kv_heads ({})",
self.num_heads,
self.num_kv_heads
);
assert_eq!(
self.compress_block % self.compress_stride,
0,
"compress_block (l={}) must be divisible by compress_stride (d={})",
self.compress_block,
self.compress_stride
);
assert_eq!(
self.select_block % self.compress_stride,
0,
"select_block (l'={}) must be divisible by compress_stride (d={})",
self.select_block,
self.compress_stride
);
assert!(
self.compress_block <= self.select_block,
"compress_block (l={}) must be <= select_block (l'={}) — paper Eq. 9 precondition",
self.compress_block,
self.select_block
);
}
#[inline]
pub fn n_rep(&self) -> usize {
self.num_heads / self.num_kv_heads
}
#[inline]
pub fn q_dim(&self) -> usize {
self.num_heads * self.head_dim
}
#[inline]
pub fn kv_dim(&self) -> usize {
self.num_kv_heads * self.head_dim
}
#[inline]
pub fn num_compress_blocks(&self, seq_len: usize) -> usize {
if seq_len < self.compress_block {
return 0;
}
(seq_len - self.compress_block) / self.compress_stride + 1
}
#[inline]
pub fn num_select_blocks(&self, seq_len: usize) -> usize {
seq_len.div_ceil(self.select_block)
}
#[inline]
pub fn compress_per_select(&self) -> usize {
self.select_block / self.compress_stride
}
#[inline]
pub fn intra_per_compress(&self) -> usize {
self.compress_block / self.compress_stride
}
#[inline]
pub fn phi_in(&self) -> usize {
self.compress_block * self.head_dim
}
}
pub struct NsaWeights {
pub phi_k_w1: Vec<f32>,
pub phi_k_b1: Vec<f32>,
pub phi_k_w2: Vec<f32>,
pub phi_k_b2: Vec<f32>,
pub phi_v_w1: Vec<f32>,
pub phi_v_b1: Vec<f32>,
pub phi_v_w2: Vec<f32>,
pub phi_v_b2: Vec<f32>,
pub k_intrablock_pos: Vec<f32>,
pub v_intrablock_pos: Vec<f32>,
pub g_proj_w: Vec<f32>,
pub g_proj_b: Vec<f32>,
}
#[derive(Default, Clone, Debug)]
pub struct NsaScratch {
phi_input: Vec<f32>,
phi_tmp1: Vec<f32>,
ck: Vec<f32>,
cv: Vec<f32>,
compress_scores: Vec<f32>,
out_cmp: Vec<f32>,
importance: Vec<f32>,
sel_indices: Vec<usize>,
sel_candidates: Vec<usize>,
sel_k: Vec<f32>,
sel_v: Vec<f32>,
sel_scores: Vec<f32>,
out_slc: Vec<f32>,
win_scores: Vec<f32>,
out_win: Vec<f32>,
gates: Vec<f32>,
}
impl NsaScratch {
pub fn reserve_for(&mut self, seq_len: usize, cfg: &NsaConfig) {
let phi_in = cfg.phi_in();
let head_dim = cfg.head_dim;
let num_heads = cfg.num_heads;
let num_kv_heads = cfg.num_kv_heads;
let max_cblocks = cfg.num_compress_blocks(seq_len);
let max_sblocks = cfg.num_select_blocks(seq_len);
let n_sel = cfg.num_selected;
let win = cfg.window;
let lp = cfg.select_block;
self.phi_input.resize(phi_in, 0.0);
self.phi_tmp1.resize(phi_in, 0.0);
self.ck.resize(num_kv_heads * max_cblocks * head_dim, 0.0);
self.cv.resize(num_kv_heads * max_cblocks * head_dim, 0.0);
self.compress_scores.resize(max_cblocks, 0.0);
self.out_cmp.resize(seq_len * num_heads * head_dim, 0.0);
self.importance
.resize(num_kv_heads * max_sblocks.max(1), 0.0);
self.sel_indices.resize(num_kv_heads * n_sel, usize::MAX);
self.sel_candidates.resize(max_sblocks, 0);
self.sel_k.resize(n_sel * lp * head_dim, 0.0);
self.sel_v.resize(n_sel * lp * head_dim, 0.0);
self.sel_scores.resize(n_sel * lp, 0.0);
self.out_slc.resize(seq_len * num_heads * head_dim, 0.0);
self.win_scores.resize(win, 0.0);
self.out_win.resize(seq_len * num_heads * head_dim, 0.0);
self.gates.resize(3 * num_heads, 0.0);
}
}
#[allow(clippy::too_many_arguments)]
pub fn apply_native_sparse_attention(
q_buf: &[f32],
q_rope_buf: &[f32],
k_cmp_buf: &[f32],
k_slc_buf: &[f32],
k_win_buf: &[f32],
v_cmp_buf: &[f32],
v_slc_buf: &[f32],
v_win_buf: &[f32],
x_buf: &[f32],
weights: &NsaWeights,
attn_out: &mut [f32],
seq_len: usize,
cfg: &NsaConfig,
scratch: &mut NsaScratch,
) {
cfg.validate();
let q_dim = cfg.q_dim();
let kv_dim = cfg.kv_dim();
let head_dim = cfg.head_dim;
let num_heads = cfg.num_heads;
let num_kv_heads = cfg.num_kv_heads;
let l = cfg.compress_block;
let d = cfg.compress_stride;
let lp = cfg.select_block;
let n_sel = cfg.num_selected;
let win = cfg.window;
let n_rep = cfg.n_rep();
let phi_in = cfg.phi_in();
assert_eq!(
q_buf.len(),
seq_len * q_dim,
"q_buf length mismatch: expected {}, got {}",
seq_len * q_dim,
q_buf.len()
);
assert_eq!(
q_rope_buf.len(),
seq_len * q_dim,
"q_rope_buf length mismatch: expected {}, got {}",
seq_len * q_dim,
q_rope_buf.len()
);
assert_eq!(
k_cmp_buf.len(),
seq_len * kv_dim,
"k_cmp_buf length mismatch: expected {}, got {}",
seq_len * kv_dim,
k_cmp_buf.len()
);
assert_eq!(
k_slc_buf.len(),
seq_len * kv_dim,
"k_slc_buf length mismatch: expected {}, got {}",
seq_len * kv_dim,
k_slc_buf.len()
);
assert_eq!(
k_win_buf.len(),
seq_len * kv_dim,
"k_win_buf length mismatch: expected {}, got {}",
seq_len * kv_dim,
k_win_buf.len()
);
assert_eq!(
v_cmp_buf.len(),
seq_len * kv_dim,
"v_cmp_buf length mismatch: expected {}, got {}",
seq_len * kv_dim,
v_cmp_buf.len()
);
assert_eq!(
v_slc_buf.len(),
seq_len * kv_dim,
"v_slc_buf length mismatch: expected {}, got {}",
seq_len * kv_dim,
v_slc_buf.len()
);
assert_eq!(
v_win_buf.len(),
seq_len * kv_dim,
"v_win_buf length mismatch: expected {}, got {}",
seq_len * kv_dim,
v_win_buf.len()
);
assert_eq!(
attn_out.len(),
seq_len * q_dim,
"attn_out length mismatch: expected {}, got {}",
seq_len * q_dim,
attn_out.len()
);
let model_dim = x_buf.len().checked_div(seq_len).unwrap_or(0);
assert_eq!(
x_buf.len(),
seq_len * model_dim,
"x_buf length must be seq_len*model_dim (model_dim inferred as {}), got {}",
model_dim,
x_buf.len()
);
assert!(
seq_len == 0 || model_dim > 0,
"x_buf must be non-empty when seq_len > 0 (model_dim inferred as 0)"
);
assert_eq!(
weights.phi_k_w1.len(),
phi_in * phi_in,
"phi_k_w1 must be [phi_in, phi_in]=[{phi_in}*{phi_in}]"
);
assert_eq!(
weights.phi_k_b1.len(),
phi_in,
"phi_k_b1 must be [phi_in={phi_in}]"
);
assert_eq!(
weights.phi_k_w2.len(),
head_dim * phi_in,
"phi_k_w2 must be [head_dim={head_dim}, phi_in={phi_in}]"
);
assert_eq!(
weights.phi_k_b2.len(),
head_dim,
"phi_k_b2 must be [head_dim={head_dim}]"
);
assert_eq!(
weights.phi_v_w1.len(),
phi_in * phi_in,
"phi_v_w1 must be [phi_in, phi_in]=[{phi_in}*{phi_in}]"
);
assert_eq!(
weights.phi_v_b1.len(),
phi_in,
"phi_v_b1 must be [phi_in={phi_in}]"
);
assert_eq!(
weights.phi_v_w2.len(),
head_dim * phi_in,
"phi_v_w2 must be [head_dim={head_dim}, phi_in={phi_in}]"
);
assert_eq!(
weights.phi_v_b2.len(),
head_dim,
"phi_v_b2 must be [head_dim={head_dim}]"
);
assert_eq!(
weights.k_intrablock_pos.len(),
num_kv_heads * l * head_dim,
"k_intrablock_pos must be [num_kv_heads={num_kv_heads}, l={l}, head_dim={head_dim}]"
);
assert_eq!(
weights.v_intrablock_pos.len(),
num_kv_heads * l * head_dim,
"v_intrablock_pos must be [num_kv_heads={num_kv_heads}, l={l}, head_dim={head_dim}]"
);
if model_dim > 0 {
let g_proj_rows = 3 * num_heads;
assert_eq!(
weights.g_proj_w.len(),
g_proj_rows * model_dim,
"g_proj_w must be [3*num_heads={g_proj_rows}, model_dim={model_dim}]"
);
}
assert_eq!(
weights.g_proj_b.len(),
3 * num_heads,
"g_proj_b must be [3*num_heads={}]",
3 * num_heads
);
if seq_len == 0 {
return;
}
scratch.reserve_for(seq_len, cfg);
attn_out.fill(0.0);
scratch.out_cmp[..seq_len * q_dim].fill(0.0);
scratch.out_slc[..seq_len * q_dim].fill(0.0);
scratch.out_win[..seq_len * q_dim].fill(0.0);
let scale = (head_dim as f32).powf(-0.5);
let max_cblocks = cfg.num_compress_blocks(seq_len);
let max_sblocks = cfg.num_select_blocks(seq_len);
let cps = cfg.compress_per_select(); let ipc = cfg.intra_per_compress();
for kv_h in 0..num_kv_heads {
let pos_enc_k = &weights.k_intrablock_pos[kv_h * l * head_dim..(kv_h + 1) * l * head_dim];
let pos_enc_v = &weights.v_intrablock_pos[kv_h * l * head_dim..(kv_h + 1) * l * head_dim];
for bi in 0..max_cblocks {
let tok_start = bi * d;
let phi_in_buf = &mut scratch.phi_input[..phi_in];
for p in 0..l {
let tok = tok_start + p;
let src = tok * kv_dim + kv_h * head_dim;
let pe = p * head_dim;
let dst = p * head_dim;
for dd in 0..head_dim {
phi_in_buf[dst + dd] = k_cmp_buf[src + dd] + pos_enc_k[pe + dd];
}
}
let tmp1 = &mut scratch.phi_tmp1[..phi_in];
tmp1.fill(0.0);
matmul_bt(phi_in_buf, &weights.phi_k_w1, tmp1, 1, phi_in, phi_in);
for i in 0..phi_in {
tmp1[i] = (tmp1[i] + weights.phi_k_b1[i]).max(0.0);
}
let ck_off = (kv_h * max_cblocks + bi) * head_dim;
let ck_slot = &mut scratch.ck[ck_off..ck_off + head_dim];
ck_slot.fill(0.0);
matmul_bt(tmp1, &weights.phi_k_w2, ck_slot, 1, phi_in, head_dim);
for i in 0..head_dim {
ck_slot[i] += weights.phi_k_b2[i];
}
let phi_in_buf = &mut scratch.phi_input[..phi_in];
for p in 0..l {
let tok = tok_start + p;
let src = tok * kv_dim + kv_h * head_dim;
let pe = p * head_dim;
let dst = p * head_dim;
for dd in 0..head_dim {
phi_in_buf[dst + dd] = v_cmp_buf[src + dd] + pos_enc_v[pe + dd];
}
}
let tmp1 = &mut scratch.phi_tmp1[..phi_in];
tmp1.fill(0.0);
matmul_bt(phi_in_buf, &weights.phi_v_w1, tmp1, 1, phi_in, phi_in);
for i in 0..phi_in {
tmp1[i] = (tmp1[i] + weights.phi_v_b1[i]).max(0.0);
}
let cv_off = (kv_h * max_cblocks + bi) * head_dim;
let cv_slot = &mut scratch.cv[cv_off..cv_off + head_dim];
cv_slot.fill(0.0);
matmul_bt(tmp1, &weights.phi_v_w2, cv_slot, 1, phi_in, head_dim);
for i in 0..head_dim {
cv_slot[i] += weights.phi_v_b2[i];
}
}
}
for qt in 0..seq_len {
scratch.importance[..num_kv_heads * max_sblocks.max(1)].fill(0.0);
for kv_h in 0..num_kv_heads {
let valid_cblocks = count_valid_compress_blocks(qt, l, d, max_cblocks);
let q_head_start = kv_h * n_rep;
for qh_local in 0..n_rep {
let qh = q_head_start + qh_local;
let out_cmp_off = qt * q_dim + qh * head_dim;
let out_cmp_slot = &mut scratch.out_cmp[out_cmp_off..out_cmp_off + head_dim];
if valid_cblocks == 0 {
out_cmp_slot.fill(0.0);
continue;
}
let q_off = qt * q_dim + qh * head_dim;
let q_head = &q_buf[q_off..q_off + head_dim];
let scores = &mut scratch.compress_scores[..valid_cblocks];
for bi in 0..valid_cblocks {
let ck_off = (kv_h * max_cblocks + bi) * head_dim;
let dot: f32 = q_head
.iter()
.zip(scratch.ck[ck_off..ck_off + head_dim].iter())
.map(|(&a, &b)| a * b)
.sum();
scores[bi] = dot * scale;
}
softmax_inplace(scores);
out_cmp_slot.fill(0.0);
for bi in 0..valid_cblocks {
let p = scores[bi];
let cv_off = (kv_h * max_cblocks + bi) * head_dim;
for dd in 0..head_dim {
out_cmp_slot[dd] += p * scratch.cv[cv_off + dd];
}
}
let valid_sblocks = count_valid_select_blocks(qt, lp, max_sblocks);
for sj in 0..valid_sblocks {
let block_score = aggregate_selection_importance(scores, sj, cps, ipc);
scratch.importance[kv_h * max_sblocks.max(1) + sj] += block_score;
}
}
}
for kv_h in 0..num_kv_heads {
let valid_sblocks = count_valid_select_blocks(qt, lp, max_sblocks);
let sel_out = &mut scratch.sel_indices[kv_h * n_sel..(kv_h + 1) * n_sel];
sel_out.fill(usize::MAX);
if valid_sblocks == 0 {
continue;
}
let n_take = n_sel.min(valid_sblocks);
let imp = &scratch.importance
[kv_h * max_sblocks.max(1)..kv_h * max_sblocks.max(1) + valid_sblocks];
let mut forced = [usize::MAX; 3];
forced[0] = 0;
let mut n_forced = 1usize;
if valid_sblocks >= 2 {
forced[n_forced] = valid_sblocks - 1;
n_forced += 1;
}
if valid_sblocks >= 3 {
forced[n_forced] = valid_sblocks - 2;
n_forced += 1;
}
forced[..n_forced].sort_unstable();
debug_assert!(
forced[..n_forced].windows(2).all(|w| w[0] != w[1]),
"forced selection blocks must be distinct by construction"
);
let n_forced = n_forced.min(n_take);
let n_extra = n_take - n_forced;
let candidates = &mut scratch.sel_candidates;
candidates.clear();
for j in 0..valid_sblocks {
if !forced[..n_forced].contains(&j) {
candidates.push(j);
}
}
let n_extra = n_extra.min(candidates.len());
for slot in 0..n_extra {
let mut best = slot;
for cand in (slot + 1)..candidates.len() {
if imp[candidates[cand]] > imp[candidates[best]] {
best = cand;
}
}
candidates.swap(slot, best);
}
let mut out_idx = 0;
for &fj in &forced[..n_forced] {
if out_idx < n_take {
sel_out[out_idx] = fj;
out_idx += 1;
}
}
for &ej in candidates[..n_extra].iter() {
if out_idx < n_take {
sel_out[out_idx] = ej;
out_idx += 1;
}
}
}
for kv_h in 0..num_kv_heads {
let valid_sblocks = count_valid_select_blocks(qt, lp, max_sblocks);
let sel_idxs = &scratch.sel_indices[kv_h * n_sel..(kv_h + 1) * n_sel];
let n_valid_sel: usize = sel_idxs
.iter()
.filter(|&&i| i != usize::MAX && i < valid_sblocks)
.count();
if n_valid_sel == 0 {
continue;
}
let max_sel_toks = n_valid_sel * lp;
let sel_k_buf = &mut scratch.sel_k[..max_sel_toks * head_dim];
let sel_v_buf = &mut scratch.sel_v[..max_sel_toks * head_dim];
let mut n_gathered = 0usize;
for &bj in sel_idxs.iter().take(n_sel) {
if bj == usize::MAX || bj >= valid_sblocks {
continue;
}
let tok_start = bj * lp;
for p in 0..lp {
let tok = tok_start + p;
if tok > qt {
continue;
}
let dst = n_gathered * head_dim;
let src = tok * kv_dim + kv_h * head_dim;
sel_k_buf[dst..dst + head_dim].copy_from_slice(&k_slc_buf[src..src + head_dim]);
sel_v_buf[dst..dst + head_dim].copy_from_slice(&v_slc_buf[src..src + head_dim]);
n_gathered += 1;
}
}
debug_assert!(n_gathered > 0, "n_valid_sel > 0 must imply n_gathered > 0");
let q_head_start = kv_h * n_rep;
for qh_local in 0..n_rep {
let qh = q_head_start + qh_local;
let q_off = qt * q_dim + qh * head_dim;
let q_head = &q_rope_buf[q_off..q_off + head_dim];
let scores = &mut scratch.sel_scores[..n_gathered];
for ti in 0..n_gathered {
let k_off = ti * head_dim;
let dot: f32 = q_head
.iter()
.zip(sel_k_buf[k_off..k_off + head_dim].iter())
.map(|(&a, &b)| a * b)
.sum();
scores[ti] = dot * scale;
}
softmax_inplace(scores);
let out_off = qt * q_dim + qh * head_dim;
let out_slot = &mut scratch.out_slc[out_off..out_off + head_dim];
out_slot.fill(0.0);
for ti in 0..n_gathered {
let p = scores[ti];
let v_off = ti * head_dim;
for dd in 0..head_dim {
out_slot[dd] += p * sel_v_buf[v_off + dd];
}
}
}
}
let win_start = qt.saturating_sub(win - 1);
let win_len = qt - win_start + 1;
for kv_h in 0..num_kv_heads {
let q_head_start = kv_h * n_rep;
for qh_local in 0..n_rep {
let qh = q_head_start + qh_local;
let q_off = qt * q_dim + qh * head_dim;
let q_head = &q_rope_buf[q_off..q_off + head_dim];
let scores = &mut scratch.win_scores[..win_len];
for (wi, tok) in (win_start..=qt).enumerate() {
let k_off = tok * kv_dim + kv_h * head_dim;
let dot: f32 = q_head
.iter()
.zip(k_win_buf[k_off..k_off + head_dim].iter())
.map(|(&a, &b)| a * b)
.sum();
scores[wi] = dot * scale;
}
softmax_inplace(scores);
let out_off = qt * q_dim + qh * head_dim;
let out_slot = &mut scratch.out_win[out_off..out_off + head_dim];
out_slot.fill(0.0);
for (wi, tok) in (win_start..=qt).enumerate() {
let p = scores[wi];
let v_off = tok * kv_dim + kv_h * head_dim;
for dd in 0..head_dim {
out_slot[dd] += p * v_win_buf[v_off + dd];
}
}
}
}
let x_t = &x_buf[qt * model_dim..(qt + 1) * model_dim];
let gates = &mut scratch.gates[..3 * num_heads];
for g_idx in 0..3 * num_heads {
let w_row = &weights.g_proj_w[g_idx * model_dim..(g_idx + 1) * model_dim];
let dot: f32 = x_t.iter().zip(w_row.iter()).map(|(&a, &b)| a * b).sum();
gates[g_idx] = sigmoid(dot + weights.g_proj_b[g_idx]);
}
for qh in 0..num_heads {
let g_cmp = gates[3 * qh];
let g_slc = gates[3 * qh + 1];
let g_win = gates[3 * qh + 2];
let off = qt * q_dim + qh * head_dim;
for dd in 0..head_dim {
attn_out[off + dd] = g_cmp * scratch.out_cmp[off + dd]
+ g_slc * scratch.out_slc[off + dd]
+ g_win * scratch.out_win[off + dd];
}
}
}
}
#[inline]
fn count_valid_compress_blocks(qt: usize, l: usize, d: usize, max_cblocks: usize) -> usize {
if qt + 1 < l {
return 0;
}
((qt + 1 - l) / d + 1).min(max_cblocks)
}
#[inline]
fn count_valid_select_blocks(qt: usize, lp: usize, max_sblocks: usize) -> usize {
(qt / lp + 1).min(max_sblocks)
}
#[inline]
fn aggregate_selection_importance(p_cmp: &[f32], sj: usize, cps: usize, ipc: usize) -> f32 {
let base = cps * sj;
let mut acc = 0.0_f32;
for m in 0..cps {
for n in 0..ipc {
let offset = m + n;
if offset <= base {
let ci = base - offset;
if ci < p_cmp.len() {
acc += p_cmp[ci];
}
}
}
}
acc
}
#[inline]
fn softmax_inplace(x: &mut [f32]) {
if x.is_empty() {
return;
}
let (max_val, any_nan) = crate::attention::softmax_row::row_max_and_any_nan(x);
if crate::attention::softmax_row::row_fails_closed_pre_exp(max_val, any_nan) {
x.fill(0.0);
return;
}
let mut sum = 0.0f32;
for v in x.iter_mut() {
*v = (*v - max_val).exp();
sum += *v;
}
crate::attention::softmax_row::finalize_row(x, sum);
}
#[inline]
fn sigmoid(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}
#[cfg(test)]
mod tests {
use super::*;
fn det_data(len: usize, seed: u64) -> Vec<f32> {
let mut state = seed.wrapping_add(0x9e37_79b9_7f4a_7c15);
let mut out = Vec::with_capacity(len);
for _ in 0..len {
state ^= state << 7;
state ^= state >> 9;
state = state.wrapping_mul(0x2545_f491_4f6c_dd1d);
let mantissa = ((state >> 41) as u32) & 0x007f_ffff;
let x = f32::from_bits(0x3f80_0000 | mantissa) - 1.5;
out.push(x);
}
out
}
fn small_cfg() -> NsaConfig {
NsaConfig {
num_heads: 2,
num_kv_heads: 1,
head_dim: 4,
compress_block: 4,
compress_stride: 2,
select_block: 4,
num_selected: 3,
window: 4,
}
}
fn make_weights(cfg: &NsaConfig, model_dim: usize, seed: u64) -> NsaWeights {
let l = cfg.compress_block;
let head_dim = cfg.head_dim;
let phi_in = cfg.phi_in();
let mut s = seed;
let mut next = |n: usize| -> Vec<f32> {
s = s.wrapping_add(0x1234_5678_9abc_def0);
det_data(n, s)
};
NsaWeights {
phi_k_w1: next(phi_in * phi_in),
phi_k_b1: next(phi_in),
phi_k_w2: next(head_dim * phi_in),
phi_k_b2: next(head_dim),
phi_v_w1: next(phi_in * phi_in),
phi_v_b1: next(phi_in),
phi_v_w2: next(head_dim * phi_in),
phi_v_b2: next(head_dim),
k_intrablock_pos: next(cfg.num_kv_heads * l * head_dim),
v_intrablock_pos: next(cfg.num_kv_heads * l * head_dim),
g_proj_w: next(3 * cfg.num_heads * model_dim),
g_proj_b: next(3 * cfg.num_heads),
}
}
fn run_nsa(cfg: &NsaConfig, seq_len: usize, seed: u64) -> Vec<f32> {
let model_dim = cfg.q_dim(); let weights = make_weights(cfg, model_dim, seed);
let q = det_data(seq_len * cfg.q_dim(), seed + 1);
let q_rope = det_data(seq_len * cfg.q_dim(), seed + 2);
let k_cmp = det_data(seq_len * cfg.kv_dim(), seed + 3);
let k_slc = det_data(seq_len * cfg.kv_dim(), seed + 4);
let k_win = det_data(seq_len * cfg.kv_dim(), seed + 5);
let v_cmp = det_data(seq_len * cfg.kv_dim(), seed + 6);
let v_slc = det_data(seq_len * cfg.kv_dim(), seed + 7);
let v_win = det_data(seq_len * cfg.kv_dim(), seed + 8);
let x = det_data(seq_len * model_dim, seed + 9);
let mut out = vec![0.0f32; seq_len * cfg.q_dim()];
let mut scratch = NsaScratch::default();
apply_native_sparse_attention(
&q,
&q_rope,
&k_cmp,
&k_slc,
&k_win,
&v_cmp,
&v_slc,
&v_win,
&x,
&weights,
&mut out,
seq_len,
cfg,
&mut scratch,
);
out
}
#[test]
fn test_nsa_config_validate_ok() {
small_cfg().validate(); }
#[test]
#[should_panic(expected = "num_heads must be > 0")]
fn test_zero_num_heads_panics() {
let mut cfg = small_cfg();
cfg.num_heads = 0;
cfg.validate();
}
#[test]
#[should_panic(expected = "num_kv_heads must be > 0")]
fn test_zero_num_kv_heads_panics() {
let mut cfg = small_cfg();
cfg.num_kv_heads = 0;
cfg.validate();
}
#[test]
#[should_panic(expected = "head_dim must be > 0")]
fn test_zero_head_dim_panics() {
let mut cfg = small_cfg();
cfg.head_dim = 0;
cfg.validate();
}
#[test]
#[should_panic(expected = "compress_block (l) must be > 0")]
fn test_zero_compress_block_panics() {
let mut cfg = small_cfg();
cfg.compress_block = 0;
cfg.validate();
}
#[test]
#[should_panic(expected = "compress_stride (d) must be > 0")]
fn test_zero_compress_stride_panics() {
let mut cfg = small_cfg();
cfg.compress_stride = 0;
cfg.validate();
}
#[test]
#[should_panic(expected = "num_heads (2) must be divisible by num_kv_heads (3)")]
fn test_non_divisible_heads_panics() {
let mut cfg = small_cfg();
cfg.num_kv_heads = 3; cfg.validate();
}
#[test]
#[should_panic(expected = "compress_block (l=5) must be divisible by compress_stride (d=2)")]
fn test_l_not_divisible_by_d_panics() {
let mut cfg = small_cfg();
cfg.compress_block = 5; cfg.validate();
}
#[test]
#[should_panic(expected = "select_block (l'=5) must be divisible by compress_stride (d=2)")]
fn test_lp_not_divisible_by_d_panics() {
let mut cfg = small_cfg();
cfg.select_block = 5; cfg.validate();
}
#[test]
#[should_panic(expected = "num_selected (n=2) must be >= 3")]
fn test_num_selected_below_3_panics() {
let mut cfg = small_cfg();
cfg.num_selected = 2; cfg.validate();
}
#[test]
#[should_panic(expected = "compress_stride (d) must be > 0")]
fn test_zero_stride_via_apply_panics() {
let cfg = NsaConfig {
num_heads: 2,
num_kv_heads: 1,
head_dim: 4,
compress_block: 4,
compress_stride: 0, select_block: 4,
num_selected: 3,
window: 4,
};
let weights = NsaWeights {
phi_k_w1: vec![],
phi_k_b1: vec![],
phi_k_w2: vec![],
phi_k_b2: vec![],
phi_v_w1: vec![],
phi_v_b1: vec![],
phi_v_w2: vec![],
phi_v_b2: vec![],
k_intrablock_pos: vec![],
v_intrablock_pos: vec![],
g_proj_w: vec![],
g_proj_b: vec![],
};
let mut out: Vec<f32> = vec![];
let mut scratch = NsaScratch::default();
apply_native_sparse_attention(
&[],
&[],
&[],
&[],
&[],
&[],
&[],
&[],
&[],
&weights,
&mut out,
1,
&cfg,
&mut scratch,
);
}
#[test]
#[should_panic(expected = "x_buf must be non-empty when seq_len > 0")]
fn test_empty_x_buf_nonempty_seq_panics() {
let cfg = small_cfg();
let seq_len = 2;
let weights = make_weights(&cfg, cfg.q_dim(), 1);
let q = det_data(seq_len * cfg.q_dim(), 1);
let q_rope = det_data(seq_len * cfg.q_dim(), 2);
let k_cmp = det_data(seq_len * cfg.kv_dim(), 3);
let k_slc = det_data(seq_len * cfg.kv_dim(), 4);
let k_win = det_data(seq_len * cfg.kv_dim(), 5);
let v_cmp = det_data(seq_len * cfg.kv_dim(), 6);
let v_slc = det_data(seq_len * cfg.kv_dim(), 7);
let v_win = det_data(seq_len * cfg.kv_dim(), 8);
let x: Vec<f32> = vec![]; let mut out = vec![0.0f32; seq_len * cfg.q_dim()];
let mut scratch = NsaScratch::default();
apply_native_sparse_attention(
&q,
&q_rope,
&k_cmp,
&k_slc,
&k_win,
&v_cmp,
&v_slc,
&v_win,
&x,
&weights,
&mut out,
seq_len,
&cfg,
&mut scratch,
);
}
#[test]
fn test_nsa_shapes_seq1() {
let cfg = small_cfg();
let out = run_nsa(&cfg, 1, 42);
assert_eq!(out.len(), cfg.q_dim());
}
#[test]
fn test_nsa_shapes_seq8() {
let cfg = small_cfg();
let out = run_nsa(&cfg, 8, 43);
assert_eq!(out.len(), 8 * cfg.q_dim());
}
#[test]
fn test_nsa_shapes_seq16() {
let cfg = small_cfg();
let out = run_nsa(&cfg, 16, 44);
assert_eq!(out.len(), 16 * cfg.q_dim());
}
#[test]
fn test_nsa_output_finite_small() {
let cfg = small_cfg();
let out = run_nsa(&cfg, 12, 100);
for (i, &v) in out.iter().enumerate() {
assert!(v.is_finite(), "output[{i}] is not finite: {v}");
}
}
#[test]
fn test_nsa_output_finite_seq1() {
let cfg = small_cfg();
let out = run_nsa(&cfg, 1, 101);
for (i, &v) in out.iter().enumerate() {
assert!(v.is_finite(), "output[{i}] is not finite: {v}");
}
}
#[test]
fn test_nsa_seq_zero() {
let cfg = small_cfg();
let model_dim = cfg.q_dim();
let weights = make_weights(&cfg, model_dim, 1);
let mut out: Vec<f32> = vec![];
let mut scratch = NsaScratch::default();
apply_native_sparse_attention(
&[],
&[],
&[],
&[],
&[],
&[],
&[],
&[],
&[],
&weights,
&mut out,
0,
&cfg,
&mut scratch,
);
}
#[test]
fn test_nsa_causal_masking() {
let cfg = small_cfg();
let seq_len = 10usize;
let model_dim = cfg.q_dim();
let weights = make_weights(&cfg, model_dim, 200);
let q = det_data(seq_len * cfg.q_dim(), 201);
let q_rope = det_data(seq_len * cfg.q_dim(), 202);
let x = det_data(seq_len * model_dim, 203);
let kv_base: [Vec<f32>; 6] =
std::array::from_fn(|i| det_data(seq_len * cfg.kv_dim(), 210 + i as u64));
let kv_perturbed: [Vec<f32>; 6] = std::array::from_fn(|i| {
let mut b = kv_base[i].clone();
for pos in 1..seq_len {
for d in 0..cfg.kv_dim() {
b[pos * cfg.kv_dim() + d] += 99_999.0;
}
}
b
});
let run = |kv: &[Vec<f32>; 6]| {
let mut out = vec![0.0f32; seq_len * cfg.q_dim()];
let mut scratch = NsaScratch::default();
apply_native_sparse_attention(
&q,
&q_rope,
&kv[0],
&kv[1],
&kv[2],
&kv[3],
&kv[4],
&kv[5],
&x,
&weights,
&mut out,
seq_len,
&cfg,
&mut scratch,
);
out
};
let out_base = run(&kv_base);
let out_perturbed = run(&kv_perturbed);
for d in 0..cfg.q_dim() {
assert_eq!(
out_base[d].to_bits(),
out_perturbed[d].to_bits(),
"pos 0 changed when only future K/V changed — causal mask leak at dim {d}"
);
}
let any_later_changed = (cfg.q_dim()..seq_len * cfg.q_dim())
.any(|i| out_base[i].to_bits() != out_perturbed[i].to_bits());
assert!(
any_later_changed,
"no later position changed after perturbing future K/V — perturbation is a no-op"
);
}
#[test]
fn test_nsa_early_tokens_finite() {
let cfg = small_cfg();
let out = run_nsa(&cfg, 3, 300);
for (i, &v) in out.iter().enumerate() {
assert!(
v.is_finite(),
"output[{i}] = {v} not finite for early-token sequence"
);
}
}
#[test]
fn test_nsa_deterministic() {
let cfg = small_cfg();
let out1 = run_nsa(&cfg, 8, 400);
let out2 = run_nsa(&cfg, 8, 400);
for (i, (&a, &b)) in out1.iter().zip(out2.iter()).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"output[{i}] differs between identical runs: {a} vs {b}"
);
}
}
#[test]
fn test_nsa_gqa_two_groups() {
let cfg = NsaConfig {
num_heads: 4,
num_kv_heads: 2,
head_dim: 4,
compress_block: 4,
compress_stride: 2,
select_block: 4,
num_selected: 3,
window: 4,
};
let out = run_nsa(&cfg, 10, 500);
assert_eq!(out.len(), 10 * cfg.q_dim());
for (i, &v) in out.iter().enumerate() {
assert!(v.is_finite(), "GQA output[{i}] not finite: {v}");
}
}
#[test]
fn test_nsa_window_only_gate() {
let cfg = small_cfg();
let seq_len = 6usize;
let model_dim = cfg.q_dim();
let mut weights = make_weights(&cfg, model_dim, 600);
weights.g_proj_w.fill(0.0);
for h in 0..cfg.num_heads {
weights.g_proj_b[3 * h] = -100.0; weights.g_proj_b[3 * h + 1] = -100.0; weights.g_proj_b[3 * h + 2] = 100.0; }
let q = det_data(seq_len * cfg.q_dim(), 601);
let q_rope = det_data(seq_len * cfg.q_dim(), 602);
let k_cmp = det_data(seq_len * cfg.kv_dim(), 603);
let k_slc = det_data(seq_len * cfg.kv_dim(), 604);
let k_win = det_data(seq_len * cfg.kv_dim(), 605);
let v_cmp = det_data(seq_len * cfg.kv_dim(), 606);
let v_slc = det_data(seq_len * cfg.kv_dim(), 607);
let v_win = det_data(seq_len * cfg.kv_dim(), 608);
let x = det_data(seq_len * model_dim, 609);
let mut out_nsa = vec![0.0f32; seq_len * cfg.q_dim()];
let mut scratch = NsaScratch::default();
apply_native_sparse_attention(
&q,
&q_rope,
&k_cmp,
&k_slc,
&k_win,
&v_cmp,
&v_slc,
&v_win,
&x,
&weights,
&mut out_nsa,
seq_len,
&cfg,
&mut scratch,
);
let mut out_win = vec![0.0f32; seq_len * cfg.q_dim()];
compute_sliding_window_reference(&q_rope, &k_win, &v_win, &mut out_win, seq_len, &cfg);
let g_win = sigmoid(100.0_f32);
for v in out_win.iter_mut() {
*v *= g_win;
}
for (i, (&a, &b)) in out_nsa.iter().zip(out_win.iter()).enumerate() {
let diff = (a - b).abs();
assert!(
diff < 1e-4,
"window-only gate: output[{i}] mismatch: nsa={a} win_ref={b} diff={diff}"
);
}
}
fn compute_sliding_window_reference(
q_rope: &[f32],
k_rope: &[f32],
v: &[f32],
out: &mut [f32],
seq_len: usize,
cfg: &NsaConfig,
) {
let head_dim = cfg.head_dim;
let q_dim = cfg.q_dim();
let kv_dim = cfg.kv_dim();
let n_rep = cfg.n_rep();
let scale = (head_dim as f32).powf(-0.5);
let win = cfg.window;
out.fill(0.0);
for qt in 0..seq_len {
let win_start = qt.saturating_sub(win - 1);
let win_len = qt - win_start + 1;
for kv_h in 0..cfg.num_kv_heads {
for qh_local in 0..n_rep {
let qh = kv_h * n_rep + qh_local;
let q_off = qt * q_dim + qh * head_dim;
let q_head = &q_rope[q_off..q_off + head_dim];
let mut scores = vec![0.0f32; win_len];
for (wi, tok) in (win_start..=qt).enumerate() {
let k_off = tok * kv_dim + kv_h * head_dim;
let dot: f32 = q_head
.iter()
.zip(k_rope[k_off..k_off + head_dim].iter())
.map(|(&a, &b)| a * b)
.sum();
scores[wi] = dot * scale;
}
softmax_inplace(&mut scores);
let out_off = qt * q_dim + qh * head_dim;
for (wi, tok) in (win_start..=qt).enumerate() {
let v_off = tok * kv_dim + kv_h * head_dim;
for dd in 0..head_dim {
out[out_off + dd] += scores[wi] * v[v_off + dd];
}
}
}
}
}
}
#[test]
fn test_num_compress_blocks() {
let cfg = small_cfg(); assert_eq!(cfg.num_compress_blocks(0), 0);
assert_eq!(cfg.num_compress_blocks(3), 0);
assert_eq!(cfg.num_compress_blocks(4), 1); assert_eq!(cfg.num_compress_blocks(5), 1); assert_eq!(cfg.num_compress_blocks(6), 2); assert_eq!(cfg.num_compress_blocks(10), 4); }
#[test]
fn test_num_select_blocks() {
let cfg = small_cfg(); assert_eq!(cfg.num_select_blocks(0), 0);
assert_eq!(cfg.num_select_blocks(3), 1); assert_eq!(cfg.num_select_blocks(4), 1);
assert_eq!(cfg.num_select_blocks(7), 2); assert_eq!(cfg.num_select_blocks(8), 2);
}
#[test]
fn test_count_valid_compress_blocks() {
assert_eq!(count_valid_compress_blocks(2, 4, 2, 5), 0);
assert_eq!(count_valid_compress_blocks(3, 4, 2, 5), 1);
assert_eq!(count_valid_compress_blocks(4, 4, 2, 5), 1);
assert_eq!(count_valid_compress_blocks(5, 4, 2, 5), 2);
assert_eq!(count_valid_compress_blocks(6, 4, 2, 5), 2);
assert_eq!(count_valid_compress_blocks(100, 4, 2, 5), 5);
}
#[test]
fn test_count_valid_select_blocks() {
assert_eq!(count_valid_select_blocks(0, 4, 5), 1);
assert_eq!(count_valid_select_blocks(3, 4, 5), 1);
assert_eq!(count_valid_select_blocks(4, 4, 5), 2);
assert_eq!(count_valid_select_blocks(8, 4, 5), 3);
assert_eq!(count_valid_select_blocks(100, 4, 5), 5);
}
#[test]
fn test_aggregate_selection_importance_hand_computed() {
let p: Vec<f32> = (0..9).map(|i| (1u32 << i) as f32).collect();
assert_eq!(aggregate_selection_importance(&p, 0, 4, 2), 1.0);
assert_eq!(aggregate_selection_importance(&p, 1, 4, 2), 45.0);
assert_eq!(aggregate_selection_importance(&p, 2, 4, 2), 720.0);
assert_eq!(aggregate_selection_importance(&p, 1, 2, 2), 9.0);
assert_eq!(aggregate_selection_importance(&p[..5], 2, 4, 2), 16.0);
}
#[test]
fn test_softmax_inplace_sums_to_one() {
let mut x = vec![1.0f32, 2.0, 3.0, 0.0, -1.0];
softmax_inplace(&mut x);
let sum: f32 = x.iter().sum();
assert!(
(sum - 1.0).abs() < 1e-6,
"softmax sum should be 1.0, got {sum}"
);
for &v in &x {
assert!(
v >= 0.0 && v.is_finite(),
"softmax output must be non-negative finite"
);
}
}
#[test]
fn test_softmax_inplace_single_element() {
let mut x = vec![42.0f32];
softmax_inplace(&mut x);
assert!((x[0] - 1.0).abs() < 1e-7);
}
#[test]
fn test_softmax_inplace_empty() {
let mut x: Vec<f32> = vec![];
softmax_inplace(&mut x); }
#[test]
fn test_softmax_preserves_relative_order() {
let mut x = vec![3.0f32, 1.0, 2.0];
let orig = x.clone();
softmax_inplace(&mut x);
assert!(x[0] > x[2], "softmax({}) > softmax({})", orig[0], orig[2]);
assert!(x[2] > x[1], "softmax({}) > softmax({})", orig[2], orig[1]);
}
#[test]
fn test_softmax_inplace_nan_fails_closed() {
let mut x = vec![1.0f32, f32::NAN, 2.0, 0.5];
softmax_inplace(&mut x);
assert!(x.iter().all(|&v| v == 0.0), "expected all-zero row: {x:?}");
}
#[test]
fn test_softmax_inplace_pos_inf_fails_closed() {
let mut x = vec![1.0f32, f32::INFINITY, 2.0];
softmax_inplace(&mut x);
assert!(x.iter().all(|&v| v == 0.0), "expected all-zero row: {x:?}");
}
#[test]
fn test_softmax_inplace_all_neg_inf_fails_closed() {
let mut x = vec![f32::NEG_INFINITY; 3];
softmax_inplace(&mut x);
assert!(x.iter().all(|&v| v == 0.0), "expected all-zero row: {x:?}");
}
#[test]
fn test_sigmoid_values() {
assert!((sigmoid(0.0) - 0.5).abs() < 1e-7);
assert!(sigmoid(100.0) > 0.999);
assert!(sigmoid(-100.0) < 0.001);
assert!(sigmoid(1.0) > sigmoid(0.0));
assert!(sigmoid(-1.0) < sigmoid(0.0));
}
#[test]
fn test_nsa_no_panic_various_seq_lens() {
let cfg = small_cfg();
for seq_len in [1, 2, 4, 5, 7, 8, 12, 16, 20] {
let out = run_nsa(&cfg, seq_len, seq_len as u64 + 700);
assert_eq!(
out.len(),
seq_len * cfg.q_dim(),
"shape mismatch at seq_len={seq_len}"
);
for (i, &v) in out.iter().enumerate() {
assert!(
v.is_finite(),
"seq_len={seq_len}: output[{i}] not finite: {v}"
);
}
}
}
}