use anyhow::{anyhow, Result};
use mlx_native::{DType, MlxBuffer, MlxDevice};
pub struct ChunkAllocsArena {
pub q_expanded_buf: MlxBuffer,
pub k_expanded_buf: MlxBuffer,
pub q_bf16_buf: MlxBuffer,
pub k_bf16_buf: MlxBuffer,
pub v_bf16_buf: MlxBuffer,
pub g_log_decay_buf: MlxBuffer,
pub o_bf16_buf: MlxBuffer,
pub seq_capacity: u32,
pub n_v_heads: u32,
pub d_k: u32,
pub d_v: u32,
}
impl ChunkAllocsArena {
pub fn new(
device: &MlxDevice,
seq_capacity: u32,
n_v_heads: u32,
d_k: u32,
d_v: u32,
) -> Result<Self> {
if seq_capacity == 0 || n_v_heads == 0 || d_k == 0 || d_v == 0 {
return Err(anyhow!(
"ChunkAllocsArena::new: zero dim \
seq_capacity={} n_v_heads={} d_k={} d_v={}",
seq_capacity,
n_v_heads,
d_k,
d_v,
));
}
let seq = seq_capacity as usize;
let nv = n_v_heads as usize;
let dk = d_k as usize;
let dv = d_v as usize;
let q_elems_exp = seq * nv * dk; let v_elems = seq * nv * dv; let g_elems = seq * nv; let out_elems_bf16 = v_elems;
let q_expanded_buf = device
.alloc_buffer(q_elems_exp * 4, DType::F32, vec![q_elems_exp])
.map_err(|e| anyhow!("ChunkAllocsArena alloc q_expanded_buf: {e}"))?;
let k_expanded_buf = device
.alloc_buffer(q_elems_exp * 4, DType::F32, vec![q_elems_exp])
.map_err(|e| anyhow!("ChunkAllocsArena alloc k_expanded_buf: {e}"))?;
let q_bf16_buf = device
.alloc_buffer(q_elems_exp * 2, DType::BF16, vec![q_elems_exp])
.map_err(|e| anyhow!("ChunkAllocsArena alloc q_bf16_buf: {e}"))?;
let k_bf16_buf = device
.alloc_buffer(q_elems_exp * 2, DType::BF16, vec![q_elems_exp])
.map_err(|e| anyhow!("ChunkAllocsArena alloc k_bf16_buf: {e}"))?;
let v_bf16_buf = device
.alloc_buffer(v_elems * 2, DType::BF16, vec![v_elems])
.map_err(|e| anyhow!("ChunkAllocsArena alloc v_bf16_buf: {e}"))?;
let g_log_decay_buf = device
.alloc_buffer(g_elems * 4, DType::F32, vec![g_elems])
.map_err(|e| anyhow!("ChunkAllocsArena alloc g_log_decay_buf: {e}"))?;
let o_bf16_buf = device
.alloc_buffer(out_elems_bf16 * 2, DType::BF16, vec![out_elems_bf16])
.map_err(|e| anyhow!("ChunkAllocsArena alloc o_bf16_buf: {e}"))?;
Ok(Self {
q_expanded_buf,
k_expanded_buf,
q_bf16_buf,
k_bf16_buf,
v_bf16_buf,
g_log_decay_buf,
o_bf16_buf,
seq_capacity,
n_v_heads,
d_k,
d_v,
})
}
pub fn validate_fits(&self, seq_len: u32, n_v_heads: u32, d_k: u32, d_v: u32) -> Result<()> {
if seq_len > self.seq_capacity {
return Err(anyhow!(
"ChunkAllocsArena::validate_fits: seq_len {} exceeds capacity {}",
seq_len,
self.seq_capacity
));
}
if n_v_heads != self.n_v_heads || d_k != self.d_k || d_v != self.d_v {
return Err(anyhow!(
"ChunkAllocsArena::validate_fits: shape mismatch — \
arena (nv={}, dk={}, dv={}) vs \
call (nv={}, dk={}, dv={})",
self.n_v_heads,
self.d_k,
self.d_v,
n_v_heads,
d_k,
d_v,
));
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn device_or_skip() -> Option<MlxDevice> {
MlxDevice::new().ok()
}
#[test]
fn test_chunk_allocs_arena_new_qwen36_35b_pp4096() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!(
"test_chunk_allocs_arena_new_qwen36_35b_pp4096: \
skipping — no Metal device"
);
return;
}
};
let (seq, nv, dk, dv) = (4096u32, 32u32, 128u32, 128u32);
let arena =
ChunkAllocsArena::new(&device, seq, nv, dk, dv).expect("chunk allocs arena pp4096");
assert_eq!(arena.seq_capacity, seq);
assert_eq!(arena.n_v_heads, nv);
assert_eq!(arena.d_k, dk);
assert_eq!(arena.d_v, dv);
let q_elems_exp = (seq as usize) * (nv as usize) * (dk as usize);
let v_elems = (seq as usize) * (nv as usize) * (dv as usize);
let g_elems = (seq as usize) * (nv as usize);
assert_eq!(
arena.q_expanded_buf.byte_len(),
q_elems_exp * 4,
"q_expanded_buf"
);
assert_eq!(
arena.k_expanded_buf.byte_len(),
q_elems_exp * 4,
"k_expanded_buf"
);
assert_eq!(arena.q_bf16_buf.byte_len(), q_elems_exp * 2, "q_bf16_buf");
assert_eq!(arena.k_bf16_buf.byte_len(), q_elems_exp * 2, "k_bf16_buf");
assert_eq!(arena.v_bf16_buf.byte_len(), v_elems * 2, "v_bf16_buf");
assert_eq!(
arena.g_log_decay_buf.byte_len(),
g_elems * 4,
"g_log_decay_buf"
);
assert_eq!(arena.o_bf16_buf.byte_len(), v_elems * 2, "o_bf16_buf");
}
#[test]
fn test_chunk_allocs_arena_new_small_shape() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_chunk_allocs_arena_new_small_shape: skipping — no Metal device");
return;
}
};
let arena =
ChunkAllocsArena::new(&device, 128, 8, 32, 32).expect("chunk allocs arena small");
assert_eq!(arena.seq_capacity, 128);
assert_eq!(arena.n_v_heads, 8);
assert_eq!(arena.d_k, 32);
assert_eq!(arena.d_v, 32);
}
#[test]
fn test_chunk_allocs_arena_new_zero_dim_rejected() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!(
"test_chunk_allocs_arena_new_zero_dim_rejected: \
skipping — no Metal device"
);
return;
}
};
assert!(ChunkAllocsArena::new(&device, 0, 8, 32, 32).is_err());
assert!(ChunkAllocsArena::new(&device, 128, 0, 32, 32).is_err());
assert!(ChunkAllocsArena::new(&device, 128, 8, 0, 32).is_err());
assert!(ChunkAllocsArena::new(&device, 128, 8, 32, 0).is_err());
}
#[test]
fn test_chunk_allocs_validate_fits_exact_match() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!(
"test_chunk_allocs_validate_fits_exact_match: skipping — no Metal device"
);
return;
}
};
let arena = ChunkAllocsArena::new(&device, 256, 8, 32, 32).expect("chunk allocs arena new");
assert!(arena.validate_fits(256, 8, 32, 32).is_ok());
assert!(arena.validate_fits(128, 8, 32, 32).is_ok());
}
#[test]
fn test_chunk_allocs_validate_fits_seq_overrun() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!(
"test_chunk_allocs_validate_fits_seq_overrun: skipping — no Metal device"
);
return;
}
};
let arena = ChunkAllocsArena::new(&device, 256, 8, 32, 32).expect("chunk allocs arena new");
assert!(arena.validate_fits(512, 8, 32, 32).is_err());
}
#[test]
fn test_chunk_allocs_validate_fits_shape_mismatch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!(
"test_chunk_allocs_validate_fits_shape_mismatch: skipping — no Metal device"
);
return;
}
};
let arena = ChunkAllocsArena::new(&device, 256, 8, 32, 32).expect("chunk allocs arena new");
assert!(arena.validate_fits(256, 4, 32, 32).is_err()); assert!(arena.validate_fits(256, 8, 16, 32).is_err()); assert!(arena.validate_fits(256, 8, 32, 16).is_err()); }
#[test]
fn test_chunk_allocs_arena_buffers_zero_initialised() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!(
"test_chunk_allocs_arena_buffers_zero_initialised: \
skipping — no Metal device"
);
return;
}
};
let arena = ChunkAllocsArena::new(&device, 64, 4, 16, 16).expect("chunk allocs arena new");
let f32_bufs: [(&MlxBuffer, &str); 3] = [
(&arena.q_expanded_buf, "q_expanded_buf"),
(&arena.k_expanded_buf, "k_expanded_buf"),
(&arena.g_log_decay_buf, "g_log_decay_buf"),
];
for (buf, name) in &f32_bufs {
let slice = buf
.as_slice::<f32>()
.unwrap_or_else(|e| panic!("{name} as_slice::<f32> failed: {e}"));
let check_len = 16.min(slice.len());
for (i, &v) in slice[..check_len].iter().enumerate() {
assert_eq!(
v, 0.0f32,
"{name}[{i}] = {v} (expected zero from device.alloc_buffer)"
);
}
}
let bf16_bufs: [(&MlxBuffer, &str); 4] = [
(&arena.q_bf16_buf, "q_bf16_buf"),
(&arena.k_bf16_buf, "k_bf16_buf"),
(&arena.v_bf16_buf, "v_bf16_buf"),
(&arena.o_bf16_buf, "o_bf16_buf"),
];
for (buf, name) in &bf16_bufs {
let slice = buf
.as_slice::<u16>()
.unwrap_or_else(|e| panic!("{name} as_slice::<u16> failed: {e}"));
let check_len = 16.min(slice.len());
for (i, &v) in slice[..check_len].iter().enumerate() {
assert_eq!(
v, 0u16,
"{name}[{i}] = {v} (expected zero from device.alloc_buffer)"
);
}
}
}
#[test]
fn test_chunk_allocs_q_k_expanded_same_size() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_chunk_allocs_q_k_expanded_same_size: skipping — no Metal device");
return;
}
};
let arena = ChunkAllocsArena::new(&device, 256, 8, 32, 32).expect("chunk allocs arena new");
assert_eq!(
arena.q_expanded_buf.byte_len(),
arena.k_expanded_buf.byte_len(),
"q_expanded and k_expanded must be the same size (both [T, H, K] F32)"
);
assert_eq!(
arena.q_bf16_buf.byte_len(),
arena.k_bf16_buf.byte_len(),
"q_bf16 and k_bf16 must be the same size (both [T, H, K] BF16)"
);
assert_eq!(
arena.v_bf16_buf.byte_len(),
arena.o_bf16_buf.byte_len(),
"v_bf16 and o_bf16 must be the same size (both [T, H, V] BF16)"
);
}
}