use crate::attention::{causal_mla_attention, causal_mla_attention_sparse, lightning_indexer_topk};
fn concat_raw_and_compressed(raw: &[f32], compressed: &[f32]) -> Vec<f32> {
let mut combined = Vec::with_capacity(raw.len() + compressed.len());
combined.extend_from_slice(raw);
combined.extend_from_slice(compressed);
combined
}
#[allow(clippy::too_many_arguments)]
pub fn hca_attention(
q: &[f32],
raw_k: &[f32],
raw_v: &[f32],
n_raw: usize,
compressed_k: &[f32],
compressed_v: &[f32],
n_compressed: usize,
n_heads: usize,
qk_head_dim: usize,
v_head_dim: usize,
) -> Vec<f32> {
let k_all = concat_raw_and_compressed(raw_k, compressed_k);
let v_all = concat_raw_and_compressed(raw_v, compressed_v);
causal_mla_attention(
q,
&k_all,
&v_all,
n_heads,
qk_head_dim,
v_head_dim,
n_raw + n_compressed,
)
}
#[allow(clippy::too_many_arguments)]
pub fn csa_attention(
q: &[f32],
raw_k: &[f32],
raw_v: &[f32],
n_raw: usize,
compressed_k: &[f32],
compressed_v: &[f32],
n_compressed: usize,
indexer_q: &[Vec<f32>],
indexer_keys: &[Vec<f32>],
indexer_weights: &[f32],
top_k: usize,
n_heads: usize,
qk_head_dim: usize,
v_head_dim: usize,
) -> Vec<f32> {
assert_eq!(indexer_keys.len(), n_compressed);
let selected = lightning_indexer_topk(indexer_q, indexer_keys, indexer_weights, top_k);
let k_all = concat_raw_and_compressed(raw_k, compressed_k);
let v_all = concat_raw_and_compressed(raw_v, compressed_v);
let mut visible: Vec<usize> = (0..n_raw).collect();
visible.extend(selected.iter().map(|&i| n_raw + i));
causal_mla_attention_sparse(
q,
&k_all,
&v_all,
n_heads,
qk_head_dim,
v_head_dim,
n_raw + n_compressed,
&visible,
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn hca_attention_with_single_raw_and_single_compressed_position_matches_full_causal_over_both()
{
let n_heads = 1;
let qk_head_dim = 2;
let v_head_dim = 2;
let q = vec![1.0, 0.0];
let raw_k = vec![1.0, 0.0];
let raw_v = vec![9.0, -3.0];
let compressed_k = vec![0.5, 0.5];
let compressed_v = vec![100.0, 200.0];
let out = hca_attention(
&q,
&raw_k,
&raw_v,
1,
&compressed_k,
&compressed_v,
1,
n_heads,
qk_head_dim,
v_head_dim,
);
let k_all = vec![1.0, 0.0, 0.5, 0.5];
let v_all = vec![9.0, -3.0, 100.0, 200.0];
let expected =
causal_mla_attention(&q, &k_all, &v_all, n_heads, qk_head_dim, v_head_dim, 2);
assert_eq!(out.len(), expected.len());
for (a, b) in out.iter().zip(expected.iter()) {
assert!((a - b).abs() < 1e-6, "a={a} b={b}");
}
}
#[test]
fn hca_attention_gives_nonzero_weight_to_compressed_positions() {
let n_heads = 1;
let qk_head_dim = 2;
let v_head_dim = 1;
let q = vec![1.0, 0.0];
let raw_k = vec![1.0, 0.0];
let raw_v = vec![5.0];
let compressed_k = vec![1.0, 0.0]; let compressed_v = vec![999.0];
let out = hca_attention(
&q,
&raw_k,
&raw_v,
1,
&compressed_k,
&compressed_v,
1,
n_heads,
qk_head_dim,
v_head_dim,
);
assert!((out[0] - 502.0).abs() < 1e-3, "out[0]={}", out[0]);
}
#[test]
fn csa_attention_with_top_k_covering_every_compressed_entry_matches_hca_attention() {
let n_heads = 1;
let qk_head_dim = 2;
let v_head_dim = 1;
let q = vec![0.3, 0.7];
let raw_k = vec![0.2, 0.4, 0.1, 0.9];
let raw_v = vec![1.0, 2.0];
let n_raw = 2;
let compressed_k = vec![0.5, 0.1, 0.05, 0.6, 0.3, 0.3];
let compressed_v = vec![10.0, 20.0, 30.0];
let n_compressed = 3;
let indexer_q = vec![vec![1.0, 0.0]];
let indexer_keys = vec![vec![0.9, 0.1], vec![0.1, 0.9], vec![0.5, 0.5]];
let indexer_weights = vec![1.0];
let csa_out = csa_attention(
&q,
&raw_k,
&raw_v,
n_raw,
&compressed_k,
&compressed_v,
n_compressed,
&indexer_q,
&indexer_keys,
&indexer_weights,
n_compressed, n_heads,
qk_head_dim,
v_head_dim,
);
let hca_out = hca_attention(
&q,
&raw_k,
&raw_v,
n_raw,
&compressed_k,
&compressed_v,
n_compressed,
n_heads,
qk_head_dim,
v_head_dim,
);
assert_eq!(csa_out.len(), hca_out.len());
for (a, b) in csa_out.iter().zip(hca_out.iter()) {
assert!((a - b).abs() < 1e-6, "a={a} b={b}");
}
}
#[test]
fn csa_attention_ignores_compressed_entries_the_indexer_does_not_select() {
let n_heads = 1;
let qk_head_dim = 2;
let v_head_dim = 1;
let q = vec![1.0, 0.0];
let raw_k = vec![1.0, 0.0];
let raw_v = vec![5.0];
let compressed_k = vec![1.0, 0.0, 1.0, 0.0]; let compressed_v = vec![7.0, 99999.0];
let indexer_q = vec![vec![1.0, 0.0]];
let indexer_keys = vec![vec![1.0, 0.0], vec![0.0, 1.0]];
let indexer_weights = vec![1.0];
let out = csa_attention(
&q,
&raw_k,
&raw_v,
1,
&compressed_k,
&compressed_v,
2,
&indexer_q,
&indexer_keys,
&indexer_weights,
1,
n_heads,
qk_head_dim,
v_head_dim,
);
assert!((out[0] - 6.0).abs() < 1e-3, "out[0]={}", out[0]);
}
}