use std::cell::Cell;
use std::collections::HashMap;
#[cfg(feature = "disk-cache")]
use std::path::Path;
use std::path::PathBuf;
use crate::CeraError;
use crate::model::{BlockType, ModelConfig};
use crate::time::Instant;
use crate::turboquant::{
CompressedKeyCache, CompressedValueCache, EncodeScratch, QueryRotationScratch, RotationState,
TurboQuantConfig,
};
pub(crate) fn try_alloc<T>(len: usize) -> Result<Vec<T>, CeraError> {
let mut v: Vec<T> = Vec::new();
v.try_reserve_exact(len)
.map_err(|_| CeraError::OutOfMemory {
requested_bytes: (len as u64).saturating_mul(std::mem::size_of::<T>() as u64),
})?;
Ok(v)
}
pub(crate) fn checked_elems<T>(count: usize, per: usize) -> Result<usize, CeraError> {
count.checked_mul(per).ok_or(CeraError::OutOfMemory {
requested_bytes: (count as u64)
.saturating_mul(per as u64)
.saturating_mul(std::mem::size_of::<T>() as u64),
})
}
pub(crate) fn zeroed<T: Clone>(len: usize, fill: T) -> Result<Vec<T>, CeraError> {
let mut v = try_alloc::<T>(len)?;
v.resize(len, fill);
Ok(v)
}
fn zeroed_f32(len: usize) -> Result<Vec<f32>, CeraError> {
zeroed(len, 0.0)
}
#[derive(Clone, Debug, Default)]
pub enum KvCompression {
#[default]
None,
TurboQuant { seed: u64, keys: bool, values: bool },
}
impl KvCompression {
pub fn turboquant(seed: u64) -> Self {
Self::TurboQuant {
seed,
keys: true,
values: true,
}
}
pub fn flags(&self) -> (bool, bool) {
match self {
Self::None => (false, false),
Self::TurboQuant { keys, values, .. } => (*keys, *values),
}
}
}
#[allow(clippy::large_enum_variant)]
pub enum LayerState {
Attention {
key_cache: Vec<f32>,
value_cache: Vec<f32>,
compressed_keys: Option<CompressedKeyCache>,
compressed_values: Option<CompressedValueCache>,
},
Conv { buffer: Vec<f32> },
}
pub struct ScratchBuffers {
pub normed: Vec<f32>,
pub ffn_input: Vec<f32>,
pub conv_proj: Vec<f32>,
pub conv_scratch: Vec<f32>,
pub q: Vec<f32>,
pub k: Vec<f32>,
pub v: Vec<f32>,
pub attn_out: Vec<f32>,
pub gate: Vec<f32>,
pub up: Vec<f32>,
pub out: Vec<f32>,
pub scores: Vec<f32>,
pub q8_scales: Vec<f32>,
pub q8_quants: Vec<i8>,
pub dequant_weight_scratch: Vec<f32>,
pub lora_tmp: Vec<f32>,
}
pub struct InferenceState {
pub layers: Vec<LayerState>,
pub seq_len: usize,
pub scratch: ScratchBuffers,
pub lora: Option<std::sync::Arc<crate::lora::LoraAdapterWeights>>,
pub tq_encode_scratch: Option<EncodeScratch>,
pub tq_query_scratch: Option<QueryRotationScratch>,
pub tq_rotations: Vec<Option<RotationState>>,
pub tq_config: Option<TurboQuantConfig>,
}
impl InferenceState {
pub fn new(num_layers: usize) -> Self {
Self {
layers: (0..num_layers)
.map(|_| LayerState::Attention {
key_cache: Vec::new(),
value_cache: Vec::new(),
compressed_keys: None,
compressed_values: None,
})
.collect(),
seq_len: 0,
scratch: ScratchBuffers {
normed: Vec::new(),
ffn_input: Vec::new(),
conv_proj: Vec::new(),
conv_scratch: Vec::new(),
q: Vec::new(),
k: Vec::new(),
v: Vec::new(),
attn_out: Vec::new(),
gate: Vec::new(),
up: Vec::new(),
out: Vec::new(),
scores: Vec::new(),
q8_scales: Vec::new(),
q8_quants: Vec::new(),
dequant_weight_scratch: Vec::new(),
lora_tmp: Vec::new(),
},
tq_encode_scratch: None,
tq_query_scratch: None,
tq_rotations: Vec::new(),
tq_config: None,
lora: None,
}
}
pub fn from_config(config: &ModelConfig) -> Result<Self, CeraError> {
Self::from_config_with_compression(config, &KvCompression::None)
}
pub fn for_prefill(config: &ModelConfig, n_tokens: usize) -> Result<Self, CeraError> {
let capacity = n_tokens.clamp(1, config.max_seq_len);
Self::from_config_capped(config, &KvCompression::None, capacity)
}
pub fn clear_for_reuse(&mut self) {
self.seq_len = 0;
for layer in &mut self.layers {
match layer {
LayerState::Attention {
key_cache,
value_cache,
..
} => {
key_cache.clear();
value_cache.clear();
}
LayerState::Conv { buffer } => buffer.iter_mut().for_each(|x| *x = 0.0),
}
}
}
pub fn from_config_with_compression(
config: &ModelConfig,
compression: &KvCompression,
) -> Result<Self, CeraError> {
Self::from_config_capped(config, compression, config.max_seq_len)
}
pub(crate) fn from_config_capped(
config: &ModelConfig,
compression: &KvCompression,
capacity: usize,
) -> Result<Self, CeraError> {
let kernel_size = config.conv_kernel_size.unwrap_or(3);
assert!(
kernel_size >= 2,
"conv_kernel_size must be at least 2, got {kernel_size}"
);
let d_conv = kernel_size - 1;
let head_dim = config.head_dim;
let q_dim = checked_elems::<f32>(config.n_heads, head_dim)?;
let max_kv_dim = checked_elems::<f32>(
config.kv_heads_per_layer.iter().copied().max().unwrap_or(0),
head_dim,
)?;
let initial_capacity = capacity;
let (compress_keys, compress_values) = compression.flags();
let tq_enabled = (compress_keys || compress_values) && head_dim.is_power_of_two();
let (compress_keys, compress_values) = if tq_enabled {
(compress_keys, compress_values)
} else {
(false, false)
};
let (tq_rotations, tq_config) = if tq_enabled {
let seed = match compression {
KvCompression::TurboQuant { seed, .. } => *seed,
KvCompression::None => 0,
};
let mut rotations = try_alloc::<Option<RotationState>>(config.block_types.len())?;
for (layer_idx, bt) in config.block_types.iter().enumerate() {
rotations.push(match bt {
BlockType::Attention => Some(RotationState::try_from_seed(
seed ^ layer_idx as u64,
head_dim,
)?),
BlockType::GatedConv => None,
});
}
(rotations, Some(TurboQuantConfig::for_head_dim(head_dim)))
} else {
(Vec::new(), None)
};
let mut layers = try_alloc::<LayerState>(config.block_types.len())?;
for layer in config.block_types.iter().enumerate().map(
|(layer_idx, bt)| -> Result<LayerState, CeraError> {
match bt {
BlockType::Attention => {
let n_kv_heads = config.kv_heads_per_layer[layer_idx];
let kv_dim = checked_elems::<f32>(n_kv_heads, head_dim)?;
let kv_capacity = checked_elems::<f32>(capacity, kv_dim)?;
let compressed_keys = if compress_keys && n_kv_heads > 0 {
Some(CompressedKeyCache::try_new(
n_kv_heads,
head_dim,
initial_capacity,
)?)
} else {
None
};
let compressed_values = if compress_values && n_kv_heads > 0 {
Some(CompressedValueCache::try_new(
n_kv_heads,
head_dim,
initial_capacity,
)?)
} else {
None
};
let key_cache = if compress_keys && n_kv_heads > 0 {
Vec::new()
} else {
try_alloc::<f32>(kv_capacity)?
};
let value_cache = if compress_values && n_kv_heads > 0 {
Vec::new()
} else {
try_alloc::<f32>(kv_capacity)?
};
Ok(LayerState::Attention {
key_cache,
value_cache,
compressed_keys,
compressed_values,
})
}
BlockType::GatedConv => Ok(LayerState::Conv {
buffer: zeroed_f32(checked_elems::<f32>(d_conv, config.hidden_size)?)?,
}),
}
},
) {
layers.push(layer?);
}
Ok(Self {
layers,
seq_len: 0,
scratch: ScratchBuffers {
normed: zeroed_f32(config.hidden_size)?,
ffn_input: zeroed_f32(config.hidden_size)?,
conv_proj: zeroed_f32(checked_elems::<f32>(3, config.hidden_size)?)?,
conv_scratch: zeroed_f32(config.hidden_size)?,
q: zeroed_f32(q_dim)?,
k: zeroed_f32(max_kv_dim)?,
v: zeroed_f32(max_kv_dim)?,
attn_out: zeroed_f32(q_dim)?,
gate: zeroed_f32(config.intermediate_size)?,
up: zeroed_f32(config.intermediate_size)?,
out: zeroed_f32(config.hidden_size)?,
scores: Vec::new(), q8_scales: Vec::new(), q8_quants: Vec::new(), dequant_weight_scratch: Vec::new(),
lora_tmp: Vec::new(),
},
tq_encode_scratch: if tq_enabled {
Some(EncodeScratch::try_new(head_dim)?)
} else {
None
},
tq_query_scratch: if tq_enabled {
Some(QueryRotationScratch::try_new(config.n_heads, head_dim)?)
} else {
None
},
tq_rotations,
tq_config,
lora: None,
})
}
pub fn append_kv(&mut self, layer: usize, k: &[f32], v: &[f32]) {
if let LayerState::Attention {
key_cache,
value_cache,
..
} = &mut self.layers[layer]
{
key_cache.extend_from_slice(k);
value_cache.extend_from_slice(v);
}
}
pub fn kv_cache(&self, layer: usize) -> (&[f32], &[f32]) {
if let LayerState::Attention {
key_cache,
value_cache,
..
} = &self.layers[layer]
{
(key_cache, value_cache)
} else {
panic!("kv_cache called on non-attention layer {layer}");
}
}
pub fn compressed_keys(&self, layer: usize) -> Option<&CompressedKeyCache> {
if let LayerState::Attention {
compressed_keys, ..
} = &self.layers[layer]
{
compressed_keys.as_ref()
} else {
None
}
}
pub fn compressed_keys_mut(&mut self, layer: usize) -> Option<&mut CompressedKeyCache> {
if let LayerState::Attention {
compressed_keys, ..
} = &mut self.layers[layer]
{
compressed_keys.as_mut()
} else {
None
}
}
pub fn is_fully_compressed(&self) -> bool {
self.layers.iter().all(|l| match l {
LayerState::Attention {
compressed_keys,
compressed_values,
..
} => compressed_keys.is_some() && compressed_values.is_some(),
LayerState::Conv { .. } => true,
})
}
pub fn is_compressed(&self) -> bool {
self.layers.iter().any(|l| {
matches!(
l,
LayerState::Attention {
compressed_keys: Some(_),
..
} | LayerState::Attention {
compressed_values: Some(_),
..
}
)
})
}
pub fn snapshot(&self) -> Option<StateSnapshot> {
let mut layers = Vec::with_capacity(self.layers.len());
for l in &self.layers {
match l {
LayerState::Attention {
key_cache,
value_cache,
compressed_keys,
compressed_values,
} => {
let snap = match (compressed_keys, compressed_values) {
(None, None) => LayerSnapshot::Attention {
k_data: bytemuck::cast_slice(key_cache).to_vec(),
v_data: bytemuck::cast_slice(value_cache).to_vec(),
},
(Some(k), Some(v)) => LayerSnapshot::AttentionCompressed {
keys: crate::turboquant::encode_compressed_keys(k),
values: crate::turboquant::encode_compressed_values(v),
},
(Some(_), None) | (None, Some(_)) => return None,
};
layers.push(snap);
}
LayerState::Conv { buffer } => layers.push(LayerSnapshot::Conv {
buffer: bytemuck::cast_slice(buffer).to_vec(),
}),
}
}
Some(StateSnapshot {
layers,
seq_len: self.seq_len,
})
}
pub fn restore(&mut self, snapshot: &StateSnapshot) {
assert_eq!(
snapshot.layers.len(),
self.layers.len(),
"snapshot layer count {} doesn't match state layer count {}",
snapshot.layers.len(),
self.layers.len()
);
fn decode_f32_into(dst: &mut Vec<f32>, src: &[u8]) {
assert!(
src.len() % 4 == 0,
"snapshot byte length {} not a multiple of 4",
src.len()
);
dst.clear();
dst.reserve(src.len() / 4);
for chunk in src.chunks_exact(4) {
dst.push(f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]));
}
}
for (layer, snap) in self.layers.iter_mut().zip(snapshot.layers.iter()) {
match (layer, snap) {
(
LayerState::Attention {
key_cache,
value_cache,
..
},
LayerSnapshot::Attention { k_data, v_data },
) => {
decode_f32_into(key_cache, k_data);
decode_f32_into(value_cache, v_data);
}
(
LayerState::Attention {
key_cache,
value_cache,
compressed_keys,
compressed_values,
},
LayerSnapshot::AttentionCompressed { keys, values },
) => {
assert!(
compressed_keys.is_some() && compressed_values.is_some(),
"AttentionCompressed snapshot restored into a live state \
missing compressed slots — caller must gate on \
`LayerSnapshot::is_compressed()` matching \
`LayerState::is_compressed()`"
);
*compressed_keys = Some(
crate::turboquant::decode_compressed_keys(keys)
.expect("invalid TQK1 blob in snapshot"),
);
*compressed_values = Some(
crate::turboquant::decode_compressed_values(values)
.expect("invalid TQV1 blob in snapshot"),
);
key_cache.clear();
value_cache.clear();
}
(LayerState::Conv { buffer }, LayerSnapshot::Conv { buffer: snap_buf }) => {
decode_f32_into(buffer, snap_buf);
}
_ => panic!("snapshot layer kind doesn't match state layer kind"),
}
}
self.seq_len = snapshot.seq_len;
}
#[allow(clippy::too_many_arguments)]
pub fn shift_kv_with_rope(
&mut self,
n_keep: usize,
shift: usize,
rope_theta: f32,
head_dim: usize,
n_kv_heads_per_layer: &[usize],
rope_type: crate::backend::cpu::RopeType,
freq_factors: Option<&[f32]>,
) {
assert!(shift > 0, "shift must be > 0");
assert!(
n_keep + shift <= self.seq_len,
"shift range out of bounds: n_keep={n_keep} + shift={shift} > seq_len={}",
self.seq_len
);
assert!(
!self.is_compressed(),
"shift_kv_with_rope called on a TurboQuant-compressed state; \
shifting compressed caches is not yet supported"
);
assert_eq!(
n_kv_heads_per_layer.len(),
self.layers.len(),
"n_kv_heads_per_layer length {} doesn't match layer count {}",
n_kv_heads_per_layer.len(),
self.layers.len(),
);
let new_seq_len = self.seq_len - shift;
for (layer_idx, layer) in self.layers.iter_mut().enumerate() {
if let LayerState::Attention {
key_cache,
value_cache,
..
} = layer
{
if key_cache.is_empty() && value_cache.is_empty() {
continue;
}
assert_eq!(
key_cache.len(),
value_cache.len(),
"KV cache length mismatch: key={} value={}",
key_cache.len(),
value_cache.len()
);
assert!(self.seq_len > 0, "attention layer has KV but seq_len is 0");
assert_eq!(
key_cache.len() % self.seq_len,
0,
"KV cache length {} not a multiple of seq_len {}",
key_cache.len(),
self.seq_len
);
let n_kv_heads = n_kv_heads_per_layer[layer_idx];
let kv_dim = key_cache.len() / self.seq_len;
assert_eq!(
n_kv_heads * head_dim,
kv_dim,
"layer {layer_idx}: n_kv_heads*head_dim ({}) != cached kv_dim ({})",
n_kv_heads * head_dim,
kv_dim
);
let drop_start = n_keep * kv_dim;
let drop_end = (n_keep + shift) * kv_dim;
key_cache.drain(drop_start..drop_end);
value_cache.drain(drop_start..drop_end);
let delta = -(shift as i32);
for t in n_keep..new_seq_len {
let row_base = t * kv_dim;
for h in 0..n_kv_heads {
let head_start = row_base + h * head_dim;
let head_end = head_start + head_dim;
let head = &mut key_cache[head_start..head_end];
match rope_type {
crate::backend::cpu::RopeType::Neox => {
crate::backend::cpu::apply_rope_delta_to_head(
head, delta, head_dim, rope_theta,
)
}
crate::backend::cpu::RopeType::Norm => {
crate::backend::cpu::apply_rope_norm_delta_to_head(
head,
delta,
head_dim,
rope_theta,
freq_factors,
)
}
}
}
}
}
}
self.seq_len = new_seq_len;
}
}
#[derive(Clone)]
pub struct StateSnapshot {
pub layers: Vec<LayerSnapshot>,
pub seq_len: usize,
}
#[derive(Clone)]
pub enum LayerSnapshot {
Attention {
k_data: Vec<u8>,
v_data: Vec<u8>,
},
AttentionCompressed {
keys: Vec<u8>,
values: Vec<u8>,
},
Conv {
buffer: Vec<u8>,
},
}
impl LayerSnapshot {
pub fn is_compressed(&self) -> bool {
matches!(self, LayerSnapshot::AttentionCompressed { .. })
}
}
impl StateSnapshot {
pub fn byte_size(&self) -> usize {
self.layers
.iter()
.map(|l| match l {
LayerSnapshot::Attention { k_data, v_data } => k_data.len() + v_data.len(),
LayerSnapshot::AttentionCompressed { keys, values } => keys.len() + values.len(),
LayerSnapshot::Conv { buffer } => buffer.len(),
})
.sum()
}
pub fn is_compressed(&self) -> bool {
self.layers.iter().any(LayerSnapshot::is_compressed)
}
}
pub struct KvCacheConfig {
pub cache_dir: Option<PathBuf>,
pub max_warm_entries: usize,
pub max_warm_bytes: u64,
pub max_cold_bytes: u64,
}
impl Default for KvCacheConfig {
fn default() -> Self {
Self {
cache_dir: None,
max_warm_entries: 32,
max_warm_bytes: 256 * 1024 * 1024,
max_cold_bytes: 10 * 1024 * 1024 * 1024,
}
}
}
struct CacheEntry {
tokens: Vec<u32>,
snapshot: StateSnapshot,
last_used: Cell<Instant>,
}
#[cfg_attr(not(feature = "disk-cache"), allow(dead_code))]
pub struct KvPrefixCache {
warm: HashMap<u64, CacheEntry>,
pub config: KvCacheConfig,
model_fingerprint: u64,
warm_bytes: u64,
}
impl KvPrefixCache {
pub fn new(config: KvCacheConfig, model_config: &ModelConfig, model_id: &str) -> Self {
Self {
warm: HashMap::new(),
model_fingerprint: model_fingerprint(model_config, model_id),
config,
warm_bytes: 0,
}
}
pub fn find_longest_prefix(&mut self, tokens: &[u32]) -> Option<(StateSnapshot, usize)> {
let warm_hit = self
.warm
.values()
.filter(|e| tokens.starts_with(&e.tokens))
.max_by_key(|e| e.tokens.len())
.map(|e| {
e.last_used.set(Instant::now());
(e.snapshot.clone(), e.tokens.len())
});
#[cfg(feature = "disk-cache")]
let cold_hit = self
.config
.cache_dir
.clone()
.and_then(|dir| self.find_cold_prefix(&dir, tokens))
.map(|snapshot| {
let len = snapshot.seq_len;
(snapshot, len)
});
#[cfg(not(feature = "disk-cache"))]
let cold_hit: Option<(StateSnapshot, usize)> = None;
let best = match (warm_hit, cold_hit) {
(Some(w), Some(c)) if c.1 > w.1 => Some(c),
(Some(w), _) => Some(w),
(None, c) => c,
};
if let Some((snapshot, len)) = &best {
if !self
.warm
.values()
.any(|e| e.tokens.len() >= *len && tokens.starts_with(&e.tokens))
{
let hash = hash_tokens(&tokens[..*len]);
let snap_bytes = snapshot.byte_size() as u64;
self.evict_warm_if_needed(snap_bytes);
if let Some(old) = self.warm.insert(
hash,
CacheEntry {
tokens: tokens[..*len].to_vec(),
snapshot: snapshot.clone(),
last_used: Cell::new(Instant::now()),
},
) {
self.warm_bytes -= old.snapshot.byte_size() as u64;
}
self.warm_bytes += snap_bytes;
}
}
best
}
pub fn insert(&mut self, tokens: &[u32], snapshot: StateSnapshot) {
if self.config.max_warm_entries == 0 && self.config.cache_dir.is_none() {
return;
}
let hash = hash_tokens(tokens);
let snap_bytes = snapshot.byte_size() as u64;
self.evict_warm_if_needed(snap_bytes);
#[cfg(feature = "disk-cache")]
if let Some(dir) = &self.config.cache_dir {
self.save_cold(dir, tokens, &snapshot);
}
if let Some(old) = self.warm.insert(
hash,
CacheEntry {
tokens: tokens.to_vec(),
snapshot,
last_used: Cell::new(Instant::now()),
},
) {
self.warm_bytes -= old.snapshot.byte_size() as u64;
}
self.warm_bytes += snap_bytes;
}
pub fn warm_bytes(&self) -> u64 {
self.warm_bytes
}
pub fn warm_count(&self) -> usize {
self.warm.len()
}
fn evict_warm_if_needed(&mut self, new_bytes: u64) {
while (self.warm.len() >= self.config.max_warm_entries
|| self.warm_bytes + new_bytes > self.config.max_warm_bytes)
&& !self.warm.is_empty()
{
let oldest = self
.warm
.iter()
.min_by_key(|(_, e)| e.last_used.get())
.map(|(k, _)| *k);
if let Some(key) = oldest {
if let Some(removed) = self.warm.remove(&key) {
self.warm_bytes -= removed.snapshot.byte_size() as u64;
}
}
}
}
#[cfg(feature = "disk-cache")]
fn cold_filename(&self, token_hash: u64) -> String {
format!(
"{:016x}_{:016x}.kvcache",
self.model_fingerprint, token_hash
)
}
#[cfg(feature = "disk-cache")]
fn save_cold(&self, dir: &Path, tokens: &[u32], snapshot: &StateSnapshot) {
if std::fs::create_dir_all(dir).is_err() {
return;
}
let mut builder =
flatbuffers::FlatBufferBuilder::with_capacity(snapshot.byte_size() + 1024);
let mut layer_offsets = Vec::with_capacity(snapshot.layers.len());
for layer in &snapshot.layers {
let (tag, k_off, v_off) = match layer {
LayerSnapshot::Attention { k_data, v_data } => {
let k = builder.create_vector(k_data);
let v = builder.create_vector(v_data);
(0u8, Some(k), Some(v))
}
LayerSnapshot::Conv { buffer } => {
let k = builder.create_vector(buffer);
(1u8, Some(k), None)
}
LayerSnapshot::AttentionCompressed { keys, values } => {
let k = builder.create_vector(keys);
let v = builder.create_vector(values);
(2u8, Some(k), Some(v))
}
};
let ld = crate::generated::cera::cache::LayerData::create(
&mut builder,
&crate::generated::cera::cache::LayerDataArgs {
type_tag: tag,
k_data: k_off,
v_data: v_off,
},
);
layer_offsets.push(ld);
}
let layers_vec = builder.create_vector(&layer_offsets);
let tokens_vec = builder.create_vector(tokens);
let entry = crate::generated::cera::cache::KvCacheEntry::create(
&mut builder,
&crate::generated::cera::cache::KvCacheEntryArgs {
model_fingerprint: self.model_fingerprint,
seq_len: snapshot.seq_len as u32,
tokens: Some(tokens_vec),
layers: Some(layers_vec),
},
);
builder.finish(entry, None);
let data = builder.finished_data();
let hash = hash_tokens(tokens);
let path = dir.join(self.cold_filename(hash));
let _ = std::fs::write(&path, data);
self.evict_cold_if_needed(dir);
}
#[cfg(feature = "disk-cache")]
fn find_cold_prefix(&self, dir: &Path, tokens: &[u32]) -> Option<StateSnapshot> {
let mut best: Option<StateSnapshot> = None;
for prefix_len in (1..=tokens.len()).rev() {
let prefix = &tokens[..prefix_len];
let hash = hash_tokens(prefix);
let path = dir.join(self.cold_filename(hash));
if path.exists() {
if let Some(snapshot) = self.load_cold_file(&path, tokens) {
best = Some(snapshot);
break; }
}
}
best
}
#[cfg(feature = "disk-cache")]
fn load_cold_file(&self, path: &Path, expected_prefix: &[u32]) -> Option<StateSnapshot> {
let data = std::fs::read(path).ok()?;
let entry = flatbuffers::root::<crate::generated::cera::cache::KvCacheEntry>(&data).ok()?;
if entry.model_fingerprint() != self.model_fingerprint {
return None;
}
let cached_tokens = entry.tokens()?;
let seq_len = entry.seq_len() as usize;
if seq_len != cached_tokens.len() {
return None;
}
if cached_tokens.len() > expected_prefix.len() {
return None;
}
for (i, ct) in cached_tokens.iter().enumerate() {
if ct != expected_prefix[i] {
return None;
}
}
let layers_fb = entry.layers()?;
let mut layers = Vec::with_capacity(layers_fb.len());
for l in layers_fb {
match l.type_tag() {
0 => {
layers.push(LayerSnapshot::Attention {
k_data: l.k_data()?.bytes().to_vec(),
v_data: l.v_data()?.bytes().to_vec(),
});
}
1 => {
layers.push(LayerSnapshot::Conv {
buffer: l.k_data()?.bytes().to_vec(),
});
}
2 => {
let keys = l.k_data()?.bytes().to_vec();
let values = l.v_data()?.bytes().to_vec();
if crate::turboquant::decode_compressed_keys(&keys).is_none()
|| crate::turboquant::decode_compressed_values(&values).is_none()
{
return None;
}
layers.push(LayerSnapshot::AttentionCompressed { keys, values });
}
_ => return None,
}
}
Some(StateSnapshot { layers, seq_len })
}
#[cfg(feature = "disk-cache")]
fn evict_cold_if_needed(&self, dir: &Path) {
let Ok(entries) = std::fs::read_dir(dir) else {
return;
};
let fp_prefix = format!("{:016x}_", self.model_fingerprint);
let mut files: Vec<(PathBuf, u64, std::time::SystemTime)> = entries
.filter_map(|e| e.ok())
.filter(|e| {
let name = e.file_name();
let name_str = name.to_string_lossy();
name_str.starts_with(&fp_prefix) && name_str.ends_with(".kvcache")
})
.filter_map(|e| {
let meta = e.metadata().ok()?;
Some((e.path(), meta.len(), meta.modified().ok()?))
})
.collect();
let total: u64 = files.iter().map(|(_, sz, _)| sz).sum();
if total <= self.config.max_cold_bytes {
return;
}
files.sort_by_key(|(_, _, t)| *t);
let mut remaining = total;
for (path, sz, _) in &files {
if remaining <= self.config.max_cold_bytes {
break;
}
let _ = std::fs::remove_file(path);
remaining -= sz;
}
}
}
fn fnv1a_u64(bytes: &[u8]) -> u64 {
const OFFSET: u64 = 0xcbf29ce484222325;
const PRIME: u64 = 0x100000001b3;
let mut h = OFFSET;
for &b in bytes {
h ^= b as u64;
h = h.wrapping_mul(PRIME);
}
h
}
fn hash_tokens(tokens: &[u32]) -> u64 {
let bytes: &[u8] = bytemuck::cast_slice(tokens);
fnv1a_u64(bytes)
}
pub fn model_fingerprint(config: &ModelConfig, model_id: &str) -> u64 {
let mut buf = Vec::with_capacity(128);
buf.extend_from_slice(model_id.as_bytes());
buf.push(0);
buf.extend_from_slice(config.architecture.as_bytes());
buf.push(0);
buf.extend_from_slice(&(config.n_layers as u64).to_le_bytes());
buf.extend_from_slice(&(config.hidden_size as u64).to_le_bytes());
buf.extend_from_slice(&(config.n_heads as u64).to_le_bytes());
for bt in &config.block_types {
buf.push(match bt {
crate::model::BlockType::Attention => 0,
crate::model::BlockType::GatedConv => 1,
});
}
for k in &config.kv_heads_per_layer {
buf.extend_from_slice(&(*k as u64).to_le_bytes());
}
fnv1a_u64(&buf)
}
#[cfg(test)]
mod tests {
use super::*;
fn tiny_config(n_layers: usize, hidden_size: usize) -> ModelConfig {
ModelConfig {
architecture: "lfm2".into(),
n_layers,
hidden_size,
intermediate_size: hidden_size * 2,
n_heads: 4,
n_kv_heads: 2,
head_dim: hidden_size / 4,
vocab_size: 256,
max_seq_len: 64,
rope_theta: 1_000_000.0,
rms_norm_eps: 1e-5,
block_types: (0..n_layers)
.map(|i| {
if i % 2 == 0 {
BlockType::Attention
} else {
BlockType::GatedConv
}
})
.collect(),
conv_kernel_size: Some(3),
kv_heads_per_layer: (0..n_layers)
.map(|i| if i % 2 == 0 { 2 } else { 0 })
.collect(),
scalars: crate::model::ScalarMultipliers::default(),
}
}
#[test]
fn for_prefill_caps_kv_capacity() {
let cfg = tiny_config(4, 16); let n = 5;
let kv_dim = cfg.kv_heads_per_layer[0] * cfg.head_dim;
let state = InferenceState::for_prefill(&cfg, n).unwrap();
if let LayerState::Attention {
key_cache,
value_cache,
..
} = &state.layers[0]
{
assert!(key_cache.capacity() >= n * kv_dim);
assert!(value_cache.capacity() >= n * kv_dim);
assert!(key_cache.capacity() < cfg.max_seq_len * kv_dim);
} else {
panic!("layer 0 should be attention");
}
let big = InferenceState::for_prefill(&cfg, cfg.max_seq_len + 100).unwrap();
if let LayerState::Attention { key_cache, .. } = &big.layers[0] {
assert!(key_cache.capacity() <= cfg.max_seq_len * kv_dim + kv_dim);
}
}
#[test]
fn clear_for_reuse_resets_but_keeps_capacity() {
let cfg = tiny_config(4, 16);
let mut state = InferenceState::for_prefill(&cfg, 8).unwrap();
state.seq_len = 3;
let (cap_k, cap_v) = if let LayerState::Attention {
key_cache,
value_cache,
..
} = &mut state.layers[0]
{
for i in 0..16 {
key_cache.push(i as f32);
value_cache.push(i as f32);
}
(key_cache.capacity(), value_cache.capacity())
} else {
panic!("layer 0 should be attention");
};
if let LayerState::Conv { buffer } = &mut state.layers[1] {
buffer.iter_mut().for_each(|x| *x = 1.0);
}
state.clear_for_reuse();
assert_eq!(state.seq_len, 0);
if let LayerState::Attention {
key_cache,
value_cache,
..
} = &state.layers[0]
{
assert!(key_cache.is_empty() && value_cache.is_empty());
assert_eq!(key_cache.capacity(), cap_k, "capacity must be retained");
assert_eq!(value_cache.capacity(), cap_v);
}
if let LayerState::Conv { buffer } = &state.layers[1] {
assert!(
buffer.iter().all(|&x| x == 0.0),
"conv buffer must be zeroed"
);
}
}
#[test]
fn snapshot_restore_round_trip_attention_and_conv() {
let cfg = tiny_config(4, 16);
let mut state = InferenceState::from_config(&cfg).unwrap();
if let LayerState::Attention {
key_cache,
value_cache,
..
} = &mut state.layers[0]
{
for i in 0..16 {
key_cache.push(i as f32 * 0.5);
value_cache.push(-(i as f32) * 0.25);
}
}
if let LayerState::Conv { buffer } = &mut state.layers[1] {
for v in buffer.iter_mut() {
*v = 0.7;
}
}
state.seq_len = 2;
let snap = state.snapshot().expect("uncompressed state must snapshot");
let mut fresh = InferenceState::from_config(&cfg).unwrap();
fresh.restore(&snap);
match (&fresh.layers[0], &state.layers[0]) {
(
LayerState::Attention {
key_cache: kr,
value_cache: vr,
..
},
LayerState::Attention {
key_cache: ko,
value_cache: vo,
..
},
) => {
assert_eq!(kr, ko, "key_cache must round-trip exactly");
assert_eq!(vr, vo, "value_cache must round-trip exactly");
}
_ => panic!("expected attention layer 0"),
}
match (&fresh.layers[1], &state.layers[1]) {
(LayerState::Conv { buffer: br }, LayerState::Conv { buffer: bo }) => {
assert_eq!(br, bo, "conv buffer must round-trip exactly")
}
_ => panic!("expected conv layer 1"),
}
assert_eq!(fresh.seq_len, state.seq_len);
}
#[test]
fn snapshot_restore_round_trip_empty_state() {
let cfg = tiny_config(2, 8);
let state = InferenceState::from_config(&cfg).unwrap();
let snap = state.snapshot().expect("empty state still snapshots");
let mut fresh = InferenceState::from_config(&cfg).unwrap();
fresh.restore(&snap);
assert_eq!(fresh.seq_len, 0);
}
#[test]
fn snapshot_compressed_state_emits_attention_compressed() {
let cfg = tiny_config(2, 8);
let mut state = InferenceState::from_config(&cfg).unwrap();
if let LayerState::Attention {
compressed_keys,
compressed_values,
..
} = &mut state.layers[0]
{
let mut keys = CompressedKeyCache::new(2, 8, 4);
let mut values = CompressedValueCache::new(2, 8, 4);
for h in 0..2 {
keys.append(h, &[0xAB, 0xCD], &[0x55], 0x1234, 0x5678);
values.append(h, &[0xEF, 0x01], 0x9ABC);
}
*compressed_keys = Some(keys);
*compressed_values = Some(values);
}
state.seq_len = 1;
assert!(state.is_compressed());
let snap = state.snapshot().expect("compressed state must snapshot");
match &snap.layers[0] {
LayerSnapshot::AttentionCompressed { keys, values } => {
assert!(keys.starts_with(b"TQK1"));
assert!(values.starts_with(b"TQV1"));
}
_ => panic!("layer 0 should be AttentionCompressed"),
}
let mut fresh = InferenceState::from_config(&cfg).unwrap();
if let LayerState::Attention {
compressed_keys,
compressed_values,
..
} = &mut fresh.layers[0]
{
*compressed_keys = Some(CompressedKeyCache::new(2, 8, 4));
*compressed_values = Some(CompressedValueCache::new(2, 8, 4));
}
fresh.restore(&snap);
match (&state.layers[0], &fresh.layers[0]) {
(
LayerState::Attention {
compressed_keys: Some(orig_k),
compressed_values: Some(orig_v),
..
},
LayerState::Attention {
compressed_keys: Some(restored_k),
compressed_values: Some(restored_v),
..
},
) => {
assert_eq!(restored_k.polar_data, orig_k.polar_data);
assert_eq!(restored_k.jl_data, orig_k.jl_data);
assert_eq!(restored_k.norms, orig_k.norms);
assert_eq!(restored_k.residual_norms, orig_k.residual_norms);
assert_eq!(restored_v.polar_data, orig_v.polar_data);
assert_eq!(restored_v.norms, orig_v.norms);
assert_eq!(restored_k.norms_f32, orig_k.norms_f32);
assert_eq!(restored_v.norms_f32, orig_v.norms_f32);
}
_ => panic!("expected both states to have populated compressed caches"),
}
assert_eq!(fresh.seq_len, state.seq_len);
}
}