use anyhow::{anyhow, Result};
use mlx_native::{DType, MlxBuffer, MlxDevice};
pub struct FaProjectionsArena {
pub x_norm_buf: MlxBuffer,
pub pre_norm_params_buf: MlxBuffer,
pub q_proj_buf: MlxBuffer,
pub k_proj_buf: MlxBuffer,
pub v_proj_buf: MlxBuffer,
pub gate_proj_buf: MlxBuffer,
pub q_normed_buf: MlxBuffer,
pub k_normed_buf: MlxBuffer,
pub qk_rms_params_buf: MlxBuffer,
pub q_rope_buf: MlxBuffer,
pub k_rope_buf: MlxBuffer,
pub gated_buf: MlxBuffer,
pub sigmoid_params_buf: MlxBuffer,
pub seq_capacity: u32,
pub hidden_size: u32,
pub n_head: u32,
pub n_kv: u32,
pub head_dim: u32,
}
impl FaProjectionsArena {
pub fn new(
device: &MlxDevice,
seq_capacity: u32,
hidden_size: u32,
n_head: u32,
n_kv: u32,
head_dim: u32,
rms_norm_eps: f32,
) -> Result<Self> {
if seq_capacity == 0 || hidden_size == 0 || n_head == 0 || n_kv == 0 || head_dim == 0 {
return Err(anyhow!(
"FaProjectionsArena::new: zero dim \
seq_capacity={} hidden_size={} n_head={} n_kv={} head_dim={}",
seq_capacity,
hidden_size,
n_head,
n_kv,
head_dim,
));
}
if n_head % n_kv != 0 {
return Err(anyhow!(
"FaProjectionsArena::new: n_head ({}) must be divisible by n_kv ({})",
n_head,
n_kv,
));
}
let seq = seq_capacity as usize;
let h = hidden_size as usize;
let nh = n_head as usize;
let nkv = n_kv as usize;
let d = head_dim as usize;
let q_total = nh * d;
let kv_total = nkv * d;
let x_norm_buf = device
.alloc_buffer(seq * h * 4, DType::F32, vec![seq, h])
.map_err(|e| anyhow!("FaProjectionsArena alloc x_norm_buf: {e}"))?;
let q_proj_buf = device
.alloc_buffer(seq * q_total * 4, DType::F32, vec![seq, q_total])
.map_err(|e| anyhow!("FaProjectionsArena alloc q_proj_buf: {e}"))?;
let k_proj_buf = device
.alloc_buffer(seq * kv_total * 4, DType::F32, vec![seq, kv_total])
.map_err(|e| anyhow!("FaProjectionsArena alloc k_proj_buf: {e}"))?;
let v_proj_buf = device
.alloc_buffer(seq * kv_total * 4, DType::F32, vec![seq, kv_total])
.map_err(|e| anyhow!("FaProjectionsArena alloc v_proj_buf: {e}"))?;
let gate_proj_buf = device
.alloc_buffer(seq * q_total * 4, DType::F32, vec![seq, q_total])
.map_err(|e| anyhow!("FaProjectionsArena alloc gate_proj_buf: {e}"))?;
let q_normed_buf = device
.alloc_buffer(seq * q_total * 4, DType::F32, vec![seq * nh, d])
.map_err(|e| anyhow!("FaProjectionsArena alloc q_normed_buf: {e}"))?;
let k_normed_buf = device
.alloc_buffer(seq * kv_total * 4, DType::F32, vec![seq * nkv, d])
.map_err(|e| anyhow!("FaProjectionsArena alloc k_normed_buf: {e}"))?;
let q_rope_buf = device
.alloc_buffer(seq * q_total * 4, DType::F32, vec![seq, nh, d])
.map_err(|e| anyhow!("FaProjectionsArena alloc q_rope_buf: {e}"))?;
let k_rope_buf = device
.alloc_buffer(seq * kv_total * 4, DType::F32, vec![seq, nkv, d])
.map_err(|e| anyhow!("FaProjectionsArena alloc k_rope_buf: {e}"))?;
let gated_buf = device
.alloc_buffer(seq * q_total * 4, DType::F32, vec![seq * q_total])
.map_err(|e| anyhow!("FaProjectionsArena alloc gated_buf: {e}"))?;
let mut pre_norm_params_buf = device
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow!("FaProjectionsArena alloc pre_norm_params_buf: {e}"))?;
{
let s = pre_norm_params_buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("pre_norm_params_buf as_mut_slice: {e}"))?;
s[0] = rms_norm_eps;
s[1] = hidden_size as f32;
}
let mut qk_rms_params_buf = device
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow!("FaProjectionsArena alloc qk_rms_params_buf: {e}"))?;
{
let s = qk_rms_params_buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("qk_rms_params_buf as_mut_slice: {e}"))?;
s[0] = rms_norm_eps;
s[1] = head_dim as f32;
}
let mut sigmoid_params_buf = device
.alloc_buffer(4, DType::U32, vec![1])
.map_err(|e| anyhow!("FaProjectionsArena alloc sigmoid_params_buf: {e}"))?;
{
let s = sigmoid_params_buf
.as_mut_slice::<u32>()
.map_err(|e| anyhow!("sigmoid_params_buf as_mut_slice: {e}"))?;
s[0] = (seq_capacity * (n_head * head_dim)) as u32;
}
Ok(Self {
x_norm_buf,
pre_norm_params_buf,
q_proj_buf,
k_proj_buf,
v_proj_buf,
gate_proj_buf,
q_normed_buf,
k_normed_buf,
qk_rms_params_buf,
q_rope_buf,
k_rope_buf,
gated_buf,
sigmoid_params_buf,
seq_capacity,
hidden_size,
n_head,
n_kv,
head_dim,
})
}
pub fn validate_fits(
&self,
seq_len: u32,
hidden_size: u32,
n_head: u32,
n_kv: u32,
head_dim: u32,
) -> Result<()> {
if seq_len > self.seq_capacity {
return Err(anyhow!(
"FaProjectionsArena::validate_fits: seq_len {} exceeds capacity {}",
seq_len,
self.seq_capacity
));
}
if hidden_size != self.hidden_size
|| n_head != self.n_head
|| n_kv != self.n_kv
|| head_dim != self.head_dim
{
return Err(anyhow!(
"FaProjectionsArena::validate_fits: shape mismatch — \
arena (h={}, n_head={}, n_kv={}, head_dim={}) vs \
call (h={}, n_head={}, n_kv={}, head_dim={})",
self.hidden_size,
self.n_head,
self.n_kv,
self.head_dim,
hidden_size,
n_head,
n_kv,
head_dim,
));
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn device_or_skip() -> Option<MlxDevice> {
MlxDevice::new().ok()
}
#[test]
fn test_fa_proj_arena_new_qwen36_pp4127() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_fa_proj_arena_new_qwen36_pp4127: skipping — no Metal device");
return;
}
};
let (seq, h, nh, nkv, d, eps) = (4127u32, 2048u32, 16u32, 2u32, 256u32, 1e-6f32);
let arena = FaProjectionsArena::new(&device, seq, h, nh, nkv, d, eps)
.expect("fa proj arena new pp4127");
assert_eq!(arena.seq_capacity, seq);
assert_eq!(arena.hidden_size, h);
assert_eq!(arena.n_head, nh);
assert_eq!(arena.n_kv, nkv);
assert_eq!(arena.head_dim, d);
let q_bytes = (seq as usize) * (nh as usize) * (d as usize) * 4;
let kv_bytes = (seq as usize) * (nkv as usize) * (d as usize) * 4;
let h_bytes = (seq as usize) * (h as usize) * 4;
assert_eq!(arena.x_norm_buf.byte_len(), h_bytes, "x_norm_buf");
assert_eq!(arena.q_proj_buf.byte_len(), q_bytes, "q_proj_buf");
assert_eq!(arena.k_proj_buf.byte_len(), kv_bytes, "k_proj_buf");
assert_eq!(arena.v_proj_buf.byte_len(), kv_bytes, "v_proj_buf");
assert_eq!(arena.gate_proj_buf.byte_len(), q_bytes, "gate_proj_buf");
assert_eq!(arena.q_normed_buf.byte_len(), q_bytes, "q_normed_buf");
assert_eq!(arena.k_normed_buf.byte_len(), kv_bytes, "k_normed_buf");
assert_eq!(arena.q_rope_buf.byte_len(), q_bytes, "q_rope_buf");
assert_eq!(arena.k_rope_buf.byte_len(), kv_bytes, "k_rope_buf");
assert_eq!(arena.gated_buf.byte_len(), q_bytes, "gated_buf");
assert_eq!(arena.pre_norm_params_buf.byte_len(), 8, "pre_norm_params");
assert_eq!(arena.qk_rms_params_buf.byte_len(), 8, "qk_rms_params");
assert_eq!(arena.sigmoid_params_buf.byte_len(), 4, "sigmoid_params");
let pn = arena
.pre_norm_params_buf
.as_slice::<f32>()
.expect("pre_norm_params_buf as_slice");
assert_eq!(pn[0], eps, "pre_norm_params[0] = eps");
assert_eq!(pn[1], h as f32, "pre_norm_params[1] = hidden_size");
let qk = arena
.qk_rms_params_buf
.as_slice::<f32>()
.expect("qk_rms_params_buf as_slice");
assert_eq!(qk[0], eps, "qk_rms_params[0] = eps");
assert_eq!(qk[1], d as f32, "qk_rms_params[1] = head_dim");
let sg = arena
.sigmoid_params_buf
.as_slice::<u32>()
.expect("sigmoid_params_buf as_slice");
assert_eq!(
sg[0],
seq * nh * d,
"sigmoid_params[0] = seq*n_head*head_dim"
);
}
#[test]
fn test_fa_proj_arena_new_small_shape() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_fa_proj_arena_new_small_shape: skipping — no Metal device");
return;
}
};
let arena =
FaProjectionsArena::new(&device, 64, 128, 4, 2, 32, 1e-5).expect("fa proj arena small");
assert_eq!(arena.seq_capacity, 64);
}
#[test]
fn test_fa_proj_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_fa_proj_arena_new_zero_dim_rejected: skipping — no Metal device");
return;
}
};
assert!(FaProjectionsArena::new(&device, 0, 128, 4, 2, 32, 1e-5).is_err());
assert!(FaProjectionsArena::new(&device, 64, 0, 4, 2, 32, 1e-5).is_err());
assert!(FaProjectionsArena::new(&device, 64, 128, 0, 2, 32, 1e-5).is_err());
assert!(FaProjectionsArena::new(&device, 64, 128, 4, 0, 32, 1e-5).is_err());
assert!(FaProjectionsArena::new(&device, 64, 128, 4, 2, 0, 1e-5).is_err());
}
#[test]
fn test_fa_proj_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_fa_proj_arena_new_gqa_divisibility_rejected: skipping — no Metal device"
);
return;
}
};
assert!(FaProjectionsArena::new(&device, 64, 128, 15, 2, 32, 1e-5).is_err());
}
#[test]
fn test_fa_proj_validate_fits_exact_match() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_fa_proj_validate_fits_exact_match: skipping — no Metal device");
return;
}
};
let arena = FaProjectionsArena::new(&device, 128, 256, 4, 2, 32, 1e-5).expect("arena");
assert!(arena.validate_fits(128, 256, 4, 2, 32).is_ok());
assert!(arena.validate_fits(64, 256, 4, 2, 32).is_ok());
}
#[test]
fn test_fa_proj_validate_fits_seq_overrun() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_fa_proj_validate_fits_seq_overrun: skipping — no Metal device");
return;
}
};
let arena = FaProjectionsArena::new(&device, 128, 256, 4, 2, 32, 1e-5).expect("arena");
assert!(arena.validate_fits(256, 256, 4, 2, 32).is_err());
}
#[test]
fn test_fa_proj_validate_fits_shape_mismatch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!("test_fa_proj_validate_fits_shape_mismatch: skipping — no Metal device");
return;
}
};
let arena = FaProjectionsArena::new(&device, 128, 256, 4, 2, 32, 1e-5).expect("arena");
assert!(arena.validate_fits(128, 128, 4, 2, 32).is_err()); assert!(arena.validate_fits(128, 256, 8, 2, 32).is_err()); assert!(arena.validate_fits(128, 256, 4, 1, 32).is_err()); assert!(arena.validate_fits(128, 256, 4, 2, 16).is_err()); }
#[test]
fn test_fa_proj_arena_buffers_zero_initialized() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match device_or_skip() {
Some(d) => d,
None => {
eprintln!(
"test_fa_proj_arena_buffers_zero_initialized: skipping — no Metal device"
);
return;
}
};
let arena = FaProjectionsArena::new(&device, 64, 128, 4, 2, 32, 1e-5).expect("arena");
let f32_outs: [(&MlxBuffer, &str); 10] = [
(&arena.x_norm_buf, "x_norm_buf"),
(&arena.q_proj_buf, "q_proj_buf"),
(&arena.k_proj_buf, "k_proj_buf"),
(&arena.v_proj_buf, "v_proj_buf"),
(&arena.gate_proj_buf, "gate_proj_buf"),
(&arena.q_normed_buf, "q_normed_buf"),
(&arena.k_normed_buf, "k_normed_buf"),
(&arena.q_rope_buf, "q_rope_buf"),
(&arena.k_rope_buf, "k_rope_buf"),
(&arena.gated_buf, "gated_buf"),
];
for (buf, name) in &f32_outs {
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)"
);
}
}
}
}