use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use cera::backend::cpu::{apply_rope_delta_to_head, apply_rope_to_head};
use cera::kv_cache::{InferenceState, LayerState};
use cera::model::{Model, ModelConfig};
fn max_abs_diff(a: &[f32], b: &[f32]) -> f32 {
a.iter()
.zip(b.iter())
.map(|(x, y)| (x - y).abs())
.fold(0.0f32, f32::max)
}
#[test]
fn rope_delta_composes_with_direct_rotation() {
let head_dim = 64;
let freq_base = 10_000.0f32;
for &p_old in &[0u32, 1, 7, 64, 2048] {
for &p_new in &[0u32, 1, 5, 63, 2000] {
let mut via_compose: Vec<f32> = (0..head_dim).map(|i| (i as f32) * 0.01).collect();
let mut via_direct = via_compose.clone();
apply_rope_to_head(&mut via_compose, p_old as usize, head_dim, freq_base);
let delta = (p_new as i32) - (p_old as i32);
apply_rope_delta_to_head(&mut via_compose, delta, head_dim, freq_base);
apply_rope_to_head(&mut via_direct, p_new as usize, head_dim, freq_base);
let err = max_abs_diff(&via_compose, &via_direct);
assert!(
err < 1e-3,
"p_old={p_old} p_new={p_new} delta={delta} max-err={err}"
);
}
}
}
#[test]
fn rope_delta_zero_is_identity() {
let head_dim = 32;
let freq_base = 10_000.0f32;
let original: Vec<f32> = (0..head_dim).map(|i| i as f32 * 0.1 + 1.0).collect();
let mut head = original.clone();
apply_rope_delta_to_head(&mut head, 0, head_dim, freq_base);
let err = max_abs_diff(&head, &original);
assert!(err < 1e-5, "delta=0 should be identity, max-err={err}");
}
fn build_state_with_rope_filled(
seq_len: usize,
head_dim: usize,
n_kv_heads_per_layer: &[usize],
freq_base: f32,
) -> InferenceState {
let mut state = InferenceState::new(n_kv_heads_per_layer.len());
state.seq_len = seq_len;
for (layer_idx, &n_kv_heads) in n_kv_heads_per_layer.iter().enumerate() {
let kv_dim = n_kv_heads * head_dim;
if let LayerState::Attention {
key_cache,
value_cache,
..
} = &mut state.layers[layer_idx]
{
key_cache.reserve(seq_len * kv_dim);
value_cache.reserve(seq_len * kv_dim);
for t in 0..seq_len {
for h in 0..n_kv_heads {
let raw: Vec<f32> = (0..head_dim)
.map(|d| (layer_idx as f32) + 0.1 * (h as f32) + 0.01 * (d as f32))
.collect();
let mut rotated = raw.clone();
apply_rope_to_head(&mut rotated, t, head_dim, freq_base);
key_cache.extend_from_slice(&rotated);
let v_row: Vec<f32> = raw.iter().map(|x| x + 0.001 * (t as f32)).collect();
value_cache.extend_from_slice(&v_row);
}
}
}
}
state
}
#[test]
fn shift_kv_with_rope_preserves_head_and_re_rotates_tail() {
let head_dim = 16;
let n_kv_heads_per_layer = vec![4usize, 2];
let seq_len = 24;
let n_keep = 5;
let shift = 7;
let freq_base = 10_000.0f32;
let mut state =
build_state_with_rope_filled(seq_len, head_dim, &n_kv_heads_per_layer, freq_base);
let head_snapshot: Vec<Vec<f32>> = n_kv_heads_per_layer
.iter()
.enumerate()
.map(|(layer_idx, &n_kv_heads)| {
let kv_dim = n_kv_heads * head_dim;
if let LayerState::Attention { key_cache, .. } = &state.layers[layer_idx] {
key_cache[..n_keep * kv_dim].to_vec()
} else {
unreachable!()
}
})
.collect();
state.shift_kv_with_rope(n_keep, shift, freq_base, head_dim, &n_kv_heads_per_layer);
assert_eq!(state.seq_len, seq_len - shift);
for (layer_idx, &n_kv_heads) in n_kv_heads_per_layer.iter().enumerate() {
let kv_dim = n_kv_heads * head_dim;
let new_seq_len = seq_len - shift;
if let LayerState::Attention {
key_cache,
value_cache,
..
} = &state.layers[layer_idx]
{
assert_eq!(key_cache.len(), new_seq_len * kv_dim);
assert_eq!(value_cache.len(), new_seq_len * kv_dim);
assert_eq!(
&key_cache[..n_keep * kv_dim],
head_snapshot[layer_idx].as_slice(),
"layer {layer_idx} head cells must be untouched"
);
for t_new in n_keep..new_seq_len {
let t_old = t_new + shift;
for h in 0..n_kv_heads {
let raw: Vec<f32> = (0..head_dim)
.map(|d| (layer_idx as f32) + 0.1 * (h as f32) + 0.01 * (d as f32))
.collect();
let mut oracle = raw.clone();
apply_rope_to_head(&mut oracle, t_new, head_dim, freq_base);
let start = t_new * kv_dim + h * head_dim;
let end = start + head_dim;
let actual = &key_cache[start..end];
let err = max_abs_diff(actual, &oracle);
assert!(
err < 1e-3,
"layer {layer_idx} head {h} t_new={t_new} t_old={t_old} max-err={err}"
);
let expected_v: Vec<f32> =
raw.iter().map(|x| x + 0.001 * (t_old as f32)).collect();
let v_actual = &value_cache[start..end];
let v_err = max_abs_diff(v_actual, &expected_v);
assert!(
v_err < 1e-6,
"V at layer {layer_idx} h={h} t_new={t_new} t_old={t_old} row identity lost: max-err={v_err}"
);
}
}
} else {
panic!("expected attention layer {layer_idx}");
}
}
}
#[test]
fn is_compressed_false_on_fresh_state() {
let state = InferenceState::new(4);
assert!(!state.is_compressed());
}
struct MockModel {
config: ModelConfig,
prefill_calls: AtomicUsize,
shift_calls: AtomicUsize,
supports_shift: bool,
}
impl MockModel {
fn new(config: ModelConfig, supports_shift: bool) -> Self {
Self {
config,
prefill_calls: AtomicUsize::new(0),
shift_calls: AtomicUsize::new(0),
supports_shift,
}
}
}
impl Model for MockModel {
fn forward(&self, _: &[u32], _: usize, _: &mut InferenceState) -> Vec<f32> {
vec![0.0; self.config.vocab_size]
}
fn forward_prefill(
&self,
tokens: &[u32],
_start_pos: usize,
state: &mut InferenceState,
) -> Vec<f32> {
self.prefill_calls.fetch_add(1, Ordering::Relaxed);
let head_dim = self.config.hidden_size / self.config.n_heads.max(1);
let kv_dim = self.config.n_kv_heads * head_dim;
for _ in tokens {
if let LayerState::Attention {
key_cache,
value_cache,
..
} = &mut state.layers[0]
{
key_cache.extend(std::iter::repeat_n(0.0f32, kv_dim));
value_cache.extend(std::iter::repeat_n(0.0f32, kv_dim));
}
state.seq_len += 1;
}
vec![0.0; self.config.vocab_size]
}
fn config(&self) -> &ModelConfig {
&self.config
}
fn supports_kv_shift(&self) -> bool {
self.supports_shift
}
fn shift_kv(&self, state: &mut InferenceState, n_keep: usize, shift: usize) {
self.shift_calls.fetch_add(1, Ordering::Relaxed);
let head_dim = self.config.hidden_size / self.config.n_heads.max(1);
state.shift_kv_with_rope(
n_keep,
shift,
self.config.rope_theta,
head_dim,
&self.config.kv_heads_per_layer,
);
}
}
fn mock_attention_config(max_seq_len: usize) -> ModelConfig {
ModelConfig {
architecture: "mock".into(),
n_layers: 1,
hidden_size: 8,
intermediate_size: 16,
n_heads: 4,
n_kv_heads: 4,
vocab_size: 8,
max_seq_len,
rope_theta: 10_000.0,
rms_norm_eps: 0.0,
block_types: vec![cera::model::BlockType::Attention],
conv_kernel_size: None,
kv_heads_per_layer: vec![4],
}
}
fn run_prefill(model: &MockModel, state: &mut InferenceState, tokens: &[u32]) -> usize {
let cancel = Arc::new(AtomicBool::new(false));
let (consumed, _) = model.forward_prefill_chunked(tokens, state.seq_len, state, 64, &cancel);
consumed
}
#[test]
fn shift_frees_capacity_when_n_keep_set() {
let max_seq_len = 32;
let n_keep = 4;
let cfg = mock_attention_config(max_seq_len);
let model = MockModel::new(cfg.clone(), true);
let mut state = InferenceState::new(cfg.block_types.len());
let first_batch: Vec<u32> = (0..28u32).collect();
assert_eq!(run_prefill(&model, &mut state, &first_batch), 28);
assert_eq!(state.seq_len, 28);
assert!(model.supports_kv_shift(), "probe must agree with field");
let shift_needed = 28 + 8 - max_seq_len;
assert_eq!(shift_needed, 4);
assert!(state.seq_len >= n_keep + shift_needed);
model.shift_kv(&mut state, n_keep, shift_needed);
assert_eq!(model.shift_calls.load(Ordering::Relaxed), 1);
assert_eq!(state.seq_len, 24);
let second_batch: Vec<u32> = (28..36u32).collect();
assert_eq!(run_prefill(&model, &mut state, &second_batch), 8);
assert_eq!(state.seq_len, 32);
assert_eq!(state.seq_len, max_seq_len);
if let LayerState::Attention { key_cache, .. } = &state.layers[0] {
let head_dim = cfg.hidden_size / cfg.n_heads;
let kv_dim = cfg.n_kv_heads * head_dim;
assert_eq!(key_cache.len(), 32 * kv_dim);
}
}
#[test]
fn can_shift_gate_all_branches() {
use cera::session::can_shift;
assert!(
can_shift(
true, 4, false,
28, 4,
),
"all conditions hold → can_shift"
);
assert!(
!can_shift(false, 4, false, 28, 4),
"supports_kv_shift=false → ContextOverflow"
);
assert!(
!can_shift(true, 0, false, 28, 4),
"n_keep=0 → ContextOverflow"
);
assert!(
!can_shift(true, 4, true, 28, 4),
"is_compressed=true → ContextOverflow"
);
assert!(
!can_shift(true, 4, false, 4, 4),
"current_pos < n_keep + shift → ContextOverflow"
);
assert!(
can_shift(true, 4, false, 8, 4),
"current_pos == n_keep + shift is allowed (inclusive)"
);
}