use anyhow::{anyhow, Result};
use mlx_native::{DType, MlxBuffer, MlxDevice};
pub struct MlxKvCache {
pub k_packed: MlxBuffer,
pub k_norms: MlxBuffer,
pub v_packed: MlxBuffer,
pub v_norms: MlxBuffer,
pub capacity: usize,
pub is_sliding: bool,
pub write_pos: usize,
pub seq_len: usize,
}
impl MlxKvCache {
pub fn trim(&mut self, n_back: usize) -> Result<usize, &'static str> {
if self.is_sliding {
return Err("trim() not yet supported on sliding cache");
}
if n_back > self.seq_len {
return Err("trim n_back exceeds seq_len");
}
self.seq_len -= n_back;
self.write_pos = self.seq_len;
Ok(self.seq_len)
}
#[inline]
pub fn visible_len(&self) -> usize {
self.seq_len
}
}
pub struct HbKvBuffers {
pub k_packed: MlxBuffer,
pub k_norms: MlxBuffer,
pub v_packed: MlxBuffer,
pub v_norms: MlxBuffer,
pub capacity: usize,
pub is_sliding: bool,
#[allow(dead_code)]
pub norms_per_pos: usize,
}
pub struct MultiSeqHbKvBuffers {
pub n_seqs: u32,
pub k_packed: MlxBuffer,
pub k_norms: MlxBuffer,
pub v_packed: MlxBuffer,
pub v_norms: MlxBuffer,
pub capacity: usize,
pub is_sliding: bool,
#[allow(dead_code)]
pub norms_per_pos: usize,
pub seq_lens: Vec<u32>,
}
pub fn alloc_hb_kv_for_layer(
dev: &MlxDevice,
layer_idx: usize,
nkv: usize,
hd: usize,
cap: usize,
is_ring: bool,
n_seqs: u32,
) -> Result<MultiSeqHbKvBuffers> {
if n_seqs == 0 {
return Err(anyhow!(
"alloc_hb_kv_for_layer L{layer_idx}: n_seqs must be > 0"
));
}
if nkv == 0 || hd == 0 || cap == 0 {
return Err(anyhow!(
"alloc_hb_kv_for_layer L{layer_idx}: nkv/hd/cap must be > 0 \
(got nkv={nkv}, hd={hd}, cap={cap})"
));
}
let norms_per_pos = (hd / 256).max(1);
let n = n_seqs as usize;
let packed_bytes = n * nkv * cap * hd; let packed_shape = vec![n, nkv, cap, hd];
let norms_elems = n * nkv * cap * norms_per_pos;
let norms_bytes = norms_elems * std::mem::size_of::<f32>();
let norms_shape = vec![n, nkv, cap, norms_per_pos];
let mut k_packed = dev
.alloc_buffer(packed_bytes, DType::U8, packed_shape.clone())
.map_err(|e| anyhow!("hb_kv L{layer_idx} K packed: {e}"))?;
let mut k_norms = dev
.alloc_buffer(norms_bytes, DType::F32, norms_shape.clone())
.map_err(|e| anyhow!("hb_kv L{layer_idx} K norms: {e}"))?;
let mut v_packed = dev
.alloc_buffer(packed_bytes, DType::U8, packed_shape)
.map_err(|e| anyhow!("hb_kv L{layer_idx} V packed: {e}"))?;
let mut v_norms = dev
.alloc_buffer(norms_bytes, DType::F32, norms_shape)
.map_err(|e| anyhow!("hb_kv L{layer_idx} V norms: {e}"))?;
if let Ok(s) = k_packed.as_mut_slice::<u8>() {
s.fill(0);
}
if let Ok(s) = v_packed.as_mut_slice::<u8>() {
s.fill(0);
}
if let Ok(s) = k_norms.as_mut_slice::<f32>() {
s.fill(0.0);
}
if let Ok(s) = v_norms.as_mut_slice::<f32>() {
s.fill(0.0);
}
Ok(MultiSeqHbKvBuffers {
n_seqs,
k_packed,
k_norms,
v_packed,
v_norms,
capacity: cap,
is_sliding: is_ring,
norms_per_pos,
seq_lens: vec![0u32; n],
})
}
pub fn layer_type_to_alloc_params(
layer_type: crate::serve::config::LayerType,
sliding_window: usize,
max_position_embeddings: usize,
) -> (bool, usize) {
use crate::serve::config::LayerType;
match layer_type {
LayerType::Sliding => (true, sliding_window),
LayerType::Full => (false, max_position_embeddings),
}
}
pub fn layer_type_to_alloc_params_per_slot(
layer_type: crate::serve::config::LayerType,
sliding_window: usize,
max_position_embeddings: usize,
max_slots: usize,
) -> (bool, usize) {
use crate::serve::config::LayerType;
match layer_type {
LayerType::Sliding => (true, sliding_window),
LayerType::Full => (false, max_position_embeddings.div_ceil(max_slots.max(1))),
}
}
impl crate::serve::multi_seq_kv::MultiSeqKvCache for MultiSeqHbKvBuffers {
fn layout(&self) -> crate::serve::multi_seq_kv::MultiSeqLayout {
crate::serve::multi_seq_kv::MultiSeqLayout::SeparateSlots
}
fn slot_count(&self) -> u32 {
self.n_seqs
}
fn seq_len(
&self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<u32, crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
Ok(self.seq_lens[slot.0 as usize])
}
fn append_for_seq(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
n_tokens: u32,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
let cur = &mut self.seq_lens[slot.0 as usize];
*cur = cur.saturating_add(n_tokens);
Ok(())
}
fn drop_seq(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
self.seq_lens[slot.0 as usize] = 0;
Ok(())
}
fn fork_seq(
&mut self,
src: crate::serve::multi_seq_kv::SlotId,
dst: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if src.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot: src,
max_slots: self.n_seqs,
});
}
if dst.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot: dst,
max_slots: self.n_seqs,
});
}
if src == dst {
return Ok(());
}
let src_idx = src.0 as usize;
let dst_idx = dst.0 as usize;
let n_seqs = self.n_seqs as usize;
gemma4_copy_buffer_slot_region(&mut self.k_packed, src_idx, dst_idx, n_seqs).map_err(
|e| crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: gemma4_leak_static_str(format!(
"fork_seq: MultiSeqHbKvBuffers k_packed copy failed ({e})"
)),
},
)?;
gemma4_copy_buffer_slot_region(&mut self.k_norms, src_idx, dst_idx, n_seqs).map_err(
|e| crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: gemma4_leak_static_str(format!(
"fork_seq: MultiSeqHbKvBuffers k_norms copy failed ({e})"
)),
},
)?;
gemma4_copy_buffer_slot_region(&mut self.v_packed, src_idx, dst_idx, n_seqs).map_err(
|e| crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: gemma4_leak_static_str(format!(
"fork_seq: MultiSeqHbKvBuffers v_packed copy failed ({e})"
)),
},
)?;
gemma4_copy_buffer_slot_region(&mut self.v_norms, src_idx, dst_idx, n_seqs).map_err(
|e| crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: gemma4_leak_static_str(format!(
"fork_seq: MultiSeqHbKvBuffers v_norms copy failed ({e})"
)),
},
)?;
self.seq_lens[dst_idx] = self.seq_lens[src_idx];
Ok(())
}
}
#[inline]
fn gemma4_leak_static_str(s: String) -> &'static str {
Box::leak(s.into_boxed_str())
}
fn gemma4_copy_buffer_slot_region(
buf: &mut MlxBuffer,
src_idx: usize,
dst_idx: usize,
n_seqs: usize,
) -> Result<()> {
anyhow::ensure!(n_seqs > 0, "fork_seq: n_seqs must be > 0");
let total_bytes = buf.byte_len();
anyhow::ensure!(
total_bytes % n_seqs == 0,
"fork_seq: total_bytes={} not divisible by n_seqs={}",
total_bytes,
n_seqs
);
let per_slot_bytes = total_bytes / n_seqs;
anyhow::ensure!(
src_idx < n_seqs && dst_idx < n_seqs,
"fork_seq: src/dst out of buffer range \
(src={src_idx}, dst={dst_idx}, n_seqs={n_seqs})"
);
if per_slot_bytes == 0 {
return Ok(());
}
let bytes = buf
.as_mut_slice::<u8>()
.map_err(|e| anyhow!("fork_seq: as_mut_slice<u8>: {e}"))?;
let src_off = src_idx * per_slot_bytes;
bytes.copy_within(src_off..src_off + per_slot_bytes, dst_idx * per_slot_bytes);
Ok(())
}
impl MultiSeqHbKvBuffers {
pub fn reset_for_slot(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
self.seq_lens[slot.0 as usize] = 0;
Ok(())
}
}
pub struct DenseKvBuffers {
pub k: MlxBuffer,
pub v: MlxBuffer,
pub capacity: usize,
pub is_sliding: bool,
pub dtype: DType,
}
impl crate::serve::kv_persist::lcp_registry::ByteSized for DenseKvBuffers {
fn byte_len(&self) -> u64 {
(self.k.byte_len() + self.v.byte_len()) as u64
}
}
pub struct HybridKvBuffers {
pub k: MlxBuffer,
pub v_packed: MlxBuffer,
pub v_norms: MlxBuffer,
pub capacity: usize,
pub is_sliding: bool,
#[allow(dead_code)]
pub norms_per_pos: usize,
pub bf16_xlen_k: Option<MlxBuffer>,
pub bf16_xlen_v: Option<MlxBuffer>,
}
impl crate::serve::kv_persist::lcp_registry::ByteSized for HybridKvBuffers {
fn byte_len(&self) -> u64 {
(self.k.byte_len() + self.v_packed.byte_len() + self.v_norms.byte_len()) as u64
}
}
pub enum GemmaLcpLayerKv {
Dense(DenseKvBuffers),
DenseAndHybrid(DenseKvBuffers, HybridKvBuffers),
}
impl std::fmt::Debug for GemmaLcpLayerKv {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Dense(d) => f
.debug_struct("Dense")
.field("capacity", &d.capacity)
.field("is_sliding", &d.is_sliding)
.field("dtype", &d.dtype)
.finish(),
Self::DenseAndHybrid(d, h) => f
.debug_struct("DenseAndHybrid")
.field("dense_capacity", &d.capacity)
.field("hybrid_capacity", &h.capacity)
.field("is_sliding", &d.is_sliding)
.finish(),
}
}
}
impl GemmaLcpLayerKv {
pub fn dense(&self) -> &DenseKvBuffers {
match self {
Self::Dense(d) => d,
Self::DenseAndHybrid(d, _) => d,
}
}
pub fn hybrid(&self) -> Option<&HybridKvBuffers> {
match self {
Self::Dense(_) => None,
Self::DenseAndHybrid(_, h) => Some(h),
}
}
}
impl crate::serve::kv_persist::lcp_registry::ByteSized for GemmaLcpLayerKv {
fn byte_len(&self) -> u64 {
match self {
Self::Dense(d) => crate::serve::kv_persist::lcp_registry::ByteSized::byte_len(d),
Self::DenseAndHybrid(d, h) => {
crate::serve::kv_persist::lcp_registry::ByteSized::byte_len(d)
+ crate::serve::kv_persist::lcp_registry::ByteSized::byte_len(h)
}
}
}
}
pub(crate) fn alloc_hybrid_kv_for_layer(
dev: &MlxDevice,
layer_idx: usize,
nkv: usize,
hd: usize,
cap: usize,
is_ring: bool,
) -> anyhow::Result<HybridKvBuffers> {
let norms_per_pos = (hd / 256).max(1);
let norms_n = nkv * cap * norms_per_pos;
let k = dev
.alloc_buffer(nkv * cap * hd * 2, DType::F16, vec![nkv, cap, hd])
.map_err(|e| anyhow!("hybrid F16 K L{layer_idx}: {e}"))?;
let full_f16_v = std::env::var("HF2Q_FULL_F16_KV")
.ok()
.map(|v| matches!(v.as_str(), "1" | "true" | "on"))
.unwrap_or(false);
let (v_packed, v_norms) = if full_f16_v {
let v_f16 = dev
.alloc_buffer(nkv * cap * hd * 2, DType::F16, vec![nkv, cap, hd])
.map_err(|e| anyhow!("hybrid F16 V L{layer_idx}: {e}"))?;
let v_norms_dummy = dev
.alloc_buffer(4, DType::F32, vec![1])
.map_err(|e| anyhow!("hybrid V norms (dummy) L{layer_idx}: {e}"))?;
(v_f16, v_norms_dummy)
} else {
let v_p = dev
.alloc_buffer(nkv * cap * hd, DType::U8, vec![nkv, cap, hd])
.map_err(|e| anyhow!("hybrid V packed L{layer_idx}: {e}"))?;
let v_n = dev
.alloc_buffer(
norms_n * 4,
DType::F32,
if norms_per_pos == 1 {
vec![nkv, cap]
} else {
vec![nkv, cap, norms_per_pos]
},
)
.map_err(|e| anyhow!("hybrid V norms L{layer_idx}: {e}"))?;
(v_p, v_n)
};
let xlen_mode = std::env::var("HF2Q_DFLASH_XLEN_SDPA").as_deref() == Ok("1");
let (bf16_xlen_k, bf16_xlen_v) = if xlen_mode {
let bk = dev
.alloc_buffer(nkv * cap * hd * 2, DType::BF16, vec![nkv, cap, hd])
.map_err(|e| anyhow!("bf16 xlen K L{layer_idx}: {e}"))?;
let bv = dev
.alloc_buffer(nkv * cap * hd * 2, DType::BF16, vec![nkv, cap, hd])
.map_err(|e| anyhow!("bf16 xlen V L{layer_idx}: {e}"))?;
(Some(bk), Some(bv))
} else {
(None, None)
};
Ok(HybridKvBuffers {
k,
v_packed,
v_norms,
capacity: cap,
is_sliding: is_ring,
norms_per_pos,
bf16_xlen_k,
bf16_xlen_v,
})
}
pub struct MultiSeqHybridKvBuffers {
pub n_seqs: u32,
pub k: MlxBuffer,
pub v_packed: MlxBuffer,
pub v_norms: MlxBuffer,
pub capacity: usize,
pub is_sliding: bool,
#[allow(dead_code)]
pub norms_per_pos: usize,
pub bf16_xlen_k: Option<MlxBuffer>,
pub bf16_xlen_v: Option<MlxBuffer>,
pub seq_lens: Vec<u32>,
}
impl crate::serve::kv_persist::lcp_registry::ByteSized for MultiSeqHybridKvBuffers {
fn byte_len(&self) -> u64 {
let mut sum =
(self.k.byte_len() + self.v_packed.byte_len() + self.v_norms.byte_len()) as u64;
if let Some(ref bk) = self.bf16_xlen_k {
sum += bk.byte_len() as u64;
}
if let Some(ref bv) = self.bf16_xlen_v {
sum += bv.byte_len() as u64;
}
sum
}
}
pub fn alloc_multi_seq_hybrid_kv_for_layer(
dev: &MlxDevice,
layer_idx: usize,
nkv: usize,
hd: usize,
cap: usize,
is_ring: bool,
n_seqs: u32,
) -> Result<MultiSeqHybridKvBuffers> {
if n_seqs == 0 {
return Err(anyhow!(
"alloc_multi_seq_hybrid_kv_for_layer L{layer_idx}: n_seqs must be > 0"
));
}
if nkv == 0 || hd == 0 || cap == 0 {
return Err(anyhow!(
"alloc_multi_seq_hybrid_kv_for_layer L{layer_idx}: nkv/hd/cap must be \
> 0 (got nkv={nkv}, hd={hd}, cap={cap})"
));
}
let norms_per_pos = (hd / 256).max(1);
let n = n_seqs as usize;
let k_elems = n * nkv * cap * hd;
let k_bytes = k_elems * 2;
let k = dev
.alloc_buffer(k_bytes, DType::F16, vec![n, nkv, cap, hd])
.map_err(|e| anyhow!("multi-seq hybrid F16 K L{layer_idx}: {e}"))?;
let full_f16_v = std::env::var("HF2Q_FULL_F16_KV")
.ok()
.map(|v| matches!(v.as_str(), "1" | "true" | "on"))
.unwrap_or(false);
let (v_packed, v_norms) = if full_f16_v {
let v_elems = n * nkv * cap * hd;
let v_bytes = v_elems * 2;
let v_f16 = dev
.alloc_buffer(v_bytes, DType::F16, vec![n, nkv, cap, hd])
.map_err(|e| anyhow!("multi-seq hybrid F16 V L{layer_idx}: {e}"))?;
let v_norms_dummy = dev
.alloc_buffer(4, DType::F32, vec![1])
.map_err(|e| anyhow!("multi-seq hybrid V norms (dummy) L{layer_idx}: {e}"))?;
(v_f16, v_norms_dummy)
} else {
let v_packed_elems = n * nkv * cap * hd;
let v_packed_bytes = v_packed_elems; let v_p = dev
.alloc_buffer(v_packed_bytes, DType::U8, vec![n, nkv, cap, hd])
.map_err(|e| anyhow!("multi-seq hybrid V packed L{layer_idx}: {e}"))?;
let v_norms_elems = n * nkv * cap * norms_per_pos;
let v_norms_bytes = v_norms_elems * std::mem::size_of::<f32>();
let v_n = dev
.alloc_buffer(v_norms_bytes, DType::F32, vec![n, nkv, cap, norms_per_pos])
.map_err(|e| anyhow!("multi-seq hybrid V norms L{layer_idx}: {e}"))?;
(v_p, v_n)
};
let xlen_mode = std::env::var("HF2Q_DFLASH_XLEN_SDPA").as_deref() == Ok("1");
let (bf16_xlen_k, bf16_xlen_v) = if xlen_mode {
let xlen_elems = n * nkv * cap * hd;
let xlen_bytes = xlen_elems * 2;
let bk = dev
.alloc_buffer(xlen_bytes, DType::BF16, vec![n, nkv, cap, hd])
.map_err(|e| anyhow!("multi-seq hybrid bf16 xlen K L{layer_idx}: {e}"))?;
let bv = dev
.alloc_buffer(xlen_bytes, DType::BF16, vec![n, nkv, cap, hd])
.map_err(|e| anyhow!("multi-seq hybrid bf16 xlen V L{layer_idx}: {e}"))?;
(Some(bk), Some(bv))
} else {
(None, None)
};
Ok(MultiSeqHybridKvBuffers {
n_seqs,
k,
v_packed,
v_norms,
capacity: cap,
is_sliding: is_ring,
norms_per_pos,
bf16_xlen_k,
bf16_xlen_v,
seq_lens: vec![0u32; n],
})
}
impl crate::serve::multi_seq_kv::MultiSeqKvCache for MultiSeqHybridKvBuffers {
fn layout(&self) -> crate::serve::multi_seq_kv::MultiSeqLayout {
crate::serve::multi_seq_kv::MultiSeqLayout::SeparateSlots
}
fn slot_count(&self) -> u32 {
self.n_seqs
}
fn seq_len(
&self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<u32, crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
Ok(self.seq_lens[slot.0 as usize])
}
fn append_for_seq(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
n_tokens: u32,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
let cur = &mut self.seq_lens[slot.0 as usize];
*cur = cur.saturating_add(n_tokens);
Ok(())
}
fn drop_seq(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
self.seq_lens[slot.0 as usize] = 0;
Ok(())
}
fn fork_seq(
&mut self,
src: crate::serve::multi_seq_kv::SlotId,
dst: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if src.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot: src,
max_slots: self.n_seqs,
});
}
if dst.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot: dst,
max_slots: self.n_seqs,
});
}
if src == dst {
return Ok(());
}
let src_idx = src.0 as usize;
let dst_idx = dst.0 as usize;
let n_seqs = self.n_seqs as usize;
gemma4_copy_buffer_slot_region(&mut self.k, src_idx, dst_idx, n_seqs).map_err(|e| {
crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: gemma4_leak_static_str(format!(
"fork_seq: MultiSeqHybridKvBuffers k copy failed ({e})"
)),
}
})?;
gemma4_copy_buffer_slot_region(&mut self.v_packed, src_idx, dst_idx, n_seqs).map_err(
|e| crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: gemma4_leak_static_str(format!(
"fork_seq: MultiSeqHybridKvBuffers v_packed copy failed ({e})"
)),
},
)?;
if self.v_norms.byte_len() >= n_seqs {
gemma4_copy_buffer_slot_region(&mut self.v_norms, src_idx, dst_idx, n_seqs).map_err(
|e| crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: gemma4_leak_static_str(format!(
"fork_seq: MultiSeqHybridKvBuffers v_norms copy failed ({e})"
)),
},
)?;
}
if let Some(ref mut bk) = self.bf16_xlen_k {
gemma4_copy_buffer_slot_region(bk, src_idx, dst_idx, n_seqs).map_err(|e| {
crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: gemma4_leak_static_str(format!(
"fork_seq: MultiSeqHybridKvBuffers bf16_xlen_k copy failed ({e})"
)),
}
})?;
}
if let Some(ref mut bv) = self.bf16_xlen_v {
gemma4_copy_buffer_slot_region(bv, src_idx, dst_idx, n_seqs).map_err(|e| {
crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: gemma4_leak_static_str(format!(
"fork_seq: MultiSeqHybridKvBuffers bf16_xlen_v copy failed ({e})"
)),
}
})?;
}
self.seq_lens[dst_idx] = self.seq_lens[src_idx];
Ok(())
}
}
impl MultiSeqHybridKvBuffers {
pub fn reset_for_slot(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
self.seq_lens[slot.0 as usize] = 0;
Ok(())
}
}
pub struct MultiSeqDenseKvBuffers {
pub n_seqs: u32,
pub k: MlxBuffer,
pub v: MlxBuffer,
pub capacity: usize,
pub is_sliding: bool,
pub dtype: DType,
pub seq_lens: Vec<u32>,
}
impl crate::serve::kv_persist::lcp_registry::ByteSized for MultiSeqDenseKvBuffers {
fn byte_len(&self) -> u64 {
(self.k.byte_len() + self.v.byte_len()) as u64
}
}
pub fn alloc_multi_seq_dense_kv_for_layer(
dev: &MlxDevice,
layer_idx: usize,
nkv: usize,
hd: usize,
cap: usize,
is_ring: bool,
dtype: DType,
n_seqs: u32,
) -> Result<MultiSeqDenseKvBuffers> {
if n_seqs == 0 {
return Err(anyhow!(
"alloc_multi_seq_dense_kv_for_layer L{layer_idx}: n_seqs must be > 0"
));
}
if nkv == 0 || hd == 0 || cap == 0 {
return Err(anyhow!(
"alloc_multi_seq_dense_kv_for_layer L{layer_idx}: nkv/hd/cap must be \
> 0 (got nkv={nkv}, hd={hd}, cap={cap})"
));
}
let n = n_seqs as usize;
let elem_bytes = dtype.size_of();
let k_elems = n * nkv * cap * hd;
let k_bytes = k_elems * elem_bytes;
let k = dev
.alloc_buffer(k_bytes, dtype, vec![n, nkv, cap, hd])
.map_err(|e| anyhow!("multi-seq dense K L{layer_idx}: {e}"))?;
let v_elems = n * nkv * cap * hd;
let v_bytes = v_elems * elem_bytes;
let v = dev
.alloc_buffer(v_bytes, dtype, vec![n, nkv, cap, hd])
.map_err(|e| anyhow!("multi-seq dense V L{layer_idx}: {e}"))?;
Ok(MultiSeqDenseKvBuffers {
n_seqs,
k,
v,
capacity: cap,
is_sliding: is_ring,
dtype,
seq_lens: vec![0u32; n],
})
}
impl crate::serve::multi_seq_kv::MultiSeqKvCache for MultiSeqDenseKvBuffers {
fn layout(&self) -> crate::serve::multi_seq_kv::MultiSeqLayout {
crate::serve::multi_seq_kv::MultiSeqLayout::SeparateSlots
}
fn slot_count(&self) -> u32 {
self.n_seqs
}
fn seq_len(
&self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<u32, crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
Ok(self.seq_lens[slot.0 as usize])
}
fn append_for_seq(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
n_tokens: u32,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
let cur = &mut self.seq_lens[slot.0 as usize];
*cur = cur.saturating_add(n_tokens);
Ok(())
}
fn drop_seq(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
self.seq_lens[slot.0 as usize] = 0;
Ok(())
}
fn fork_seq(
&mut self,
src: crate::serve::multi_seq_kv::SlotId,
dst: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if src.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot: src,
max_slots: self.n_seqs,
});
}
if dst.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot: dst,
max_slots: self.n_seqs,
});
}
if src == dst {
return Ok(());
}
let src_idx = src.0 as usize;
let dst_idx = dst.0 as usize;
let n_seqs = self.n_seqs as usize;
gemma4_copy_buffer_slot_region(&mut self.k, src_idx, dst_idx, n_seqs).map_err(|e| {
crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: gemma4_leak_static_str(format!(
"fork_seq: MultiSeqDenseKvBuffers k copy failed ({e})"
)),
}
})?;
gemma4_copy_buffer_slot_region(&mut self.v, src_idx, dst_idx, n_seqs).map_err(|e| {
crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: gemma4_leak_static_str(format!(
"fork_seq: MultiSeqDenseKvBuffers v copy failed ({e})"
)),
}
})?;
self.seq_lens[dst_idx] = self.seq_lens[src_idx];
Ok(())
}
}
impl MultiSeqDenseKvBuffers {
pub fn reset_for_slot(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
self.seq_lens[slot.0 as usize] = 0;
Ok(())
}
}
pub struct MultiSeqMlxKvCache {
pub n_seqs: u32,
pub k_packed: MlxBuffer,
pub k_norms: MlxBuffer,
pub v_packed: MlxBuffer,
pub v_norms: MlxBuffer,
pub capacity: usize,
pub is_sliding: bool,
pub norms_per_pos: usize,
pub seq_lens: Vec<u32>,
}
impl crate::serve::kv_persist::lcp_registry::ByteSized for MultiSeqMlxKvCache {
fn byte_len(&self) -> u64 {
(self.k_packed.byte_len()
+ self.k_norms.byte_len()
+ self.v_packed.byte_len()
+ self.v_norms.byte_len()) as u64
}
}
pub fn alloc_multi_seq_mlx_kv_for_layer(
dev: &MlxDevice,
layer_idx: usize,
nkv: usize,
hd: usize,
cap: usize,
is_ring: bool,
norms_per_pos: usize,
n_seqs: u32,
) -> Result<MultiSeqMlxKvCache> {
if n_seqs == 0 {
return Err(anyhow!(
"alloc_multi_seq_mlx_kv_for_layer L{layer_idx}: n_seqs must be > 0"
));
}
if nkv == 0 || hd == 0 || cap == 0 || norms_per_pos == 0 {
return Err(anyhow!(
"alloc_multi_seq_mlx_kv_for_layer L{layer_idx}: nkv/hd/cap/norms_per_pos \
must be > 0 (got nkv={nkv}, hd={hd}, cap={cap}, norms_per_pos={norms_per_pos})"
));
}
if hd % 2 != 0 {
return Err(anyhow!(
"alloc_multi_seq_mlx_kv_for_layer L{layer_idx}: hd must be even for 4-bit \
nibble-packed K/V (hd/2 stride; got hd={hd})"
));
}
let n = n_seqs as usize;
let hd_half = hd / 2;
let k_packed_elems = n * nkv * cap * hd_half;
let k_packed_bytes = k_packed_elems; let k_packed = dev
.alloc_buffer(k_packed_bytes, DType::U8, vec![n, nkv, cap, hd_half])
.map_err(|e| anyhow!("multi-seq MLX K packed L{layer_idx}: {e}"))?;
let k_norms_elems = n * nkv * cap * norms_per_pos;
let k_norms_bytes = k_norms_elems * 4;
let k_norms_shape: Vec<usize> = if norms_per_pos == 1 {
vec![n, nkv, cap]
} else {
vec![n, nkv, cap, norms_per_pos]
};
let k_norms = dev
.alloc_buffer(k_norms_bytes, DType::F32, k_norms_shape)
.map_err(|e| anyhow!("multi-seq MLX K norms L{layer_idx}: {e}"))?;
let v_packed_elems = n * nkv * cap * hd_half;
let v_packed_bytes = v_packed_elems;
let v_packed = dev
.alloc_buffer(v_packed_bytes, DType::U8, vec![n, nkv, cap, hd_half])
.map_err(|e| anyhow!("multi-seq MLX V packed L{layer_idx}: {e}"))?;
let v_norms_elems = n * nkv * cap * norms_per_pos;
let v_norms_bytes = v_norms_elems * 4;
let v_norms_shape: Vec<usize> = if norms_per_pos == 1 {
vec![n, nkv, cap]
} else {
vec![n, nkv, cap, norms_per_pos]
};
let v_norms = dev
.alloc_buffer(v_norms_bytes, DType::F32, v_norms_shape)
.map_err(|e| anyhow!("multi-seq MLX V norms L{layer_idx}: {e}"))?;
Ok(MultiSeqMlxKvCache {
n_seqs,
k_packed,
k_norms,
v_packed,
v_norms,
capacity: cap,
is_sliding: is_ring,
norms_per_pos,
seq_lens: vec![0u32; n],
})
}
impl crate::serve::multi_seq_kv::MultiSeqKvCache for MultiSeqMlxKvCache {
fn layout(&self) -> crate::serve::multi_seq_kv::MultiSeqLayout {
crate::serve::multi_seq_kv::MultiSeqLayout::SeparateSlots
}
fn slot_count(&self) -> u32 {
self.n_seqs
}
fn seq_len(
&self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<u32, crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
Ok(self.seq_lens[slot.0 as usize])
}
fn append_for_seq(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
n_tokens: u32,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
let cur = &mut self.seq_lens[slot.0 as usize];
*cur = cur.saturating_add(n_tokens);
Ok(())
}
fn drop_seq(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
self.seq_lens[slot.0 as usize] = 0;
Ok(())
}
fn fork_seq(
&mut self,
src: crate::serve::multi_seq_kv::SlotId,
dst: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if src.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot: src,
max_slots: self.n_seqs,
});
}
if dst.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot: dst,
max_slots: self.n_seqs,
});
}
if src == dst {
return Ok(());
}
let src_idx = src.0 as usize;
let dst_idx = dst.0 as usize;
let n_seqs = self.n_seqs as usize;
gemma4_copy_buffer_slot_region(&mut self.k_packed, src_idx, dst_idx, n_seqs).map_err(
|e| crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: gemma4_leak_static_str(format!(
"fork_seq: MultiSeqMlxKvCache k_packed copy failed ({e})"
)),
},
)?;
gemma4_copy_buffer_slot_region(&mut self.k_norms, src_idx, dst_idx, n_seqs).map_err(
|e| crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: gemma4_leak_static_str(format!(
"fork_seq: MultiSeqMlxKvCache k_norms copy failed ({e})"
)),
},
)?;
gemma4_copy_buffer_slot_region(&mut self.v_packed, src_idx, dst_idx, n_seqs).map_err(
|e| crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: gemma4_leak_static_str(format!(
"fork_seq: MultiSeqMlxKvCache v_packed copy failed ({e})"
)),
},
)?;
gemma4_copy_buffer_slot_region(&mut self.v_norms, src_idx, dst_idx, n_seqs).map_err(
|e| crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: gemma4_leak_static_str(format!(
"fork_seq: MultiSeqMlxKvCache v_norms copy failed ({e})"
)),
},
)?;
self.seq_lens[dst_idx] = self.seq_lens[src_idx];
Ok(())
}
}
impl MultiSeqMlxKvCache {
pub fn reset_for_slot(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= self.n_seqs {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: self.n_seqs,
});
}
self.seq_lens[slot.0 as usize] = 0;
Ok(())
}
}
impl crate::serve::multi_seq_kv::MultiSeqKvCache for DenseKvBuffers {
fn layout(&self) -> crate::serve::multi_seq_kv::MultiSeqLayout {
crate::serve::multi_seq_kv::MultiSeqLayout::SeparateSlots
}
fn slot_count(&self) -> u32 {
1
}
fn seq_len(
&self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<u32, crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= 1 {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: 1,
});
}
Ok(0)
}
fn append_for_seq(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
_n_tokens: u32,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= 1 {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: 1,
});
}
Err(
crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: "DenseKvBuffers::append_for_seq (legacy single-seq path; full multi-seq lift shipped in ADR-040 Phase A3b iter-2 — use MultiSeqDenseKvBuffers via alloc_multi_seq_dense_kv_for_layer)",
},
)
}
fn drop_seq(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= 1 {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: 1,
});
}
Err(
crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: "DenseKvBuffers::drop_seq (legacy single-seq path; full multi-seq lift shipped in ADR-040 Phase A3b iter-2 — use MultiSeqDenseKvBuffers via alloc_multi_seq_dense_kv_for_layer)",
},
)
}
fn fork_seq(
&mut self,
src: crate::serve::multi_seq_kv::SlotId,
dst: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if src.0 >= 1 {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot: src,
max_slots: 1,
});
}
if dst.0 >= 1 {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot: dst,
max_slots: 1,
});
}
Ok(())
}
}
impl crate::serve::multi_seq_kv::MultiSeqKvCache for MlxKvCache {
fn layout(&self) -> crate::serve::multi_seq_kv::MultiSeqLayout {
crate::serve::multi_seq_kv::MultiSeqLayout::SeparateSlots
}
fn slot_count(&self) -> u32 {
1
}
fn seq_len(
&self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<u32, crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= 1 {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: 1,
});
}
Ok(u32::try_from(self.seq_len).unwrap_or(u32::MAX))
}
fn append_for_seq(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
_n_tokens: u32,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= 1 {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: 1,
});
}
Err(
crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: "MlxKvCache::append_for_seq (legacy 4-bit single-seq path; full multi-seq lift shipped in ADR-040 Phase A3b iter-3 — use MultiSeqMlxKvCache via alloc_multi_seq_mlx_kv_for_layer)",
},
)
}
fn drop_seq(
&mut self,
slot: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if slot.0 >= 1 {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot,
max_slots: 1,
});
}
Err(
crate::serve::multi_seq_kv::MultiSeqError::CapabilityUnsupported {
capability: "MlxKvCache::drop_seq (legacy 4-bit single-seq path; full multi-seq lift shipped in ADR-040 Phase A3b iter-3 — use MultiSeqMlxKvCache via alloc_multi_seq_mlx_kv_for_layer)",
},
)
}
fn fork_seq(
&mut self,
src: crate::serve::multi_seq_kv::SlotId,
dst: crate::serve::multi_seq_kv::SlotId,
) -> Result<(), crate::serve::multi_seq_kv::MultiSeqError> {
if src.0 >= 1 {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot: src,
max_slots: 1,
});
}
if dst.0 >= 1 {
return Err(crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange {
slot: dst,
max_slots: 1,
});
}
Ok(())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum DecodeRegime {
#[default]
Default,
#[allow(dead_code)]
ForceTq,
#[allow(dead_code)]
ForceDense,
}
#[cfg(test)]
mod tests {
use super::*;
fn skip_dev() -> Option<MlxDevice> {
match MlxDevice::new() {
Ok(d) => Some(d),
Err(_) => {
eprintln!("skip: no MlxDevice");
None
}
}
}
#[test]
fn mlx_kv_cache_trim_linear_decrements_seq_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let buf = || dev.alloc_buffer(4, DType::F32, vec![1]).unwrap();
let mut cache = MlxKvCache {
k_packed: buf(),
k_norms: buf(),
v_packed: buf(),
v_norms: buf(),
capacity: 16,
is_sliding: false,
write_pos: 8,
seq_len: 8,
};
let new_len = cache.trim(3).unwrap();
assert_eq!(new_len, 5);
assert_eq!(cache.seq_len, 5);
assert_eq!(cache.write_pos, 5);
}
#[test]
fn mlx_kv_cache_trim_sliding_errors() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let buf = || dev.alloc_buffer(4, DType::F32, vec![1]).unwrap();
let mut cache = MlxKvCache {
k_packed: buf(),
k_norms: buf(),
v_packed: buf(),
v_norms: buf(),
capacity: 16,
is_sliding: true,
write_pos: 4,
seq_len: 4,
};
assert!(cache.trim(1).is_err());
}
#[test]
fn mlx_kv_cache_trim_overflow_errors() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let buf = || dev.alloc_buffer(4, DType::F32, vec![1]).unwrap();
let mut cache = MlxKvCache {
k_packed: buf(),
k_norms: buf(),
v_packed: buf(),
v_norms: buf(),
capacity: 16,
is_sliding: false,
write_pos: 3,
seq_len: 3,
};
assert!(cache.trim(10).is_err());
}
#[test]
fn mlx_kv_cache_visible_len_eq_seq_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let buf = || dev.alloc_buffer(4, DType::F32, vec![1]).unwrap();
let cache = MlxKvCache {
k_packed: buf(),
k_norms: buf(),
v_packed: buf(),
v_norms: buf(),
capacity: 32,
is_sliding: false,
write_pos: 7,
seq_len: 7,
};
assert_eq!(cache.visible_len(), cache.seq_len);
}
#[test]
fn decode_regime_default_via_default_trait() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let r: DecodeRegime = Default::default();
assert_eq!(r, DecodeRegime::Default);
}
#[test]
fn decode_regime_variants_distinct() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
assert_ne!(DecodeRegime::Default, DecodeRegime::ForceTq);
assert_ne!(DecodeRegime::Default, DecodeRegime::ForceDense);
assert_ne!(DecodeRegime::ForceTq, DecodeRegime::ForceDense);
}
#[test]
fn hybrid_kv_buffers_byte_len_sums_fields() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2;
let cap = 4;
let hd = 256;
let k = dev
.alloc_buffer(nkv * cap * hd * 2, DType::F16, vec![nkv, cap, hd])
.unwrap();
let v_packed = dev
.alloc_buffer(nkv * cap * hd, DType::U8, vec![nkv, cap, hd])
.unwrap();
let v_norms = dev
.alloc_buffer(nkv * cap * 4, DType::F32, vec![nkv, cap])
.unwrap();
let k_bytes = k.byte_len();
let vp_bytes = v_packed.byte_len();
let vn_bytes = v_norms.byte_len();
let buf = HybridKvBuffers {
k,
v_packed,
v_norms,
capacity: cap,
is_sliding: false,
norms_per_pos: 1,
bf16_xlen_k: None,
bf16_xlen_v: None,
};
use crate::serve::kv_persist::lcp_registry::ByteSized;
assert_eq!(buf.byte_len(), (k_bytes + vp_bytes + vn_bytes) as u64);
}
#[test]
fn dense_kv_buffers_byte_len_sums_k_plus_v() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2;
let cap = 8;
let hd = 256;
let k = dev
.alloc_buffer(nkv * cap * hd * 4, DType::F32, vec![nkv, cap, hd])
.unwrap();
let v = dev
.alloc_buffer(nkv * cap * hd * 4, DType::F32, vec![nkv, cap, hd])
.unwrap();
let kb = k.byte_len();
let vb = v.byte_len();
let buf = DenseKvBuffers {
k,
v,
capacity: cap,
is_sliding: false,
dtype: DType::F32,
};
use crate::serve::kv_persist::lcp_registry::ByteSized;
assert_eq!(buf.byte_len(), (kb + vb) as u64);
}
#[test]
fn alloc_hybrid_kv_for_layer_no_xlen_no_full_f16() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
std::env::remove_var("HF2Q_FULL_F16_KV");
std::env::remove_var("HF2Q_DFLASH_XLEN_SDPA");
let buf = alloc_hybrid_kv_for_layer(&dev, 0, 2, 256, 8, false).unwrap();
assert!(buf.bf16_xlen_k.is_none());
assert!(buf.bf16_xlen_v.is_none());
assert_eq!(buf.capacity, 8);
assert!(!buf.is_sliding);
assert_eq!(buf.norms_per_pos, 1);
}
#[test]
fn alloc_hybrid_kv_for_layer_full_f16_v_allocates_f16() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
std::env::set_var("HF2Q_FULL_F16_KV", "1");
std::env::remove_var("HF2Q_DFLASH_XLEN_SDPA");
let buf = alloc_hybrid_kv_for_layer(&dev, 1, 2, 256, 4, true).unwrap();
assert_eq!(buf.v_norms.byte_len(), 4);
assert!(buf.is_sliding);
std::env::remove_var("HF2Q_FULL_F16_KV");
}
#[test]
fn alloc_hybrid_kv_for_layer_xlen_allocates_bf16_buffers() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
std::env::remove_var("HF2Q_FULL_F16_KV");
std::env::set_var("HF2Q_DFLASH_XLEN_SDPA", "1");
let buf = alloc_hybrid_kv_for_layer(&dev, 2, 2, 256, 4, false).unwrap();
assert!(buf.bf16_xlen_k.is_some());
assert!(buf.bf16_xlen_v.is_some());
std::env::remove_var("HF2Q_DFLASH_XLEN_SDPA");
}
#[test]
fn alloc_hybrid_kv_for_layer_norms_per_pos_d256_d512() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
std::env::remove_var("HF2Q_FULL_F16_KV");
std::env::remove_var("HF2Q_DFLASH_XLEN_SDPA");
let buf256 = alloc_hybrid_kv_for_layer(&dev, 0, 2, 256, 4, false).unwrap();
assert_eq!(buf256.norms_per_pos, 1);
let buf512 = alloc_hybrid_kv_for_layer(&dev, 0, 2, 512, 4, false).unwrap();
assert_eq!(buf512.norms_per_pos, 2);
}
use crate::serve::multi_seq_kv::{MultiSeqError, MultiSeqKvCache as _, MultiSeqLayout, SlotId};
#[test]
fn h6_hb_kv_buffers_n_seqs_4_byte_scale() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 8usize;
let baseline =
alloc_hb_kv_for_layer(&dev, 0, nkv, hd, cap, false, 1).expect("H6: alloc at n_seqs=1");
let lifted =
alloc_hb_kv_for_layer(&dev, 0, nkv, hd, cap, false, 4).expect("H6: alloc at n_seqs=4");
assert_eq!(baseline.n_seqs, 1, "H6: baseline n_seqs=1");
assert_eq!(lifted.n_seqs, 4, "H6: lifted n_seqs=4");
assert_eq!(
lifted.k_packed.byte_len(),
baseline.k_packed.byte_len() * 4,
"H6 FALSIFIED: k_packed does not scale 4× ({} != {} * 4 = {})",
lifted.k_packed.byte_len(),
baseline.k_packed.byte_len(),
baseline.k_packed.byte_len() * 4
);
assert_eq!(
lifted.v_packed.byte_len(),
baseline.v_packed.byte_len() * 4,
"H6 FALSIFIED: v_packed does not scale 4× ({} != {} * 4 = {})",
lifted.v_packed.byte_len(),
baseline.v_packed.byte_len(),
baseline.v_packed.byte_len() * 4
);
assert_eq!(
lifted.k_norms.byte_len(),
baseline.k_norms.byte_len() * 4,
"H6 FALSIFIED: k_norms does not scale 4× ({} != {})",
lifted.k_norms.byte_len(),
baseline.k_norms.byte_len() * 4
);
assert_eq!(
lifted.v_norms.byte_len(),
baseline.v_norms.byte_len() * 4,
"H6 FALSIFIED: v_norms does not scale 4×"
);
assert_eq!(baseline.seq_lens.len(), 1, "H6: baseline seq_lens.len()");
assert_eq!(lifted.seq_lens.len(), 4, "H6: lifted seq_lens.len()");
assert!(
baseline.seq_lens.iter().all(|&x| x == 0),
"H6: baseline seq_lens initialized to 0"
);
assert!(
lifted.seq_lens.iter().all(|&x| x == 0),
"H6: lifted seq_lens initialized to 0"
);
}
#[test]
fn h7_hb_kv_sliding_per_slot_isolation() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 4usize; let mut cache = alloc_hb_kv_for_layer(&dev, 0, nkv, hd, cap, true, 2)
.expect("H7: alloc n_seqs=2 sliding");
assert!(cache.is_sliding, "H7: sliding flag propagated");
assert_eq!(cache.n_seqs, 2);
let slot_packed = nkv * cap * hd;
let total_packed = cache.k_packed.byte_len();
assert_eq!(
total_packed,
2 * slot_packed,
"H7 fixture sanity: total bytes = 2 * slot_packed"
);
{
let k_slice = cache
.k_packed
.as_mut_slice::<u8>()
.expect("k_packed u8 mut");
for (i, b) in k_slice[slot_packed..2 * slot_packed].iter_mut().enumerate() {
*b = ((i % 251) + 1) as u8;
}
}
{
let v_slice = cache
.v_packed
.as_mut_slice::<u8>()
.expect("v_packed u8 mut");
for (i, b) in v_slice[slot_packed..2 * slot_packed].iter_mut().enumerate() {
*b = ((i % 253) + 1) as u8;
}
}
let k_slot1_before: Vec<u8> = cache.k_packed.as_slice::<u8>().expect("k_packed u8")
[slot_packed..2 * slot_packed]
.to_vec();
let v_slot1_before: Vec<u8> = cache.v_packed.as_slice::<u8>().expect("v_packed u8")
[slot_packed..2 * slot_packed]
.to_vec();
let k_slot0_before: Vec<u8> =
cache.k_packed.as_slice::<u8>().expect("k_packed u8")[0..slot_packed].to_vec();
assert!(
k_slot0_before.iter().all(|&b| b == 0),
"H7 fixture sanity: slot 0 K region zero-init"
);
cache
.append_for_seq(SlotId(0), 3)
.expect("H7: append slot 0 cursor");
assert_eq!(cache.seq_lens[0], 3, "H7: slot 0 cursor advanced");
assert_eq!(cache.seq_lens[1], 0, "H7: slot 1 cursor untouched");
let k_slot1_after: Vec<u8> = cache.k_packed.as_slice::<u8>().expect("k_packed u8")
[slot_packed..2 * slot_packed]
.to_vec();
let v_slot1_after: Vec<u8> = cache.v_packed.as_slice::<u8>().expect("v_packed u8")
[slot_packed..2 * slot_packed]
.to_vec();
assert_eq!(
k_slot1_before, k_slot1_after,
"H7 FALSIFIED: slot 1's k_packed bytes changed after slot-0 \
cursor advance — per-slot isolation invariant broken"
);
assert_eq!(
v_slot1_before, v_slot1_after,
"H7 FALSIFIED: slot 1's v_packed bytes changed after slot-0 \
cursor advance"
);
}
#[test]
fn h8_alloc_hb_kv_for_layer_byte_equivalent_to_pre_refactor() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 8usize;
let norms_per_pos = (hd / 256).max(1);
let expected_packed_bytes = nkv * cap * hd; let expected_norms_bytes = nkv * cap * norms_per_pos * std::mem::size_of::<f32>();
let helper =
alloc_hb_kv_for_layer(&dev, 0, nkv, hd, cap, false, 1).expect("H8: helper at n_seqs=1");
assert_eq!(
helper.k_packed.byte_len(),
expected_packed_bytes,
"H8 FALSIFIED: k_packed bytes diverge from inline formula \
({} != {})",
helper.k_packed.byte_len(),
expected_packed_bytes
);
assert_eq!(
helper.v_packed.byte_len(),
expected_packed_bytes,
"H8 FALSIFIED: v_packed bytes diverge from inline formula"
);
assert_eq!(
helper.k_norms.byte_len(),
expected_norms_bytes,
"H8 FALSIFIED: k_norms bytes diverge from inline formula \
({} != {})",
helper.k_norms.byte_len(),
expected_norms_bytes
);
assert_eq!(
helper.v_norms.byte_len(),
expected_norms_bytes,
"H8 FALSIFIED: v_norms bytes diverge from inline formula"
);
assert_eq!(
helper.k_packed.shape(),
&[1, nkv, cap, hd],
"H8: helper k_packed shape includes leading n_seqs=1 axis"
);
assert_eq!(helper.norms_per_pos, norms_per_pos);
assert_eq!(helper.capacity, cap);
}
#[test]
fn gemma4_hb_kv_n_seqs_outermost_axis() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let cache_1 = alloc_hb_kv_for_layer(&dev, 0, 2, 256, 8, false, 1).expect("alloc n_seqs=1");
let cache_4 = alloc_hb_kv_for_layer(&dev, 0, 2, 256, 8, false, 4).expect("alloc n_seqs=4");
for (name, b1, b4) in [
("k_packed", &cache_1.k_packed, &cache_4.k_packed),
("v_packed", &cache_1.v_packed, &cache_4.v_packed),
("k_norms", &cache_1.k_norms, &cache_4.k_norms),
("v_norms", &cache_1.v_norms, &cache_4.v_norms),
] {
let s1 = b1.shape().to_vec();
let s4 = b4.shape().to_vec();
assert_eq!(
s1.len(),
4,
"M5: {name} (n_seqs=1) must be 4-D; got {:?}",
s1
);
assert_eq!(
s4.len(),
4,
"M5: {name} (n_seqs=4) must be 4-D; got {:?}",
s4
);
assert_eq!(
s1[0], 1,
"M5: {name} baseline shape[0] must be n_seqs=1; got {:?}",
s1
);
assert_eq!(
s4[0], 4,
"M5 FALSIFIED: {name} shape[0] must be n_seqs=4 \
(n_seqs landed on wrong axis); got {:?}",
s4
);
assert_eq!(
&s4[1..],
&s1[1..],
"M5 FALSIFIED: {name} non-n_seqs dims diverge between \
n_seqs=1 ({:?}) and n_seqs=4 ({:?})",
s1,
s4
);
}
}
#[test]
fn gemma4_hb_kv_slot_count_matches_n_seqs() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let c1 = alloc_hb_kv_for_layer(&dev, 0, 2, 256, 8, false, 1).expect("alloc 1");
let c4 = alloc_hb_kv_for_layer(&dev, 0, 2, 256, 8, false, 4).expect("alloc 4");
assert_eq!(c1.slot_count(), 1);
assert_eq!(c4.slot_count(), 4);
}
#[test]
fn gemma4_hb_kv_layout_is_separate_slots() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let c = alloc_hb_kv_for_layer(&dev, 0, 2, 256, 8, false, 4).expect("alloc");
assert_eq!(c.layout(), MultiSeqLayout::SeparateSlots);
}
#[test]
fn gemma4_hb_kv_slot_out_of_range_errors_named() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let mut c = alloc_hb_kv_for_layer(&dev, 0, 2, 256, 8, false, 4).expect("alloc");
let err = c.seq_len(SlotId(4)).expect_err("slot 4 OOR for n_seqs=4");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(4),
max_slots: 4
}
);
let err = c.seq_len(SlotId(99)).expect_err("slot 99 OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(99),
max_slots: 4
}
);
let err = c.append_for_seq(SlotId(4), 1).expect_err("append OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(4),
max_slots: 4
}
);
let err = c.drop_seq(SlotId(4)).expect_err("drop OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(4),
max_slots: 4
}
);
let err = c
.fork_seq(SlotId(4), SlotId(5))
.expect_err("fork: src OOR first");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(4),
max_slots: 4
}
);
let err = c.fork_seq(SlotId(0), SlotId(4)).expect_err("fork: dst OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(4),
max_slots: 4
}
);
}
#[test]
fn gemma4_hb_kv_append_advances_target_slot_only() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let mut c = alloc_hb_kv_for_layer(&dev, 0, 2, 256, 8, false, 4).expect("alloc");
for s in 0..4 {
assert_eq!(c.seq_len(SlotId(s)).expect("seq_len in range"), 0);
}
c.append_for_seq(SlotId(0), 5).expect("append slot 0");
c.append_for_seq(SlotId(2), 3).expect("append slot 2");
assert_eq!(c.seq_len(SlotId(0)).unwrap(), 5);
assert_eq!(c.seq_len(SlotId(1)).unwrap(), 0, "slot 1 untouched");
assert_eq!(c.seq_len(SlotId(2)).unwrap(), 3);
assert_eq!(c.seq_len(SlotId(3)).unwrap(), 0, "slot 3 untouched");
}
#[test]
fn gemma4_hb_kv_drop_resets_seq_len_for_target_slot_only() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let mut c = alloc_hb_kv_for_layer(&dev, 0, 2, 256, 8, false, 4).expect("alloc");
c.append_for_seq(SlotId(0), 10).unwrap();
c.append_for_seq(SlotId(1), 20).unwrap();
c.append_for_seq(SlotId(2), 30).unwrap();
c.append_for_seq(SlotId(3), 40).unwrap();
c.drop_seq(SlotId(2)).expect("drop slot 2");
assert_eq!(c.seq_len(SlotId(0)).unwrap(), 10);
assert_eq!(c.seq_len(SlotId(1)).unwrap(), 20);
assert_eq!(c.seq_len(SlotId(2)).unwrap(), 0, "slot 2 reset");
assert_eq!(c.seq_len(SlotId(3)).unwrap(), 40);
assert_eq!(c.seq_lens[2], 0, "underlying cursor wiped");
assert_eq!(c.seq_lens[0], 10, "untouched cursors preserved");
}
#[test]
fn gemma4_hb_kv_drop_does_not_zero_k_packed_buffer() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 4usize;
let mut c = alloc_hb_kv_for_layer(&dev, 0, nkv, hd, cap, false, 2).expect("alloc n_seqs=2");
let slot_packed = nkv * cap * hd;
{
let k = c.k_packed.as_mut_slice::<u8>().expect("k_packed u8 mut");
for (i, b) in k[..slot_packed].iter_mut().enumerate() {
*b = (((i * 7) % 251) + 1) as u8;
}
}
{
let v = c.v_packed.as_mut_slice::<u8>().expect("v_packed u8 mut");
for (i, b) in v[..slot_packed].iter_mut().enumerate() {
*b = (((i * 11) % 253) + 1) as u8;
}
}
c.append_for_seq(SlotId(0), 2).expect("append slot 0");
let k_before: Vec<u8> =
c.k_packed.as_slice::<u8>().expect("k_packed u8")[..slot_packed].to_vec();
let v_before: Vec<u8> =
c.v_packed.as_slice::<u8>().expect("v_packed u8")[..slot_packed].to_vec();
assert!(
k_before.iter().any(|&b| b != 0),
"M4-G fixture sanity: deterministic upload must produce \
non-zero bytes (else test is vacuous)"
);
c.drop_seq(SlotId(0)).expect("drop slot 0");
assert_eq!(c.seq_lens[0], 0, "cursor reset");
let k_after: Vec<u8> =
c.k_packed.as_slice::<u8>().expect("k_packed u8 after")[..slot_packed].to_vec();
let v_after: Vec<u8> =
c.v_packed.as_slice::<u8>().expect("v_packed u8 after")[..slot_packed].to_vec();
assert_eq!(
k_before, k_after,
"M4-G FALSIFIED: drop_seq mutated k_packed contents for \
slot 0. Per Phase A3a contract, drop_seq is cursor-only; \
buffer-content reset is kernel-dispatcher-owned at Phase \
B4c. An in-place zero here would break the next \
admission's buffer-reuse correctness."
);
assert_eq!(
v_before, v_after,
"M4-G FALSIFIED: drop_seq mutated v_packed contents for slot 0."
);
}
#[test]
fn gemma4_hb_kv_fork_to_self_is_noop_ok() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let mut c = alloc_hb_kv_for_layer(&dev, 0, 2, 256, 8, false, 4).expect("alloc");
c.append_for_seq(SlotId(2), 9).unwrap();
c.fork_seq(SlotId(2), SlotId(2)).expect("fork self ok");
assert_eq!(c.seq_len(SlotId(2)).unwrap(), 9);
assert_eq!(c.seq_len(SlotId(0)).unwrap(), 0);
assert_eq!(c.seq_len(SlotId(1)).unwrap(), 0);
assert_eq!(c.seq_len(SlotId(3)).unwrap(), 0);
}
#[test]
fn historical_gemma4_hb_kv_fork_cross_slot_closure_at_phase_a3c() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let mut c = alloc_hb_kv_for_layer(&dev, 0, 2, 256, 8, false, 4).expect("alloc");
c.append_for_seq(SlotId(0), 7).unwrap();
c.fork_seq(SlotId(0), SlotId(1)).expect(
"iter-A3c closure: cross-slot fork must return Ok(()) — \
was previously CapabilityUnsupported per A3a typed-clamp",
);
assert_eq!(
c.seq_len(SlotId(1)).unwrap(),
7,
"iter-A3c closure: fork_seq must copy src's seq_len to dst"
);
assert_eq!(
c.seq_len(SlotId(0)).unwrap(),
7,
"iter-A3c closure: fork_seq must NOT modify src's seq_len (sub-pin for H163)"
);
}
#[test]
fn a3a_mixed_layer_alloc_full_sliding_byte_isolation() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
use crate::serve::config::LayerType;
let layer_types: Vec<LayerType> = vec![
LayerType::Full,
LayerType::Sliding,
LayerType::Full,
LayerType::Sliding,
];
let n_seqs: u32 = 2;
let nkv: usize = 2;
let hd: usize = 256;
let max_seq_len: usize = 32;
let sliding_window: usize = 8;
for (layer_idx, lt) in layer_types.iter().enumerate() {
let (is_ring, cap) =
super::layer_type_to_alloc_params(*lt, sliding_window, max_seq_len);
let buf = alloc_hb_kv_for_layer(&dev, layer_idx, nkv, hd, cap, is_ring, n_seqs)
.unwrap_or_else(|e| {
panic!("L{layer_idx} ({lt:?}): alloc_hb_kv_for_layer must succeed; got {e}")
});
assert_eq!(
buf.is_sliding, is_ring,
"L{layer_idx} ({lt:?}): is_sliding={} does NOT match \
expected={is_ring} (layer-type plumbing broken)",
buf.is_sliding,
);
let cap_label = if is_ring {
"sliding_window"
} else {
"max_seq_len"
};
assert_eq!(
buf.capacity, cap,
"L{layer_idx} ({lt:?}): capacity={} does NOT match \
expected={cap} ({cap_label})",
buf.capacity,
);
assert_eq!(
buf.seq_lens.len(),
n_seqs as usize,
"L{layer_idx}: seq_lens.len() must equal n_seqs"
);
assert!(
buf.seq_lens.iter().all(|&x| x == 0),
"L{layer_idx}: seq_lens must be zero-initialised"
);
assert_eq!(
buf.n_seqs, n_seqs,
"L{layer_idx}: n_seqs must propagate from call site"
);
let expected_packed_bytes = (n_seqs as usize) * nkv * cap * hd;
assert_eq!(
buf.k_packed.byte_len(),
expected_packed_bytes,
"L{layer_idx} ({lt:?}): k_packed byte_len mismatch",
);
assert_eq!(
buf.v_packed.byte_len(),
expected_packed_bytes,
"L{layer_idx} ({lt:?}): v_packed byte_len mismatch",
);
}
}
#[test]
fn a5c_layer_type_to_alloc_params_mapping_pinned() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::serve::config::LayerType;
let sliding_window: usize = 4_096;
let max_pos: usize = 131_072;
let (is_ring_s, cap_s) =
super::layer_type_to_alloc_params(LayerType::Sliding, sliding_window, max_pos);
assert!(is_ring_s, "Sliding MUST map to is_ring=true (ring buffer)");
assert_eq!(
cap_s, sliding_window,
"Sliding MUST map to capacity=sliding_window={sliding_window}"
);
let (is_ring_f, cap_f) =
super::layer_type_to_alloc_params(LayerType::Full, sliding_window, max_pos);
assert!(!is_ring_f, "Full MUST map to is_ring=false (linear buffer)");
assert_eq!(
cap_f, max_pos,
"Full MUST map to capacity=max_position_embeddings={max_pos}"
);
assert_ne!(
cap_s, cap_f,
"Sliding + Full MUST yield distinct capacities in a realistic \
production config (sliding_window != max_position_embeddings); \
a swap of the two arms in `layer_type_to_alloc_params` would \
make these equal and break the assertion above"
);
}
#[test]
fn a5c_production_gemma4_model_routes_through_layer_type_helper() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::serve::config::LayerType;
let sliding_window: usize = 1_024;
let max_pos: usize = 131_072;
let (is_ring, cap) =
super::layer_type_to_alloc_params(LayerType::Sliding, sliding_window, max_pos);
assert!(
is_ring && cap == sliding_window,
"production Sliding layer alloc shape: (ring=true, cap=1024); \
got (ring={is_ring}, cap={cap}) — gemma4/model.rs:1247-1257 \
would allocate the wrong shape if this mapping drifts"
);
let (is_ring, cap) =
super::layer_type_to_alloc_params(LayerType::Full, sliding_window, max_pos);
assert!(
!is_ring && cap == max_pos,
"production Full layer alloc shape: (ring=false, cap=131072); \
got (ring={is_ring}, cap={cap})"
);
}
#[test]
fn iter_f_kvcap_per_slot_alloc_params_mapping_pinned() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::serve::config::LayerType;
let sliding_window: usize = 1_024;
let max_pos: usize = 262_144;
for &n in &[1usize, 2, 4, 8] {
let (is_ring, cap) = super::layer_type_to_alloc_params_per_slot(
LayerType::Sliding,
sliding_window,
max_pos,
n,
);
assert!(is_ring, "Sliding MUST stay a ring buffer (max_slots={n})");
assert_eq!(
cap, sliding_window,
"Sliding capacity MUST stay sliding_window regardless of \
max_slots — dividing the ring window would corrupt its \
semantics (max_slots={n})"
);
}
let (is_ring_f8, cap_f8) =
super::layer_type_to_alloc_params_per_slot(LayerType::Full, sliding_window, max_pos, 8);
assert!(!is_ring_f8, "Full MUST stay linear (is_ring=false)");
assert_eq!(
cap_f8,
max_pos / 8,
"Full per-slot capacity at max_slots=8 MUST be max_position_\
embeddings/8 = {} (the literal '8×32k')",
max_pos / 8
);
let (_, cap_single) =
super::layer_type_to_alloc_params(LayerType::Full, sliding_window, max_pos);
let (_, cap_n1) =
super::layer_type_to_alloc_params_per_slot(LayerType::Full, sliding_window, max_pos, 1);
assert_eq!(
cap_n1, cap_single,
"max_slots=1 MUST be identity with the single-seq helper \
(no N=1 / SerialFifo regression): got {cap_n1} vs {cap_single}"
);
for &n in &[1usize, 2, 4, 8] {
let (_, per_slot) = super::layer_type_to_alloc_params_per_slot(
LayerType::Full,
sliding_window,
max_pos,
n,
);
let total = per_slot * n;
assert!(
total >= max_pos && total < max_pos + n,
"total Full KV (per_slot {per_slot} × {n} slots = {total}) MUST \
stay ≈ one full context ({max_pos}), not grow linearly in N"
);
}
}
#[test]
fn a3a_layer_type_variants_are_full_and_sliding_only() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::serve::config::LayerType;
fn name(lt: LayerType) -> &'static str {
match lt {
LayerType::Full => "Full",
LayerType::Sliding => "Sliding",
}
}
assert_eq!(name(LayerType::Full), "Full");
assert_eq!(name(LayerType::Sliding), "Sliding");
}
#[test]
fn h11_multi_seq_hybrid_kv_n_seqs_4_byte_scale() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
std::env::remove_var("HF2Q_FULL_F16_KV");
std::env::remove_var("HF2Q_DFLASH_XLEN_SDPA");
let nkv = 2usize;
let hd = 256usize;
let cap = 8usize;
let baseline = alloc_multi_seq_hybrid_kv_for_layer(&dev, 0, nkv, hd, cap, false, 1)
.expect("H11: alloc at n_seqs=1");
let lifted = alloc_multi_seq_hybrid_kv_for_layer(&dev, 0, nkv, hd, cap, false, 4)
.expect("H11: alloc at n_seqs=4");
assert_eq!(baseline.n_seqs, 1, "H11: baseline n_seqs=1");
assert_eq!(lifted.n_seqs, 4, "H11: lifted n_seqs=4");
assert_eq!(
lifted.k.byte_len(),
baseline.k.byte_len() * 4,
"H11 FALSIFIED: F16 K does not scale 4× ({} != {} * 4 = {})",
lifted.k.byte_len(),
baseline.k.byte_len(),
baseline.k.byte_len() * 4
);
assert_eq!(
lifted.v_packed.byte_len(),
baseline.v_packed.byte_len() * 4,
"H11 FALSIFIED: V packed does not scale 4× ({} != {} * 4 = {})",
lifted.v_packed.byte_len(),
baseline.v_packed.byte_len(),
baseline.v_packed.byte_len() * 4
);
assert_eq!(
lifted.v_norms.byte_len(),
baseline.v_norms.byte_len() * 4,
"H11 FALSIFIED: V norms does not scale 4× ({} != {})",
lifted.v_norms.byte_len(),
baseline.v_norms.byte_len() * 4
);
assert_eq!(baseline.seq_lens.len(), 1, "H11: baseline seq_lens.len()");
assert_eq!(lifted.seq_lens.len(), 4, "H11: lifted seq_lens.len()");
assert!(
baseline.seq_lens.iter().all(|&x| x == 0),
"H11: baseline seq_lens zero-init"
);
assert!(
lifted.seq_lens.iter().all(|&x| x == 0),
"H11: lifted seq_lens zero-init"
);
for (name, b) in [
("k", &lifted.k),
("v_packed", &lifted.v_packed),
("v_norms", &lifted.v_norms),
] {
let s = b.shape().to_vec();
assert_eq!(s.len(), 4, "H11 M5: {name} must be 4-D; got {:?}", s);
assert_eq!(
s[0], 4,
"H11 M5 FALSIFIED: {name} shape[0] must be n_seqs=4 (n_seqs landed \
on wrong axis); got {:?}",
s
);
}
let expected_k_bytes = 4usize * nkv * cap * hd * 2; let expected_v_packed_bytes = 4usize * nkv * cap * hd; let expected_v_norms_bytes = 4usize * nkv * cap * 1 * 4; let expected_total = expected_k_bytes + expected_v_packed_bytes + expected_v_norms_bytes;
assert_eq!(
lifted.k.byte_len(),
expected_k_bytes,
"H11 EXACT FORMULA FALSIFIED: K F16 bytes ({}) != n*nkv*cap*hd*2 ({})",
lifted.k.byte_len(),
expected_k_bytes
);
assert_eq!(
lifted.v_packed.byte_len(),
expected_v_packed_bytes,
"H11 EXACT FORMULA FALSIFIED: V packed U8 bytes ({}) != n*nkv*cap*hd ({})",
lifted.v_packed.byte_len(),
expected_v_packed_bytes
);
assert_eq!(
lifted.v_norms.byte_len(),
expected_v_norms_bytes,
"H11 EXACT FORMULA FALSIFIED: V norms F32 bytes ({}) != n*nkv*cap*1*4 ({})",
lifted.v_norms.byte_len(),
expected_v_norms_bytes
);
let actual_total =
lifted.k.byte_len() + lifted.v_packed.byte_len() + lifted.v_norms.byte_len();
assert_eq!(
actual_total, expected_total,
"H11 EXACT FORMULA FALSIFIED: composition K+V_packed+V_norms = {} != {}",
actual_total, expected_total
);
assert_eq!(
actual_total, 49408,
"H11 EXACT FORMULA FALSIFIED at concrete value: expected 49408 bytes \
for n_seqs=4 nkv=2 cap=8 hd=256 default (no full-F16, no xlen), got {}",
actual_total
);
assert!(
lifted.bf16_xlen_k.is_none() && lifted.bf16_xlen_v.is_none(),
"H11 EXACT FORMULA: xlen buffers must be None on default path"
);
}
#[test]
fn h11r_multi_seq_hybrid_kv_realistic_sliding_shape_byte_formula() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
std::env::remove_var("HF2Q_FULL_F16_KV");
std::env::remove_var("HF2Q_DFLASH_XLEN_SDPA");
let nkv = 8usize;
let hd = 256usize;
let cap = 512usize;
let n_seqs = 4u32;
let lifted = alloc_multi_seq_hybrid_kv_for_layer(&dev, 0, nkv, hd, cap, false, n_seqs)
.expect("H11r: realistic shape alloc");
let n = n_seqs as usize;
let expected_k = n * nkv * cap * hd * 2; let expected_v = n * nkv * cap * hd; let expected_norms = n * nkv * cap * 1 * 4; let expected_total = expected_k + expected_v + expected_norms;
assert_eq!(lifted.k.byte_len(), expected_k, "H11r: K F16");
assert_eq!(lifted.v_packed.byte_len(), expected_v, "H11r: V packed U8");
assert_eq!(
lifted.v_norms.byte_len(),
expected_norms,
"H11r: V norms F32"
);
let actual = lifted.k.byte_len() + lifted.v_packed.byte_len() + lifted.v_norms.byte_len();
assert_eq!(actual, expected_total, "H11r: composition");
assert_eq!(
actual, 12_648_448,
"H11r CONCRETE FALSIFIED: realistic Gemma 4 sliding shape n=4 nkv=8 \
hd=256 cap=512 should sum to 12_648_448 bytes (~12 MB); got {}",
actual
);
}
#[test]
fn h12_multi_seq_hybrid_kv_per_slot_byte_isolation() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
std::env::remove_var("HF2Q_FULL_F16_KV");
std::env::remove_var("HF2Q_DFLASH_XLEN_SDPA");
let nkv = 2usize;
let hd = 256usize;
let cap = 4usize;
let mut cache = alloc_multi_seq_hybrid_kv_for_layer(&dev, 0, nkv, hd, cap, false, 2)
.expect("H12: alloc n_seqs=2");
assert_eq!(cache.n_seqs, 2);
let slot_k_bytes = nkv * cap * hd * 2;
let slot_v_bytes = nkv * cap * hd;
let slot_vn_bytes = nkv * cap * 1 * 4;
assert_eq!(
cache.k.byte_len(),
2 * slot_k_bytes,
"H12 fixture sanity: K total = 2 * slot_k_bytes"
);
assert_eq!(
cache.v_packed.byte_len(),
2 * slot_v_bytes,
"H12 fixture sanity: V packed total = 2 * slot_v_bytes"
);
assert_eq!(
cache.v_norms.byte_len(),
2 * slot_vn_bytes,
"H12 fixture sanity: V norms total = 2 * slot_vn_bytes"
);
{
let k_slice = cache.k.as_mut_slice::<u8>().expect("k F16 as u8 mut");
for (i, b) in k_slice[..slot_k_bytes].iter_mut().enumerate() {
*b = (((i * 7) % 251) + 1) as u8;
}
}
{
let v_slice = cache
.v_packed
.as_mut_slice::<u8>()
.expect("v_packed u8 mut");
for (i, b) in v_slice[..slot_v_bytes].iter_mut().enumerate() {
*b = (((i * 11) % 253) + 1) as u8;
}
}
{
let vn_slice = cache
.v_norms
.as_mut_slice::<f32>()
.expect("v_norms f32 mut");
let slot_vn_f32 = nkv * cap * 1; for (i, f) in vn_slice[..slot_vn_f32].iter_mut().enumerate() {
*f = (i as f32) * 0.123_45;
}
}
let k_slot1_before: Vec<u8> =
cache.k.as_slice::<u8>().expect("k F16 as u8")[slot_k_bytes..2 * slot_k_bytes].to_vec();
let v_slot1_before: Vec<u8> = cache.v_packed.as_slice::<u8>().expect("v_packed u8")
[slot_v_bytes..2 * slot_v_bytes]
.to_vec();
let vn_slot1_before: Vec<f32> = cache.v_norms.as_slice::<f32>().expect("v_norms f32")
[(nkv * cap * 1)..2 * (nkv * cap * 1)]
.to_vec();
assert!(
k_slot1_before.iter().all(|&b| b == 0),
"H12 fixture sanity: slot 1 K zero-init"
);
assert!(
v_slot1_before.iter().all(|&b| b == 0),
"H12 fixture sanity: slot 1 V packed zero-init"
);
assert!(
vn_slot1_before.iter().all(|&f| f == 0.0),
"H12 fixture sanity: slot 1 V norms zero-init"
);
cache
.append_for_seq(SlotId(0), 3)
.expect("H12: append slot 0");
assert_eq!(cache.seq_lens[0], 3);
assert_eq!(cache.seq_lens[1], 0);
let k_slot1_after: Vec<u8> =
cache.k.as_slice::<u8>().expect("k F16 as u8")[slot_k_bytes..2 * slot_k_bytes].to_vec();
let v_slot1_after: Vec<u8> = cache.v_packed.as_slice::<u8>().expect("v_packed u8")
[slot_v_bytes..2 * slot_v_bytes]
.to_vec();
let vn_slot1_after: Vec<f32> = cache.v_norms.as_slice::<f32>().expect("v_norms f32")
[(nkv * cap * 1)..2 * (nkv * cap * 1)]
.to_vec();
assert_eq!(
k_slot1_before, k_slot1_after,
"H12 FALSIFIED: slot 1 K bytes changed after slot-0 write"
);
assert_eq!(
v_slot1_before, v_slot1_after,
"H12 FALSIFIED: slot 1 V packed bytes changed"
);
assert_eq!(
vn_slot1_before, vn_slot1_after,
"H12 FALSIFIED: slot 1 V norms changed"
);
}
#[test]
fn h13_multi_seq_hybrid_kv_cursor_independence() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
std::env::remove_var("HF2Q_FULL_F16_KV");
std::env::remove_var("HF2Q_DFLASH_XLEN_SDPA");
let mut c =
alloc_multi_seq_hybrid_kv_for_layer(&dev, 0, 2, 256, 8, false, 4).expect("alloc");
for s in 0..4 {
assert_eq!(c.seq_len(SlotId(s)).expect("seq_len in range"), 0);
}
c.append_for_seq(SlotId(0), 5).expect("append slot 0");
c.append_for_seq(SlotId(2), 3).expect("append slot 2");
assert_eq!(c.seq_len(SlotId(0)).unwrap(), 5);
assert_eq!(
c.seq_len(SlotId(1)).unwrap(),
0,
"H13 FALSIFIED: slot 1 cursor touched by slot 0/2 append"
);
assert_eq!(c.seq_len(SlotId(2)).unwrap(), 3);
assert_eq!(
c.seq_len(SlotId(3)).unwrap(),
0,
"H13 FALSIFIED: slot 3 cursor touched by slot 0/2 append"
);
c.drop_seq(SlotId(0)).expect("drop slot 0");
assert_eq!(c.seq_len(SlotId(0)).unwrap(), 0, "H13: slot 0 reset");
assert_eq!(
c.seq_len(SlotId(2)).unwrap(),
3,
"H13: slot 2 preserved through slot 0 drop"
);
}
#[test]
fn h14_multi_seq_hybrid_kv_xlen_optional_coexists_with_u8_v() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
std::env::remove_var("HF2Q_FULL_F16_KV");
std::env::set_var("HF2Q_DFLASH_XLEN_SDPA", "1");
let xlen_on = alloc_multi_seq_hybrid_kv_for_layer(&dev, 0, 2, 256, 4, false, 3)
.expect("H14: alloc xlen on");
assert_eq!(xlen_on.n_seqs, 3);
assert!(
xlen_on.bf16_xlen_k.is_some(),
"H14 FALSIFIED: xlen K must be Some when env gate set"
);
assert!(
xlen_on.bf16_xlen_v.is_some(),
"H14 FALSIFIED: xlen V must be Some when env gate set"
);
let bk = xlen_on.bf16_xlen_k.as_ref().unwrap();
let bv = xlen_on.bf16_xlen_v.as_ref().unwrap();
let expected_xlen_bytes = 3 * 2 * 4 * 256 * 2; assert_eq!(
bk.byte_len(),
expected_xlen_bytes,
"H14 FALSIFIED: xlen K bytes wrong ({} != {})",
bk.byte_len(),
expected_xlen_bytes
);
assert_eq!(
bv.byte_len(),
expected_xlen_bytes,
"H14 FALSIFIED: xlen V bytes wrong"
);
assert_eq!(
bk.shape(),
&[3, 2, 4, 256],
"H14: xlen K shape n_seqs outermost"
);
assert_eq!(
bv.shape(),
&[3, 2, 4, 256],
"H14: xlen V shape n_seqs outermost"
);
assert_eq!(
xlen_on.v_packed.byte_len(),
3 * 2 * 4 * 256,
"H14: U8 V packed coexists with xlen"
);
assert_eq!(
xlen_on.v_norms.byte_len(),
3 * 2 * 4 * 1 * 4,
"H14: F32 V norms coexists with xlen"
);
std::env::remove_var("HF2Q_DFLASH_XLEN_SDPA");
let xlen_off = alloc_multi_seq_hybrid_kv_for_layer(&dev, 0, 2, 256, 4, false, 3)
.expect("H14: alloc xlen off");
assert!(
xlen_off.bf16_xlen_k.is_none(),
"H14 FALSIFIED: xlen K must be None when env gate unset"
);
assert!(
xlen_off.bf16_xlen_v.is_none(),
"H14 FALSIFIED: xlen V must be None when env gate unset"
);
assert_eq!(xlen_off.v_packed.byte_len(), 3 * 2 * 4 * 256);
assert_eq!(xlen_off.v_norms.byte_len(), 3 * 2 * 4 * 1 * 4);
}
#[test]
fn h15_dense_kv_buffers_typed_clamp_slot_count_one() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::serve::multi_seq_kv::MultiSeqKvCache;
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2;
let cap = 8;
let hd = 256;
let k = dev
.alloc_buffer(nkv * cap * hd * 4, DType::F32, vec![nkv, cap, hd])
.unwrap();
let v = dev
.alloc_buffer(nkv * cap * hd * 4, DType::F32, vec![nkv, cap, hd])
.unwrap();
let mut buf = DenseKvBuffers {
k,
v,
capacity: cap,
is_sliding: false,
dtype: DType::F32,
};
assert_eq!(
buf.slot_count(),
1,
"H15 FALSIFIED: DenseKvBuffers slot_count must be 1"
);
assert_eq!(buf.layout(), MultiSeqLayout::SeparateSlots);
assert_eq!(buf.seq_len(SlotId(0)).unwrap(), 0);
let err = buf.seq_len(SlotId(1)).expect_err("slot 1 OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(1),
max_slots: 1
},
"H15 FALSIFIED: SlotOutOfRange shape wrong; got {err:?}"
);
let err = buf
.append_for_seq(SlotId(2), 1)
.expect_err("append slot 2 OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(2),
max_slots: 1
}
);
let err = buf
.append_for_seq(SlotId(0), 1)
.expect_err("append clamped to iter-A3b-2");
match err {
MultiSeqError::CapabilityUnsupported { capability } => {
assert!(
capability.contains("DenseKvBuffers"),
"label must name struct: {capability}"
);
assert!(
capability.contains("A3b iter-2"),
"label must name deferral: {capability}"
);
}
other => panic!("H15: expected CapabilityUnsupported; got {other:?}"),
}
let err = buf.drop_seq(SlotId(5)).expect_err("drop slot 5 OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(5),
max_slots: 1
}
);
let err = buf.drop_seq(SlotId(0)).expect_err("drop clamped");
assert!(matches!(err, MultiSeqError::CapabilityUnsupported { .. }));
buf.fork_seq(SlotId(0), SlotId(0))
.expect("self-fork ok no-op");
let err = buf
.fork_seq(SlotId(1), SlotId(0))
.expect_err("fork src OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(1),
max_slots: 1
}
);
let err = buf
.fork_seq(SlotId(0), SlotId(2))
.expect_err("fork dst OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(2),
max_slots: 1
}
);
}
#[test]
fn h16_mlx_kv_cache_typed_clamp_slot_count_one() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::serve::multi_seq_kv::MultiSeqKvCache;
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let buf = || dev.alloc_buffer(4, DType::F32, vec![1]).unwrap();
let mut cache = MlxKvCache {
k_packed: buf(),
k_norms: buf(),
v_packed: buf(),
v_norms: buf(),
capacity: 16,
is_sliding: false,
write_pos: 5,
seq_len: 5,
};
assert_eq!(
cache.slot_count(),
1,
"H16 FALSIFIED: MlxKvCache slot_count must be 1"
);
assert_eq!(cache.layout(), MultiSeqLayout::SeparateSlots);
assert_eq!(
cache.seq_len(SlotId(0)).unwrap(),
5,
"H16: seq_len(0) reports legacy cursor (was 5)"
);
let err = cache.seq_len(SlotId(1)).expect_err("slot 1 OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(1),
max_slots: 1
},
"H16 FALSIFIED: SlotOutOfRange shape wrong; got {err:?}"
);
let err = cache.append_for_seq(SlotId(3), 1).expect_err("append OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(3),
max_slots: 1
}
);
let err = cache
.append_for_seq(SlotId(0), 1)
.expect_err("append clamped");
match err {
MultiSeqError::CapabilityUnsupported { capability } => {
assert!(
capability.contains("MlxKvCache"),
"label must name struct: {capability}"
);
assert!(
capability.contains("A3b iter-3"),
"label must name deferral: {capability}"
);
assert!(
capability.contains("legacy 4-bit"),
"label must name legacy path: {capability}"
);
}
other => panic!("H16: expected CapabilityUnsupported; got {other:?}"),
}
let err = cache.drop_seq(SlotId(7)).expect_err("drop OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(7),
max_slots: 1
}
);
let err = cache.drop_seq(SlotId(0)).expect_err("drop clamped");
assert!(matches!(err, MultiSeqError::CapabilityUnsupported { .. }));
cache
.fork_seq(SlotId(0), SlotId(0))
.expect("self-fork no-op");
}
#[test]
fn h10_post_falsification_hybrid_kv_default_is_on() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
std::env::remove_var("HF2Q_HYBRID_KV");
let env = crate::debug::investigation_env::InvestigationEnv::from_env();
assert!(
env.hybrid_kv,
"H10 post-falsification FALSIFIED: HF2Q_HYBRID_KV default flipped to OFF; \
A3b iter-1 assumed default-ON per ADR-029 iter-13 + \
`src/debug/investigation_env.rs:878` (env_default_true). \
If this flip is intentional, update the A3b iter-1 block \
comment at kv_cache.rs and re-examine whether HybridKvBuffers \
remains the production-default variant for Gemma 4."
);
}
#[test]
fn iter_b4c_kernel_iter1_multi_seq_hb_kv_reset_for_slot_per_slot_isolation_2026_05_30() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::serve::multi_seq_kv::SlotId;
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 4usize;
let n_seqs: u32 = 4;
let mut cache =
alloc_hb_kv_for_layer(&dev, 0, nkv, hd, cap, false, n_seqs).expect("alloc n_seqs=4");
for s in 0..(n_seqs as usize) {
cache.seq_lens[s] = (s as u32) + 11;
}
let slot_packed = nkv * cap * hd;
{
let k_slice = cache.k_packed.as_mut_slice::<u8>().expect("k_packed u8");
for s in 0..(n_seqs as usize) {
let start = s * slot_packed;
for (i, b) in k_slice[start..start + slot_packed].iter_mut().enumerate() {
*b = (((s * 17 + i) % 251) + 1) as u8;
}
}
}
{
let v_slice = cache.v_packed.as_mut_slice::<u8>().expect("v_packed u8");
for s in 0..(n_seqs as usize) {
let start = s * slot_packed;
for (i, b) in v_slice[start..start + slot_packed].iter_mut().enumerate() {
*b = (((s * 19 + i) % 253) + 1) as u8;
}
}
}
let k_before: Vec<u8> = cache
.k_packed
.as_slice::<u8>()
.expect("k_packed read")
.to_vec();
let v_before: Vec<u8> = cache
.v_packed
.as_slice::<u8>()
.expect("v_packed read")
.to_vec();
cache
.reset_for_slot(SlotId(1))
.expect("reset_for_slot(1) on n_seqs=4");
for s in 0..(n_seqs as usize) {
if s == 1 {
assert_eq!(
cache.seq_lens[s], 0,
"iter-B4c-kernel iter-1: slot 1 cursor must be 0 after reset_for_slot(1)"
);
} else {
assert_eq!(
cache.seq_lens[s],
(s as u32) + 11,
"iter-B4c-kernel iter-1: slot {s} cursor must be untouched"
);
}
}
let k_after: Vec<u8> = cache
.k_packed
.as_slice::<u8>()
.expect("k_packed read 2")
.to_vec();
let v_after: Vec<u8> = cache
.v_packed
.as_slice::<u8>()
.expect("v_packed read 2")
.to_vec();
assert_eq!(
k_before, k_after,
"iter-B4c-kernel iter-1: reset_for_slot must NOT zero K packed bytes \
(cursor-masked discipline; matches drop_seq invariant)"
);
assert_eq!(
v_before, v_after,
"iter-B4c-kernel iter-1: reset_for_slot must NOT zero V packed bytes"
);
}
#[test]
fn iter_b4c_kernel_iter1_multi_seq_hb_kv_reset_for_slot_bounds_typed_2026_05_30() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::serve::multi_seq_kv::{MultiSeqError, SlotId};
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let mut cache =
alloc_hb_kv_for_layer(&dev, 0, 2, 256, 4, false, 4).expect("alloc n_seqs=4");
let err = cache.reset_for_slot(SlotId(4)).expect_err("slot 4 OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(4),
max_slots: 4
},
"iter-B4c-kernel iter-1: OOR must surface SlotOutOfRange"
);
let mut cache1 =
alloc_hb_kv_for_layer(&dev, 0, 2, 256, 4, false, 1).expect("alloc n_seqs=1");
cache1.seq_lens[0] = 7;
cache1
.reset_for_slot(SlotId(0))
.expect("SlotId(0) at n_seqs=1 must succeed");
assert_eq!(cache1.seq_lens[0], 0);
}
#[test]
fn iter_b4c_kernel_iter1_multi_seq_hybrid_kv_reset_for_slot_per_slot_isolation_2026_05_30() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::serve::multi_seq_kv::SlotId;
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
std::env::remove_var("HF2Q_FULL_F16_KV");
std::env::remove_var("HF2Q_DFLASH_XLEN_SDPA");
let mut cache = alloc_multi_seq_hybrid_kv_for_layer(&dev, 0, 2, 256, 4, false, 4)
.expect("alloc multi-seq hybrid n_seqs=4");
for s in 0..4 {
cache.seq_lens[s] = (s as u32) * 5 + 3;
}
cache.reset_for_slot(SlotId(2)).expect("reset_for_slot(2)");
assert_eq!(cache.seq_lens[0], 3);
assert_eq!(cache.seq_lens[1], 8);
assert_eq!(
cache.seq_lens[2], 0,
"iter-B4c-kernel iter-1: slot 2 cursor must be 0 after reset"
);
assert_eq!(cache.seq_lens[3], 18);
let err = cache.reset_for_slot(SlotId(99)).expect_err("slot 99 OOR");
assert!(matches!(
err,
crate::serve::multi_seq_kv::MultiSeqError::SlotOutOfRange { .. }
));
}
#[test]
fn h144_multi_seq_dense_kv_buffers_sibling_struct_exists() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 8usize;
let n_seqs = 3u32;
let buf_f32_lin = alloc_multi_seq_dense_kv_for_layer(
&dev,
0,
nkv,
hd,
cap,
false,
DType::F32,
n_seqs,
)
.expect("H144: alloc F32 linear");
assert_eq!(buf_f32_lin.n_seqs, n_seqs, "H144: n_seqs propagation");
assert_eq!(buf_f32_lin.dtype, DType::F32, "H144: dtype propagation F32");
assert!(
!buf_f32_lin.is_sliding,
"H144: is_sliding=false propagation"
);
assert_eq!(buf_f32_lin.capacity, cap, "H144: capacity propagation");
assert_eq!(
buf_f32_lin.seq_lens.len(),
n_seqs as usize,
"H144 FALSIFIED: seq_lens.len() must equal n_seqs"
);
assert!(
buf_f32_lin.seq_lens.iter().all(|&x| x == 0),
"H144 FALSIFIED: seq_lens zero-init"
);
assert_eq!(
buf_f32_lin.k.shape(),
&[n_seqs as usize, nkv, cap, hd],
"H144 FALSIFIED: K shape n_seqs outermost"
);
assert_eq!(
buf_f32_lin.v.shape(),
&[n_seqs as usize, nkv, cap, hd],
"H144 FALSIFIED: V shape n_seqs outermost"
);
let buf_f16_ring = alloc_multi_seq_dense_kv_for_layer(
&dev,
7,
nkv,
hd,
cap,
true,
DType::F16,
n_seqs,
)
.expect("H144: alloc F16 sliding");
assert_eq!(
buf_f16_ring.dtype,
DType::F16,
"H144: dtype propagation F16"
);
assert!(buf_f16_ring.is_sliding, "H144: is_sliding=true propagation");
}
#[test]
fn h145_alloc_multi_seq_dense_kv_for_layer_preflight_errors() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
assert!(
alloc_multi_seq_dense_kv_for_layer(&dev, 0, 2, 256, 8, false, DType::F32, 0).is_err(),
"H145 FALSIFIED: n_seqs=0 must error"
);
assert!(
alloc_multi_seq_dense_kv_for_layer(&dev, 0, 0, 256, 8, false, DType::F32, 1).is_err(),
"H145 FALSIFIED: nkv=0 must error"
);
assert!(
alloc_multi_seq_dense_kv_for_layer(&dev, 0, 2, 0, 8, false, DType::F32, 1).is_err(),
"H145 FALSIFIED: hd=0 must error"
);
assert!(
alloc_multi_seq_dense_kv_for_layer(&dev, 0, 2, 256, 0, false, DType::F32, 1).is_err(),
"H145 FALSIFIED: cap=0 must error"
);
}
#[test]
fn h146_multi_seq_dense_kv_n_seqs_4_byte_scale_exact_formula() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 8usize;
let f32_baseline =
alloc_multi_seq_dense_kv_for_layer(&dev, 0, nkv, hd, cap, false, DType::F32, 1)
.expect("H146: F32 alloc n_seqs=1");
let f32_lifted =
alloc_multi_seq_dense_kv_for_layer(&dev, 0, nkv, hd, cap, false, DType::F32, 4)
.expect("H146: F32 alloc n_seqs=4");
assert_eq!(f32_baseline.n_seqs, 1);
assert_eq!(f32_lifted.n_seqs, 4);
assert_eq!(
f32_lifted.k.byte_len(),
f32_baseline.k.byte_len() * 4,
"H146 FALSIFIED: F32 K not 4× scale"
);
assert_eq!(
f32_lifted.v.byte_len(),
f32_baseline.v.byte_len() * 4,
"H146 FALSIFIED: F32 V not 4× scale"
);
let expected_f32_k = 4usize * nkv * cap * hd * 4; let expected_f32_v = expected_f32_k;
let expected_f32_total = expected_f32_k + expected_f32_v;
assert_eq!(
f32_lifted.k.byte_len(),
expected_f32_k,
"H146 EXACT FORMULA FALSIFIED: F32 K bytes != n*nkv*cap*hd*4"
);
assert_eq!(
f32_lifted.v.byte_len(),
expected_f32_v,
"H146 EXACT FORMULA FALSIFIED: F32 V bytes != n*nkv*cap*hd*4"
);
let actual_f32_total = f32_lifted.k.byte_len() + f32_lifted.v.byte_len();
assert_eq!(
actual_f32_total, expected_f32_total,
"H146 EXACT FORMULA FALSIFIED: F32 composition"
);
assert_eq!(
actual_f32_total, 131_072,
"H146 EXACT FORMULA FALSIFIED at concrete value: F32 expected 131072 \
bytes for n_seqs=4 nkv=2 cap=8 hd=256; got {}",
actual_f32_total
);
let f16_baseline =
alloc_multi_seq_dense_kv_for_layer(&dev, 0, nkv, hd, cap, false, DType::F16, 1)
.expect("H146: F16 alloc n_seqs=1");
let f16_lifted =
alloc_multi_seq_dense_kv_for_layer(&dev, 0, nkv, hd, cap, false, DType::F16, 4)
.expect("H146: F16 alloc n_seqs=4");
assert_eq!(
f16_lifted.k.byte_len(),
f16_baseline.k.byte_len() * 4,
"H146 FALSIFIED: F16 K not 4× scale"
);
assert_eq!(
f16_lifted.v.byte_len(),
f16_baseline.v.byte_len() * 4,
"H146 FALSIFIED: F16 V not 4× scale"
);
let expected_f16_k = 4usize * nkv * cap * hd * 2; let expected_f16_v = expected_f16_k;
let expected_f16_total = expected_f16_k + expected_f16_v;
assert_eq!(
f16_lifted.k.byte_len(),
expected_f16_k,
"H146 EXACT FORMULA FALSIFIED: F16 K bytes != n*nkv*cap*hd*2"
);
let actual_f16_total = f16_lifted.k.byte_len() + f16_lifted.v.byte_len();
assert_eq!(
actual_f16_total, expected_f16_total,
"H146 EXACT FORMULA FALSIFIED: F16 composition"
);
assert_eq!(
actual_f16_total, 65_536,
"H146 EXACT FORMULA FALSIFIED at concrete value: F16 expected 65536 \
bytes for n_seqs=4 nkv=2 cap=8 hd=256; got {}",
actual_f16_total
);
assert_eq!(f32_baseline.seq_lens.len(), 1);
assert_eq!(f32_lifted.seq_lens.len(), 4);
assert_eq!(f16_lifted.seq_lens.len(), 4);
assert!(f32_lifted.seq_lens.iter().all(|&x| x == 0));
for (name, b) in [("k", &f32_lifted.k), ("v", &f32_lifted.v)] {
let s = b.shape().to_vec();
assert_eq!(s.len(), 4, "H146 M5: {name} must be 4-D; got {:?}", s);
assert_eq!(
s[0], 4,
"H146 M5 FALSIFIED: {name} shape[0] must be n_seqs=4 (n_seqs \
landed on wrong axis); got {:?}",
s
);
}
}
#[test]
fn h147_multi_seq_dense_kv_per_slot_byte_isolation() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 4usize;
let mut cache =
alloc_multi_seq_dense_kv_for_layer(&dev, 0, nkv, hd, cap, false, DType::F32, 2)
.expect("H147: alloc n_seqs=2");
assert_eq!(cache.n_seqs, 2);
let slot_k_bytes = nkv * cap * hd * 4;
let slot_v_bytes = nkv * cap * hd * 4;
assert_eq!(cache.k.byte_len(), 2 * slot_k_bytes, "H147: K total");
assert_eq!(cache.v.byte_len(), 2 * slot_v_bytes, "H147: V total");
{
let k_slice = cache.k.as_mut_slice::<u8>().expect("K F32 as u8 mut");
for (i, b) in k_slice[..slot_k_bytes].iter_mut().enumerate() {
*b = (((i * 7) % 251) + 1) as u8;
}
}
{
let v_slice = cache.v.as_mut_slice::<u8>().expect("V F32 as u8 mut");
for (i, b) in v_slice[..slot_v_bytes].iter_mut().enumerate() {
*b = (((i * 11) % 253) + 1) as u8;
}
}
let k_slot1_before: Vec<u8> =
cache.k.as_slice::<u8>().expect("K read")[slot_k_bytes..2 * slot_k_bytes].to_vec();
let v_slot1_before: Vec<u8> =
cache.v.as_slice::<u8>().expect("V read")[slot_v_bytes..2 * slot_v_bytes].to_vec();
assert!(
k_slot1_before.iter().all(|&b| b == 0),
"H147 fixture sanity: slot 1 K zero-init"
);
assert!(
v_slot1_before.iter().all(|&b| b == 0),
"H147 fixture sanity: slot 1 V zero-init"
);
cache
.append_for_seq(SlotId(0), 3)
.expect("H147: append slot 0");
assert_eq!(cache.seq_lens[0], 3);
assert_eq!(cache.seq_lens[1], 0);
let k_slot1_after: Vec<u8> =
cache.k.as_slice::<u8>().expect("K read 2")[slot_k_bytes..2 * slot_k_bytes].to_vec();
let v_slot1_after: Vec<u8> =
cache.v.as_slice::<u8>().expect("V read 2")[slot_v_bytes..2 * slot_v_bytes].to_vec();
assert_eq!(
k_slot1_before, k_slot1_after,
"H147 FALSIFIED: slot 1 K bytes changed after slot-0 write"
);
assert_eq!(
v_slot1_before, v_slot1_after,
"H147 FALSIFIED: slot 1 V bytes changed after slot-0 write"
);
}
#[test]
fn h148_multi_seq_dense_kv_n_seqs_1_byte_equivalent_to_legacy() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 8usize;
for &dtype in &[DType::F32, DType::F16] {
let multi = alloc_multi_seq_dense_kv_for_layer(&dev, 0, nkv, hd, cap, false, dtype, 1)
.expect("H148: alloc multi-seq n_seqs=1");
let legacy_bytes = nkv * cap * hd * dtype.size_of();
assert_eq!(
multi.k.byte_len(),
legacy_bytes,
"H148 FALSIFIED ({:?}): K bytes {} != legacy {}",
dtype,
multi.k.byte_len(),
legacy_bytes
);
assert_eq!(
multi.v.byte_len(),
legacy_bytes,
"H148 FALSIFIED ({:?}): V bytes {} != legacy {}",
dtype,
multi.v.byte_len(),
legacy_bytes
);
let legacy_total = 2 * legacy_bytes;
use crate::serve::kv_persist::lcp_registry::ByteSized;
assert_eq!(
multi.byte_len(),
legacy_total as u64,
"H148 FALSIFIED ({:?}): total byte_len {} != legacy K+V {}",
dtype,
multi.byte_len(),
legacy_total
);
}
}
#[test]
fn h149_multi_seq_dense_kv_multi_seq_kv_cache_impl() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let n_seqs = 4u32;
let mut cache =
alloc_multi_seq_dense_kv_for_layer(&dev, 0, 2, 256, 8, false, DType::F32, n_seqs)
.expect("H149: alloc n_seqs=4");
assert_eq!(
cache.slot_count(),
n_seqs,
"H149 FALSIFIED: slot_count must equal n_seqs={n_seqs}"
);
assert_eq!(cache.layout(), MultiSeqLayout::SeparateSlots);
for s in 0..n_seqs {
assert_eq!(
cache.seq_len(SlotId(s)).expect("seq_len in range"),
0,
"H149: slot {s} starts at cursor 0"
);
}
let err = cache.seq_len(SlotId(n_seqs)).expect_err("OOR n_seqs");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(n_seqs),
max_slots: n_seqs,
},
"H149 FALSIFIED: seq_len OOR shape; got {err:?}"
);
cache.append_for_seq(SlotId(0), 5).expect("append slot 0");
cache.append_for_seq(SlotId(2), 3).expect("append slot 2");
assert_eq!(cache.seq_len(SlotId(0)).unwrap(), 5);
assert_eq!(
cache.seq_len(SlotId(1)).unwrap(),
0,
"H149 FALSIFIED: slot 1 cursor touched by slot 0/2 append"
);
assert_eq!(cache.seq_len(SlotId(2)).unwrap(), 3);
assert_eq!(
cache.seq_len(SlotId(3)).unwrap(),
0,
"H149 FALSIFIED: slot 3 cursor touched by slot 0/2 append"
);
let err = cache
.append_for_seq(SlotId(n_seqs + 1), 1)
.expect_err("append OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(n_seqs + 1),
max_slots: n_seqs,
}
);
cache.drop_seq(SlotId(0)).expect("drop slot 0");
assert_eq!(cache.seq_len(SlotId(0)).unwrap(), 0, "H149: slot 0 reset");
assert_eq!(
cache.seq_len(SlotId(2)).unwrap(),
3,
"H149 FALSIFIED: slot 2 preserved through slot 0 drop"
);
let err = cache.drop_seq(SlotId(99)).expect_err("drop OOR");
assert!(matches!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(99),
max_slots: 4
}
));
cache
.fork_seq(SlotId(1), SlotId(1))
.expect("self-fork no-op");
let err = cache
.fork_seq(SlotId(99), SlotId(0))
.expect_err("fork src OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(99),
max_slots: 4
}
);
let err = cache
.fork_seq(SlotId(0), SlotId(99))
.expect_err("fork dst OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(99),
max_slots: 4
}
);
cache
.append_for_seq(SlotId(0), 5)
.expect("re-seed slot 0 for fork");
let slot0_before = cache.seq_len(SlotId(0)).unwrap();
let slot2_before = cache.seq_len(SlotId(2)).unwrap();
cache
.fork_seq(SlotId(0), SlotId(1))
.expect("iter-A3c closure: cross-slot fork must return Ok(())");
assert_eq!(
cache.seq_len(SlotId(1)).unwrap(),
slot0_before,
"H149 closure: fork must copy src's seq_len to dst"
);
assert_eq!(
cache.seq_len(SlotId(0)).unwrap(),
slot0_before,
"H149 closure: fork must NOT mutate src's seq_len"
);
assert_eq!(
cache.seq_len(SlotId(2)).unwrap(),
slot2_before,
"H149 closure: fork must NOT mutate non-src non-dst slots"
);
}
#[test]
fn h150_multi_seq_dense_kv_reset_for_slot_per_slot_isolation_and_bounds() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 4usize;
let n_seqs = 4u32;
let mut cache =
alloc_multi_seq_dense_kv_for_layer(&dev, 0, nkv, hd, cap, false, DType::F32, n_seqs)
.expect("alloc n_seqs=4");
for s in 0..(n_seqs as usize) {
cache.seq_lens[s] = (s as u32) + 11;
}
let slot_k_bytes = nkv * cap * hd * 4; let slot_v_bytes = nkv * cap * hd * 4;
{
let k_slice = cache.k.as_mut_slice::<u8>().expect("k u8");
for s in 0..(n_seqs as usize) {
let start = s * slot_k_bytes;
for (i, b) in k_slice[start..start + slot_k_bytes].iter_mut().enumerate() {
*b = (((s * 17 + i) % 251) + 1) as u8;
}
}
}
{
let v_slice = cache.v.as_mut_slice::<u8>().expect("v u8");
for s in 0..(n_seqs as usize) {
let start = s * slot_v_bytes;
for (i, b) in v_slice[start..start + slot_v_bytes].iter_mut().enumerate() {
*b = (((s * 19 + i) % 253) + 1) as u8;
}
}
}
let k_before: Vec<u8> = cache.k.as_slice::<u8>().expect("k read").to_vec();
let v_before: Vec<u8> = cache.v.as_slice::<u8>().expect("v read").to_vec();
cache
.reset_for_slot(SlotId(1))
.expect("reset_for_slot(1) on n_seqs=4");
for s in 0..(n_seqs as usize) {
if s == 1 {
assert_eq!(
cache.seq_lens[s], 0,
"H150 FALSIFIED: slot 1 cursor must be 0 after reset_for_slot(1)"
);
} else {
assert_eq!(
cache.seq_lens[s],
(s as u32) + 11,
"H150 FALSIFIED: slot {s} cursor must be untouched"
);
}
}
let k_after: Vec<u8> = cache.k.as_slice::<u8>().expect("k read 2").to_vec();
let v_after: Vec<u8> = cache.v.as_slice::<u8>().expect("v read 2").to_vec();
assert_eq!(
k_before, k_after,
"H150 FALSIFIED: reset_for_slot must NOT zero K bytes \
(cursor-masked discipline; matches drop_seq invariant)"
);
assert_eq!(
v_before, v_after,
"H150 FALSIFIED: reset_for_slot must NOT zero V bytes"
);
let err = cache.reset_for_slot(SlotId(99)).expect_err("slot 99 OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(99),
max_slots: 4
},
"H150 FALSIFIED: reset OOR shape; got {err:?}"
);
let mut cache1 =
alloc_multi_seq_dense_kv_for_layer(&dev, 0, 2, 256, 4, false, DType::F32, 1)
.expect("alloc n_seqs=1");
cache1.seq_lens[0] = 7;
cache1
.reset_for_slot(SlotId(0))
.expect("SlotId(0) at n_seqs=1 must succeed (byte-equivalence case)");
assert_eq!(cache1.seq_lens[0], 0);
}
#[test]
fn h151_multi_seq_mlx_kv_cache_sibling_struct_exists() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 8usize;
let n_seqs = 3u32;
let buf_lin = alloc_multi_seq_mlx_kv_for_layer(
&dev, 0, nkv, hd, cap, false, 1, n_seqs,
)
.expect("H151: alloc norms_per_pos=1 linear");
assert_eq!(buf_lin.n_seqs, n_seqs, "H151: n_seqs propagation");
assert_eq!(buf_lin.norms_per_pos, 1, "H151: norms_per_pos propagation");
assert!(!buf_lin.is_sliding, "H151: is_sliding=false propagation");
assert_eq!(buf_lin.capacity, cap, "H151: capacity propagation");
assert_eq!(
buf_lin.seq_lens.len(),
n_seqs as usize,
"H151 FALSIFIED: seq_lens.len() must equal n_seqs"
);
assert!(
buf_lin.seq_lens.iter().all(|&x| x == 0),
"H151 FALSIFIED: seq_lens zero-init"
);
assert_eq!(
buf_lin.k_packed.shape(),
&[n_seqs as usize, nkv, cap, hd / 2],
"H151 FALSIFIED: k_packed shape n_seqs outermost"
);
assert_eq!(
buf_lin.v_packed.shape(),
&[n_seqs as usize, nkv, cap, hd / 2],
"H151 FALSIFIED: v_packed shape n_seqs outermost"
);
assert_eq!(
buf_lin.k_norms.shape(),
&[n_seqs as usize, nkv, cap],
"H151 FALSIFIED: k_norms shape (norms_per_pos=1) n_seqs outermost"
);
assert_eq!(
buf_lin.v_norms.shape(),
&[n_seqs as usize, nkv, cap],
"H151 FALSIFIED: v_norms shape (norms_per_pos=1) n_seqs outermost"
);
let hd_big = 512usize;
let buf_ring = alloc_multi_seq_mlx_kv_for_layer(
&dev, 7, nkv, hd_big, cap, true, 2, n_seqs,
)
.expect("H151: alloc norms_per_pos=2 sliding");
assert_eq!(
buf_ring.norms_per_pos, 2,
"H151: norms_per_pos=2 propagation"
);
assert!(buf_ring.is_sliding, "H151: is_sliding=true propagation");
assert_eq!(
buf_ring.k_norms.shape(),
&[n_seqs as usize, nkv, cap, 2],
"H151 FALSIFIED: k_norms shape (norms_per_pos=2) must be 4-D"
);
assert_eq!(
buf_ring.v_norms.shape(),
&[n_seqs as usize, nkv, cap, 2],
"H151 FALSIFIED: v_norms shape (norms_per_pos=2) must be 4-D"
);
}
#[test]
fn h152_alloc_multi_seq_mlx_kv_for_layer_preflight_errors() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
assert!(
alloc_multi_seq_mlx_kv_for_layer(&dev, 0, 2, 256, 8, false, 1, 0).is_err(),
"H152 FALSIFIED: n_seqs=0 must error"
);
assert!(
alloc_multi_seq_mlx_kv_for_layer(&dev, 0, 0, 256, 8, false, 1, 1).is_err(),
"H152 FALSIFIED: nkv=0 must error"
);
assert!(
alloc_multi_seq_mlx_kv_for_layer(&dev, 0, 2, 0, 8, false, 1, 1).is_err(),
"H152 FALSIFIED: hd=0 must error"
);
assert!(
alloc_multi_seq_mlx_kv_for_layer(&dev, 0, 2, 256, 0, false, 1, 1).is_err(),
"H152 FALSIFIED: cap=0 must error"
);
assert!(
alloc_multi_seq_mlx_kv_for_layer(&dev, 0, 2, 256, 8, false, 0, 1).is_err(),
"H152 FALSIFIED: norms_per_pos=0 must error"
);
assert!(
alloc_multi_seq_mlx_kv_for_layer(&dev, 0, 2, 257, 8, false, 1, 1).is_err(),
"H152 FALSIFIED: odd hd must error (4-bit nibble-packed)"
);
}
#[test]
fn h153_multi_seq_mlx_kv_n_seqs_4_byte_scale_exact_formula() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 8usize;
let baseline_1 = alloc_multi_seq_mlx_kv_for_layer(&dev, 0, nkv, hd, cap, false, 1, 1)
.expect("H153: alloc norms_per_pos=1 n_seqs=1");
let lifted_1 = alloc_multi_seq_mlx_kv_for_layer(&dev, 0, nkv, hd, cap, false, 1, 4)
.expect("H153: alloc norms_per_pos=1 n_seqs=4");
assert_eq!(baseline_1.n_seqs, 1);
assert_eq!(lifted_1.n_seqs, 4);
assert_eq!(
lifted_1.k_packed.byte_len(),
baseline_1.k_packed.byte_len() * 4,
"H153 FALSIFIED: k_packed not 4× scale"
);
assert_eq!(
lifted_1.v_packed.byte_len(),
baseline_1.v_packed.byte_len() * 4,
"H153 FALSIFIED: v_packed not 4× scale"
);
assert_eq!(
lifted_1.k_norms.byte_len(),
baseline_1.k_norms.byte_len() * 4,
"H153 FALSIFIED: k_norms not 4× scale"
);
assert_eq!(
lifted_1.v_norms.byte_len(),
baseline_1.v_norms.byte_len() * 4,
"H153 FALSIFIED: v_norms not 4× scale"
);
let expected_k_packed = 4usize * nkv * cap * (hd / 2); let expected_v_packed = expected_k_packed;
let expected_k_norms = 4usize * nkv * cap * 1 * 4; let expected_v_norms = expected_k_norms;
let expected_total =
expected_k_packed + expected_v_packed + expected_k_norms + expected_v_norms;
assert_eq!(
lifted_1.k_packed.byte_len(),
expected_k_packed,
"H153 EXACT FORMULA FALSIFIED: k_packed bytes != n*nkv*cap*hd/2"
);
assert_eq!(
lifted_1.v_packed.byte_len(),
expected_v_packed,
"H153 EXACT FORMULA FALSIFIED: v_packed bytes != n*nkv*cap*hd/2"
);
assert_eq!(
lifted_1.k_norms.byte_len(),
expected_k_norms,
"H153 EXACT FORMULA FALSIFIED: k_norms bytes != n*nkv*cap*1*4"
);
assert_eq!(
lifted_1.v_norms.byte_len(),
expected_v_norms,
"H153 EXACT FORMULA FALSIFIED: v_norms bytes != n*nkv*cap*1*4"
);
use crate::serve::kv_persist::lcp_registry::ByteSized;
let actual_total = lifted_1.byte_len() as usize;
assert_eq!(
actual_total, expected_total,
"H153 EXACT FORMULA FALSIFIED: total composition"
);
assert_eq!(
actual_total, 16_896,
"H153 EXACT FORMULA FALSIFIED at concrete value: norms_per_pos=1 \
expected 16896 bytes for n_seqs=4 nkv=2 cap=8 hd=256; got {}",
actual_total
);
let hd_big = 512usize;
let baseline_2 = alloc_multi_seq_mlx_kv_for_layer(&dev, 0, nkv, hd_big, cap, false, 2, 1)
.expect("H153: alloc norms_per_pos=2 n_seqs=1");
let lifted_2 = alloc_multi_seq_mlx_kv_for_layer(&dev, 0, nkv, hd_big, cap, false, 2, 4)
.expect("H153: alloc norms_per_pos=2 n_seqs=4");
assert_eq!(
lifted_2.k_packed.byte_len(),
baseline_2.k_packed.byte_len() * 4,
"H153 FALSIFIED: k_packed norms_per_pos=2 not 4× scale"
);
assert_eq!(
lifted_2.k_norms.byte_len(),
baseline_2.k_norms.byte_len() * 4,
"H153 FALSIFIED: k_norms norms_per_pos=2 not 4× scale"
);
let expected_k_packed_2 = 4usize * nkv * cap * (hd_big / 2); let expected_k_norms_2 = 4usize * nkv * cap * 2 * 4; let expected_total_2 = 2 * expected_k_packed_2 + 2 * expected_k_norms_2;
let actual_total_2 = lifted_2.byte_len() as usize;
assert_eq!(
actual_total_2, expected_total_2,
"H153 EXACT FORMULA FALSIFIED: norms_per_pos=2 composition"
);
assert_eq!(
actual_total_2, 33_792,
"H153 EXACT FORMULA FALSIFIED at concrete value: norms_per_pos=2 \
expected 33792 bytes for n_seqs=4 nkv=2 cap=8 hd=512; got {}",
actual_total_2
);
assert_eq!(baseline_1.seq_lens.len(), 1);
assert_eq!(lifted_1.seq_lens.len(), 4);
assert_eq!(lifted_2.seq_lens.len(), 4);
assert!(lifted_1.seq_lens.iter().all(|&x| x == 0));
for (name, b) in [
("k_packed", &lifted_1.k_packed),
("v_packed", &lifted_1.v_packed),
("k_norms", &lifted_1.k_norms),
("v_norms", &lifted_1.v_norms),
] {
let s = b.shape().to_vec();
assert!(s.len() >= 3, "H153: {name} must be ≥3-D; got {:?}", s);
assert_eq!(
s[0], 4,
"H153 FALSIFIED: {name} shape[0] must be n_seqs=4 (n_seqs \
landed on wrong axis); got {:?}",
s
);
}
}
#[test]
fn h154_multi_seq_mlx_kv_per_slot_byte_isolation() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 4usize;
let mut cache = alloc_multi_seq_mlx_kv_for_layer(
&dev, 0, nkv, hd, cap, false, 1, 2,
)
.expect("H154: alloc n_seqs=2");
assert_eq!(cache.n_seqs, 2);
let slot_kp_bytes = nkv * cap * (hd / 2);
let slot_vp_bytes = nkv * cap * (hd / 2);
let slot_kn_bytes = nkv * cap * 1 * 4;
let slot_vn_bytes = nkv * cap * 1 * 4;
assert_eq!(
cache.k_packed.byte_len(),
2 * slot_kp_bytes,
"H154: k_packed total"
);
assert_eq!(
cache.v_packed.byte_len(),
2 * slot_vp_bytes,
"H154: v_packed total"
);
assert_eq!(
cache.k_norms.byte_len(),
2 * slot_kn_bytes,
"H154: k_norms total"
);
assert_eq!(
cache.v_norms.byte_len(),
2 * slot_vn_bytes,
"H154: v_norms total"
);
{
let s = cache.k_packed.as_mut_slice::<u8>().expect("kp u8 mut");
for (i, b) in s[..slot_kp_bytes].iter_mut().enumerate() {
*b = (((i * 7) % 251) + 1) as u8;
}
}
{
let s = cache.v_packed.as_mut_slice::<u8>().expect("vp u8 mut");
for (i, b) in s[..slot_vp_bytes].iter_mut().enumerate() {
*b = (((i * 11) % 251) + 1) as u8;
}
}
{
let s = cache.k_norms.as_mut_slice::<u8>().expect("kn u8 mut");
for (i, b) in s[..slot_kn_bytes].iter_mut().enumerate() {
*b = (((i * 13) % 251) + 1) as u8;
}
}
{
let s = cache.v_norms.as_mut_slice::<u8>().expect("vn u8 mut");
for (i, b) in s[..slot_vn_bytes].iter_mut().enumerate() {
*b = (((i * 17) % 251) + 1) as u8;
}
}
let kp_slot1_before: Vec<u8> = cache.k_packed.as_slice::<u8>().expect("kp r")
[slot_kp_bytes..2 * slot_kp_bytes]
.to_vec();
let vp_slot1_before: Vec<u8> = cache.v_packed.as_slice::<u8>().expect("vp r")
[slot_vp_bytes..2 * slot_vp_bytes]
.to_vec();
let kn_slot1_before: Vec<u8> = cache.k_norms.as_slice::<u8>().expect("kn r")
[slot_kn_bytes..2 * slot_kn_bytes]
.to_vec();
let vn_slot1_before: Vec<u8> = cache.v_norms.as_slice::<u8>().expect("vn r")
[slot_vn_bytes..2 * slot_vn_bytes]
.to_vec();
assert!(
kp_slot1_before.iter().all(|&b| b == 0),
"H154 sanity: kp slot1 zero"
);
assert!(
vp_slot1_before.iter().all(|&b| b == 0),
"H154 sanity: vp slot1 zero"
);
assert!(
kn_slot1_before.iter().all(|&b| b == 0),
"H154 sanity: kn slot1 zero"
);
assert!(
vn_slot1_before.iter().all(|&b| b == 0),
"H154 sanity: vn slot1 zero"
);
cache
.append_for_seq(SlotId(0), 3)
.expect("H154: append slot 0");
assert_eq!(cache.seq_lens[0], 3);
assert_eq!(cache.seq_lens[1], 0);
let kp_slot1_after: Vec<u8> = cache.k_packed.as_slice::<u8>().expect("kp r2")
[slot_kp_bytes..2 * slot_kp_bytes]
.to_vec();
let vp_slot1_after: Vec<u8> = cache.v_packed.as_slice::<u8>().expect("vp r2")
[slot_vp_bytes..2 * slot_vp_bytes]
.to_vec();
let kn_slot1_after: Vec<u8> = cache.k_norms.as_slice::<u8>().expect("kn r2")
[slot_kn_bytes..2 * slot_kn_bytes]
.to_vec();
let vn_slot1_after: Vec<u8> = cache.v_norms.as_slice::<u8>().expect("vn r2")
[slot_vn_bytes..2 * slot_vn_bytes]
.to_vec();
assert_eq!(
kp_slot1_before, kp_slot1_after,
"H154 FALSIFIED: slot 1 k_packed bytes changed after slot-0 write"
);
assert_eq!(
vp_slot1_before, vp_slot1_after,
"H154 FALSIFIED: slot 1 v_packed bytes changed after slot-0 write"
);
assert_eq!(
kn_slot1_before, kn_slot1_after,
"H154 FALSIFIED: slot 1 k_norms bytes changed after slot-0 write"
);
assert_eq!(
vn_slot1_before, vn_slot1_after,
"H154 FALSIFIED: slot 1 v_norms bytes changed after slot-0 write"
);
}
#[test]
fn h155_multi_seq_mlx_kv_n_seqs_1_byte_equivalent_to_legacy() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let cap = 8usize;
for &(hd, norms_per_pos) in &[(256usize, 1usize), (512usize, 2usize)] {
let multi =
alloc_multi_seq_mlx_kv_for_layer(&dev, 0, nkv, hd, cap, false, norms_per_pos, 1)
.expect("H155: alloc multi-seq n_seqs=1");
let legacy_packed_bytes = nkv * cap * (hd / 2);
let legacy_norms_bytes = nkv * cap * norms_per_pos * 4;
assert_eq!(
multi.k_packed.byte_len(),
legacy_packed_bytes,
"H155 FALSIFIED (hd={hd} norms_per_pos={norms_per_pos}): \
k_packed bytes {} != legacy {}",
multi.k_packed.byte_len(),
legacy_packed_bytes
);
assert_eq!(
multi.v_packed.byte_len(),
legacy_packed_bytes,
"H155 FALSIFIED (hd={hd} norms_per_pos={norms_per_pos}): \
v_packed bytes {} != legacy {}",
multi.v_packed.byte_len(),
legacy_packed_bytes
);
assert_eq!(
multi.k_norms.byte_len(),
legacy_norms_bytes,
"H155 FALSIFIED (hd={hd} norms_per_pos={norms_per_pos}): \
k_norms bytes {} != legacy {}",
multi.k_norms.byte_len(),
legacy_norms_bytes
);
assert_eq!(
multi.v_norms.byte_len(),
legacy_norms_bytes,
"H155 FALSIFIED (hd={hd} norms_per_pos={norms_per_pos}): \
v_norms bytes {} != legacy {}",
multi.v_norms.byte_len(),
legacy_norms_bytes
);
let legacy_total = 2 * legacy_packed_bytes + 2 * legacy_norms_bytes;
use crate::serve::kv_persist::lcp_registry::ByteSized;
assert_eq!(
multi.byte_len(),
legacy_total as u64,
"H155 FALSIFIED (hd={hd} norms_per_pos={norms_per_pos}): \
total byte_len {} != legacy 4-buffer sum {}",
multi.byte_len(),
legacy_total
);
}
}
#[test]
fn h156_multi_seq_mlx_kv_multi_seq_kv_cache_impl() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let n_seqs = 4u32;
let mut cache = alloc_multi_seq_mlx_kv_for_layer(&dev, 0, 2, 256, 8, false, 1, n_seqs)
.expect("H156: alloc n_seqs=4");
assert_eq!(
cache.slot_count(),
n_seqs,
"H156 FALSIFIED: slot_count must equal n_seqs={n_seqs}"
);
assert_eq!(cache.layout(), MultiSeqLayout::SeparateSlots);
for s in 0..n_seqs {
assert_eq!(
cache.seq_len(SlotId(s)).expect("seq_len in range"),
0,
"H156: slot {s} starts at cursor 0"
);
}
let err = cache.seq_len(SlotId(n_seqs)).expect_err("OOR n_seqs");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(n_seqs),
max_slots: n_seqs,
},
"H156 FALSIFIED: seq_len OOR shape; got {err:?}"
);
cache.append_for_seq(SlotId(0), 5).expect("append slot 0");
cache.append_for_seq(SlotId(2), 3).expect("append slot 2");
assert_eq!(cache.seq_len(SlotId(0)).unwrap(), 5);
assert_eq!(
cache.seq_len(SlotId(1)).unwrap(),
0,
"H156 FALSIFIED: slot 1 cursor touched by slot 0/2 append"
);
assert_eq!(cache.seq_len(SlotId(2)).unwrap(), 3);
assert_eq!(
cache.seq_len(SlotId(3)).unwrap(),
0,
"H156 FALSIFIED: slot 3 cursor touched by slot 0/2 append"
);
let err = cache
.append_for_seq(SlotId(n_seqs + 1), 1)
.expect_err("append OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(n_seqs + 1),
max_slots: n_seqs,
}
);
cache.drop_seq(SlotId(0)).expect("drop slot 0");
assert_eq!(cache.seq_len(SlotId(0)).unwrap(), 0, "H156: slot 0 reset");
assert_eq!(
cache.seq_len(SlotId(2)).unwrap(),
3,
"H156 FALSIFIED: slot 2 preserved through slot 0 drop"
);
let err = cache.drop_seq(SlotId(99)).expect_err("drop OOR");
assert!(matches!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(99),
max_slots: 4
}
));
cache
.fork_seq(SlotId(1), SlotId(1))
.expect("self-fork no-op");
let err = cache
.fork_seq(SlotId(99), SlotId(0))
.expect_err("fork src OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(99),
max_slots: 4
}
);
let err = cache
.fork_seq(SlotId(0), SlotId(99))
.expect_err("fork dst OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(99),
max_slots: 4
}
);
cache
.append_for_seq(SlotId(0), 5)
.expect("re-seed slot 0 for fork");
let slot0_before = cache.seq_len(SlotId(0)).unwrap();
let slot2_before = cache.seq_len(SlotId(2)).unwrap();
cache
.fork_seq(SlotId(0), SlotId(1))
.expect("iter-A3c closure: cross-slot fork must return Ok(())");
assert_eq!(
cache.seq_len(SlotId(1)).unwrap(),
slot0_before,
"H156 closure: fork must copy src's seq_len to dst"
);
assert_eq!(
cache.seq_len(SlotId(0)).unwrap(),
slot0_before,
"H156 closure: fork must NOT mutate src's seq_len"
);
assert_eq!(
cache.seq_len(SlotId(2)).unwrap(),
slot2_before,
"H156 closure: fork must NOT mutate non-src non-dst slots"
);
}
#[test]
fn h157_multi_seq_mlx_kv_reset_for_slot_per_slot_isolation_and_bounds() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 4usize;
let n_seqs = 4u32;
let mut cache = alloc_multi_seq_mlx_kv_for_layer(&dev, 0, nkv, hd, cap, false, 1, n_seqs)
.expect("alloc n_seqs=4");
for s in 0..(n_seqs as usize) {
cache.seq_lens[s] = (s as u32) + 11;
}
let slot_kp_bytes = nkv * cap * (hd / 2);
let slot_vp_bytes = nkv * cap * (hd / 2);
let slot_kn_bytes = nkv * cap * 1 * 4;
let slot_vn_bytes = nkv * cap * 1 * 4;
{
let s = cache.k_packed.as_mut_slice::<u8>().expect("kp u8");
for slot in 0..(n_seqs as usize) {
let start = slot * slot_kp_bytes;
for (i, b) in s[start..start + slot_kp_bytes].iter_mut().enumerate() {
*b = (((slot * 17 + i) % 251) + 1) as u8;
}
}
}
{
let s = cache.v_packed.as_mut_slice::<u8>().expect("vp u8");
for slot in 0..(n_seqs as usize) {
let start = slot * slot_vp_bytes;
for (i, b) in s[start..start + slot_vp_bytes].iter_mut().enumerate() {
*b = (((slot * 19 + i) % 251) + 1) as u8;
}
}
}
{
let s = cache.k_norms.as_mut_slice::<u8>().expect("kn u8");
for slot in 0..(n_seqs as usize) {
let start = slot * slot_kn_bytes;
for (i, b) in s[start..start + slot_kn_bytes].iter_mut().enumerate() {
*b = (((slot * 23 + i) % 251) + 1) as u8;
}
}
}
{
let s = cache.v_norms.as_mut_slice::<u8>().expect("vn u8");
for slot in 0..(n_seqs as usize) {
let start = slot * slot_vn_bytes;
for (i, b) in s[start..start + slot_vn_bytes].iter_mut().enumerate() {
*b = (((slot * 29 + i) % 251) + 1) as u8;
}
}
}
let kp_before: Vec<u8> = cache.k_packed.as_slice::<u8>().expect("kp r").to_vec();
let vp_before: Vec<u8> = cache.v_packed.as_slice::<u8>().expect("vp r").to_vec();
let kn_before: Vec<u8> = cache.k_norms.as_slice::<u8>().expect("kn r").to_vec();
let vn_before: Vec<u8> = cache.v_norms.as_slice::<u8>().expect("vn r").to_vec();
cache
.reset_for_slot(SlotId(1))
.expect("reset_for_slot(1) on n_seqs=4");
for s in 0..(n_seqs as usize) {
if s == 1 {
assert_eq!(
cache.seq_lens[s], 0,
"H157 FALSIFIED: slot 1 cursor must be 0 after reset_for_slot(1)"
);
} else {
assert_eq!(
cache.seq_lens[s],
(s as u32) + 11,
"H157 FALSIFIED: slot {s} cursor must be untouched"
);
}
}
let kp_after: Vec<u8> = cache.k_packed.as_slice::<u8>().expect("kp r2").to_vec();
let vp_after: Vec<u8> = cache.v_packed.as_slice::<u8>().expect("vp r2").to_vec();
let kn_after: Vec<u8> = cache.k_norms.as_slice::<u8>().expect("kn r2").to_vec();
let vn_after: Vec<u8> = cache.v_norms.as_slice::<u8>().expect("vn r2").to_vec();
assert_eq!(
kp_before, kp_after,
"H157 FALSIFIED: reset_for_slot must NOT zero k_packed bytes \
(cursor-masked discipline; matches drop_seq invariant)"
);
assert_eq!(
vp_before, vp_after,
"H157 FALSIFIED: reset_for_slot must NOT zero v_packed bytes"
);
assert_eq!(
kn_before, kn_after,
"H157 FALSIFIED: reset_for_slot must NOT zero k_norms bytes"
);
assert_eq!(
vn_before, vn_after,
"H157 FALSIFIED: reset_for_slot must NOT zero v_norms bytes"
);
let err = cache.reset_for_slot(SlotId(99)).expect_err("slot 99 OOR");
assert_eq!(
err,
MultiSeqError::SlotOutOfRange {
slot: SlotId(99),
max_slots: 4
},
"H157 FALSIFIED: reset OOR shape; got {err:?}"
);
let mut cache1 = alloc_multi_seq_mlx_kv_for_layer(&dev, 0, 2, 256, 4, false, 1, 1)
.expect("alloc n_seqs=1");
cache1.seq_lens[0] = 7;
cache1
.reset_for_slot(SlotId(0))
.expect("SlotId(0) at n_seqs=1 must succeed (byte-equivalence case)");
assert_eq!(cache1.seq_lens[0], 0);
}
#[test]
fn h159_multi_seq_hb_kv_fork_seq_cross_slot_copies_all_buffers() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 8usize;
let n_seqs = 4u32;
let mut c =
alloc_hb_kv_for_layer(&dev, 0, nkv, hd, cap, false, n_seqs).expect("alloc n_seqs=4");
let slot_kp = nkv * cap * hd;
let slot_vp = nkv * cap * hd;
let slot_kn = nkv * cap * 1 * 4;
let slot_vn = nkv * cap * 1 * 4;
{
let s = c.k_packed.as_mut_slice::<u8>().expect("kp u8");
for (i, b) in s[..slot_kp].iter_mut().enumerate() {
*b = (((i * 7) % 251) + 1) as u8;
}
}
{
let s = c.v_packed.as_mut_slice::<u8>().expect("vp u8");
for (i, b) in s[..slot_vp].iter_mut().enumerate() {
*b = (((i * 11) % 253) + 1) as u8;
}
}
{
let s = c.k_norms.as_mut_slice::<u8>().expect("kn u8");
for (i, b) in s[..slot_kn].iter_mut().enumerate() {
*b = (((i * 13) % 251) + 1) as u8;
}
}
{
let s = c.v_norms.as_mut_slice::<u8>().expect("vn u8");
for (i, b) in s[..slot_vn].iter_mut().enumerate() {
*b = (((i * 17) % 251) + 1) as u8;
}
}
c.append_for_seq(SlotId(0), 5).unwrap();
let src_kp = c.k_packed.as_slice::<u8>().unwrap()[..slot_kp].to_vec();
let src_vp = c.v_packed.as_slice::<u8>().unwrap()[..slot_vp].to_vec();
let src_kn = c.k_norms.as_slice::<u8>().unwrap()[..slot_kn].to_vec();
let src_vn = c.v_norms.as_slice::<u8>().unwrap()[..slot_vn].to_vec();
c.fork_seq(SlotId(0), SlotId(2))
.expect("H159: fork must succeed post-A3c");
let dst_kp = c.k_packed.as_slice::<u8>().unwrap()[2 * slot_kp..3 * slot_kp].to_vec();
let dst_vp = c.v_packed.as_slice::<u8>().unwrap()[2 * slot_vp..3 * slot_vp].to_vec();
let dst_kn = c.k_norms.as_slice::<u8>().unwrap()[2 * slot_kn..3 * slot_kn].to_vec();
let dst_vn = c.v_norms.as_slice::<u8>().unwrap()[2 * slot_vn..3 * slot_vn].to_vec();
assert_eq!(src_kp, dst_kp, "H159 FALSIFIED: k_packed dst != src");
assert_eq!(src_vp, dst_vp, "H159 FALSIFIED: v_packed dst != src");
assert_eq!(src_kn, dst_kn, "H159 FALSIFIED: k_norms dst != src");
assert_eq!(src_vn, dst_vn, "H159 FALSIFIED: v_norms dst != src");
assert_eq!(c.seq_len(SlotId(2)).unwrap(), 5, "H159: cursor copied");
assert_eq!(
c.seq_len(SlotId(0)).unwrap(),
5,
"H159: src cursor unchanged"
);
let src_kp_after = c.k_packed.as_slice::<u8>().unwrap()[..slot_kp].to_vec();
assert_eq!(src_kp, src_kp_after, "H159 FALSIFIED: src k_packed mutated");
}
#[test]
fn h160_multi_seq_hybrid_kv_fork_seq_cross_slot_copies_all_buffers() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 8usize;
let n_seqs = 4u32;
let mut c = alloc_multi_seq_hybrid_kv_for_layer(&dev, 0, nkv, hd, cap, false, n_seqs)
.expect("alloc n_seqs=4");
let slot_k_bytes = nkv * cap * hd * 2;
let slot_vp_bytes = nkv * cap * hd;
let slot_vn_bytes = nkv * cap * 1 * 4;
{
let s = c.k.as_mut_slice::<u8>().expect("k u8");
let start = 1 * slot_k_bytes;
for (i, b) in s[start..start + slot_k_bytes].iter_mut().enumerate() {
*b = (((i * 23) % 251) + 1) as u8;
}
}
{
let s = c.v_packed.as_mut_slice::<u8>().expect("vp u8");
let start = 1 * slot_vp_bytes;
for (i, b) in s[start..start + slot_vp_bytes].iter_mut().enumerate() {
*b = (((i * 29) % 253) + 1) as u8;
}
}
{
let s = c.v_norms.as_mut_slice::<u8>().expect("vn u8");
let start = 1 * slot_vn_bytes;
for (i, b) in s[start..start + slot_vn_bytes].iter_mut().enumerate() {
*b = (((i * 31) % 251) + 1) as u8;
}
}
c.append_for_seq(SlotId(1), 9).unwrap();
let src_k = c.k.as_slice::<u8>().unwrap()[slot_k_bytes..2 * slot_k_bytes].to_vec();
let src_vp =
c.v_packed.as_slice::<u8>().unwrap()[slot_vp_bytes..2 * slot_vp_bytes].to_vec();
let src_vn = c.v_norms.as_slice::<u8>().unwrap()[slot_vn_bytes..2 * slot_vn_bytes].to_vec();
c.fork_seq(SlotId(1), SlotId(3))
.expect("H160: fork must succeed post-A3c");
let dst_k = c.k.as_slice::<u8>().unwrap()[3 * slot_k_bytes..4 * slot_k_bytes].to_vec();
let dst_vp =
c.v_packed.as_slice::<u8>().unwrap()[3 * slot_vp_bytes..4 * slot_vp_bytes].to_vec();
let dst_vn =
c.v_norms.as_slice::<u8>().unwrap()[3 * slot_vn_bytes..4 * slot_vn_bytes].to_vec();
assert_eq!(src_k, dst_k, "H160 FALSIFIED: k dst != src");
assert_eq!(src_vp, dst_vp, "H160 FALSIFIED: v_packed dst != src");
assert_eq!(src_vn, dst_vn, "H160 FALSIFIED: v_norms dst != src");
assert_eq!(c.seq_len(SlotId(3)).unwrap(), 9, "H160: cursor copied");
assert_eq!(
c.seq_len(SlotId(1)).unwrap(),
9,
"H160: src cursor unchanged"
);
let src_k_after = c.k.as_slice::<u8>().unwrap()[slot_k_bytes..2 * slot_k_bytes].to_vec();
assert_eq!(src_k, src_k_after, "H160 FALSIFIED: src k mutated");
}
#[test]
fn h161_multi_seq_dense_kv_fork_seq_cross_slot_copies_all_buffers() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 4usize;
let n_seqs = 4u32;
let mut c =
alloc_multi_seq_dense_kv_for_layer(&dev, 0, nkv, hd, cap, false, DType::F32, n_seqs)
.expect("alloc n_seqs=4");
let slot_k_bytes = nkv * cap * hd * 4;
let slot_v_bytes = nkv * cap * hd * 4;
{
let s = c.k.as_mut_slice::<u8>().expect("k u8");
let start = 2 * slot_k_bytes;
for (i, b) in s[start..start + slot_k_bytes].iter_mut().enumerate() {
*b = (((i * 37) % 251) + 1) as u8;
}
}
{
let s = c.v.as_mut_slice::<u8>().expect("v u8");
let start = 2 * slot_v_bytes;
for (i, b) in s[start..start + slot_v_bytes].iter_mut().enumerate() {
*b = (((i * 41) % 253) + 1) as u8;
}
}
c.append_for_seq(SlotId(2), 3).unwrap();
let src_k = c.k.as_slice::<u8>().unwrap()[2 * slot_k_bytes..3 * slot_k_bytes].to_vec();
let src_v = c.v.as_slice::<u8>().unwrap()[2 * slot_v_bytes..3 * slot_v_bytes].to_vec();
c.fork_seq(SlotId(2), SlotId(0))
.expect("H161: fork must succeed post-A3c");
let dst_k = c.k.as_slice::<u8>().unwrap()[..slot_k_bytes].to_vec();
let dst_v = c.v.as_slice::<u8>().unwrap()[..slot_v_bytes].to_vec();
assert_eq!(src_k, dst_k, "H161 FALSIFIED: k dst != src");
assert_eq!(src_v, dst_v, "H161 FALSIFIED: v dst != src");
assert_eq!(c.seq_len(SlotId(0)).unwrap(), 3, "H161: cursor copied");
assert_eq!(
c.seq_len(SlotId(2)).unwrap(),
3,
"H161: src cursor unchanged"
);
let src_k_after =
c.k.as_slice::<u8>().unwrap()[2 * slot_k_bytes..3 * slot_k_bytes].to_vec();
assert_eq!(src_k, src_k_after, "H161 FALSIFIED: src k mutated");
}
#[test]
fn h162_multi_seq_mlx_kv_fork_seq_cross_slot_copies_all_buffers() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = match skip_dev() {
Some(d) => d,
None => return,
};
let nkv = 2usize;
let hd = 256usize;
let cap = 4usize;
let n_seqs = 4u32;
let mut c = alloc_multi_seq_mlx_kv_for_layer(&dev, 0, nkv, hd, cap, false, 1, n_seqs)
.expect("alloc n_seqs=4");
let slot_kp = nkv * cap * (hd / 2);
let slot_vp = nkv * cap * (hd / 2);
let slot_kn = nkv * cap * 1 * 4;
let slot_vn = nkv * cap * 1 * 4;
{
let s = c.k_packed.as_mut_slice::<u8>().expect("kp u8");
for (i, b) in s[..slot_kp].iter_mut().enumerate() {
*b = (((i * 43) % 251) + 1) as u8;
}
}
{
let s = c.v_packed.as_mut_slice::<u8>().expect("vp u8");
for (i, b) in s[..slot_vp].iter_mut().enumerate() {
*b = (((i * 47) % 253) + 1) as u8;
}
}
{
let s = c.k_norms.as_mut_slice::<u8>().expect("kn u8");
for (i, b) in s[..slot_kn].iter_mut().enumerate() {
*b = (((i * 53) % 251) + 1) as u8;
}
}
{
let s = c.v_norms.as_mut_slice::<u8>().expect("vn u8");
for (i, b) in s[..slot_vn].iter_mut().enumerate() {
*b = (((i * 59) % 251) + 1) as u8;
}
}
c.append_for_seq(SlotId(0), 4).unwrap();
let src_kp = c.k_packed.as_slice::<u8>().unwrap()[..slot_kp].to_vec();
let src_vp = c.v_packed.as_slice::<u8>().unwrap()[..slot_vp].to_vec();
let src_kn = c.k_norms.as_slice::<u8>().unwrap()[..slot_kn].to_vec();
let src_vn = c.v_norms.as_slice::<u8>().unwrap()[..slot_vn].to_vec();
c.fork_seq(SlotId(0), SlotId(3))
.expect("H162: fork must succeed post-A3c");
let dst_kp = c.k_packed.as_slice::<u8>().unwrap()[3 * slot_kp..4 * slot_kp].to_vec();
let dst_vp = c.v_packed.as_slice::<u8>().unwrap()[3 * slot_vp..4 * slot_vp].to_vec();
let dst_kn = c.k_norms.as_slice::<u8>().unwrap()[3 * slot_kn..4 * slot_kn].to_vec();
let dst_vn = c.v_norms.as_slice::<u8>().unwrap()[3 * slot_vn..4 * slot_vn].to_vec();
assert_eq!(src_kp, dst_kp, "H162 FALSIFIED: k_packed dst != src");
assert_eq!(src_vp, dst_vp, "H162 FALSIFIED: v_packed dst != src");
assert_eq!(src_kn, dst_kn, "H162 FALSIFIED: k_norms dst != src");
assert_eq!(src_vn, dst_vn, "H162 FALSIFIED: v_norms dst != src");
assert_eq!(c.seq_len(SlotId(3)).unwrap(), 4, "H162: cursor copied");
assert_eq!(
c.seq_len(SlotId(0)).unwrap(),
4,
"H162: src cursor unchanged"
);
let src_kp_after = c.k_packed.as_slice::<u8>().unwrap()[..slot_kp].to_vec();
assert_eq!(src_kp, src_kp_after, "H162 FALSIFIED: src k_packed mutated");
}
#[test]
fn gemma_lcp_layer_kv_bytesized_sums_both_legs() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let dev = MlxDevice::new().expect("device");
let (nkv, cap, hd) = (2usize, 8usize, 4usize);
let dense = DenseKvBuffers {
k: dev
.alloc_buffer(nkv * cap * hd * 4, DType::F32, vec![nkv, cap, hd])
.unwrap(),
v: dev
.alloc_buffer(nkv * cap * hd * 4, DType::F32, vec![nkv, cap, hd])
.unwrap(),
capacity: cap,
is_sliding: false,
dtype: DType::F32,
};
let hybrid = HybridKvBuffers {
k: dev
.alloc_buffer(nkv * cap * hd * 2, DType::F16, vec![nkv, cap, hd])
.unwrap(),
v_packed: dev
.alloc_buffer(nkv * cap * hd, DType::U8, vec![nkv, cap, hd])
.unwrap(),
v_norms: dev
.alloc_buffer(nkv * cap * 4, DType::F32, vec![nkv, cap])
.unwrap(),
capacity: cap,
is_sliding: false,
norms_per_pos: 1,
bf16_xlen_k: None,
bf16_xlen_v: None,
};
use crate::serve::kv_persist::lcp_registry::ByteSized;
let d_bytes = ByteSized::byte_len(&dense);
let h_bytes = ByteSized::byte_len(&hybrid);
assert_eq!(
d_bytes,
(nkv * cap * hd * 4 * 2) as u64,
"dense = 2 F32 buffers"
);
assert_eq!(
h_bytes,
(nkv * cap * hd * 2 + nkv * cap * hd + nkv * cap * 4) as u64
);
let dense_only = GemmaLcpLayerKv::Dense(DenseKvBuffers {
k: dev
.alloc_buffer(nkv * cap * hd * 4, DType::F32, vec![nkv, cap, hd])
.unwrap(),
v: dev
.alloc_buffer(nkv * cap * hd * 4, DType::F32, vec![nkv, cap, hd])
.unwrap(),
capacity: cap,
is_sliding: false,
dtype: DType::F32,
});
assert_eq!(ByteSized::byte_len(&dense_only), d_bytes);
let both = GemmaLcpLayerKv::DenseAndHybrid(
DenseKvBuffers {
k: dev
.alloc_buffer(nkv * cap * hd * 4, DType::F32, vec![nkv, cap, hd])
.unwrap(),
v: dev
.alloc_buffer(nkv * cap * hd * 4, DType::F32, vec![nkv, cap, hd])
.unwrap(),
capacity: cap,
is_sliding: false,
dtype: DType::F32,
},
hybrid,
);
assert_eq!(
ByteSized::byte_len(&both),
d_bytes + h_bytes,
"DenseAndHybrid byte_len must sum dense + hybrid legs exactly"
);
assert!(both.hybrid().is_some());
assert!(dense_only.hybrid().is_none());
assert_eq!(both.dense().capacity, cap);
}
}