use anyhow::{anyhow, Result};
use mlx_native::{DType, MlxBuffer, MlxDevice};
pub struct DnPrefillArena {
pub qkv_conv_buf: MlxBuffer,
pub ssm_params_buf: MlxBuffer,
pub q_scaled_buf: MlxBuffer,
pub g_buf: MlxBuffer,
pub beta_buf: MlxBuffer,
pub g_params_buf: MlxBuffer,
pub gated_buf: MlxBuffer,
pub op8_params_buf: MlxBuffer,
pub attn_out_buf: MlxBuffer,
pub gdn_params_buf: MlxBuffer,
pub q_split_buf: MlxBuffer,
pub k_split_buf: MlxBuffer,
pub v_split_buf: MlxBuffer,
pub x_norm_buf: MlxBuffer,
pub pre_norm_params_buf: MlxBuffer,
pub qkv_raw_buf: MlxBuffer,
pub z_buf: MlxBuffer,
pub alpha_logit_buf: MlxBuffer,
pub beta_logit_buf: MlxBuffer,
pub q_l2_buf: MlxBuffer,
pub k_normed_buf: MlxBuffer,
pub l2_params_q_buf: MlxBuffer,
pub l2_params_k_buf: MlxBuffer,
pub seq_capacity: u32,
pub hidden_size: u32,
pub n_k_heads: u32,
pub n_v_heads: u32,
pub d_k: u32,
pub d_v: u32,
}
impl DnPrefillArena {
#[allow(clippy::too_many_arguments)]
pub fn new(
device: &MlxDevice,
seq_capacity: u32,
hidden_size: u32,
n_k_heads: u32,
n_v_heads: u32,
d_k: u32,
d_v: u32,
) -> Result<Self> {
if seq_capacity == 0
|| hidden_size == 0
|| n_k_heads == 0
|| n_v_heads == 0
|| d_k == 0
|| d_v == 0
{
return Err(anyhow!(
"DnPrefillArena::new: zero dim \
seq_capacity={} hidden_size={} n_k_heads={} n_v_heads={} d_k={} d_v={}",
seq_capacity,
hidden_size,
n_k_heads,
n_v_heads,
d_k,
d_v,
));
}
let seq = seq_capacity as usize;
let h = hidden_size as usize;
let nk = n_k_heads as usize;
let nv = n_v_heads as usize;
let dk = d_k as usize;
let dv = d_v as usize;
let qkv_channels = 2 * nk * dk + nv * dv;
let z_channels = nv * dv;
let q_sp = nk * dk;
let k_sp = nk * dk;
let v_sp = nv * dv;
let n_q_elems = seq * nk * dk;
let g_n = seq * nv;
let gated_elems = seq * nv * dv;
let attn_out_elems = nv * seq * dv;
let qkv_conv_buf = device
.alloc_buffer(seq * qkv_channels * 4, DType::F32, vec![seq, qkv_channels])
.map_err(|e| anyhow!("DnPrefillArena alloc qkv_conv_buf: {e}"))?;
let ssm_params_buf = device
.alloc_buffer(4 * 4, DType::U32, vec![4])
.map_err(|e| anyhow!("DnPrefillArena alloc ssm_params_buf: {e}"))?;
let q_scaled_buf = device
.alloc_buffer(n_q_elems * 4, DType::F32, vec![n_q_elems])
.map_err(|e| anyhow!("DnPrefillArena alloc q_scaled_buf: {e}"))?;
let g_buf = device
.alloc_buffer(g_n * 4, DType::F32, vec![g_n])
.map_err(|e| anyhow!("DnPrefillArena alloc g_buf: {e}"))?;
let beta_buf = device
.alloc_buffer(g_n * 4, DType::F32, vec![g_n])
.map_err(|e| anyhow!("DnPrefillArena alloc beta_buf: {e}"))?;
let g_params_buf = device
.alloc_buffer(8, DType::U32, vec![2])
.map_err(|e| anyhow!("DnPrefillArena alloc g_params_buf: {e}"))?;
let gated_buf = device
.alloc_buffer(gated_elems * 4, DType::F32, vec![gated_elems])
.map_err(|e| anyhow!("DnPrefillArena alloc gated_buf: {e}"))?;
let op8_params_buf = device
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow!("DnPrefillArena alloc op8_params_buf: {e}"))?;
let attn_out_buf = device
.alloc_buffer(attn_out_elems * 4, DType::F32, vec![attn_out_elems])
.map_err(|e| anyhow!("DnPrefillArena alloc attn_out_buf: {e}"))?;
let gdn_params_buf = device
.alloc_buffer(9 * 4, DType::U32, vec![9])
.map_err(|e| anyhow!("DnPrefillArena alloc gdn_params_buf: {e}"))?;
let q_split_buf = device
.alloc_buffer(seq * q_sp * 4, DType::F32, vec![seq, q_sp])
.map_err(|e| anyhow!("DnPrefillArena alloc q_split_buf: {e}"))?;
let k_split_buf = device
.alloc_buffer(seq * k_sp * 4, DType::F32, vec![seq, k_sp])
.map_err(|e| anyhow!("DnPrefillArena alloc k_split_buf: {e}"))?;
let v_split_buf = device
.alloc_buffer(seq * v_sp * 4, DType::F32, vec![seq, v_sp])
.map_err(|e| anyhow!("DnPrefillArena alloc v_split_buf: {e}"))?;
let x_norm_buf = device
.alloc_buffer(seq * h * 4, DType::F32, vec![seq, h])
.map_err(|e| anyhow!("DnPrefillArena alloc x_norm_buf: {e}"))?;
let pre_norm_params_buf = device
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow!("DnPrefillArena alloc pre_norm_params_buf: {e}"))?;
let qkv_raw_buf = device
.alloc_buffer(seq * qkv_channels * 4, DType::F32, vec![seq, qkv_channels])
.map_err(|e| anyhow!("DnPrefillArena alloc qkv_raw_buf: {e}"))?;
let z_buf = device
.alloc_buffer(seq * z_channels * 4, DType::F32, vec![seq, z_channels])
.map_err(|e| anyhow!("DnPrefillArena alloc z_buf: {e}"))?;
let alpha_logit_buf = device
.alloc_buffer(seq * nv * 4, DType::F32, vec![seq, nv])
.map_err(|e| anyhow!("DnPrefillArena alloc alpha_logit_buf: {e}"))?;
let beta_logit_buf = device
.alloc_buffer(seq * nv * 4, DType::F32, vec![seq, nv])
.map_err(|e| anyhow!("DnPrefillArena alloc beta_logit_buf: {e}"))?;
let q_l2_rows = seq * nk;
let q_l2_buf = device
.alloc_buffer(q_l2_rows * dk * 4, DType::F32, vec![q_l2_rows, dk])
.map_err(|e| anyhow!("DnPrefillArena alloc q_l2_buf: {e}"))?;
let k_normed_buf = device
.alloc_buffer(q_l2_rows * dk * 4, DType::F32, vec![q_l2_rows, dk])
.map_err(|e| anyhow!("DnPrefillArena alloc k_normed_buf: {e}"))?;
let l2_params_q_buf = device
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow!("DnPrefillArena alloc l2_params_q_buf: {e}"))?;
let l2_params_k_buf = device
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow!("DnPrefillArena alloc l2_params_k_buf: {e}"))?;
Ok(Self {
qkv_conv_buf,
ssm_params_buf,
q_scaled_buf,
g_buf,
beta_buf,
g_params_buf,
gated_buf,
op8_params_buf,
attn_out_buf,
gdn_params_buf,
q_split_buf,
k_split_buf,
v_split_buf,
x_norm_buf,
pre_norm_params_buf,
qkv_raw_buf,
z_buf,
alpha_logit_buf,
beta_logit_buf,
q_l2_buf,
k_normed_buf,
l2_params_q_buf,
l2_params_k_buf,
seq_capacity,
hidden_size,
n_k_heads,
n_v_heads,
d_k,
d_v,
})
}
pub fn validate_fits(
&self,
seq_len: u32,
hidden_size: u32,
n_k_heads: u32,
n_v_heads: u32,
d_k: u32,
d_v: u32,
) -> Result<()> {
if seq_len > self.seq_capacity {
return Err(anyhow!(
"DnPrefillArena::validate_fits: seq_len {} exceeds capacity {}",
seq_len,
self.seq_capacity
));
}
if hidden_size != self.hidden_size
|| n_k_heads != self.n_k_heads
|| n_v_heads != self.n_v_heads
|| d_k != self.d_k
|| d_v != self.d_v
{
return Err(anyhow!(
"DnPrefillArena::validate_fits: shape mismatch — \
arena (h={}, nk={}, nv={}, dk={}, dv={}) vs \
call (h={}, nk={}, nv={}, dk={}, dv={})",
self.hidden_size,
self.n_k_heads,
self.n_v_heads,
self.d_k,
self.d_v,
hidden_size,
n_k_heads,
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_dn_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_dn_arena_new_qwen36_35b_pp4096: skipping — no Metal device");
return;
}
};
let (seq, h, nk, nv, dk, dv) = (4096u32, 2048u32, 16u32, 32u32, 128u32, 128u32);
let arena = DnPrefillArena::new(&device, seq, h, nk, nv, dk, dv).expect("dn arena pp4096");
assert_eq!(arena.seq_capacity, seq);
assert_eq!(arena.hidden_size, h);
assert_eq!(arena.n_k_heads, nk);
assert_eq!(arena.n_v_heads, nv);
assert_eq!(arena.d_k, dk);
assert_eq!(arena.d_v, dv);
let qkv_channels = (2 * nk * dk + nv * dv) as usize;
let z_channels = (nv * dv) as usize;
assert_eq!(
arena.qkv_conv_buf.byte_len(),
(seq as usize) * qkv_channels * 4,
"qkv_conv_buf"
);
assert_eq!(arena.ssm_params_buf.byte_len(), 16, "ssm_params_buf");
assert_eq!(
arena.q_scaled_buf.byte_len(),
(seq as usize) * (nk as usize) * (dk as usize) * 4,
"q_scaled_buf"
);
assert_eq!(
arena.g_buf.byte_len(),
(seq as usize) * (nv as usize) * 4,
"g_buf"
);
assert_eq!(
arena.beta_buf.byte_len(),
(seq as usize) * (nv as usize) * 4,
"beta_buf"
);
assert_eq!(arena.g_params_buf.byte_len(), 8, "g_params_buf");
assert_eq!(
arena.gated_buf.byte_len(),
(seq as usize) * (nv as usize) * (dv as usize) * 4,
"gated_buf"
);
assert_eq!(arena.op8_params_buf.byte_len(), 8, "op8_params_buf");
assert_eq!(
arena.attn_out_buf.byte_len(),
(nv as usize) * (seq as usize) * (dv as usize) * 4,
"attn_out_buf"
);
assert_eq!(arena.gdn_params_buf.byte_len(), 36, "gdn_params_buf");
let q_sp = (nk * dk) as usize;
let k_sp = (nk * dk) as usize;
let v_sp = (nv * dv) as usize;
assert_eq!(
arena.q_split_buf.byte_len(),
(seq as usize) * q_sp * 4,
"q_split_buf"
);
assert_eq!(
arena.k_split_buf.byte_len(),
(seq as usize) * k_sp * 4,
"k_split_buf"
);
assert_eq!(
arena.v_split_buf.byte_len(),
(seq as usize) * v_sp * 4,
"v_split_buf"
);
assert_eq!(
arena.x_norm_buf.byte_len(),
(seq as usize) * (h as usize) * 4,
"x_norm_buf"
);
assert_eq!(
arena.pre_norm_params_buf.byte_len(),
8,
"pre_norm_params_buf"
);
assert_eq!(
arena.qkv_raw_buf.byte_len(),
(seq as usize) * qkv_channels * 4,
"qkv_raw_buf"
);
assert_eq!(
arena.z_buf.byte_len(),
(seq as usize) * z_channels * 4,
"z_buf"
);
assert_eq!(
arena.alpha_logit_buf.byte_len(),
(seq as usize) * (nv as usize) * 4,
"alpha_logit_buf"
);
assert_eq!(
arena.beta_logit_buf.byte_len(),
(seq as usize) * (nv as usize) * 4,
"beta_logit_buf"
);
let q_l2_rows = (seq as usize) * (nk as usize);
assert_eq!(
arena.q_l2_buf.byte_len(),
q_l2_rows * (dk as usize) * 4,
"q_l2_buf"
);
assert_eq!(
arena.k_normed_buf.byte_len(),
q_l2_rows * (dk as usize) * 4,
"k_normed_buf"
);
assert_eq!(arena.l2_params_q_buf.byte_len(), 8, "l2_params_q_buf");
assert_eq!(arena.l2_params_k_buf.byte_len(), 8, "l2_params_k_buf");
}
#[test]
fn test_dn_arena_new_small_shape() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_dn_arena_new_small_shape: skipping — no Metal device");
return;
}
};
let arena = DnPrefillArena::new(&device, 64, 128, 4, 8, 16, 16).expect("dn arena small");
assert_eq!(arena.seq_capacity, 64);
assert_eq!(arena.hidden_size, 128);
assert_eq!(arena.n_k_heads, 4);
assert_eq!(arena.n_v_heads, 8);
assert_eq!(arena.d_k, 16);
assert_eq!(arena.d_v, 16);
}
#[test]
fn test_dn_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_dn_arena_new_zero_dim_rejected: skipping — no Metal device");
return;
}
};
assert!(DnPrefillArena::new(&device, 0, 128, 4, 8, 16, 16).is_err());
assert!(DnPrefillArena::new(&device, 64, 0, 4, 8, 16, 16).is_err());
assert!(DnPrefillArena::new(&device, 64, 128, 0, 8, 16, 16).is_err());
assert!(DnPrefillArena::new(&device, 64, 128, 4, 0, 16, 16).is_err());
assert!(DnPrefillArena::new(&device, 64, 128, 4, 8, 0, 16).is_err());
assert!(DnPrefillArena::new(&device, 64, 128, 4, 8, 16, 0).is_err());
}
#[test]
fn test_dn_validate_fits_exact_match() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_dn_validate_fits_exact_match: skipping — no Metal device");
return;
}
};
let arena = DnPrefillArena::new(&device, 128, 256, 4, 8, 32, 32).expect("dn arena new");
assert!(arena.validate_fits(128, 256, 4, 8, 32, 32).is_ok());
assert!(arena.validate_fits(64, 256, 4, 8, 32, 32).is_ok());
}
#[test]
fn test_dn_validate_fits_seq_overrun() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_dn_validate_fits_seq_overrun: skipping — no Metal device");
return;
}
};
let arena = DnPrefillArena::new(&device, 128, 256, 4, 8, 32, 32).expect("dn arena new");
assert!(arena.validate_fits(256, 256, 4, 8, 32, 32).is_err());
}
#[test]
fn test_dn_validate_fits_shape_mismatch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_dn_validate_fits_shape_mismatch: skipping — no Metal device");
return;
}
};
let arena = DnPrefillArena::new(&device, 128, 256, 4, 8, 32, 32).expect("dn arena new");
assert!(arena.validate_fits(128, 128, 4, 8, 32, 32).is_err()); assert!(arena.validate_fits(128, 256, 8, 8, 32, 32).is_err()); assert!(arena.validate_fits(128, 256, 4, 4, 32, 32).is_err()); assert!(arena.validate_fits(128, 256, 4, 8, 16, 32).is_err()); assert!(arena.validate_fits(128, 256, 4, 8, 32, 16).is_err()); }
#[test]
fn test_dn_arena_buffers_zero_initialised() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_dn_arena_buffers_zero_initialised: skipping — no Metal device");
return;
}
};
let arena = DnPrefillArena::new(&device, 64, 128, 4, 8, 16, 16).expect("dn arena new");
let f32_bufs: [(&MlxBuffer, &str); 14] = [
(&arena.qkv_conv_buf, "qkv_conv_buf"),
(&arena.q_scaled_buf, "q_scaled_buf"),
(&arena.g_buf, "g_buf"),
(&arena.beta_buf, "beta_buf"),
(&arena.gated_buf, "gated_buf"),
(&arena.attn_out_buf, "attn_out_buf"),
(&arena.q_split_buf, "q_split_buf"),
(&arena.k_split_buf, "k_split_buf"),
(&arena.v_split_buf, "v_split_buf"),
(&arena.x_norm_buf, "x_norm_buf"),
(&arena.qkv_raw_buf, "qkv_raw_buf"),
(&arena.z_buf, "z_buf"),
(&arena.alpha_logit_buf, "alpha_logit_buf"),
(&arena.beta_logit_buf, "beta_logit_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 f32_params: [(&MlxBuffer, &str); 4] = [
(&arena.op8_params_buf, "op8_params_buf"),
(&arena.pre_norm_params_buf, "pre_norm_params_buf"),
(&arena.l2_params_q_buf, "l2_params_q_buf"),
(&arena.l2_params_k_buf, "l2_params_k_buf"),
];
for (buf, name) in &f32_params {
let slice = buf
.as_slice::<f32>()
.unwrap_or_else(|e| panic!("{name} as_slice::<f32> failed: {e}"));
assert_eq!(slice[0], 0.0f32, "{name}[0] should be zero");
}
let u32_params: [(&MlxBuffer, &str); 3] = [
(&arena.ssm_params_buf, "ssm_params_buf"),
(&arena.g_params_buf, "g_params_buf"),
(&arena.gdn_params_buf, "gdn_params_buf"),
];
for (buf, name) in &u32_params {
let slice = buf
.as_slice::<u32>()
.unwrap_or_else(|e| panic!("{name} as_slice::<u32> failed: {e}"));
assert_eq!(slice[0], 0u32, "{name}[0] should be zero");
}
}
#[test]
fn test_dn_arena_q_k_split_buf_same_size() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_dn_arena_q_k_split_buf_same_size: skipping — no Metal device");
return;
}
};
let arena = DnPrefillArena::new(&device, 128, 256, 4, 8, 32, 32).expect("dn arena new");
assert_eq!(
arena.q_split_buf.byte_len(),
arena.k_split_buf.byte_len(),
"q_split and k_split must be the same size (q_sp == k_sp)"
);
}
}