pub fn kv_cache_formats() -> (&'static str, &'static str) {
static F: std::sync::OnceLock<(&'static str, &'static str)> = std::sync::OnceLock::new();
*F.get_or_init(|| {
let k = match std::env::var("MEMRA_KV_K").as_deref() {
Ok("fp8") => "fp8",
Ok("q8_0") | Ok("") | Err(_) => "q8_0",
Ok(o) => panic!("MEMRA_KV_K={o} unsupported (q8_0 | fp8)"),
};
let v = match std::env::var("MEMRA_KV_V").as_deref() {
Ok("q4_0") => "q4_0",
Ok("fp8") => "fp8",
Ok("q5_1") | Ok("") | Err(_) => "q5_1",
Ok(o) => panic!("MEMRA_KV_V={o} unsupported (q5_1 | q4_0 | fp8)"),
};
if (k, v) != ("q8_0", "q5_1") {
eprintln!("[memra] KV cache format: K={k} V={v} (non-default — new numeric config)");
}
(k, v)
})
}
pub fn kv_blk_bytes() -> (usize, usize) {
let (k, v) = kv_cache_formats();
let kb = match k { "fp8" => 32, _ => 34 };
let vb = match v { "q4_0" => 18, "fp8" => 32, _ => 24 };
(kb, vb)
}
pub fn gkv_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_GEMMA_GKV").map(|v| v != "0").unwrap_or(true))
}
pub fn wkv_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_GEMMA_WKV").map(|v| v != "0")
.unwrap_or_else(|_| std::env::var("MEMRA_DRAFT").is_err()))
}
pub static KV_FP8_FORCE: std::sync::atomic::AtomicI8 = std::sync::atomic::AtomicI8::new(-1);
pub fn kv_fp8_on() -> bool {
static ENV: std::sync::OnceLock<Option<bool>> = std::sync::OnceLock::new();
if let Some(v) = *ENV.get_or_init(|| std::env::var("MEMRA_KV_FP8").ok()
.map(|v| v == "1")) { return v; }
matches!(KV_FP8_FORCE.load(std::sync::atomic::Ordering::Relaxed), 1)
}
pub fn swa_ring_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_SWA_RING").as_deref() == Ok("1"))
}
pub const PRIME_CHUNK_MAX_TOKENS: usize = 4096;
const SWA_VIEW_ALIGNMENT_ROWS: usize = 32;
pub fn swa_ring_rows(window: usize, max_ctx: usize) -> usize {
max_ctx.min(window + PRIME_CHUNK_MAX_TOKENS + (SWA_VIEW_ALIGNMENT_ROWS - 1))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct KvRing {
rows: usize,
window: usize,
base: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum KvRingAppend {
Contiguous {
write_row: usize,
},
Rebase {
src_row: usize,
keep_rows: usize,
new_base: usize,
write_row: usize,
},
}
impl KvRing {
pub fn new(rows: usize, window: usize) -> Self {
assert!(window > 0 && rows > 0, "invalid SWA ring geometry");
Self { rows, window, base: 0 }
}
pub fn rows(&self) -> usize { self.rows }
pub fn base(&self) -> usize { self.base }
pub fn window(&self) -> usize { self.window }
pub fn append_plan(
&self,
len: usize,
retain_from: usize,
append_rows: usize,
) -> Result<KvRingAppend, String> {
if len < self.base || retain_from < self.base || retain_from > len {
return Err(format!(
"SWA ring lapped required rows (base {}, retain {retain_from}, len {len})",
self.base
));
}
let used = len - self.base;
if used > self.rows {
return Err(format!("SWA ring state exceeds capacity ({used} > {})", self.rows));
}
if used.saturating_add(append_rows) <= self.rows {
return Ok(KvRingAppend::Contiguous {
write_row: used % self.rows,
});
}
let keep_rows = len - retain_from;
if keep_rows.saturating_add(append_rows) > self.rows {
return Err(format!(
"SWA ring append does not fit (keep {keep_rows} + append {append_rows} > {})",
self.rows
));
}
Ok(KvRingAppend::Rebase {
src_row: retain_from - self.base,
keep_rows,
new_base: retain_from,
write_row: keep_rows,
})
}
pub fn apply_rebase(&mut self, new_base: usize) {
debug_assert!(new_base >= self.base);
self.base = new_base;
}
pub fn physical_range(
&self,
start: usize,
end: usize,
) -> Result<std::ops::Range<usize>, String> {
if start < self.base || end < start || end - self.base > self.rows {
return Err(format!(
"SWA ring view [{start},{end}) is outside resident [{},{})",
self.base,
self.base + self.rows
));
}
let start_row = (start - self.base) % self.rows;
let len = end - start;
debug_assert!(start_row + len <= self.rows, "ring view must be contiguous after rebase");
Ok(start_row..start_row + len)
}
pub fn can_rewind_to(&self, len: usize) -> bool {
let raw = len.saturating_sub(self.window - 1);
let view_start = raw & !(SWA_VIEW_ALIGNMENT_ROWS - 1);
view_start >= self.base
}
}
pub trait KvDev {
fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>>;
fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>>;
fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>>;
fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>>;
fn clone_dtod(&self, src: &CudaSlice<f32>) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>>;
fn copy_into(&self, dst: &mut CudaSlice<f32>, off: usize, src: &CudaSlice<f32>, len: usize)
-> Result<(), Box<dyn std::error::Error>>;
fn set_i32_one(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>>;
}
use memra_gguf::config::{LayerKind, ModelConfig};
use cudarc::driver::CudaSlice;
pub struct KvLayer {
pub k: CudaSlice<u8>, pub v: CudaSlice<u8>, pub kv_dim_k: usize, pub kv_dim_v: usize, pub k_tok_bytes: usize, pub v_tok_bytes: usize, pub len: usize,
pub ring: Option<KvRing>,
pub len_d: CudaSlice<i32>,
}
impl KvLayer {
pub fn physical_rows(
&self,
start: usize,
end: usize,
) -> Result<std::ops::Range<usize>, String> {
match &self.ring {
Some(ring) => ring.physical_range(start, end),
None => Ok(start..end),
}
}
}
pub struct RecurLayer {
pub conv_state: CudaSlice<f32>, pub ssm_state: CudaSlice<f32>, pub ssm_state_alt: CudaSlice<f32>,
}
pub struct Cache {
pub kv: Vec<Option<KvLayer>>,
pub recur: Vec<Option<RecurLayer>>,
pub pos: usize,
pub max_ctx: usize,
pub last_logits_dev: Option<CudaSlice<f32>>,
pub dflash_taps: Option<DflashTapSink>,
}
fn full_attention_kv_layout(cfg: &ModelConfig, il: u32) -> (usize, usize, usize, usize) {
debug_assert_eq!(cfg.layer_kind(il), LayerKind::FullAttention);
let n_head_kv = cfg.n_head_kv as usize;
let (kv_dim_k, kv_dim_v) = match &cfg.gemma4 {
Some(g) => {
let hd = if g.swa_pattern[il as usize] {
g.key_length_swa
} else {
g.key_length_global
} as usize;
let d = match g.head_count_kv.get(il as usize) {
Some(n) => hd * *n as usize,
None => hd * n_head_kv,
};
(d, d)
}
None => (
cfg.head_dim_k as usize * n_head_kv,
cfg.head_dim_v as usize * n_head_kv,
),
};
assert!(
kv_dim_k % 32 == 0 && kv_dim_v % 32 == 0,
"KVQUANT requires per-layer kv_dim_k%32==0 && kv_dim_v%32==0 \
(layer {il}: k={kv_dim_k} v={kv_dim_v})"
);
let (kbb, vbb) = kv_blk_bytes();
let g4_global_fp8 = gkv_on()
&& cfg
.gemma4
.as_ref()
.is_some_and(|g| !g.swa_pattern[il as usize]);
let g4_windowed_fp8 = wkv_on()
&& cfg
.gemma4
.as_ref()
.is_some_and(|g| g.swa_pattern[il as usize]);
let qwen_fp8 = kv_fp8_on() && cfg.gemma4.is_none();
let (kbb_l, vbb_l) = if g4_global_fp8 || g4_windowed_fp8 || qwen_fp8 {
(32, 32)
} else {
(kbb, vbb)
};
(kv_dim_k, kv_dim_v, kbb_l, vbb_l)
}
fn kv_plane_allocation_bytes(rows: usize, token_bytes: usize) -> usize {
rows * token_bytes + 8
}
pub fn cache_bytes_per_token(cfg: &ModelConfig) -> usize {
let shared = cfg.gemma4.as_ref().map(|g| g.shared_kv_layers).unwrap_or(0);
(0..cfg.n_layer)
.filter(|&il| cfg.layer_kind(il) == LayerKind::FullAttention)
.filter(|&il| shared == 0 || il < cfg.n_layer - shared)
.map(|il| {
let (kv_dim_k, kv_dim_v, kbb, vbb) = full_attention_kv_layout(cfg, il);
(kv_dim_k / 32) * kbb + (kv_dim_v / 32) * vbb
})
.sum()
}
pub fn cache_ring_bytes_per_token(cfg: &ModelConfig) -> usize {
if !swa_ring_on() || !cfg.arch.is_step35() {
return 0;
}
let shared = cfg.gemma4.as_ref().map(|g| g.shared_kv_layers).unwrap_or(0);
(0..cfg.n_layer)
.filter(|&il| cfg.layer_kind(il) == LayerKind::FullAttention)
.filter(|&il| shared == 0 || il < cfg.n_layer - shared)
.filter(|&il| cfg.layer_geometry(il).is_some_and(|geometry| geometry.window.is_some()))
.map(|il| {
let (kv_dim_k, kv_dim_v, kbb, vbb) = full_attention_kv_layout(cfg, il);
(kv_dim_k / 32) * kbb + (kv_dim_v / 32) * vbb
})
.sum()
}
pub fn cache_ring_row_cap(cfg: &ModelConfig) -> usize {
if !swa_ring_on() || !cfg.arch.is_step35() {
return 0;
}
cfg.geometry
.as_ref()
.and_then(|table| table.classes().iter().find_map(|geometry| geometry.window))
.map(|window| swa_ring_rows(window as usize, usize::MAX))
.unwrap_or(0)
}
pub struct DflashTapSink {
pub layer_ids: Vec<usize>,
pub buf: CudaSlice<f32>,
pub hidden: usize,
pub t: usize,
}
pub struct CacheSnapshot {
pub kv_len: Vec<Option<usize>>, pub conv: Vec<Option<CudaSlice<f32>>>, pub ssm: Vec<Option<CudaSlice<f32>>>,
pub pos: usize,
}
impl Cache {
pub fn new(
e: &impl KvDev,
cfg: &ModelConfig,
max_ctx: usize,
) -> Result<Self, Box<dyn std::error::Error>> {
Self::new_inner(&|_| e, cfg, max_ctx)
}
pub fn new_pp2(
dev0: &dyn KvDev,
dev1: &dyn KvDev,
split: usize,
cfg: &ModelConfig,
max_ctx: usize,
) -> Result<Self, Box<dyn std::error::Error>> {
Self::new_inner(&|il| if il < split { dev0 } else { dev1 }, cfg, max_ctx)
}
pub fn new_ppn<'a>(
devs: &[&'a dyn KvDev],
fence: &[usize],
cfg: &ModelConfig,
max_ctx: usize,
) -> Result<Self, Box<dyn std::error::Error>> {
assert_eq!(devs.len() + 1, fence.len(), "ppn cache: devs vs fence mismatch");
let pick = |il: usize| -> &dyn KvDev {
let s = match fence[1..fence.len() - 1].binary_search(&il) {
Ok(k) => k + 1,
Err(k) => k,
};
devs[s.min(devs.len() - 1)]
};
Self::new_inner(&pick, cfg, max_ctx)
}
fn new_inner<'a>(
pick: &dyn Fn(usize) -> &'a dyn KvDev,
cfg: &ModelConfig,
max_ctx: usize,
) -> Result<Self, Box<dyn std::error::Error>> {
let n = cfg.n_layer as usize;
let mut kv = Vec::with_capacity(n);
let mut recur = Vec::with_capacity(n);
let head_dim_k = cfg.head_dim_k as usize;
let head_dim_v = cfg.head_dim_v as usize;
assert!(head_dim_k % 32 == 0 && head_dim_v % 32 == 0,
"KVQUANT requires head_dim_k%32==0 && head_dim_v%32==0 (got k={head_dim_k} v={head_dim_v})");
let (conv_dim, d_state, num_v, d_conv) = if let Some(s) = &cfg.ssm {
let num_k = s.group_count as usize;
let num_v = s.time_step_rank as usize;
let ds = s.state_size as usize;
(
ds * num_k * 2 + ds * num_v,
ds,
num_v,
s.conv_kernel as usize,
)
} else {
(0, 0, 0, 0)
};
for il in 0..cfg.n_layer {
let e = pick(il as usize);
let g4_shared = cfg.gemma4.as_ref().map(|g| g.shared_kv_layers).unwrap_or(0);
if g4_shared > 0 && il >= cfg.n_layer - g4_shared {
kv.push(None);
recur.push(None);
continue;
}
match cfg.layer_kind(il) {
LayerKind::FullAttention => {
let (kv_dim_k, kv_dim_v, kbb_l, vbb_l) =
full_attention_kv_layout(cfg, il);
let k_tok_bytes = (kv_dim_k / 32) * kbb_l;
let v_tok_bytes = (kv_dim_v / 32) * vbb_l;
let ring = if swa_ring_on() && cfg.arch.is_step35() {
cfg.layer_geometry(il)
.and_then(|geometry| geometry.window)
.map(|window| {
let window = window as usize;
KvRing::new(swa_ring_rows(window, max_ctx), window)
})
} else {
None
};
let alloc_rows = ring.as_ref().map(KvRing::rows).unwrap_or(max_ctx);
kv.push(Some(KvLayer {
k: e.alloc_u8(kv_plane_allocation_bytes(alloc_rows, k_tok_bytes))?,
v: e.alloc_u8(kv_plane_allocation_bytes(alloc_rows, v_tok_bytes))?,
kv_dim_k,
kv_dim_v,
k_tok_bytes,
v_tok_bytes,
len: 0,
ring,
len_d: e.htod_i32(&[0])?,
}));
recur.push(None);
}
LayerKind::LinearAttention => {
kv.push(None);
recur.push(Some(RecurLayer {
conv_state: e.zeros(conv_dim * (d_conv - 1))?,
ssm_state: e.zeros(d_state * d_state * num_v)?,
ssm_state_alt: e.zeros(d_state * d_state * num_v)?,
}));
}
}
}
Ok(Cache { kv, recur, pos: 0, max_ctx, dflash_taps: None, last_logits_dev: None })
}
pub fn has_swa_ring(&self) -> bool {
self.kv.iter().flatten().any(|layer| layer.ring.is_some())
}
pub fn can_rollback(&self, snap: &CacheSnapshot, accept_len: usize) -> bool {
self.kv.iter().zip(&snap.kv_len).all(|(layer, saved)| {
match (layer, saved) {
(Some(layer), Some(saved)) => layer
.ring
.as_ref()
.is_none_or(|ring| ring.can_rewind_to(saved + accept_len)),
_ => true,
}
})
}
pub fn snapshot(&self, e: &impl KvDev) -> Result<CacheSnapshot, Box<dyn std::error::Error>> {
let n = self.kv.len();
let mut kv_len = Vec::with_capacity(n);
let mut conv = Vec::with_capacity(n);
let mut ssm = Vec::with_capacity(n);
for il in 0..n {
match &self.kv[il] {
Some(kvl) => kv_len.push(Some(kvl.len)),
None => kv_len.push(None),
}
match &self.recur[il] {
Some(rl) => {
conv.push(Some(e.clone_dtod(&rl.conv_state)?));
ssm.push(Some(e.clone_dtod(&rl.ssm_state)?));
}
None => {
conv.push(None);
ssm.push(None);
}
}
}
Ok(CacheSnapshot {
kv_len,
conv,
ssm,
pos: self.pos,
})
}
pub fn snapshot_into(
&self,
e: &impl KvDev,
snap: &mut CacheSnapshot,
) -> Result<(), Box<dyn std::error::Error>> {
let n = self.kv.len();
for il in 0..n {
snap.kv_len[il] = self.kv[il].as_ref().map(|kvl| kvl.len);
if let Some(rl) = &self.recur[il] {
let dc = snap.conv[il]
.as_mut()
.expect("snapshot_into: shape mismatch (conv)");
let ds = snap.ssm[il]
.as_mut()
.expect("snapshot_into: shape mismatch (ssm)");
let (cn, sn) = (rl.conv_state.len(), rl.ssm_state.len());
e.copy_into(dc, 0, &rl.conv_state, cn)?;
e.copy_into(ds, 0, &rl.ssm_state, sn)?;
}
}
snap.pos = self.pos;
Ok(())
}
pub fn rollback(
&mut self,
e: &impl KvDev,
snap: &CacheSnapshot,
accept_len: usize,
) -> Result<(), Box<dyn std::error::Error>> {
if !self.can_rollback(snap, accept_len) {
return Err("SWA ring rewind checkpoint has been lapped; full re-prime required".into());
}
for il in 0..self.kv.len() {
if let (Some(kvl), Some(saved)) = (self.kv[il].as_mut(), snap.kv_len[il]) {
kvl.len = saved + accept_len;
e.set_i32_one(&mut kvl.len_d, kvl.len as i32)?;
}
if let Some(rl) = self.recur[il].as_mut() {
if let Some(c) = &snap.conv[il] {
e.copy_into(&mut rl.conv_state, 0, c, c.len())?;
}
if let Some(s) = &snap.ssm[il] {
e.copy_into(&mut rl.ssm_state, 0, s, s.len())?;
}
}
}
self.pos = snap.pos;
Ok(())
}
}
#[cfg(test)]
mod swa_ring_tests {
use super::{kv_plane_allocation_bytes, swa_ring_rows, KvRing, KvRingAppend};
#[test]
fn allocation_rows_cover_window_max_prime_and_alignment_slack() {
assert_eq!(swa_ring_rows(512, 262_144), 512 + 4096 + 31);
assert_eq!(swa_ring_rows(512, 4096), 4096);
assert_eq!(
kv_plane_allocation_bytes(4639, 1088),
4639 * 1088 + 8,
"the Step35 session plane allocates ring rows plus the existing tail pad",
);
}
#[test]
fn ring_matches_flat_bytes_before_wrap() {
let ring = KvRing::new(swa_ring_rows(512, 262_144), 512);
let flat: Vec<u32> = (0..1024).collect();
let mut physical = vec![u32::MAX; ring.rows()];
let KvRingAppend::Contiguous { write_row } = ring.append_plan(0, 0, flat.len()).unwrap()
else { panic!("first append unexpectedly wrapped") };
physical[write_row..write_row + flat.len()].copy_from_slice(&flat);
let view = ring.physical_range(0, flat.len()).unwrap();
assert_eq!(&physical[view], flat.as_slice());
}
#[test]
fn wrap_rebases_the_exact_aligned_prime_view() {
let mut ring = KvRing::new(swa_ring_rows(512, 262_144), 512);
let flat: Vec<u32> = (0..8192).collect();
let mut physical = vec![u32::MAX; ring.rows()];
let KvRingAppend::Contiguous { write_row } = ring.append_plan(0, 0, 4096).unwrap()
else { panic!("first prime chunk unexpectedly wrapped") };
physical[write_row..write_row + 4096].copy_from_slice(&flat[..4096]);
let off = (4096usize - (512 - 1)) & !31usize;
let KvRingAppend::Rebase {
src_row,
keep_rows,
new_base,
write_row,
} = ring.append_plan(4096, off, 4096).unwrap()
else { panic!("second prime chunk did not wrap") };
let retained = physical[src_row..src_row + keep_rows].to_vec();
physical[..keep_rows].copy_from_slice(&retained);
ring.apply_rebase(new_base);
physical[write_row..write_row + 4096].copy_from_slice(&flat[4096..8192]);
let view = ring.physical_range(off, 8192).unwrap();
assert_eq!(&physical[view], &flat[off..8192]);
assert_eq!(ring.base(), off);
}
#[test]
fn rewind_declines_once_the_required_window_was_lapped() {
let mut ring = KvRing::new(swa_ring_rows(512, 262_144), 512);
let KvRingAppend::Rebase { new_base, .. } =
ring.append_plan(4096, 3584, 4096).unwrap()
else { panic!("expected wrap") };
ring.apply_rebase(new_base);
assert!(ring.can_rewind_to(4095));
assert!(!ring.can_rewind_to(4094));
assert!(!ring.can_rewind_to(0));
}
}