use anyhow::{anyhow, Result};
use mlx_native::{DType, MlxBuffer, MlxDevice};
pub struct FaPrefillArena {
pub q_bf16_seq: MlxBuffer,
pub q_bf16_hm: MlxBuffer,
pub k_bf16_seq: MlxBuffer,
pub k_bf16_hm: MlxBuffer,
pub v_bf16_seq: MlxBuffer,
pub v_bf16_hm: MlxBuffer,
pub out_bf16_hm: MlxBuffer,
pub seq_capacity: u32,
pub n_heads: u32,
pub n_kv_heads: u32,
pub head_dim: u32,
}
impl FaPrefillArena {
pub fn new(
device: &MlxDevice,
seq_capacity: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
) -> Result<Self> {
if seq_capacity == 0 || n_heads == 0 || n_kv_heads == 0 || head_dim == 0 {
return Err(anyhow!(
"FaPrefillArena::new: zero dim \
seq_capacity={} n_heads={} n_kv_heads={} head_dim={}",
seq_capacity,
n_heads,
n_kv_heads,
head_dim
));
}
if n_heads % n_kv_heads != 0 {
return Err(anyhow!(
"FaPrefillArena::new: n_heads ({}) must be divisible by n_kv_heads ({})",
n_heads,
n_kv_heads
));
}
let seq = seq_capacity as usize;
let nh = n_heads as usize;
let nkv = n_kv_heads as usize;
let d = head_dim as usize;
let q_elems = seq * nh * d;
let k_elems = seq * nkv * d;
let v_elems = seq * nkv * d; let out_elems = seq * nh * d;
let q_bf16_seq = device
.alloc_buffer(q_elems * 2, DType::BF16, vec![seq, nh, d])
.map_err(|e| anyhow!("FaPrefillArena::new alloc q_bf16_seq: {e}"))?;
let q_bf16_hm = device
.alloc_buffer(q_elems * 2, DType::BF16, vec![1, nh, seq, d])
.map_err(|e| anyhow!("FaPrefillArena::new alloc q_bf16_hm: {e}"))?;
let k_bf16_seq = device
.alloc_buffer(k_elems * 2, DType::BF16, vec![seq, nkv, d])
.map_err(|e| anyhow!("FaPrefillArena::new alloc k_bf16_seq: {e}"))?;
let k_bf16_hm = device
.alloc_buffer(k_elems * 2, DType::BF16, vec![1, nkv, seq, d])
.map_err(|e| anyhow!("FaPrefillArena::new alloc k_bf16_hm: {e}"))?;
let v_bf16_seq = device
.alloc_buffer(v_elems * 2, DType::BF16, vec![seq, nkv, d])
.map_err(|e| anyhow!("FaPrefillArena::new alloc v_bf16_seq: {e}"))?;
let v_bf16_hm = device
.alloc_buffer(v_elems * 2, DType::BF16, vec![1, nkv, seq, d])
.map_err(|e| anyhow!("FaPrefillArena::new alloc v_bf16_hm: {e}"))?;
let out_bf16_hm = device
.alloc_buffer(out_elems * 2, DType::BF16, vec![1, nh, seq, d])
.map_err(|e| anyhow!("FaPrefillArena::new alloc out_bf16_hm: {e}"))?;
Ok(Self {
q_bf16_seq,
q_bf16_hm,
k_bf16_seq,
k_bf16_hm,
v_bf16_seq,
v_bf16_hm,
out_bf16_hm,
seq_capacity,
n_heads,
n_kv_heads,
head_dim,
})
}
pub fn validate_fits(
&self,
seq_len: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
) -> Result<()> {
if seq_len > self.seq_capacity {
return Err(anyhow!(
"FaPrefillArena::validate_fits: seq_len {} exceeds capacity {}",
seq_len,
self.seq_capacity
));
}
if n_heads != self.n_heads || n_kv_heads != self.n_kv_heads || head_dim != self.head_dim {
return Err(anyhow!(
"FaPrefillArena::validate_fits: shape mismatch — \
arena (n_heads={}, n_kv_heads={}, head_dim={}) vs \
call (n_heads={}, n_kv_heads={}, head_dim={})",
self.n_heads,
self.n_kv_heads,
self.head_dim,
n_heads,
n_kv_heads,
head_dim,
));
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn device_or_skip() -> Option<MlxDevice> {
MlxDevice::new().ok()
}
#[test]
fn test_arena_new_qwen35_pp101() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_arena_new_qwen35_pp101: skipping — no Metal device");
return;
}
};
let (seq, nh, nkv, d) = (101usize, 16usize, 2usize, 256usize);
let arena = FaPrefillArena::new(&device, seq as u32, nh as u32, nkv as u32, d as u32)
.expect("arena new pp101");
assert_eq!(arena.seq_capacity, seq as u32);
assert_eq!(arena.n_heads, nh as u32);
assert_eq!(arena.n_kv_heads, nkv as u32);
assert_eq!(arena.head_dim, d as u32);
let q_bytes = seq * nh * d * 2;
let k_bytes = seq * nkv * d * 2;
assert_eq!(arena.q_bf16_seq.byte_len(), q_bytes, "q_bf16_seq byte_len");
assert_eq!(arena.q_bf16_hm.byte_len(), q_bytes, "q_bf16_hm byte_len");
assert_eq!(arena.k_bf16_seq.byte_len(), k_bytes, "k_bf16_seq byte_len");
assert_eq!(arena.k_bf16_hm.byte_len(), k_bytes, "k_bf16_hm byte_len");
assert_eq!(arena.v_bf16_seq.byte_len(), k_bytes, "v_bf16_seq byte_len");
assert_eq!(arena.v_bf16_hm.byte_len(), k_bytes, "v_bf16_hm byte_len");
assert_eq!(
arena.out_bf16_hm.byte_len(),
q_bytes,
"out_bf16_hm byte_len"
);
}
#[test]
fn test_arena_new_qwen35_pp4096() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_arena_new_qwen35_pp4096: skipping — no Metal device");
return;
}
};
let arena = FaPrefillArena::new(&device, 4096, 16, 2, 256).expect("arena new pp4096");
assert_eq!(arena.seq_capacity, 4096);
}
#[test]
fn test_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_arena_new_zero_dim_rejected: skipping — no Metal device");
return;
}
};
assert!(
FaPrefillArena::new(&device, 0, 16, 2, 256).is_err(),
"seq=0 should be rejected"
);
assert!(
FaPrefillArena::new(&device, 101, 0, 2, 256).is_err(),
"n_heads=0 should be rejected"
);
assert!(
FaPrefillArena::new(&device, 101, 16, 0, 256).is_err(),
"n_kv_heads=0 should be rejected"
);
assert!(
FaPrefillArena::new(&device, 101, 16, 2, 0).is_err(),
"head_dim=0 should be rejected"
);
}
#[test]
fn test_arena_new_gqa_divisibility_rejected() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_arena_new_gqa_divisibility_rejected: skipping — no Metal device");
return;
}
};
assert!(
FaPrefillArena::new(&device, 101, 15, 2, 256).is_err(),
"n_heads=15, n_kv_heads=2 (15 % 2 != 0) should be rejected"
);
}
#[test]
fn test_validate_fits_seq_overrun() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_validate_fits_seq_overrun: skipping — no Metal device");
return;
}
};
let arena = FaPrefillArena::new(&device, 128, 16, 2, 256).expect("arena new seq128");
assert!(
arena.validate_fits(256, 16, 2, 256).is_err(),
"seq_len=256 > capacity=128 should be rejected"
);
}
#[test]
fn test_validate_fits_shape_mismatch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_validate_fits_shape_mismatch: skipping — no Metal device");
return;
}
};
let arena = FaPrefillArena::new(&device, 128, 16, 2, 256).expect("arena new nh16");
assert!(
arena.validate_fits(128, 8, 2, 256).is_err(),
"n_heads=8 vs arena n_heads=16 should be rejected"
);
assert!(
arena.validate_fits(128, 16, 4, 256).is_err(),
"n_kv_heads=4 vs arena n_kv_heads=2 should be rejected"
);
assert!(
arena.validate_fits(128, 16, 2, 128).is_err(),
"head_dim=128 vs arena head_dim=256 should be rejected"
);
}
#[test]
fn test_validate_fits_exact_match() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_validate_fits_exact_match: skipping — no Metal device");
return;
}
};
let arena = FaPrefillArena::new(&device, 128, 16, 2, 256).expect("arena new seq128");
assert!(
arena.validate_fits(128, 16, 2, 256).is_ok(),
"exact-match validate_fits should return Ok"
);
assert!(
arena.validate_fits(64, 16, 2, 256).is_ok(),
"seq_len=64 <= capacity=128 should be Ok"
);
}
#[test]
fn test_arena_buffers_zero_initialized() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_arena_buffers_zero_initialized: skipping — no Metal device");
return;
}
};
let arena = FaPrefillArena::new(&device, 64, 16, 2, 256).expect("arena new seq64");
let bufs: [(&MlxBuffer, &str); 7] = [
(&arena.q_bf16_seq, "q_bf16_seq"),
(&arena.q_bf16_hm, "q_bf16_hm"),
(&arena.k_bf16_seq, "k_bf16_seq"),
(&arena.k_bf16_hm, "k_bf16_hm"),
(&arena.v_bf16_seq, "v_bf16_seq"),
(&arena.v_bf16_hm, "v_bf16_hm"),
(&arena.out_bf16_hm, "out_bf16_hm"),
];
for (buf, name) in &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:#06x} (BF16 0x0000 = +0.0), expected zero"
);
}
}
}
}