pub const MINIMAX_BLOCK_SIZE: usize = 128;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct BlockSparseConfig {
pub block_size: usize,
pub top_blocks: usize,
pub init_blocks: usize,
pub local_blocks: usize,
}
impl Default for BlockSparseConfig {
fn default() -> Self {
BlockSparseConfig {
block_size: MINIMAX_BLOCK_SIZE,
top_blocks: 32,
init_blocks: 1,
local_blocks: 2,
}
}
}
impl BlockSparseConfig {
pub fn causal_blocks(&self, query_pos: usize) -> usize {
if self.block_size == 0 {
return 0;
}
query_pos / self.block_size + 1
}
}
pub fn block_sparse_select(
index_q: &[Vec<f32>],
index_k: &[Vec<f32>],
query_pos: usize,
cfg: &BlockSparseConfig,
) -> Vec<Vec<usize>> {
let n_blocks = cfg.causal_blocks(query_pos).min(
if cfg.block_size == 0 {
0
} else {
index_k.len().div_ceil(cfg.block_size)
},
);
if n_blocks == 0 {
return vec![Vec::new(); index_q.len()];
}
index_q
.iter()
.map(|q| select_for_head(q, index_k, query_pos, n_blocks, cfg))
.collect()
}
fn select_for_head(
q: &[f32],
index_k: &[Vec<f32>],
query_pos: usize,
n_blocks: usize,
cfg: &BlockSparseConfig,
) -> Vec<usize> {
let mut forced = vec![false; n_blocks];
for slot in forced.iter_mut().take(cfg.init_blocks.min(n_blocks)) {
*slot = true;
}
for slot in forced
.iter_mut()
.skip(n_blocks.saturating_sub(cfg.local_blocks))
{
*slot = true;
}
let budget = cfg.top_blocks.max(forced.iter().filter(|f| **f).count());
let mut chosen: Vec<usize> = (0..n_blocks).filter(|&b| forced[b]).collect();
if chosen.len() >= budget || chosen.len() == n_blocks {
chosen.sort_unstable();
return chosen;
}
let mut scored: Vec<(usize, f32)> = (0..n_blocks)
.filter(|&b| !forced[b])
.map(|b| {
let start = b * cfg.block_size;
let end = ((b + 1) * cfg.block_size)
.min(index_k.len())
.min(query_pos + 1);
let best = (start..end)
.map(|p| dot(q, &index_k[p]))
.fold(f32::NEG_INFINITY, f32::max);
(b, best)
})
.collect();
scored.sort_by(|a, b| b.1.total_cmp(&a.1).then(a.0.cmp(&b.0)));
for (b, _) in scored.into_iter().take(budget - chosen.len()) {
chosen.push(b);
}
chosen.sort_unstable();
chosen
}
fn dot(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
}
pub fn positions_of_blocks(
blocks: &[usize],
query_pos: usize,
n_positions: usize,
cfg: &BlockSparseConfig,
) -> Vec<usize> {
let mut out = Vec::new();
for &b in blocks {
let start = b * cfg.block_size;
let end = ((b + 1) * cfg.block_size)
.min(n_positions)
.min(query_pos + 1);
out.extend(start..end);
}
out.sort_unstable();
out.dedup();
out
}
#[cfg(test)]
mod tests {
use super::*;
fn keys(n: usize, hot: &[usize]) -> Vec<Vec<f32>> {
(0..n)
.map(|p| {
if hot.contains(&p) {
vec![1.0, 0.0]
} else {
vec![0.0, 0.01]
}
})
.collect()
}
fn cfg(block: usize, top: usize, init: usize, local: usize) -> BlockSparseConfig {
BlockSparseConfig {
block_size: block,
top_blocks: top,
init_blocks: init,
local_blocks: local,
}
}
#[test]
fn a_selection_is_never_empty_however_poor_the_scores() {
let c = cfg(4, 1, 1, 1);
let k: Vec<Vec<f32>> = (0..64).map(|_| vec![0.0, 0.0]).collect();
let q = vec![vec![1.0, 0.0]];
for pos in [0usize, 1, 3, 4, 17, 63] {
let sel = block_sparse_select(&q, &k, pos, &c);
assert!(
!sel[0].is_empty(),
"position {pos} selected nothing, which is a NaN downstream"
);
let covered = positions_of_blocks(&sel[0], pos, k.len(), &c);
assert!(!covered.is_empty(), "position {pos} covers no positions");
assert!(
covered.iter().all(|&p| p <= pos),
"position {pos} saw the future"
);
}
}
#[test]
fn a_block_is_scored_by_its_best_position_not_its_average() {
let c = cfg(4, 3, 1, 1);
let k = keys(16, &[5]);
let q = vec![vec![1.0, 0.0]];
let sel = block_sparse_select(&q, &k, 15, &c);
assert!(
sel[0].contains(&1),
"the block holding the one hot key must be chosen, got {:?}",
sel[0]
);
}
#[test]
fn each_kv_head_selects_for_itself() {
let c = cfg(4, 3, 1, 1);
let mut k: Vec<Vec<f32>> = (0..16).map(|_| vec![0.0, 0.0]).collect();
k[5] = vec![1.0, 0.0]; k[9] = vec![0.0, 1.0]; let q = vec![vec![1.0, 0.0], vec![0.0, 1.0]];
let sel = block_sparse_select(&q, &k, 15, &c);
assert_eq!(sel.len(), 2, "one selection per KV head");
assert!(sel[0].contains(&1), "head 0: {:?}", sel[0]);
assert!(sel[1].contains(&2), "head 1: {:?}", sel[1]);
assert_ne!(sel[0], sel[1], "the heads must not be reduced together");
}
#[test]
fn the_first_and_newest_blocks_are_included_before_scoring() {
let c = cfg(4, 3, 1, 1);
let k = keys(16, &[4, 5, 8, 9]);
let q = vec![vec![1.0, 0.0]];
let sel = block_sparse_select(&q, &k, 15, &c);
assert!(sel[0].contains(&0), "init block missing: {:?}", sel[0]);
assert!(sel[0].contains(&3), "local block missing: {:?}", sel[0]);
}
#[test]
fn the_selection_depends_only_on_the_order_of_the_scores() {
let c = cfg(4, 3, 1, 1);
let k = keys(16, &[5]);
let scaled: Vec<Vec<f32>> = k
.iter()
.map(|row| row.iter().map(|v| v * 1000.0).collect())
.collect();
let q = vec![vec![1.0, 0.0]];
assert_eq!(
block_sparse_select(&q, &k, 15, &c),
block_sparse_select(&q, &scaled, 15, &c)
);
}
#[test]
fn selection_is_causal_and_the_current_block_is_clipped() {
let c = cfg(4, 8, 1, 1);
let k = keys(16, &[]);
let q = vec![vec![1.0, 0.0]];
let sel = block_sparse_select(&q, &k, 6, &c);
assert!(
sel[0].iter().all(|&b| b <= 1),
"block 2 starts at position 8, after the query at 6: {:?}",
sel[0]
);
let covered = positions_of_blocks(&sel[0], 6, k.len(), &c);
assert_eq!(covered, vec![0, 1, 2, 3, 4, 5, 6]);
}
#[test]
fn the_budget_bounds_the_selection_including_the_forced_blocks() {
let c = cfg(4, 3, 1, 1);
let k = keys(64, &[]);
let q = vec![vec![1.0, 0.0]];
let sel = block_sparse_select(&q, &k, 63, &c);
assert_eq!(sel[0].len(), 3, "budget of 3: {:?}", sel[0]);
let tight = cfg(4, 1, 2, 2);
let sel = block_sparse_select(&q, &k, 63, &tight);
assert_eq!(sel[0], vec![0, 1, 14, 15]);
}
#[test]
fn ties_break_deterministically_toward_the_lower_block() {
let c = cfg(4, 3, 0, 0);
let k: Vec<Vec<f32>> = (0..16).map(|_| vec![1.0, 0.0]).collect();
let q = vec![vec![1.0, 0.0]];
let first = block_sparse_select(&q, &k, 15, &c);
assert_eq!(first[0], vec![0, 1, 2]);
for _ in 0..8 {
assert_eq!(block_sparse_select(&q, &k, 15, &c), first);
}
}
#[test]
fn a_query_beyond_the_keys_is_clipped_to_what_exists() {
let c = cfg(4, 8, 1, 1);
let k = keys(6, &[]);
let q = vec![vec![1.0, 0.0]];
let sel = block_sparse_select(&q, &k, 100, &c);
assert_eq!(sel[0], vec![0, 1], "only two blocks of keys exist");
let covered = positions_of_blocks(&sel[0], 100, k.len(), &c);
assert_eq!(covered, vec![0, 1, 2, 3, 4, 5]);
}
#[test]
fn the_real_block_size_is_the_kv_page_size() {
assert_eq!(MINIMAX_BLOCK_SIZE, 128);
let c = BlockSparseConfig::default();
assert_eq!(c.block_size, MINIMAX_BLOCK_SIZE);
assert_eq!(c.causal_blocks(0), 1);
assert_eq!(c.causal_blocks(127), 1);
assert_eq!(c.causal_blocks(128), 2);
}
}