#![allow(clippy::needless_range_loop)]
use super::tensor::Mat;
use crate::error::{FocrError, FocrResult};
use crate::quant::recipe::is_truthy;
use rayon::prelude::*;
use std::sync::atomic::{AtomicU64, Ordering};
pub const NUM_HEADS: usize = 10;
pub const HEAD_DIM: usize = 128;
pub const RING_WINDOW: usize = 128;
static NEXT_RING_CACHE_ID: AtomicU64 = AtomicU64::new(1);
fn next_ring_cache_id() -> u64 {
let mut current = NEXT_RING_CACHE_ID.load(Ordering::Relaxed);
loop {
let next = current
.checked_add(1)
.expect("rswa: exhausted all u64 RingCache identities");
match NEXT_RING_CACHE_ID.compare_exchange_weak(
current,
next,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => return current,
Err(observed) => current = observed,
}
}
}
#[inline]
#[must_use]
fn scale() -> f32 {
1.0 / (HEAD_DIM as f32).sqrt()
}
const ATTN_GEMM_ENV: &str = "FOCR_ATTN_GEMM";
const INT8_KV_ENV: &str = "FOCR_INT8_KV";
const PARALLEL_HEADS_ENV: &str = "FOCR_RSWA_PARALLEL_ATTN";
fn risky_opt_in_enabled_for(value: Option<&str>) -> bool {
value.is_some_and(is_truthy)
}
fn attn_gemm_enabled() -> bool {
use std::sync::OnceLock;
static FLAG: OnceLock<bool> = OnceLock::new();
*FLAG.get_or_init(|| {
let value = std::env::var(ATTN_GEMM_ENV).ok();
risky_opt_in_enabled_for(value.as_deref())
})
}
fn int8_kv_enabled() -> bool {
use std::sync::OnceLock;
static FLAG: OnceLock<bool> = OnceLock::new();
*FLAG.get_or_init(|| {
let value = std::env::var(INT8_KV_ENV).ok();
risky_opt_in_enabled_for(value.as_deref())
})
}
fn parallel_heads_enabled() -> bool {
use std::sync::OnceLock;
static FLAG: OnceLock<bool> = OnceLock::new();
*FLAG.get_or_init(|| {
match std::env::var(PARALLEL_HEADS_ENV)
.ok()
.map(|value| value.trim().to_ascii_lowercase())
.as_deref()
{
Some("0" | "off" | "false" | "no") => false,
Some("1" | "on" | "true" | "yes") => true,
_ => true,
}
})
}
fn checked_cache_region_elems(rows: usize) -> usize {
rows.checked_mul(HEAD_DIM)
.expect("rswa: cache rows*HEAD_DIM overflow")
}
fn checked_head_major_layout(seq: usize, label: &str) -> FocrResult<(usize, usize)> {
let stride = seq.checked_mul(HEAD_DIM).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"rswa: {label} seq*HEAD_DIM overflow (seq={seq}, HEAD_DIM={HEAD_DIM})"
))
})?;
let expect = NUM_HEADS.checked_mul(stride).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"rswa: {label} NUM_HEADS*seq*HEAD_DIM overflow \
(NUM_HEADS={NUM_HEADS}, seq={seq}, HEAD_DIM={HEAD_DIM})"
))
})?;
Ok((stride, expect))
}
#[derive(Debug)]
pub struct RingCache {
cache_id: u64,
ref_capacity: usize,
ref_k: Vec<Vec<f32>>,
ref_v: Vec<Vec<f32>>,
ring_k: Vec<Vec<f32>>,
ring_v: Vec<Vec<f32>>,
prefill_len: Option<usize>,
ring_len: usize,
ring_pos: usize,
decode_writes: usize,
lineage_epoch: u64,
int8: Option<Int8Kv>,
}
impl Clone for RingCache {
fn clone(&self) -> Self {
Self {
cache_id: next_ring_cache_id(),
ref_capacity: self.ref_capacity,
ref_k: self.ref_k.clone(),
ref_v: self.ref_v.clone(),
ring_k: self.ring_k.clone(),
ring_v: self.ring_v.clone(),
prefill_len: self.prefill_len,
ring_len: self.ring_len,
ring_pos: self.ring_pos,
decode_writes: self.decode_writes,
lineage_epoch: self.lineage_epoch,
int8: self.int8.clone(),
}
}
}
#[derive(Debug, Clone)]
struct Int8Kv {
ref_k: Vec<Vec<i8>>,
ref_k_scale: Vec<Vec<f32>>,
ref_v: Vec<Vec<i8>>,
ref_v_scale: Vec<Vec<f32>>,
ring_k: Vec<Vec<i8>>,
ring_k_scale: Vec<Vec<f32>>,
ring_v: Vec<Vec<i8>>,
ring_v_scale: Vec<Vec<f32>>,
}
impl Int8Kv {
fn new(ref_capacity: usize) -> Self {
let ref_elems = checked_cache_region_elems(ref_capacity);
let ring_elems = checked_cache_region_elems(RING_WINDOW);
let i8_ref = || (0..NUM_HEADS).map(|_| vec![0i8; ref_elems]).collect();
let sc_ref = || (0..NUM_HEADS).map(|_| vec![0.0f32; ref_capacity]).collect();
let i8_ring = || (0..NUM_HEADS).map(|_| vec![0i8; ring_elems]).collect();
let sc_ring = || (0..NUM_HEADS).map(|_| vec![0.0f32; RING_WINDOW]).collect();
Self {
ref_k: i8_ref(),
ref_k_scale: sc_ref(),
ref_v: i8_ref(),
ref_v_scale: sc_ref(),
ring_k: i8_ring(),
ring_k_scale: sc_ring(),
ring_v: i8_ring(),
ring_v_scale: sc_ring(),
}
}
}
fn quantize_row_i8(row: &[f32], q: &mut [i8]) -> f32 {
let maxabs = row.iter().fold(0.0f32, |m, &x| m.max(x.abs()));
let scale = if maxabs > 0.0 { maxabs / 127.0 } else { 1.0 };
for (qd, &x) in q.iter_mut().zip(row.iter()) {
*qd = (x / scale).round_ties_even().clamp(-127.0, 127.0) as i8;
}
scale
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RingCheckpoint {
cache_id: u64,
prefill_len: Option<usize>,
ring_len: usize,
ring_pos: usize,
decode_writes: usize,
lineage_epoch: u64,
}
impl RingCheckpoint {
#[must_use]
pub fn prefill_len(&self) -> Option<usize> {
self.prefill_len
}
#[must_use]
pub fn ring_len(&self) -> usize {
self.ring_len
}
#[must_use]
pub fn ring_pos(&self) -> usize {
self.ring_pos
}
#[must_use]
pub fn effective_len(&self) -> usize {
self.prefill_len.unwrap_or(0) + self.ring_len
}
}
impl RingCache {
#[must_use]
pub fn new(prefill_capacity: usize) -> Self {
Self::new_inner(prefill_capacity, int8_kv_enabled())
}
#[must_use]
fn new_inner(prefill_capacity: usize, with_int8: bool) -> Self {
let ref_elems = checked_cache_region_elems(prefill_capacity);
let ring_elems = checked_cache_region_elems(RING_WINDOW);
Self {
cache_id: next_ring_cache_id(),
ref_capacity: prefill_capacity,
ref_k: (0..NUM_HEADS).map(|_| vec![0.0f32; ref_elems]).collect(),
ref_v: (0..NUM_HEADS).map(|_| vec![0.0f32; ref_elems]).collect(),
ring_k: (0..NUM_HEADS).map(|_| vec![0.0f32; ring_elems]).collect(),
ring_v: (0..NUM_HEADS).map(|_| vec![0.0f32; ring_elems]).collect(),
prefill_len: None,
ring_len: 0,
ring_pos: 0,
decode_writes: 0,
lineage_epoch: 0,
int8: with_int8.then(|| Int8Kv::new(prefill_capacity)),
}
}
#[must_use]
pub fn ref_capacity(&self) -> usize {
self.ref_capacity
}
#[must_use]
pub fn prefill_len(&self) -> Option<usize> {
self.prefill_len
}
#[must_use]
pub fn reference_k(&self, h: usize) -> &[f32] {
&self.ref_k[h][..self.prefill_len.unwrap_or(0) * HEAD_DIM]
}
#[must_use]
pub fn reference_v(&self, h: usize) -> &[f32] {
&self.ref_v[h][..self.prefill_len.unwrap_or(0) * HEAD_DIM]
}
#[must_use]
pub fn ring_len(&self) -> usize {
self.ring_len
}
#[must_use]
pub fn ring_pos(&self) -> usize {
self.ring_pos
}
#[must_use]
pub fn is_warm(&self) -> bool {
self.ring_len >= RING_WINDOW
}
#[must_use]
pub fn effective_len(&self) -> usize {
self.prefill_len.unwrap_or(0) + self.ring_len
}
pub fn record_prefill(&mut self, k: &[f32], v: &[f32], seq: usize) -> FocrResult<()> {
if seq > self.ref_capacity {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: prefill seq {seq} exceeds ref_capacity {}",
self.ref_capacity
)));
}
let (stride, expect) = checked_head_major_layout(seq, "prefill")?;
if k.len() != expect || v.len() != expect {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: prefill k/v len {}/{} != NUM_HEADS*seq*HEAD_DIM {}",
k.len(),
v.len(),
expect
)));
}
let next_lineage_epoch = self.lineage_epoch.checked_add(1).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"rswa: checkpoint lineage epoch exhausted during record_prefill"
))
})?;
for h in 0..NUM_HEADS {
let src = &k[h * stride..(h + 1) * stride];
self.ref_k[h][..stride].copy_from_slice(src);
let src = &v[h * stride..(h + 1) * stride];
self.ref_v[h][..stride].copy_from_slice(src);
}
if let Some(i8kv) = self.int8.as_mut() {
for h in 0..NUM_HEADS {
for r in 0..seq {
let off = r * HEAD_DIM;
i8kv.ref_k_scale[h][r] = quantize_row_i8(
&self.ref_k[h][off..off + HEAD_DIM],
&mut i8kv.ref_k[h][off..off + HEAD_DIM],
);
i8kv.ref_v_scale[h][r] = quantize_row_i8(
&self.ref_v[h][off..off + HEAD_DIM],
&mut i8kv.ref_v[h][off..off + HEAD_DIM],
);
}
}
}
self.prefill_len = Some(seq);
self.ring_len = 0;
self.ring_pos = 0;
self.decode_writes = 0;
self.lineage_epoch = next_lineage_epoch;
Ok(())
}
pub fn write_decode_step(&mut self, k_step: &[f32], v_step: &[f32]) -> FocrResult<usize> {
if self.prefill_len.is_none() {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: write_decode_step before record_prefill"
)));
}
let expect = NUM_HEADS * HEAD_DIM;
if k_step.len() != expect || v_step.len() != expect {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: decode step k/v len {}/{} != NUM_HEADS*HEAD_DIM {}",
k_step.len(),
v_step.len(),
expect
)));
}
let next_decode_writes = self.decode_writes.checked_add(1).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"rswa: decode-write counter exhausted before ring mutation"
))
})?;
let slot = if self.ring_len < RING_WINDOW {
let slot = self.ring_len;
for h in 0..NUM_HEADS {
let off = slot * HEAD_DIM;
let src = &k_step[h * HEAD_DIM..(h + 1) * HEAD_DIM];
self.ring_k[h][off..off + HEAD_DIM].copy_from_slice(src);
let src = &v_step[h * HEAD_DIM..(h + 1) * HEAD_DIM];
self.ring_v[h][off..off + HEAD_DIM].copy_from_slice(src);
}
self.ring_len += 1;
self.ring_pos = self.ring_len % RING_WINDOW;
slot
} else {
let slot = self.ring_pos;
for h in 0..NUM_HEADS {
let off = slot * HEAD_DIM;
let src = &k_step[h * HEAD_DIM..(h + 1) * HEAD_DIM];
self.ring_k[h][off..off + HEAD_DIM].copy_from_slice(src);
let src = &v_step[h * HEAD_DIM..(h + 1) * HEAD_DIM];
self.ring_v[h][off..off + HEAD_DIM].copy_from_slice(src);
}
self.ring_pos = (self.ring_pos + 1) % RING_WINDOW;
slot
};
if let Some(i8kv) = self.int8.as_mut() {
let off = slot * HEAD_DIM;
for h in 0..NUM_HEADS {
i8kv.ring_k_scale[h][slot] = quantize_row_i8(
&self.ring_k[h][off..off + HEAD_DIM],
&mut i8kv.ring_k[h][off..off + HEAD_DIM],
);
i8kv.ring_v_scale[h][slot] = quantize_row_i8(
&self.ring_v[h][off..off + HEAD_DIM],
&mut i8kv.ring_v[h][off..off + HEAD_DIM],
);
}
}
self.decode_writes = next_decode_writes;
Ok(slot)
}
#[must_use]
pub fn checkpoint(&self) -> RingCheckpoint {
RingCheckpoint {
cache_id: self.cache_id,
prefill_len: self.prefill_len,
ring_len: self.ring_len,
ring_pos: self.ring_pos,
decode_writes: self.decode_writes,
lineage_epoch: self.lineage_epoch,
}
}
pub fn rollback_to(&mut self, cp: &RingCheckpoint) -> FocrResult<()> {
if cp.cache_id != self.cache_id {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: rollback_to checkpoint belongs to a different cache instance (cp.cache_id={} != cache.id={})",
cp.cache_id,
self.cache_id
)));
}
if cp.prefill_len != self.prefill_len {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: rollback_to checkpoint prefill_len {:?} != cache {:?}",
cp.prefill_len,
self.prefill_len
)));
}
if cp.lineage_epoch != self.lineage_epoch {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: rollback_to stale checkpoint lineage (cp.epoch={} != cache.epoch={}); the checkpoint belongs to an abandoned rollback branch or an earlier prefill",
cp.lineage_epoch,
self.lineage_epoch
)));
}
if cp.decode_writes > self.decode_writes {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: rollback_to checkpoint is from the future \
(cp.decode_writes={} > cache={})",
cp.decode_writes,
self.decode_writes
)));
}
let discarded = self.decode_writes - cp.decode_writes;
if cp.ring_len > RING_WINDOW || discarded > RING_WINDOW - cp.ring_len {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: rollback_to is not lossless — {discarded} discarded steps \
wrapped the ring past checkpoint.ring_len={} (W={RING_WINDOW}); an \
evicted slot's prior K/V cannot be restored from cursors alone",
cp.ring_len
)));
}
let next_lineage_epoch = if discarded == 0 {
self.lineage_epoch
} else {
self.lineage_epoch.checked_add(1).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"rswa: checkpoint lineage epoch exhausted before rollback"
))
})?
};
self.ring_len = cp.ring_len;
self.ring_pos = cp.ring_pos;
self.decode_writes = cp.decode_writes;
self.lineage_epoch = next_lineage_epoch;
Ok(())
}
#[inline]
fn ref_k_row(&self, h: usize, r: usize) -> &[f32] {
let off = r * HEAD_DIM;
&self.ref_k[h][off..off + HEAD_DIM]
}
#[inline]
fn ref_v_row(&self, h: usize, r: usize) -> &[f32] {
let off = r * HEAD_DIM;
&self.ref_v[h][off..off + HEAD_DIM]
}
#[inline]
fn ring_k_row(&self, h: usize, r: usize) -> &[f32] {
let off = r * HEAD_DIM;
&self.ring_k[h][off..off + HEAD_DIM]
}
#[inline]
fn ring_v_row(&self, h: usize, r: usize) -> &[f32] {
let off = r * HEAD_DIM;
&self.ring_v[h][off..off + HEAD_DIM]
}
}
#[inline]
fn dot(a: &[f32], b: &[f32]) -> f32 {
let mut acc = 0.0f32;
for i in 0..HEAD_DIM {
acc += a[i] * b[i];
}
acc
}
#[derive(Debug)]
pub struct BatchedRingCache {
streams: Vec<Vec<RingCache>>,
n_layers: usize,
}
impl BatchedRingCache {
#[must_use]
pub fn new(prefill_caps: &[usize], n_layers: usize) -> Self {
assert!(n_layers > 0, "BatchedRingCache: n_layers must be > 0");
assert!(
!prefill_caps.is_empty(),
"BatchedRingCache: needs at least one stream"
);
let streams = prefill_caps
.iter()
.map(|&cap| (0..n_layers).map(|_| RingCache::new(cap)).collect())
.collect();
Self { streams, n_layers }
}
pub fn from_streams(streams: Vec<Vec<RingCache>>) -> FocrResult<Self> {
let Some(first) = streams.first() else {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: BatchedRingCache::from_streams needs at least one stream"
)));
};
let n_layers = first.len();
if n_layers == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: BatchedRingCache::from_streams stream 0 has zero layers"
)));
}
for (s, layers) in streams.iter().enumerate() {
if layers.len() != n_layers {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: BatchedRingCache::from_streams stream {s} has {} layers != {n_layers}",
layers.len()
)));
}
}
Ok(Self { streams, n_layers })
}
#[must_use]
pub fn num_streams(&self) -> usize {
self.streams.len()
}
#[must_use]
pub fn num_layers(&self) -> usize {
self.n_layers
}
#[must_use]
pub fn stream(&self, s: usize) -> &[RingCache] {
&self.streams[s]
}
#[must_use]
pub fn layer(&self, s: usize, layer: usize) -> &RingCache {
&self.streams[s][layer]
}
pub fn layer_mut(&mut self, s: usize, layer: usize) -> &mut RingCache {
&mut self.streams[s][layer]
}
pub fn record_prefill(
&mut self,
s: usize,
layer: usize,
k: &[f32],
v: &[f32],
seq: usize,
) -> FocrResult<()> {
self.streams[s][layer].record_prefill(k, v, seq)
}
pub fn write_decode_step(
&mut self,
s: usize,
layer: usize,
k_step: &[f32],
v_step: &[f32],
) -> FocrResult<usize> {
self.streams[s][layer].write_decode_step(k_step, v_step)
}
#[must_use]
pub fn checkpoint(&self, s: usize, layer: usize) -> RingCheckpoint {
self.streams[s][layer].checkpoint()
}
pub fn rollback_to(&mut self, s: usize, layer: usize, cp: &RingCheckpoint) -> FocrResult<()> {
self.streams[s][layer].rollback_to(cp)
}
#[must_use]
pub fn kv_f32_bytes(&self) -> usize {
const F32: usize = core::mem::size_of::<f32>();
let mut bytes = 0usize;
for stream in &self.streams {
for cache in stream {
let rows = cache.ref_capacity() + RING_WINDOW;
bytes += 2 * NUM_HEADS * rows * HEAD_DIM * F32;
}
}
bytes
}
}
pub fn decode_attention(cache: &RingCache, q: &[f32]) -> FocrResult<Mat> {
if int8_kv_enabled() && cache.int8.is_some() {
decode_attention_int8(cache, q)
} else if attn_gemm_enabled() {
decode_attention_gemm(cache, q)
} else if parallel_heads_enabled() {
decode_attention_scalar_parallel(cache, q)
} else {
decode_attention_scalar(cache, q)
}
}
fn decode_dims(cache: &RingCache, q: &[f32]) -> FocrResult<(usize, usize)> {
let Some(prefill_len) = cache.prefill_len else {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: decode_attention before record_prefill"
)));
};
let expect = NUM_HEADS * HEAD_DIM;
if q.len() != expect {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: query len {} != NUM_HEADS*HEAD_DIM {}",
q.len(),
expect
)));
}
let ring_len = cache.ring_len;
if prefill_len + ring_len == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: empty attention key set (prefill_len=0, ring_len=0)"
)));
}
Ok((prefill_len, ring_len))
}
fn decode_attention_scalar(cache: &RingCache, q: &[f32]) -> FocrResult<Mat> {
let (prefill_len, ring_len) = decode_dims(cache, q)?;
let s = scale();
let mut out = vec![0.0f32; NUM_HEADS * HEAD_DIM];
for h in 0..NUM_HEADS {
let qh = &q[h * HEAD_DIM..(h + 1) * HEAD_DIM];
let dst = &mut out[h * HEAD_DIM..(h + 1) * HEAD_DIM];
decode_attention_scalar_head(cache, h, qh, prefill_len, ring_len, s, dst);
}
Ok(Mat::from_vec(1, NUM_HEADS * HEAD_DIM, out))
}
fn decode_attention_scalar_parallel(cache: &RingCache, q: &[f32]) -> FocrResult<Mat> {
let (prefill_len, ring_len) = decode_dims(cache, q)?;
let s = scale();
let mut out = vec![0.0f32; NUM_HEADS * HEAD_DIM];
out.par_chunks_mut(HEAD_DIM)
.enumerate()
.for_each(|(h, dst)| {
let qh = &q[h * HEAD_DIM..(h + 1) * HEAD_DIM];
decode_attention_scalar_head(cache, h, qh, prefill_len, ring_len, s, dst);
});
Ok(Mat::from_vec(1, NUM_HEADS * HEAD_DIM, out))
}
#[inline]
fn decode_attention_scalar_head(
cache: &RingCache,
h: usize,
qh: &[f32],
prefill_len: usize,
ring_len: usize,
s: f32,
dst: &mut [f32],
) {
debug_assert_eq!(qh.len(), HEAD_DIM);
debug_assert_eq!(dst.len(), HEAD_DIM);
let mut run_max = f32::NEG_INFINITY;
let mut run_den = 0.0f32;
let mut acc = [0.0f32; HEAD_DIM];
for r in 0..prefill_len {
let score = dot(qh, cache.ref_k_row(h, r)) * s;
fold(
&mut run_max,
&mut run_den,
&mut acc,
score,
cache.ref_v_row(h, r),
);
}
for r in 0..ring_len {
let score = dot(qh, cache.ring_k_row(h, r)) * s;
fold(
&mut run_max,
&mut run_den,
&mut acc,
score,
cache.ring_v_row(h, r),
);
}
let inv = if run_den > 0.0 { 1.0 / run_den } else { 0.0 };
for i in 0..HEAD_DIM {
dst[i] = acc[i] * inv;
}
}
#[inline]
fn block_scores(q: &[f32], keys: &[f32], rows: usize, scale: f32, out: &mut [f32]) {
for r in 0..rows {
let krow = &keys[r * HEAD_DIM..(r + 1) * HEAD_DIM];
let mut acc = 0.0f32;
for d in 0..HEAD_DIM {
acc += q[d] * krow[d];
}
out[r] = acc * scale;
}
}
#[inline]
fn block_accumulate(probs: &[f32], vals: &[f32], rows: usize, acc: &mut [f32; HEAD_DIM]) {
for r in 0..rows {
let vrow = &vals[r * HEAD_DIM..(r + 1) * HEAD_DIM];
let p = probs[r];
for d in 0..HEAD_DIM {
acc[d] += p * vrow[d];
}
}
}
#[inline]
fn softmax_inplace(scores: &mut [f32]) -> f32 {
let mut mx = f32::NEG_INFINITY;
for &sc in scores.iter() {
if sc > mx {
mx = sc;
}
}
let mut den = 0.0f32;
for sc in scores.iter_mut() {
let w = (*sc - mx).exp();
*sc = w;
den += w;
}
den
}
fn decode_attention_gemm(cache: &RingCache, q: &[f32]) -> FocrResult<Mat> {
let (prefill_len, ring_len) = decode_dims(cache, q)?;
let total = prefill_len + ring_len;
let s = scale();
let mut out = vec![0.0f32; NUM_HEADS * HEAD_DIM];
let mut scores = vec![0.0f32; total];
for h in 0..NUM_HEADS {
let qh = &q[h * HEAD_DIM..(h + 1) * HEAD_DIM];
block_scores(
qh,
&cache.ref_k[h][..prefill_len * HEAD_DIM],
prefill_len,
s,
&mut scores[..prefill_len],
);
block_scores(
qh,
&cache.ring_k[h][..ring_len * HEAD_DIM],
ring_len,
s,
&mut scores[prefill_len..total],
);
let den = softmax_inplace(&mut scores[..total]);
let mut acc = [0.0f32; HEAD_DIM];
block_accumulate(
&scores[..prefill_len],
&cache.ref_v[h][..prefill_len * HEAD_DIM],
prefill_len,
&mut acc,
);
block_accumulate(
&scores[prefill_len..total],
&cache.ring_v[h][..ring_len * HEAD_DIM],
ring_len,
&mut acc,
);
let inv = if den > 0.0 { 1.0 / den } else { 0.0 };
let dst = &mut out[h * HEAD_DIM..(h + 1) * HEAD_DIM];
for i in 0..HEAD_DIM {
dst[i] = acc[i] * inv;
}
}
Ok(Mat::from_vec(1, NUM_HEADS * HEAD_DIM, out))
}
fn decode_attention_int8(cache: &RingCache, q: &[f32]) -> FocrResult<Mat> {
let (prefill_len, ring_len) = decode_dims(cache, q)?;
let Some(i8kv) = cache.int8.as_ref() else {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: decode_attention_int8 without an int8 KV mirror (FOCR_INT8_KV)"
)));
};
let total = prefill_len + ring_len;
let s = scale();
let mut out = vec![0.0f32; NUM_HEADS * HEAD_DIM];
let mut scores = vec![0.0f32; total];
let mut qi8 = [0i8; HEAD_DIM];
let mut acc_i32 = vec![0i32; total];
for h in 0..NUM_HEADS {
let qh = &q[h * HEAD_DIM..(h + 1) * HEAD_DIM];
let qscale = quantize_row_i8(qh, &mut qi8);
if prefill_len > 0 {
let dst = &mut acc_i32[..prefill_len];
dst.fill(0);
crate::simd::igemm_s8s8(
&qi8,
&i8kv.ref_k[h][..prefill_len * HEAD_DIM],
1,
HEAD_DIM,
prefill_len,
dst,
);
let k_scale = &i8kv.ref_k_scale[h];
for r in 0..prefill_len {
scores[r] = acc_i32[r] as f32 * qscale * k_scale[r] * s;
}
}
if ring_len > 0 {
let dst = &mut acc_i32[..ring_len];
dst.fill(0);
crate::simd::igemm_s8s8(
&qi8,
&i8kv.ring_k[h][..ring_len * HEAD_DIM],
1,
HEAD_DIM,
ring_len,
dst,
);
let k_scale = &i8kv.ring_k_scale[h];
for r in 0..ring_len {
scores[prefill_len + r] = acc_i32[r] as f32 * qscale * k_scale[r] * s;
}
}
let den = softmax_inplace(&mut scores[..total]);
let mut acc = [0.0f32; HEAD_DIM];
for r in 0..prefill_len {
let vrow = &i8kv.ref_v[h][r * HEAD_DIM..(r + 1) * HEAD_DIM];
let pw = scores[r] * i8kv.ref_v_scale[h][r];
for d in 0..HEAD_DIM {
acc[d] += pw * f32::from(vrow[d]);
}
}
for r in 0..ring_len {
let vrow = &i8kv.ring_v[h][r * HEAD_DIM..(r + 1) * HEAD_DIM];
let pw = scores[prefill_len + r] * i8kv.ring_v_scale[h][r];
for d in 0..HEAD_DIM {
acc[d] += pw * f32::from(vrow[d]);
}
}
let inv = if den > 0.0 { 1.0 / den } else { 0.0 };
let dst = &mut out[h * HEAD_DIM..(h + 1) * HEAD_DIM];
for i in 0..HEAD_DIM {
dst[i] = acc[i] * inv;
}
}
Ok(Mat::from_vec(1, NUM_HEADS * HEAD_DIM, out))
}
fn validate_decode_step_mat(label: &str, mat: &Mat, expect: usize) -> FocrResult<()> {
if mat.rows != 1 || mat.cols != expect {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: attention expects {label} [1, {expect}], got [{},{}]",
mat.rows,
mat.cols
)));
}
if mat.data.len() != expect {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: attention {label} data len {} != NUM_HEADS*HEAD_DIM {}",
mat.data.len(),
expect
)));
}
Ok(())
}
#[inline]
fn fold(run_max: &mut f32, run_den: &mut f32, acc: &mut [f32; HEAD_DIM], score: f32, v: &[f32]) {
if score > *run_max {
let correction = if run_max.is_finite() {
(*run_max - score).exp()
} else {
0.0
};
*run_den *= correction;
for a in acc.iter_mut() {
*a *= correction;
}
*run_max = score;
}
let w = (score - *run_max).exp();
*run_den += w;
for i in 0..HEAD_DIM {
acc[i] += w * v[i];
}
}
pub fn attention(
cache: &mut RingCache,
q: &Mat,
k: &Mat,
v: &Mat,
position_ids: &[usize],
) -> FocrResult<Mat> {
let expect = NUM_HEADS * HEAD_DIM;
validate_decode_step_mat("q", q, expect)?;
validate_decode_step_mat("k", k, expect)?;
validate_decode_step_mat("v", v, expect)?;
if position_ids.len() != 1 {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: attention expects a single decode position_id, got {}",
position_ids.len()
)));
}
cache.write_decode_step(&k.data, &v.data)?;
decode_attention(cache, &q.data)
}
pub fn verify_attention(
cache: &RingCache,
draft_q: &[&[f32]],
draft_k: &[&[f32]],
draft_v: &[&[f32]],
) -> FocrResult<Vec<Mat>> {
let k = draft_q.len();
if draft_k.len() != k || draft_v.len() != k {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: verify_attention draft q/k/v counts disagree ({k}, {}, {})",
draft_k.len(),
draft_v.len()
)));
}
let expect = NUM_HEADS * HEAD_DIM;
for (i, q) in draft_q.iter().enumerate() {
if q.len() != expect {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: verify_attention draft_q[{i}] len {} != NUM_HEADS*HEAD_DIM {expect}",
q.len()
)));
}
}
let mut work = cache.clone();
let mut out = Vec::with_capacity(k);
for i in 0..k {
work.write_decode_step(draft_k[i], draft_v[i])?;
out.push(decode_attention(&work, draft_q[i])?);
}
Ok(out)
}
pub fn batched_decode_attention(
cache: &BatchedRingCache,
layer: usize,
queries: &[f32],
) -> FocrResult<Vec<Mat>> {
if layer >= cache.num_layers() {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: batched_decode_attention layer {layer} >= num_layers {}",
cache.num_layers()
)));
}
let b = cache.num_streams();
let per_stream = NUM_HEADS * HEAD_DIM;
let expect = b * per_stream;
if queries.len() != expect {
return Err(FocrError::Other(anyhow::anyhow!(
"rswa: batched_decode_attention queries len {} != B*NUM_HEADS*HEAD_DIM {} \
(B={b})",
queries.len(),
expect
)));
}
let mut contexts = Vec::with_capacity(b);
for s in 0..b {
let q_s = &queries[s * per_stream..(s + 1) * per_stream];
contexts.push(decode_attention(cache.layer(s, layer), q_s)?);
}
Ok(contexts)
}
#[cfg(test)]
mod tests {
use super::*;
fn assert_err_contains<T>(res: FocrResult<T>, needle: &str) {
let message = match res {
Ok(_) => String::from("<ok>"),
Err(err) => err.to_string(),
};
assert!(
message.contains(needle),
"error {message:?} did not contain {needle:?}"
);
}
fn fill_head_major(seq: usize, f: impl Fn(usize, usize) -> f32) -> Vec<f32> {
let mut out = vec![0.0f32; NUM_HEADS * seq * HEAD_DIM];
for h in 0..NUM_HEADS {
for r in 0..seq {
for d in 0..HEAD_DIM {
out[(h * seq + r) * HEAD_DIM + d] = f(r, d);
}
}
}
out
}
fn one_token(f: impl Fn(usize, usize) -> f32) -> Vec<f32> {
let mut out = vec![0.0f32; NUM_HEADS * HEAD_DIM];
for h in 0..NUM_HEADS {
for d in 0..HEAD_DIM {
out[h * HEAD_DIM + d] = f(h, d);
}
}
out
}
#[test]
fn constants_match_spec() {
assert_eq!(NUM_HEADS, 10);
assert_eq!(HEAD_DIM, 128);
assert_eq!(RING_WINDOW, 128);
assert!((scale() - 1.0 / (128.0f32).sqrt()).abs() < 1e-9);
}
#[test]
fn new_allocates_worst_case() {
let cache = RingCache::new(32768 - 128);
assert_eq!(cache.ref_capacity(), 32768 - 128);
assert_eq!(cache.prefill_len(), None);
assert_eq!(cache.ring_len(), 0);
assert_eq!(cache.ring_pos(), 0);
assert!(!cache.is_warm());
assert_eq!(cache.ring_k.len(), NUM_HEADS);
assert_eq!(cache.ring_k[0].len(), RING_WINDOW * HEAD_DIM);
}
#[test]
#[should_panic(expected = "cache rows*HEAD_DIM overflow")]
fn new_rejects_ref_capacity_shape_overflow_before_allocating() {
let _ = RingCache::new(usize::MAX / HEAD_DIM + 1);
}
#[test]
fn head_major_layout_rejects_stride_overflow() {
let err = checked_head_major_layout(usize::MAX / HEAD_DIM + 1, "test")
.expect_err("seq*HEAD_DIM should overflow");
assert!(err.to_string().contains("seq*HEAD_DIM overflow"));
}
#[test]
fn head_major_layout_rejects_total_overflow() {
let seq = (usize::MAX / HEAD_DIM) / NUM_HEADS + 1;
let err = checked_head_major_layout(seq, "test")
.expect_err("NUM_HEADS*seq*HEAD_DIM should overflow");
assert!(err.to_string().contains("NUM_HEADS*seq*HEAD_DIM overflow"));
}
#[test]
fn record_prefill_sets_boundary() {
let mut cache = RingCache::new(64);
let k = fill_head_major(8, |r, _| r as f32);
let v = fill_head_major(8, |r, _| (r * 2) as f32);
cache.record_prefill(&k, &v, 8).unwrap();
assert_eq!(cache.prefill_len(), Some(8));
assert_eq!(cache.effective_len(), 8);
assert_eq!(cache.ref_k_row(0, 3)[0], 3.0);
assert_eq!(cache.ref_v_row(0, 3)[0], 6.0);
}
#[test]
fn record_prefill_rejects_overflow() {
let mut cache = RingCache::new(4);
let k = fill_head_major(8, |_, _| 1.0);
let v = fill_head_major(8, |_, _| 1.0);
assert!(cache.record_prefill(&k, &v, 8).is_err());
}
#[test]
fn decode_single_key_returns_value() {
let mut cache = RingCache::new(8);
let k = fill_head_major(1, |_, d| if d == 0 { 1.0 } else { 0.0 });
let v = fill_head_major(1, |_, _| 7.0);
cache.record_prefill(&k, &v, 1).unwrap();
let q = one_token(|_, d| if d == 0 { 5.0 } else { 0.0 });
let out = decode_attention(&cache, &q).unwrap();
assert_eq!(out.shape(), (1, NUM_HEADS * HEAD_DIM));
for &x in &out.data {
assert!((x - 7.0).abs() < 1e-5);
}
}
#[test]
fn decode_equal_scores_averages_values() {
let mut cache = RingCache::new(8);
let k = fill_head_major(2, |_, d| if d == 0 { 1.0 } else { 0.0 });
let v = fill_head_major(2, |r, _| if r == 0 { 2.0 } else { 4.0 });
cache.record_prefill(&k, &v, 2).unwrap();
let q = one_token(|_, d| if d == 0 { 3.0 } else { 0.0 });
let out = decode_attention(&cache, &q).unwrap();
for &x in &out.data {
assert!((x - 3.0).abs() < 1e-5, "got {x}");
}
}
#[test]
fn online_matches_naive_softmax() {
let mut cache = RingCache::new(16);
let m = 5usize;
let k = fill_head_major(m, |r, d| {
((r + 1) as f32) * (if d == 0 { 1.0 } else { 0.0 })
});
let v = fill_head_major(m, |r, d| (r as f32) + (d as f32) * 0.01);
cache.record_prefill(&k, &v, m).unwrap();
let q = one_token(|_, d| if d == 0 { 0.5 } else { 0.0 });
let out = decode_attention(&cache, &q).unwrap();
let s = scale();
let mut scores = vec![0.0f32; m];
for (r, sc) in scores.iter_mut().enumerate() {
let mut d0 = 0.0f32;
#[allow(clippy::needless_range_loop)]
for d in 0..HEAD_DIM {
d0 += q[d] * cache.ref_k_row(0, r)[d];
}
*sc = d0 * s;
}
let mx = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let exps: Vec<f32> = scores.iter().map(|&x| (x - mx).exp()).collect();
let den: f32 = exps.iter().sum();
let mut expect0 = 0.0f32;
#[allow(clippy::needless_range_loop)]
for r in 0..m {
expect0 += (exps[r] / den) * cache.ref_v_row(0, r)[0];
}
assert!(
(out.data[0] - expect0).abs() < 1e-4,
"online {} naive {}",
out.data[0],
expect0
);
}
#[test]
fn warmup_appends_without_eviction() {
let mut cache = RingCache::new(8);
let k = fill_head_major(4, |_, _| 1.0);
let v = fill_head_major(4, |_, _| 1.0);
cache.record_prefill(&k, &v, 4).unwrap();
for step in 0..RING_WINDOW {
let kt = one_token(|_, _| step as f32);
let vt = one_token(|_, _| step as f32);
let slot = cache.write_decode_step(&kt, &vt).unwrap();
assert_eq!(slot, step, "warm-up writes are append-in-order");
assert_eq!(cache.ring_len(), step + 1);
assert!(cache.effective_len() == 4 + step + 1);
}
assert!(cache.is_warm());
assert_eq!(cache.ring_len(), RING_WINDOW);
assert_eq!(cache.ring_pos(), 0);
}
#[test]
fn steady_state_overwrites_modulo_w() {
let mut cache = RingCache::new(4);
let k = fill_head_major(2, |_, _| 0.0);
let v = fill_head_major(2, |_, _| 0.0);
cache.record_prefill(&k, &v, 2).unwrap();
for _ in 0..RING_WINDOW {
let t = one_token(|_, _| 0.0);
cache.write_decode_step(&t, &t).unwrap();
}
assert!(cache.is_warm());
let kt = one_token(|_, _| 99.0);
let slot0 = cache.write_decode_step(&kt, &kt).unwrap();
assert_eq!(slot0, 0);
assert_eq!(cache.ring_pos(), 1);
let slot1 = cache.write_decode_step(&kt, &kt).unwrap();
assert_eq!(slot1, 1);
assert_eq!(cache.ring_pos(), 2);
assert_eq!(cache.ring_len(), RING_WINDOW);
assert_eq!(cache.ring_k_row(0, 0)[0], 99.0);
}
#[test]
fn ring_pos_wraps_modulo_w() {
let mut cache = RingCache::new(2);
let k = fill_head_major(1, |_, _| 0.0);
cache.record_prefill(&k, &k, 1).unwrap();
for _ in 0..RING_WINDOW {
let t = one_token(|_, _| 0.0);
cache.write_decode_step(&t, &t).unwrap();
}
for expected_slot in 0..RING_WINDOW {
let t = one_token(|_, _| 0.0);
let slot = cache.write_decode_step(&t, &t).unwrap();
assert_eq!(slot, expected_slot);
}
assert_eq!(cache.ring_pos(), 0);
}
#[test]
fn attention_entry_writes_then_attends() {
let mut cache = RingCache::new(8);
let k = fill_head_major(1, |_, d| if d == 1 { 1.0 } else { 0.0 });
let v = fill_head_major(1, |_, _| 1.0);
cache.record_prefill(&k, &v, 1).unwrap();
let q = Mat::from_vec(
1,
NUM_HEADS * HEAD_DIM,
one_token(|_, d| if d == 0 { 10.0 } else { 0.0 }),
);
let kt = Mat::from_vec(
1,
NUM_HEADS * HEAD_DIM,
one_token(|_, d| if d == 0 { 1.0 } else { 0.0 }),
);
let vt = Mat::from_vec(1, NUM_HEADS * HEAD_DIM, one_token(|_, _| 5.0));
let out = attention(&mut cache, &q, &kt, &vt, &[42]).unwrap();
assert_eq!(out.shape(), (1, NUM_HEADS * HEAD_DIM));
assert_eq!(cache.ring_len(), 1);
let ring_w = (10.0f32 / (HEAD_DIM as f32).sqrt()).exp();
let expect = (ring_w * 5.0 + 1.0) / (ring_w + 1.0);
for &x in &out.data {
assert!((x - expect).abs() < 1e-4, "got {x}, expected {expect}");
assert!(x > 3.0, "ring token should dominate the reference, got {x}");
}
}
#[test]
fn decode_before_prefill_errors() {
let cache = RingCache::new(4);
let q = one_token(|_, _| 1.0);
assert!(decode_attention(&cache, &q).is_err());
}
#[test]
fn write_step_before_prefill_errors() {
let mut cache = RingCache::new(4);
let t = one_token(|_, _| 1.0);
assert!(cache.write_decode_step(&t, &t).is_err());
}
#[test]
fn attention_rejects_multi_row_query() {
let mut cache = RingCache::new(4);
let k = fill_head_major(1, |_, _| 0.0);
cache.record_prefill(&k, &k, 1).unwrap();
let q = Mat::zeros(2, NUM_HEADS * HEAD_DIM);
let kt = Mat::zeros(1, NUM_HEADS * HEAD_DIM);
let vt = Mat::zeros(1, NUM_HEADS * HEAD_DIM);
assert!(attention(&mut cache, &q, &kt, &vt, &[0]).is_err());
}
#[test]
fn attention_rejects_malformed_query_without_mutating_cache() {
let mut cache = RingCache::new(4);
let k = fill_head_major(1, |_, _| 0.0);
cache.record_prefill(&k, &k, 1).unwrap();
let q = Mat {
rows: 1,
cols: NUM_HEADS * HEAD_DIM,
data: vec![0.0; NUM_HEADS * HEAD_DIM - 1],
};
let kt = Mat::from_vec(1, NUM_HEADS * HEAD_DIM, one_token(|_, _| 1.0));
let vt = Mat::from_vec(1, NUM_HEADS * HEAD_DIM, one_token(|_, _| 2.0));
assert_err_contains(
attention(&mut cache, &q, &kt, &vt, &[0]),
"attention q data len",
);
assert_eq!(cache.ring_len(), 0, "malformed q must not write K/V");
}
#[test]
fn attention_rejects_kv_logical_shape_mismatch_before_mutating_cache() {
let mut cache = RingCache::new(4);
let prefill = fill_head_major(1, |_, _| 0.0);
cache.record_prefill(&prefill, &prefill, 1).unwrap();
let q = Mat::from_vec(1, NUM_HEADS * HEAD_DIM, one_token(|_, _| 1.0));
let kt = Mat {
rows: 2,
cols: (NUM_HEADS * HEAD_DIM) / 2,
data: one_token(|_, _| 1.0),
};
let vt = Mat::from_vec(1, NUM_HEADS * HEAD_DIM, one_token(|_, _| 2.0));
assert_err_contains(
attention(&mut cache, &q, &kt, &vt, &[0]),
"attention expects k [1",
);
assert_eq!(cache.ring_len(), 0, "malformed k must not write K/V");
}
#[test]
fn risky_decode_attention_opt_ins_require_truthy_values() {
for value in [
None,
Some(""),
Some("0"),
Some("off"),
Some("false"),
Some("no"),
] {
assert!(
!risky_opt_in_enabled_for(value),
"{value:?} must keep accuracy-risky attention experiments disabled"
);
}
for value in ["1", "on", "true", "yes", " TRUE "] {
assert!(
risky_opt_in_enabled_for(Some(value)),
"{value:?} must explicitly enable the requested experiment"
);
}
}
fn max_abs_diff(a: &Mat, b: &Mat) -> f32 {
assert_eq!(a.shape(), b.shape(), "shape mismatch in max_abs_diff");
a.data
.iter()
.zip(b.data.iter())
.map(|(x, y)| (x - y).abs())
.fold(0.0f32, f32::max)
}
fn build_cache(
pf: usize,
ring: usize,
int8: bool,
kf: impl Fn(usize, usize) -> f32,
vf: impl Fn(usize, usize) -> f32,
rk: impl Fn(usize, usize, usize) -> f32,
rv: impl Fn(usize, usize, usize) -> f32,
) -> RingCache {
let mut cache = RingCache::new_inner(pf + ring + 8, int8);
let k = fill_head_major(pf, &kf);
let v = fill_head_major(pf, &vf);
cache.record_prefill(&k, &v, pf).unwrap();
for step in 0..ring {
let kt = one_token(|h, d| rk(step, h, d));
let vt = one_token(|h, d| rv(step, h, d));
cache.write_decode_step(&kt, &vt).unwrap();
}
cache
}
#[test]
fn int8_row_quantization_uses_ties_to_even() {
let mut row = vec![0.0f32; HEAD_DIM];
row[..9].copy_from_slice(&[127.0, 0.5, 1.5, 2.5, 3.5, -0.5, -1.5, -2.5, -3.5]);
let mut quantized = [0i8; HEAD_DIM];
let scale = quantize_row_i8(&row, &mut quantized);
assert_eq!(scale.to_bits(), 1.0f32.to_bits());
assert_eq!(&quantized[..9], &[127, 0, 2, 2, 4, 0, -2, -2, -4]);
}
#[test]
fn int8_row_quantization_divides_instead_of_multiplying_by_reciprocal() {
let mut row = vec![0.0f32; HEAD_DIM];
row[0] = f32::from_bits(0x1e2a_010a);
row[1] = f32::from_bits(0x9e26_a853);
let mut quantized = [0i8; HEAD_DIM];
let scale = quantize_row_i8(&row, &mut quantized);
let reciprocal_result = (row[1] * (1.0 / scale)).round().clamp(-127.0, 127.0) as i8;
assert_eq!(quantized[0], 127);
assert_eq!(quantized[1], -124);
assert_eq!(
reciprocal_result, -125,
"fixture must distinguish the two formulas"
);
}
#[test]
fn int8_row_quantization_matches_scalar_contract() {
let mut state = 0xd1b5_4a32_d192_ed03u64;
for case in 0..256 {
let mut row = vec![0.0f32; HEAD_DIM];
let magnitude = 2.0f32.powi((case % 41) - 20);
for x in &mut row {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
let unit = ((state >> 40) as u32) as f32 / ((1u32 << 24) - 1) as f32;
*x = (unit * 2.0 - 1.0) * magnitude;
}
if case == 0 {
row.fill(0.0);
}
let expected_scale = row.iter().fold(0.0f32, |m, &x| m.max(x.abs()));
let expected_scale = if expected_scale > 0.0 {
expected_scale / 127.0
} else {
1.0
};
let expected = row
.iter()
.map(|&x| (x / expected_scale).round_ties_even().clamp(-127.0, 127.0) as i8)
.collect::<Vec<_>>();
let mut actual = [0i8; HEAD_DIM];
let actual_scale = quantize_row_i8(&row, &mut actual);
assert_eq!(
actual_scale.to_bits(),
expected_scale.to_bits(),
"case {case}"
);
assert_eq!(actual.as_slice(), expected.as_slice(), "case {case}");
}
}
#[test]
fn parallel_heads_are_bit_identical_to_serial() {
let cache = build_cache(
23,
RING_WINDOW + 3,
false,
|r, d| ((r * 29 + d * 11) % 37) as f32 * 0.03125 - 0.5,
|r, d| ((r * 17 + d * 7) % 31) as f32 * 0.046875 - 0.7,
|s, h, d| ((s * 13 + h * 5 + d * 3) % 41) as f32 * 0.0234375 - 0.4,
|s, h, d| ((s * 19 + h * 7 + d * 2) % 43) as f32 * 0.01953125 - 0.3,
);
let q = one_token(|h, d| ((h * 23 + d * 13) % 47) as f32 * 0.02734375 - 0.6);
let serial = decode_attention_scalar(&cache, &q).unwrap();
let serial_bits = serial.data.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
for threads in [1, 2, 8, 10] {
let pool = rayon::ThreadPoolBuilder::new()
.num_threads(threads)
.build()
.unwrap();
let parallel = pool
.install(|| decode_attention_scalar_parallel(&cache, &q))
.unwrap();
assert_eq!(
serial_bits,
parallel
.data
.iter()
.map(|x| x.to_bits())
.collect::<Vec<_>>(),
"parallel result changed with {threads} worker threads"
);
}
}
#[test]
fn gemm_attention_matches_scalar_reference() {
let cache = build_cache(
7,
5,
false,
|r, d| ((r * 13 + d * 7) % 17) as f32 * 0.11 - 0.9,
|r, d| ((r * 5 + d * 3) % 11) as f32 * 0.07 - 0.3,
|s, h, d| ((h * 3 + d * 2 + s) % 13) as f32 * 0.05 - 0.31,
|s, h, d| ((h + d * 4 + s) % 9) as f32 * 0.06 - 0.2,
);
let q = one_token(|h, d| ((h * 2 + d) % 7) as f32 * 0.2 - 0.6);
let gemm = decode_attention_gemm(&cache, &q).unwrap();
let scalar = decode_attention_scalar(&cache, &q).unwrap();
let max_abs = max_abs_diff(&gemm, &scalar);
assert!(max_abs <= 2.0e-6, "gemm vs scalar max_abs={max_abs}");
}
#[test]
fn default_dispatch_is_bit_exact_scalar() {
let cache = build_cache(
4,
3,
false,
|r, d| ((r + d) % 5) as f32 * 0.13 - 0.3,
|r, d| ((r * 2 + d) % 6) as f32 * 0.09 - 0.2,
|s, h, d| ((h + d + s) % 7) as f32 * 0.04 - 0.1,
|s, h, d| ((h * 2 + d + s) % 5) as f32 * 0.05 - 0.1,
);
let q = one_token(|h, d| ((h + d) % 9) as f32 * 0.1 - 0.4);
let public = decode_attention(&cache, &q).unwrap();
let scalar = decode_attention_scalar(&cache, &q).unwrap();
assert_eq!(public.data, scalar.data);
}
#[test]
fn int8_qk_i32_accumulation_cannot_overflow() {
let worst = 127i64 * 127 * HEAD_DIM as i64;
assert_eq!(worst, 2_064_512);
assert!(
worst <= i64::from(i32::MAX),
"worst-case int8 QK accumulation {worst} overflows i32"
);
}
#[test]
fn int8_mirror_allocated_only_when_enabled() {
assert!(RingCache::new_inner(8, false).int8.is_none());
let c = RingCache::new_inner(8, true);
let i8kv = c.int8.as_ref().expect("int8 mirror present");
assert_eq!(i8kv.ref_k.len(), NUM_HEADS);
assert_eq!(i8kv.ring_k[0].len(), RING_WINDOW * HEAD_DIM);
assert_eq!(i8kv.ref_k_scale[0].len(), 8);
}
#[test]
fn int8_kv_attention_matches_gemm_when_losslessly_quantizable() {
let anchor = |base: usize| {
move |r: usize, d: usize| -> f32 {
if d == 0 {
127.0
} else {
(((r * base + d * 3) % 7) as i32 - 3) as f32
}
}
};
let ranchor = |base: usize| {
move |s: usize, h: usize, d: usize| -> f32 {
if d == 0 {
127.0
} else {
(((h * base + d * 5 + s * 2) % 7) as i32 - 3) as f32
}
}
};
let cache = build_cache(6, 4, true, anchor(31), anchor(17), ranchor(13), ranchor(11));
let q = one_token(|_h, d| {
if d == 0 {
127.0
} else {
((d % 7) as i32 - 3) as f32
}
});
let int8 = decode_attention_int8(&cache, &q).unwrap();
let gemm = decode_attention_gemm(&cache, &q).unwrap();
let scalar = decode_attention_scalar(&cache, &q).unwrap();
let i8_vs_gemm = max_abs_diff(&int8, &gemm);
assert!(i8_vs_gemm <= 1.0e-6, "int8 vs gemm max_abs={i8_vs_gemm}");
let i8_vs_scalar = max_abs_diff(&int8, &scalar);
let out_mag = scalar.data.iter().fold(0.0f32, |m, &x| m.max(x.abs()));
let rel = i8_vs_scalar / out_mag.max(1.0);
assert!(
rel <= 2.0e-6,
"int8 vs scalar rel={rel} (abs={i8_vs_scalar}, out_mag={out_mag})"
);
}
#[test]
fn int8_kv_attention_runs_on_lossy_inputs() {
let cache = build_cache(
5,
3,
true,
|r, d| ((r * 7 + d * 3) % 19) as f32 * 0.013 - 0.12,
|r, d| ((r * 11 + d) % 23) as f32 * 0.009 - 0.1,
|s, h, d| ((h * 5 + d * 2 + s) % 17) as f32 * 0.011 - 0.09,
|s, h, d| ((h + d * 3 + s) % 13) as f32 * 0.012 - 0.07,
);
let q = one_token(|h, d| ((h * 3 + d) % 11) as f32 * 0.02 - 0.1);
let int8 = decode_attention_int8(&cache, &q).unwrap();
assert_eq!(int8.shape(), (1, NUM_HEADS * HEAD_DIM));
assert!(int8.data.iter().all(|x| x.is_finite()));
}
#[test]
fn int8_path_without_mirror_errors() {
let cache = build_cache(
3,
0,
false, |_, _| 0.5,
|_, _| 0.5,
|_, _, _| 0.0,
|_, _, _| 0.0,
);
let q = one_token(|_, _| 0.5);
assert_err_contains(
decode_attention_int8(&cache, &q),
"without an int8 KV mirror",
);
}
#[test]
fn checkpoint_rollback_restores_ring_rows_byte_identical() {
let pf = 5usize;
let accept = 9usize;
let accept_tok =
|t: usize| one_token(|h, d| ((h * 7 + d * 3 + t * 5) % 13) as f32 * 0.1 - 0.6);
let build = || {
let mut c = RingCache::new(64);
let k = fill_head_major(pf, |r, d| ((r * 3 + d) % 11) as f32 * 0.07 - 0.3);
let v = fill_head_major(pf, |r, d| ((r + d * 2) % 9) as f32 * 0.05 - 0.2);
c.record_prefill(&k, &v, pf).unwrap();
for t in 0..accept {
let kt = accept_tok(t);
c.write_decode_step(&kt, &kt).unwrap();
}
c
};
let control = build();
let mut test = build();
let cp = test.checkpoint();
assert_eq!(cp.ring_len(), accept);
assert_eq!(cp.ring_pos(), accept);
assert_eq!(cp.prefill_len(), Some(pf));
assert_eq!(cp.effective_len(), pf + accept);
for t in 0..6 {
let kt = one_token(|h, d| ((h + d + t) % 5) as f32 * 0.3 + 0.9);
test.write_decode_step(&kt, &kt).unwrap();
}
assert_eq!(test.ring_len(), accept + 6);
test.rollback_to(&cp).unwrap();
assert_eq!(test.ring_len(), accept);
assert_eq!(test.ring_pos(), accept);
assert_eq!(test.decode_writes, accept);
assert_eq!(test.effective_len(), pf + accept);
assert!(!test.is_warm());
for h in 0..NUM_HEADS {
for r in 0..test.ring_len() {
assert_eq!(
test.ring_k_row(h, r),
control.ring_k_row(h, r),
"ring_k h{h} r{r}"
);
assert_eq!(
test.ring_v_row(h, r),
control.ring_v_row(h, r),
"ring_v h{h} r{r}"
);
}
for r in 0..pf {
assert_eq!(
test.ref_k_row(h, r),
control.ref_k_row(h, r),
"ref_k h{h} r{r}"
);
assert_eq!(
test.ref_v_row(h, r),
control.ref_v_row(h, r),
"ref_v h{h} r{r}"
);
}
}
assert_ne!(
test.ring_k_row(0, accept),
control.ring_k_row(0, accept),
"dead slot is intentionally left stale (never read)"
);
let kt = accept_tok(accept);
assert_eq!(test.write_decode_step(&kt, &kt).unwrap(), accept);
assert_eq!(test.ring_len(), accept + 1);
}
#[test]
fn rollback_rejects_eviction_in_steady_state_without_mutation() {
let mut c = RingCache::new(8);
let k = fill_head_major(3, |_, _| 0.25);
c.record_prefill(&k, &k, 3).unwrap();
for _ in 0..RING_WINDOW {
let t = one_token(|_, _| 0.5);
c.write_decode_step(&t, &t).unwrap();
}
assert!(c.is_warm());
let cp = c.checkpoint();
let t = one_token(|_, _| 1.0);
c.write_decode_step(&t, &t).unwrap();
let current = c.checkpoint();
let ring_k = c.ring_k.clone();
let ring_v = c.ring_v.clone();
let error = c
.rollback_to(&cp)
.expect_err("steady-state overwrite cannot be restored from cursors");
assert!(error.to_string().contains("not lossless"));
assert_eq!(c.checkpoint(), current, "failed rollback changed cursors");
assert_eq!(c.ring_k, ring_k, "failed rollback changed ring K");
assert_eq!(c.ring_v, ring_v, "failed rollback changed ring V");
}
#[test]
fn rollback_rejects_abandoned_branch_checkpoint_without_mutation() {
let mut c = RingCache::new(16);
let prefill = fill_head_major(4, |r, d| ((r * 7 + d) % 13) as f32 * 0.1);
c.record_prefill(&prefill, &prefill, 4).unwrap();
let accepted = one_token(|_, _| 0.25);
c.write_decode_step(&accepted, &accepted).unwrap();
let branch_point = c.checkpoint();
let abandoned = one_token(|_, _| 1.0);
c.write_decode_step(&abandoned, &abandoned).unwrap();
let stale = c.checkpoint();
let abandoned_ring_k = c.ring_k.clone();
c.rollback_to(&branch_point).unwrap();
let replacement = one_token(|_, _| -2.0);
c.write_decode_step(&replacement, &replacement).unwrap();
assert_eq!(stale.ring_len, c.ring_len);
assert_eq!(stale.ring_pos, c.ring_pos);
assert_eq!(stale.decode_writes, c.decode_writes);
assert_ne!(stale.lineage_epoch, c.lineage_epoch);
assert_ne!(abandoned_ring_k, c.ring_k);
let current = c.checkpoint();
let ring_k = c.ring_k.clone();
let ring_v = c.ring_v.clone();
let error = c
.rollback_to(&stale)
.expect_err("checkpoint from abandoned branch must be stale");
assert!(error.to_string().contains("stale checkpoint lineage"));
assert_eq!(c.checkpoint(), current, "refusal changed cursors");
assert_eq!(c.ring_k, ring_k, "refusal changed replacement ring K");
assert_eq!(c.ring_v, ring_v, "refusal changed replacement ring V");
}
#[test]
fn rollback_rejects_different_cache_checkpoint_without_mutation() {
let prefill = fill_head_major(4, |r, d| ((r * 5 + d) % 17) as f32 * 0.1);
let mut issuer = RingCache::new(16);
let mut target = RingCache::new(16);
issuer.record_prefill(&prefill, &prefill, 4).unwrap();
target.record_prefill(&prefill, &prefill, 4).unwrap();
for value in [1.0, 2.0] {
let step = one_token(|_, _| value);
issuer.write_decode_step(&step, &step).unwrap();
}
let foreign = issuer.checkpoint();
for value in [-1.0, -2.0, -3.0] {
let step = one_token(|_, _| value);
target.write_decode_step(&step, &step).unwrap();
}
assert_ne!(foreign.cache_id, target.cache_id);
assert_eq!(foreign.prefill_len, target.prefill_len);
assert_eq!(foreign.lineage_epoch, target.lineage_epoch);
assert!(foreign.decode_writes < target.decode_writes);
let current = target.checkpoint();
let ring_k = target.ring_k.clone();
let ring_v = target.ring_v.clone();
let error = target
.rollback_to(&foreign)
.expect_err("checkpoint from independent cache must be refused");
assert!(error.to_string().contains("different cache instance"));
assert_eq!(target.checkpoint(), current, "refusal changed cursors");
assert_eq!(target.ring_k, ring_k, "refusal changed target ring K");
assert_eq!(target.ring_v, ring_v, "refusal changed target ring V");
}
#[test]
fn rollback_noop_is_identity() {
let mut c = RingCache::new(16);
let k = fill_head_major(4, |r, d| ((r + d) % 7) as f32 * 0.1);
c.record_prefill(&k, &k, 4).unwrap();
let t = one_token(|_, _| 0.3);
c.write_decode_step(&t, &t).unwrap();
let cp = c.checkpoint();
c.rollback_to(&cp).unwrap();
assert_eq!(c.ring_len(), 1);
assert_eq!(c.ring_pos(), 1);
assert_eq!(c.decode_writes, 1);
}
}
#[cfg(test)]
mod batched_ring_tests {
use super::{BatchedRingCache, HEAD_DIM, NUM_HEADS, RING_WINDOW, RingCache, decode_attention};
fn val(stream: usize, layer: usize, h: usize, r: usize, d: usize) -> f32 {
let mut x = (stream as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15)
^ (layer as u64).wrapping_mul(0xC2B2_AE3D_27D4_EB4F)
^ (h as u64).wrapping_mul(0x1656_67B1_9E37_79F9)
^ (r as u64).wrapping_mul(0xD6E8_FEB8_6659_FD93)
^ (d as u64).wrapping_mul(0x27D4_EB2F_1656_67C5);
x ^= x >> 33;
x = x.wrapping_mul(0xFF51_AFD7_ED55_8CCD);
x ^= x >> 33;
((x >> 40) as f32 / (1u64 << 24) as f32) * 2.0 - 1.0
}
fn build_prefill(stream: usize, layer: usize, seq: usize) -> (Vec<f32>, Vec<f32>) {
let mut k = vec![0.0f32; NUM_HEADS * seq * HEAD_DIM];
let mut v = vec![0.0f32; NUM_HEADS * seq * HEAD_DIM];
for h in 0..NUM_HEADS {
for r in 0..seq {
for d in 0..HEAD_DIM {
let idx = h * seq * HEAD_DIM + r * HEAD_DIM + d;
k[idx] = val(stream, layer, h, r, d);
v[idx] = val(stream, layer, h, r, d + 1);
}
}
}
(k, v)
}
fn build_step(stream: usize, layer: usize, t: usize) -> (Vec<f32>, Vec<f32>) {
let mut k = vec![0.0f32; NUM_HEADS * HEAD_DIM];
let mut v = vec![0.0f32; NUM_HEADS * HEAD_DIM];
for h in 0..NUM_HEADS {
for d in 0..HEAD_DIM {
let idx = h * HEAD_DIM + d;
k[idx] = val(stream, layer, h, 1_000 + t, d);
v[idx] = val(stream, layer, h, 1_000 + t, d + 1);
}
}
(k, v)
}
fn build_q(stream: usize, layer: usize, t: usize) -> Vec<f32> {
let mut q = vec![0.0f32; NUM_HEADS * HEAD_DIM];
for h in 0..NUM_HEADS {
for d in 0..HEAD_DIM {
q[h * HEAD_DIM + d] = val(stream, layer, h, 2_000 + t, d);
}
}
q
}
#[test]
fn per_stream_independent_and_bit_exact() {
let n_layers = 2usize;
let caps = [5usize, 9, 16, 4, 7, 20, 3, 11]; let mut bc = BatchedRingCache::new(&caps, n_layers);
assert_eq!(bc.num_streams(), caps.len());
assert_eq!(bc.num_layers(), n_layers);
for (s, &cap) in caps.iter().enumerate() {
for l in 0..n_layers {
let (k, v) = build_prefill(s, l, cap);
bc.record_prefill(s, l, &k, &v, cap).expect("prefill");
}
}
let st = 5usize;
let mut standalone: Vec<RingCache> =
(0..n_layers).map(|_| RingCache::new(caps[st])).collect();
for (l, cache) in standalone.iter_mut().enumerate() {
let (k, v) = build_prefill(st, l, caps[st]);
cache
.record_prefill(&k, &v, caps[st])
.expect("standalone prefill");
}
let steps = 200usize; for t in 0..steps {
for (s, _) in caps.iter().enumerate() {
for l in 0..n_layers {
let (k, v) = build_step(s, l, t);
bc.write_decode_step(s, l, &k, &v).expect("batched step");
}
}
for (l, cache) in standalone.iter_mut().enumerate() {
let (k, v) = build_step(st, l, t);
cache.write_decode_step(&k, &v).expect("standalone step");
}
if t == 0 || t == RING_WINDOW - 1 || t == RING_WINDOW || t == steps - 1 {
for l in 0..n_layers {
let bcl = bc.layer(st, l);
let sal = &standalone[l];
assert_eq!(
bcl.prefill_len(),
sal.prefill_len(),
"prefill_len l{l} t{t}"
);
assert_eq!(bcl.ring_len(), sal.ring_len(), "ring_len l{l} t{t}");
assert_eq!(bcl.ring_pos(), sal.ring_pos(), "ring_pos l{l} t{t}");
assert_eq!(
bcl.effective_len(),
sal.effective_len(),
"eff_len l{l} t{t}"
);
assert_eq!(bcl.is_warm(), sal.is_warm(), "is_warm l{l} t{t}");
let q = build_q(st, l, t);
let a = decode_attention(bcl, &q).expect("batched attn");
let b = decode_attention(sal, &q).expect("standalone attn");
assert_eq!(a.rows, b.rows);
assert_eq!(a.cols, b.cols);
assert_eq!(
a.data, b.data,
"stream {st} layer {l} step {t}: attention differs"
);
}
}
}
}
#[test]
fn large_batch_invariants_bounded_and_reference_unwritten() {
let b = 256usize;
let seq = 4usize;
let caps = vec![seq; b];
let mut bc = BatchedRingCache::new(&caps, 1);
assert_eq!(bc.num_streams(), b);
for s in 0..b {
let (k, v) = build_prefill(s, 0, seq);
bc.record_prefill(s, 0, &k, &v, seq).expect("prefill");
}
let steps = 200usize; for t in 0..steps {
for s in 0..b {
let (k, v) = build_step(s, 0, t);
bc.write_decode_step(s, 0, &k, &v).expect("step");
}
}
for s in 0..b {
let c = bc.layer(s, 0);
assert_eq!(
c.prefill_len(),
Some(seq),
"stream {s}: reference block must be untouched by decode"
);
assert_eq!(c.ring_len(), RING_WINDOW, "stream {s}: ring saturates");
assert!(c.is_warm(), "stream {s}: warm after {steps} steps");
assert_eq!(c.effective_len(), seq + RING_WINDOW, "stream {s}");
assert!(
c.effective_len() <= c.prefill_len().expect("prefilled") + RING_WINDOW,
"stream {s}: generated-token KV is bounded"
);
}
}
#[test]
fn kv_f32_bytes_matches_budget_arithmetic() {
let caps = [10usize, 20, 30];
let n_layers = 12usize;
let bc = BatchedRingCache::new(&caps, n_layers);
let f32sz = core::mem::size_of::<f32>();
let expect: usize = caps
.iter()
.map(|&cap| n_layers * 2 * NUM_HEADS * (cap + RING_WINDOW) * HEAD_DIM * f32sz)
.sum();
assert_eq!(bc.kv_f32_bytes(), expect);
assert!(bc.kv_f32_bytes() > 0);
}
#[test]
fn from_streams_adopts_prebuilt_rings_bit_exact() {
let n_layers = 2usize;
let caps = [6usize, 13];
let mut built: Vec<Vec<RingCache>> = Vec::new();
for (s, &cap) in caps.iter().enumerate() {
let mut layers: Vec<RingCache> = (0..n_layers).map(|_| RingCache::new(cap)).collect();
for (l, cache) in layers.iter_mut().enumerate() {
let (k, v) = build_prefill(s, l, cap);
cache.record_prefill(&k, &v, cap).expect("prefill");
}
built.push(layers);
}
let st = 1usize;
let mut standalone: Vec<RingCache> =
(0..n_layers).map(|_| RingCache::new(caps[st])).collect();
for (l, cache) in standalone.iter_mut().enumerate() {
let (k, v) = build_prefill(st, l, caps[st]);
cache
.record_prefill(&k, &v, caps[st])
.expect("mirror prefill");
}
let mut bc = BatchedRingCache::from_streams(built).expect("adopt");
assert_eq!(bc.num_streams(), caps.len());
assert_eq!(bc.num_layers(), n_layers);
for t in 0..5usize {
for (l, cache) in standalone.iter_mut().enumerate() {
let (k, v) = build_step(st, l, t);
cache.write_decode_step(&k, &v).expect("mirror step");
}
for l in 0..n_layers {
let (k, v) = build_step(st, l, t);
bc.write_decode_step(st, l, &k, &v).expect("adopted step");
}
for l in 0..n_layers {
let q = build_q(st, l, t);
let a = decode_attention(bc.layer(st, l), &q).expect("adopted attn");
let b = decode_attention(&standalone[l], &q).expect("mirror attn");
assert_eq!(
a.data, b.data,
"adopted stream {st} layer {l} step {t} differs"
);
}
}
}
#[test]
fn from_streams_rejects_empty_and_ragged() {
assert!(BatchedRingCache::from_streams(Vec::new()).is_err());
let s0: Vec<RingCache> = vec![RingCache::new(4), RingCache::new(4)];
let s1: Vec<RingCache> = vec![RingCache::new(4)]; assert!(BatchedRingCache::from_streams(vec![s0, s1]).is_err());
}
}