use anyhow::{anyhow, Context, Result};
use mlx_native::ops::dense_gemv_bf16::dense_gemv_bf16_f32;
use mlx_native::ops::dense_mm_bf16::{dense_matmul_bf16_f32_tensor, DenseMmBf16F32Params};
use mlx_native::ops::elementwise::{cast, elementwise_add, CastDirection};
use mlx_native::ops::flash_attn_prefill::{
dispatch_flash_attn_prefill_bf16_d256, dispatch_flash_attn_prefill_bf16_d256_resume,
FlashAttnPrefillParams, FlashAttnPrefillResumeParams,
};
use mlx_native::ops::flash_attn_vec::{
flash_attn_vec, tmp_buffer_bytes as flash_attn_vec_tmp_bytes,
tmp_buffer_bytes_with_qL as flash_attn_vec_tmp_bytes_with_qL, FlashAttnVecParams,
};
use mlx_native::ops::kv_cache_copy::dispatch_kv_cache_copy_seq_f32_dual;
use mlx_native::ops::quantized_matmul_ggml::{
quantized_matmul_ggml, GgmlQuantizedMatmulParams, GgmlType,
};
use mlx_native::ops::rms_norm;
use mlx_native::ops::rope_multi::{dispatch_rope_multi_cached, RopeMultiMode, RopeMultiParams};
use mlx_native::ops::sdpa::{sdpa, SdpaParams};
use mlx_native::ops::sdpa_decode::dispatch_sdpa_decode;
use mlx_native::ops::sigmoid_mul::dispatch_sigmoid_mul;
use mlx_native::ops::silu_mul::dispatch_silu_mul;
use mlx_native::ops::transpose::{permute_021_bf16, permute_021_bf16_to_f32, permute_021_f32};
use mlx_native::ops::tree_attention::{self as tree_attn_ops, TreeAttentionParams};
use mlx_native::{DType, KernelRegistry, MlxBuffer, MlxDevice};
use super::encoder_stage::LayerEncoder;
use super::full_attn::FullAttnLayerWeights;
use super::kv_cache::FullAttnKvSlot;
use crate::serve::multi_seq_kv::SlotId;
#[inline]
fn slot_k_v_region_for_full_attn(
slot_id: SlotId,
n_kv_heads: u32,
max_seq_len: u32,
head_dim: u32,
) -> (u64, usize) {
let n_elements = (n_kv_heads as usize) * (max_seq_len as usize) * (head_dim as usize);
let byte_offset = (slot_id.0 as u64)
.checked_mul(n_elements as u64)
.and_then(|e| e.checked_mul(std::mem::size_of::<f32>() as u64))
.expect("slot K/V byte offset overflow (slot_id * n_kv * max_seq * head_dim * 4)");
(byte_offset, n_elements)
}
fn write_kv_with_optional_tq_encode(
enc: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
k_seq_major: &MlxBuffer,
v_seq_major: &MlxBuffer,
slot: &mut FullAttnKvSlot,
n_kv_heads: u32,
head_dim: u32,
max_seq_len: u32,
cur_len: u32,
n_tokens: u32,
slot_id: SlotId,
) -> Result<()> {
if let (Some(dst_k), Some(dst_v)) = (slot.k.as_ref(), slot.v.as_ref()) {
let (byte_offset, n_elements) =
slot_k_v_region_for_full_attn(slot_id, n_kv_heads, max_seq_len, head_dim);
let dst_k_view = dst_k.slice_view(byte_offset, n_elements);
let dst_v_view = dst_v.slice_view(byte_offset, n_elements);
dispatch_kv_cache_copy_seq_f32_dual(
enc,
registry,
device.metal_device(),
k_seq_major,
v_seq_major,
&dst_k_view,
&dst_v_view,
n_kv_heads,
head_dim,
max_seq_len,
cur_len,
n_tokens,
0,
)
.context("kv_cache_copy_seq_f32_dual (write_kv_with_optional_tq_encode)")?;
}
if slot.tq.is_some() && n_tokens > 0 && (head_dim == 256 || head_dim == 512) {
enc.memory_barrier();
let codebook_bits = crate::debug::INVESTIGATION_ENV.tq_codebook_bits;
let cb_bits = if matches!(codebook_bits, 5 | 6 | 8) {
codebook_bits
} else {
8
};
slot.encode_seq_tokens_to_tq_for_slot(
k_seq_major,
true,
n_tokens,
n_kv_heads,
head_dim,
max_seq_len,
cur_len,
0,
false,
1.0,
cb_bits,
slot_id,
enc,
registry,
device,
)
.context("TQ encode K (write_kv_with_optional_tq_encode)")?;
slot.encode_seq_tokens_to_tq_for_slot(
v_seq_major,
false,
n_tokens,
n_kv_heads,
head_dim,
max_seq_len,
cur_len,
0,
false,
1.0,
cb_bits,
slot_id,
enc,
registry,
device,
)
.context("TQ encode V (write_kv_with_optional_tq_encode)")?;
}
Ok(())
}
fn dispatch_decode_sdpa_with_optional_tq(
enc: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
q_seq_major: &MlxBuffer,
slot: &FullAttnKvSlot,
out_buf: &MlxBuffer,
fa_tmp: &MlxBuffer,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
kv_seq_len: u32,
max_seq_len: u32,
slot_id: SlotId,
) -> Result<()> {
let scale = 1.0_f32 / (head_dim as f32).sqrt();
if slot.tq.is_some() && (head_dim == 256 || head_dim == 512) {
let codebook_bits = crate::debug::INVESTIGATION_ENV.tq_codebook_bits;
let cb_bits = if matches!(codebook_bits, 5 | 6 | 8) {
codebook_bits
} else {
8
};
mlx_native::ops::fwht_standalone::dispatch_fwht_sign_premult_f32(
enc,
registry,
device.metal_device(),
q_seq_major,
n_heads,
head_dim,
)
.context("dispatch_fwht_sign_premult_f32 (TQ decode pre-rotation)")?;
enc.memory_barrier();
let tq_params = super::kv_cache::Qwen35TqSdpaParams {
num_heads: n_heads,
num_kv_heads: n_kv_heads,
head_dim,
kv_seq_len,
kv_capacity: max_seq_len,
scale,
mask_type: 0, sliding_window: 0,
softcap: 0.0,
ring_start: 0,
scale_factor_d512: 1.0,
codebook_bits: cb_bits,
};
slot.dispatch_tq_sdpa_for_slot(
q_seq_major,
out_buf,
fa_tmp,
&tq_params,
slot_id,
enc,
registry,
device,
)
.context("dispatch_tq_sdpa (TQ decode SDPA)")?;
enc.memory_barrier();
mlx_native::ops::fwht_standalone::dispatch_fwht_sign_undo_f32(
enc,
registry,
device.metal_device(),
out_buf,
n_heads,
head_dim,
)
.context("dispatch_fwht_sign_undo_f32 (TQ decode post-rotation)")?;
} else {
let fa_params = FlashAttnVecParams {
num_heads: n_heads,
num_kv_heads: n_kv_heads,
head_dim,
kv_seq_len,
kv_capacity: max_seq_len,
scale,
mask_type: 0,
sliding_window: 0,
softcap: 0.0,
q_seq_len: FlashAttnVecParams::DEFAULT_Q_SEQ_LEN,
};
let kbuf = slot.k.as_ref().expect(
"flash_attn_vec F32 fallback: slot.k is None but TQ branch \
not taken — iter-34 alloc/SDPA gating invariant regressed \
(tq_kv_active=true ⇒ slot.tq=Some ⇒ TQ chain above runs).",
);
let vbuf = slot
.v
.as_ref()
.expect("flash_attn_vec F32: slot.v is None (see slot.k)");
let (byte_offset, n_elements) =
slot_k_v_region_for_full_attn(slot_id, n_kv_heads, max_seq_len, head_dim);
let kbuf_view = kbuf.slice_view(byte_offset, n_elements);
let vbuf_view = vbuf.slice_view(byte_offset, n_elements);
flash_attn_vec(
enc,
registry,
device,
q_seq_major,
&kbuf_view,
&vbuf_view,
out_buf,
fa_tmp,
&fa_params,
)
.context("flash_attn_vec (legacy F32 decode)")?;
}
Ok(())
}
pub struct FullAttnWeightsGpu {
pub attn_norm: MlxBuffer,
pub post_attn_norm: MlxBuffer,
pub wq: MlxBuffer,
pub wk: MlxBuffer,
pub wv: MlxBuffer,
pub w_gate: MlxBuffer,
pub attn_q_norm: MlxBuffer,
pub attn_k_norm: MlxBuffer,
pub wo: MlxBuffer,
}
impl FullAttnWeightsGpu {
pub fn from_cpu(weights: &FullAttnLayerWeights, device: &MlxDevice) -> Result<Self> {
Ok(Self {
attn_norm: upload_f32_weight(&weights.attn_norm, device)?,
post_attn_norm: upload_f32_weight(&weights.post_attn_norm, device)?,
wq: upload_q4_0_from_f32(&weights.wq, device)?,
wk: upload_q4_0_from_f32(&weights.wk, device)?,
wv: upload_q4_0_from_f32(&weights.wv, device)?,
w_gate: upload_q4_0_from_f32(&weights.w_gate, device)?,
attn_q_norm: upload_f32_weight(&weights.attn_q_norm, device)?,
attn_k_norm: upload_f32_weight(&weights.attn_k_norm, device)?,
wo: upload_q4_0_from_f32(&weights.wo, device)?,
})
}
#[cfg(test)]
pub fn from_cpu_f32(weights: &FullAttnLayerWeights, device: &MlxDevice) -> Result<Self> {
Ok(Self {
attn_norm: upload_f32(&weights.attn_norm, device)?,
post_attn_norm: upload_f32(&weights.post_attn_norm, device)?,
wq: upload_f32(&weights.wq, device)?,
wk: upload_f32(&weights.wk, device)?,
wv: upload_f32(&weights.wv, device)?,
w_gate: upload_f32(&weights.w_gate, device)?,
attn_q_norm: upload_f32(&weights.attn_q_norm, device)?,
attn_k_norm: upload_f32(&weights.attn_k_norm, device)?,
wo: upload_f32(&weights.wo, device)?,
})
}
}
#[inline(always)]
fn f32_to_bf16_rne(v: f32) -> u16 {
let bits = v.to_bits();
if (bits & 0x7FFF_FFFF) > 0x7F80_0000 {
return ((bits >> 16) | 0x0040) as u16; }
let rounding_bias = 0x7FFF_u32 + ((bits >> 16) & 1);
((bits + rounding_bias) >> 16) as u16
}
pub fn upload_bf16_from_f32(data: &[f32], device: &MlxDevice) -> Result<MlxBuffer> {
let n = data.len();
let byte_len = n * 2; let mut buf = device
.alloc_buffer(byte_len, DType::BF16, vec![n])
.map_err(|e| anyhow!("alloc bf16 buffer len={n}: {e}"))?;
{
let slice = buf
.as_mut_slice::<u16>()
.map_err(|e| anyhow!("mut_slice bf16: {e}"))?;
for (i, &v) in data.iter().enumerate() {
slice[i] = f32_to_bf16_rne(v);
}
}
super::weight_pool::register_weight_buffer(device, &buf)
.map_err(|e| anyhow!("register_weight_buffer bf16 len={n}: {e}"))?;
Ok(buf)
}
pub fn encode_q4_0_blocks(vals: &[f32]) -> Vec<u8> {
use half::f16;
const QK: usize = 32;
let n = vals.len();
assert_eq!(
n % QK,
0,
"encode_q4_0_blocks: n={n} must be divisible by QK=32"
);
let n_blocks = n / QK;
let mut out = vec![0u8; n_blocks * 18];
for b in 0..n_blocks {
let block = &vals[b * QK..(b + 1) * QK];
let amax = block.iter().cloned().map(f32::abs).fold(0.0f32, f32::max);
let d = if amax > 0.0 { amax / 7.0 } else { 1.0 };
let d_f16 = f16::from_f32(d);
let off = b * 18;
out[off..off + 2].copy_from_slice(&d_f16.to_le_bytes());
for j in 0..16 {
let q0 = ((block[j] / d).round().clamp(-8.0, 7.0) as i8 + 8) as u8;
let q1 = ((block[j + 16] / d).round().clamp(-8.0, 7.0) as i8 + 8) as u8;
out[off + 2 + j] = (q0 & 0x0F) | ((q1 & 0x0F) << 4);
}
}
out
}
pub fn upload_q4_0_from_f32(data: &[f32], device: &MlxDevice) -> Result<MlxBuffer> {
let blocks = encode_q4_0_blocks(data);
let byte_len = blocks.len();
let mut buf = device
.alloc_buffer(byte_len, DType::U8, vec![byte_len])
.map_err(|e| anyhow!("alloc q4_0 buffer len={byte_len}: {e}"))?;
{
let slice = buf
.as_mut_slice::<u8>()
.map_err(|e| anyhow!("mut_slice q4_0: {e}"))?;
slice.copy_from_slice(&blocks);
}
super::weight_pool::register_weight_buffer(device, &buf)
.map_err(|e| anyhow!("register_weight_buffer q4_0 len={byte_len}: {e}"))?;
Ok(buf)
}
pub fn upload_f32(data: &[f32], device: &MlxDevice) -> Result<MlxBuffer> {
let byte_len = data.len() * 4;
let mut buf = device
.alloc_buffer(byte_len, DType::F32, vec![data.len()])
.map_err(|e| anyhow!("alloc f32 buffer len={}: {e}", data.len()))?;
{
let slice = buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("mut_slice: {e}"))?;
slice.copy_from_slice(data);
}
Ok(buf)
}
pub fn upload_f32_weight(data: &[f32], device: &MlxDevice) -> Result<MlxBuffer> {
let buf = upload_f32(data, device)?;
super::weight_pool::register_weight_buffer(device, &buf)
.map_err(|e| anyhow!("register_weight_buffer f32 len={}: {e}", data.len()))?;
Ok(buf)
}
pub fn upload_f32_into(data: &[f32], buf: &mut MlxBuffer) -> Result<()> {
anyhow::ensure!(
buf.dtype() == DType::F32,
"upload_f32_into: expected F32 buffer, got {:?}",
buf.dtype()
);
anyhow::ensure!(
buf.element_count() >= data.len(),
"upload_f32_into: buf too small (cap={} < data={})",
buf.element_count(),
data.len()
);
let slice = buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("mut_slice: {e}"))?;
slice[..data.len()].copy_from_slice(data);
Ok(())
}
pub fn download_f32(buf: &MlxBuffer) -> Result<Vec<f32>> {
if buf.dtype() != DType::F32 {
return Err(anyhow!("download_f32: buffer dtype {} != f32", buf.dtype()));
}
let slice: &[f32] = buf.as_slice().map_err(|e| anyhow!("as_slice: {e}"))?;
Ok(slice.to_vec())
}
pub fn apply_q_or_k_per_head_rms_norm(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
norm_weight: &MlxBuffer,
seq_len: u32,
n_heads: u32,
head_dim: u32,
eps: f32,
) -> Result<MlxBuffer> {
let rows = seq_len * n_heads;
let dim = head_dim;
let out = super::decode_pool::pooled_alloc_buffer(
device,
(rows * dim) as usize * 4,
DType::F32,
vec![rows as usize, dim as usize],
)
.map_err(|e| anyhow!("alloc out: {e}"))?;
let mut params = super::decode_pool::pooled_alloc_buffer(device, 8, DType::F32, vec![2])
.map_err(|e| anyhow!("alloc params: {e}"))?;
{
let s = params
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("mut_slice: {e}"))?;
s[0] = eps;
s[1] = dim as f32;
}
rms_norm::dispatch_rms_norm(
encoder,
registry,
device.metal_device(),
input,
norm_weight,
&out,
¶ms,
rows,
dim,
)
.context("dispatch_rms_norm per-head")?;
Ok(out)
}
pub fn apply_imrope(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
positions: &MlxBuffer,
seq_len: u32,
n_heads: u32,
head_dim: u32,
rotary_dim: u32,
freq_base: f32,
mrope_section: [u32; 4],
) -> Result<MlxBuffer> {
let params = RopeMultiParams {
head_dim,
rope_dim: rotary_dim,
n_heads,
seq_len,
freq_base,
mode: RopeMultiMode::Imrope,
sections: mrope_section,
};
let out = super::decode_pool::pooled_alloc_buffer(
device,
(seq_len * n_heads * head_dim) as usize * 4,
DType::F32,
vec![seq_len as usize, n_heads as usize, head_dim as usize],
)
.map_err(|e| anyhow!("alloc imrope out (pooled): {e}"))?;
dispatch_rope_multi_cached(encoder, registry, device, input, &out, positions, params)
.context("dispatch_rope_multi_cached")?;
Ok(out)
}
pub fn apply_sigmoid_gate_multiply(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
attn_out: &MlxBuffer,
gate: &MlxBuffer,
n_elements: u32,
) -> Result<MlxBuffer> {
let out = super::decode_pool::pooled_alloc_buffer(
device,
n_elements as usize * 4,
DType::F32,
vec![n_elements as usize],
)
.map_err(|e| anyhow!("alloc sigmoid-mul out: {e}"))?;
let mut params = super::decode_pool::pooled_alloc_buffer(device, 4, DType::U32, vec![1])
.map_err(|e| anyhow!("alloc params: {e}"))?;
params
.as_mut_slice::<u32>()
.map_err(|e| anyhow!("mut_slice: {e}"))?[0] = n_elements;
dispatch_sigmoid_mul(
encoder,
registry,
device.metal_device(),
attn_out,
gate,
&out,
¶ms,
n_elements,
)
.context("dispatch_sigmoid_mul")?;
Ok(out)
}
pub fn apply_pre_attn_rms_norm(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weights_gpu: &FullAttnWeightsGpu,
seq_len: u32,
hidden_size: u32,
eps: f32,
) -> Result<MlxBuffer> {
let out = super::decode_pool::pooled_alloc_buffer(
device,
(seq_len * hidden_size) as usize * 4,
DType::F32,
vec![seq_len as usize, hidden_size as usize],
)
.map_err(|e| anyhow!("alloc out: {e}"))?;
let mut params = super::decode_pool::pooled_alloc_buffer(device, 8, DType::F32, vec![2])
.map_err(|e| anyhow!("alloc params: {e}"))?;
{
let s = params
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("mut_slice: {e}"))?;
s[0] = eps;
s[1] = hidden_size as f32;
}
rms_norm::dispatch_rms_norm(
encoder,
registry,
device.metal_device(),
input,
&weights_gpu.attn_norm,
&out,
¶ms,
seq_len,
hidden_size,
)
.context("dispatch_rms_norm")?;
Ok(out)
}
pub fn apply_linear_projection_f32(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
seq_len: u32,
in_features: u32,
out_features: u32,
) -> Result<MlxBuffer> {
debug_assert_eq!(
input.dtype(),
DType::F32,
"apply_linear_projection_f32: input must be F32 (kernel paths assume F32); got {}",
input.dtype()
);
let out_bytes = (seq_len * out_features) as usize * 4;
let mut dst = device
.alloc_buffer(
out_bytes,
DType::F32,
vec![seq_len as usize, out_features as usize],
)
.map_err(|e| anyhow!("alloc projection output: {e}"))?;
match weight.dtype() {
DType::U8 => {
let params = GgmlQuantizedMatmulParams {
m: seq_len,
n: out_features,
k: in_features,
ggml_type: GgmlType::Q4_0,
};
quantized_matmul_ggml(encoder, registry, device, input, weight, &mut dst, ¶ms)
.context("quantized_matmul_ggml Q4_0")?;
}
DType::BF16 => {
let params = DenseMmBf16F32Params {
m: seq_len,
n: out_features,
k: in_features,
src0_batch: 1,
src1_batch: 1,
};
if seq_len == 1 {
dense_gemv_bf16_f32(encoder, registry, device, weight, input, &mut dst, ¶ms)
.context("dense_gemv_bf16_f32 (M=1)")?;
} else {
dense_matmul_bf16_f32_tensor(
encoder, registry, device, weight, input, &mut dst, ¶ms,
)
.context("dense_matmul_bf16_f32_tensor")?;
}
}
DType::F32 => {
let n_w = (out_features * in_features) as usize;
let weight_bf16 = super::decode_pool::pooled_alloc_buffer(
device,
n_w * 2,
DType::BF16,
vec![out_features as usize, in_features as usize],
)
.map_err(|e| anyhow!("alloc weight_bf16 (pooled): {e}"))?;
cast(
encoder,
registry,
device.metal_device(),
weight,
&weight_bf16,
n_w,
CastDirection::F32ToBF16,
)
.context("cast weight F32→BF16")?;
encoder.memory_barrier();
let params = DenseMmBf16F32Params {
m: seq_len,
n: out_features,
k: in_features,
src0_batch: 1,
src1_batch: 1,
};
dense_matmul_bf16_f32_tensor(
encoder,
registry,
device,
&weight_bf16,
input,
&mut dst,
¶ms,
)
.context("dense_matmul_bf16_f32_tensor (F32 legacy)")?;
}
other => {
return Err(anyhow!(
"apply_linear_projection_f32: unsupported weight dtype {:?}",
other
));
}
}
Ok(dst)
}
pub fn apply_linear_projection_f32_qweight(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
qweight: &crate::serve::forward_mlx_shared::MlxQWeight,
seq_len: u32,
in_features: u32,
out_features: u32,
) -> Result<MlxBuffer> {
debug_assert_eq!(
input.dtype(),
DType::F32,
"apply_linear_projection_f32_qweight: input must be F32; got {}",
input.dtype()
);
let out_bytes = (seq_len * out_features) as usize * 4;
let mut dst = device
.alloc_buffer(
out_bytes,
DType::F32,
vec![seq_len as usize, out_features as usize],
)
.map_err(|e| anyhow!("alloc projection output: {e}"))?;
match qweight.buffer.dtype() {
DType::U8 => {
let params = GgmlQuantizedMatmulParams {
m: seq_len,
n: out_features,
k: in_features,
ggml_type: qweight.info.ggml_dtype,
};
quantized_matmul_ggml(
encoder,
registry,
device,
input,
&qweight.buffer,
&mut dst,
¶ms,
)
.with_context(|| {
format!(
"quantized_matmul_ggml ggml_type={:?}",
qweight.info.ggml_dtype
)
})?;
}
DType::BF16 => {
let params = DenseMmBf16F32Params {
m: seq_len,
n: out_features,
k: in_features,
src0_batch: 1,
src1_batch: 1,
};
if seq_len == 1 {
dense_gemv_bf16_f32(
encoder,
registry,
device,
&qweight.buffer,
input,
&mut dst,
¶ms,
)
.context("dense_gemv_bf16_f32 (M=1)")?;
} else {
dense_matmul_bf16_f32_tensor(
encoder,
registry,
device,
&qweight.buffer,
input,
&mut dst,
¶ms,
)
.context("dense_matmul_bf16_f32_tensor")?;
}
}
DType::F32 => {
let n_w = (out_features * in_features) as usize;
let weight_bf16 = super::decode_pool::pooled_alloc_buffer(
device,
n_w * 2,
DType::BF16,
vec![out_features as usize, in_features as usize],
)
.map_err(|e| anyhow!("alloc weight_bf16 (pooled): {e}"))?;
cast(
encoder,
registry,
device.metal_device(),
&qweight.buffer,
&weight_bf16,
n_w,
CastDirection::F32ToBF16,
)
.context("cast weight F32→BF16")?;
encoder.memory_barrier();
let params = DenseMmBf16F32Params {
m: seq_len,
n: out_features,
k: in_features,
src0_batch: 1,
src1_batch: 1,
};
dense_matmul_bf16_f32_tensor(
encoder,
registry,
device,
&weight_bf16,
input,
&mut dst,
¶ms,
)
.context("dense_matmul_bf16_f32_tensor (F32 legacy)")?;
}
other => {
return Err(anyhow!(
"apply_linear_projection_f32_qweight: unsupported weight dtype {:?}",
other
));
}
}
Ok(dst)
}
pub fn apply_linear_projection_f32_into(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
dst: &mut MlxBuffer,
seq_len: u32,
in_features: u32,
out_features: u32,
) -> Result<()> {
debug_assert_eq!(
input.dtype(),
DType::F32,
"apply_linear_projection_f32_into: input must be F32; got {}",
input.dtype()
);
debug_assert_eq!(
dst.dtype(),
DType::F32,
"apply_linear_projection_f32_into: dst must be F32; got {}",
dst.dtype()
);
match weight.dtype() {
DType::U8 => {
let params = GgmlQuantizedMatmulParams {
m: seq_len,
n: out_features,
k: in_features,
ggml_type: GgmlType::Q4_0,
};
quantized_matmul_ggml(encoder, registry, device, input, weight, dst, ¶ms)
.context("quantized_matmul_ggml Q4_0 (into)")?;
}
DType::BF16 => {
let params = DenseMmBf16F32Params {
m: seq_len,
n: out_features,
k: in_features,
src0_batch: 1,
src1_batch: 1,
};
if seq_len == 1 {
dense_gemv_bf16_f32(encoder, registry, device, weight, input, dst, ¶ms)
.context("dense_gemv_bf16_f32 (M=1, into)")?;
} else {
dense_matmul_bf16_f32_tensor(
encoder, registry, device, weight, input, dst, ¶ms,
)
.context("dense_matmul_bf16_f32_tensor (into)")?;
}
}
DType::F32 => {
let n_w = (out_features * in_features) as usize;
let weight_bf16 = super::decode_pool::pooled_alloc_buffer(
device,
n_w * 2,
DType::BF16,
vec![out_features as usize, in_features as usize],
)
.map_err(|e| anyhow!("alloc weight_bf16 (pooled, into): {e}"))?;
cast(
encoder,
registry,
device.metal_device(),
weight,
&weight_bf16,
n_w,
CastDirection::F32ToBF16,
)
.context("cast weight F32→BF16 (into)")?;
encoder.memory_barrier();
let params = DenseMmBf16F32Params {
m: seq_len,
n: out_features,
k: in_features,
src0_batch: 1,
src1_batch: 1,
};
dense_matmul_bf16_f32_tensor(
encoder,
registry,
device,
&weight_bf16,
input,
dst,
¶ms,
)
.context("dense_matmul_bf16_f32_tensor F32 legacy (into)")?;
}
other => {
return Err(anyhow!(
"apply_linear_projection_f32_into: unsupported weight dtype {:?}",
other
));
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn apply_linear_projection_f32_pooled(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
seq_len: u32,
in_features: u32,
out_features: u32,
) -> Result<MlxBuffer> {
if seq_len != 1 {
return apply_linear_projection_f32(
encoder,
registry,
device,
input,
weight,
seq_len,
in_features,
out_features,
);
}
let out_bytes = (seq_len * out_features) as usize * 4;
let mut dst = super::decode_pool::pooled_alloc_buffer(
device,
out_bytes,
DType::F32,
vec![seq_len as usize, out_features as usize],
)
.map_err(|e| anyhow!("alloc projection output (pooled): {e}"))?;
apply_linear_projection_f32_into(
encoder,
registry,
device,
input,
weight,
&mut dst,
seq_len,
in_features,
out_features,
)?;
Ok(dst)
}
pub fn apply_pre_attn_rms_norm_into(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weights_gpu: &FullAttnWeightsGpu,
out: &MlxBuffer,
params: &MlxBuffer,
seq_len: u32,
hidden_size: u32,
) -> Result<()> {
rms_norm::dispatch_rms_norm(
encoder,
registry,
device.metal_device(),
input,
&weights_gpu.attn_norm,
out,
params,
seq_len,
hidden_size,
)
.context("dispatch_rms_norm (arena into)")?;
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn apply_q_or_k_per_head_rms_norm_into(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
norm_weight: &MlxBuffer,
out: &MlxBuffer,
params: &MlxBuffer,
seq_len: u32,
n_heads: u32,
head_dim: u32,
) -> Result<()> {
let rows = seq_len * n_heads;
let dim = head_dim;
rms_norm::dispatch_rms_norm(
encoder,
registry,
device.metal_device(),
input,
norm_weight,
out,
params,
rows,
dim,
)
.context("dispatch_rms_norm per-head (arena into)")?;
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn apply_imrope_into(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
out: &MlxBuffer,
positions: &MlxBuffer,
seq_len: u32,
n_heads: u32,
head_dim: u32,
rotary_dim: u32,
freq_base: f32,
mrope_section: [u32; 4],
) -> Result<()> {
let params = RopeMultiParams {
head_dim,
rope_dim: rotary_dim,
n_heads,
seq_len,
freq_base,
mode: RopeMultiMode::Imrope,
sections: mrope_section,
};
dispatch_rope_multi_cached(encoder, registry, device, input, out, positions, params)
.context("dispatch_rope_multi_cached (arena into)")?;
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn apply_sigmoid_gate_multiply_into(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
attn_out: &MlxBuffer,
gate: &MlxBuffer,
out: &MlxBuffer,
params: &MlxBuffer,
n_elements: u32,
) -> Result<()> {
dispatch_sigmoid_mul(
encoder,
registry,
device.metal_device(),
attn_out,
gate,
out,
params,
n_elements,
)
.context("dispatch_sigmoid_mul (arena into)")?;
Ok(())
}
pub fn permute_seq_head_dim_to_head_seq_dim_cpu(
data: &[f32],
seq_len: usize,
n_heads: usize,
head_dim: usize,
) -> Vec<f32> {
let mut out = vec![0.0f32; seq_len * n_heads * head_dim];
for h in 0..n_heads {
for t in 0..seq_len {
let src_off = (t * n_heads + h) * head_dim;
let dst_off = (h * seq_len + t) * head_dim;
out[dst_off..dst_off + head_dim].copy_from_slice(&data[src_off..src_off + head_dim]);
}
}
out
}
pub fn apply_sdpa_causal(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
q_head_major: &MlxBuffer,
k_head_major: &MlxBuffer,
v_head_major: &MlxBuffer,
seq_len: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
) -> Result<MlxBuffer> {
let out = super::decode_pool::pooled_alloc_buffer(
device,
(n_heads * seq_len * head_dim) as usize * 4,
DType::F32,
vec![1, n_heads as usize, seq_len as usize, head_dim as usize],
)
.map_err(|e| anyhow!("alloc sdpa output: {e}"))?;
let params = SdpaParams {
n_heads,
n_kv_heads,
head_dim,
seq_len,
kv_seq_len: seq_len,
scale: 1.0 / (head_dim as f32).sqrt(),
kv_capacity: 0, do_causal: true,
};
sdpa(
encoder,
registry,
device,
q_head_major,
k_head_major,
v_head_major,
&out,
¶ms,
1,
)
.context("sdpa")?;
Ok(out)
}
pub fn apply_sdpa_causal_from_seq_major(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
q_seq_major: &MlxBuffer,
k_seq_major: &MlxBuffer,
v_seq_major: &MlxBuffer,
seq_len: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
) -> Result<MlxBuffer> {
let seq = seq_len as usize;
let nh = n_heads as usize;
let nkv = n_kv_heads as usize;
let d = head_dim as usize;
encoder
.commit_and_wait()
.context("commit before sdpa permute")?;
let q_cpu = download_f32(q_seq_major)?;
let k_cpu = download_f32(k_seq_major)?;
let v_cpu = download_f32(v_seq_major)?;
let q_hm = permute_seq_head_dim_to_head_seq_dim_cpu(&q_cpu, seq, nh, d);
let k_hm = permute_seq_head_dim_to_head_seq_dim_cpu(&k_cpu, seq, nkv, d);
let v_hm = permute_seq_head_dim_to_head_seq_dim_cpu(&v_cpu, seq, nkv, d);
let q_gpu = upload_f32(&q_hm, device)?;
let k_gpu = upload_f32(&k_hm, device)?;
let v_gpu = upload_f32(&v_hm, device)?;
let mut enc2 = device.command_encoder().context("new encoder for sdpa")?;
let out_hm = apply_sdpa_causal(
&mut enc2, registry, device, &q_gpu, &k_gpu, &v_gpu, seq_len, n_heads, n_kv_heads, head_dim,
)?;
enc2.commit_and_wait().context("sdpa commit")?;
let out_hm_cpu = download_f32(&out_hm)?;
let mut out_sm = vec![0.0f32; seq * nh * d];
for h in 0..nh {
for t in 0..seq {
let src = (h * seq + t) * d;
let dst = (t * nh + h) * d;
out_sm[dst..dst + d].copy_from_slice(&out_hm_cpu[src..src + d]);
}
}
upload_f32(&out_sm, device)
}
#[allow(clippy::too_many_arguments)]
pub fn apply_flash_attn_prefill_seq_major_into(
enc: &mut mlx_native::CommandEncoder,
device: &MlxDevice,
registry: &mut KernelRegistry,
q_seq_major: &MlxBuffer,
k_seq_major: &MlxBuffer,
v_seq_major: &MlxBuffer,
out_seq: &MlxBuffer,
seq_len: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
arena: &mut crate::inference::models::qwen35::FaPrefillArena,
) -> Result<()> {
if head_dim != 256 {
return Err(anyhow!(
"apply_flash_attn_prefill_seq_major_into: head_dim must be 256 \
(D=256 dispatcher); got {head_dim}. Other head_dims need a \
different mlx-native dispatcher (D=64 / D=512) or a new port."
));
}
let seq = seq_len 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;
arena
.validate_fits(seq_len, n_heads, n_kv_heads, head_dim)
.context("FA bridge: arena validate_fits")?;
cast(
enc,
registry,
device.metal_device(),
q_seq_major,
&arena.q_bf16_seq,
q_elems,
CastDirection::F32ToBF16,
)
.context("FA bridge: cast Q F32→BF16")?;
enc.memory_barrier();
permute_021_bf16(
enc,
registry,
device.metal_device(),
&arena.q_bf16_seq,
&arena.q_bf16_hm,
seq,
nh,
d,
)
.context("FA bridge: permute_021 Q [seq, nh, d] → [nh, seq, d]")?;
cast(
enc,
registry,
device.metal_device(),
k_seq_major,
&arena.k_bf16_seq,
k_elems,
CastDirection::F32ToBF16,
)
.context("FA bridge: cast K F32→BF16")?;
enc.memory_barrier();
permute_021_bf16(
enc,
registry,
device.metal_device(),
&arena.k_bf16_seq,
&arena.k_bf16_hm,
seq,
nkv,
d,
)
.context("FA bridge: permute_021 K [seq, nkv, d] → [nkv, seq, d]")?;
cast(
enc,
registry,
device.metal_device(),
v_seq_major,
&arena.v_bf16_seq,
v_elems,
CastDirection::F32ToBF16,
)
.context("FA bridge: cast V F32→BF16")?;
enc.memory_barrier();
permute_021_bf16(
enc,
registry,
device.metal_device(),
&arena.v_bf16_seq,
&arena.v_bf16_hm,
seq,
nkv,
d,
)
.context("FA bridge: permute_021 V [seq, nkv, d] → [nkv, seq, d]")?;
enc.memory_barrier();
let scale = 1.0 / (d as f32).sqrt();
dispatch_flash_attn_prefill_bf16_d256(
enc,
device,
registry,
&arena.q_bf16_hm,
&arena.k_bf16_hm,
&arena.v_bf16_hm,
None,
&mut arena.out_bf16_hm,
&FlashAttnPrefillParams {
n_heads,
n_kv_heads,
head_dim,
seq_len_q: seq_len,
seq_len_k: seq_len,
batch: 1,
scale,
do_causal: true,
},
)
.context("FA bridge: dispatch_flash_attn_prefill_bf16_d256")?;
enc.memory_barrier();
permute_021_bf16_to_f32(
enc,
registry,
device.metal_device(),
&arena.out_bf16_hm,
out_seq,
nh,
seq,
d,
)
.context("FA bridge: permute_021_bf16_to_f32 out [nh, seq, d] → [seq, nh, d] F32")?;
Ok(())
}
#[derive(Debug, Clone, Copy)]
pub struct Qwen35TreeVerifyParams {
pub num_q_heads: u32,
pub num_kv_heads: u32,
pub head_dim: u32,
pub q_seq_len: u32,
pub kv_seq_len: u32,
pub kv_capacity: u32,
pub mask_stride: u32,
pub scale: f32,
}
#[allow(clippy::too_many_arguments)]
pub fn dispatch_qwen35_tree_verify_attention(
enc: &mut mlx_native::CommandEncoder,
device: &MlxDevice,
registry: &mut KernelRegistry,
q_head_outer: &MlxBuffer,
k_head_outer: &MlxBuffer,
v_head_outer: &MlxBuffer,
tree_mask: &MlxBuffer,
params: Qwen35TreeVerifyParams,
) -> Result<MlxBuffer> {
if params.head_dim != 128 {
return Err(anyhow!(
"dispatch_qwen35_tree_verify_attention: head_dim must be 128 \
(Qwen3.5 tree-verify dispatcher); got {}. Other head_dims need \
a different target-model wrapper.",
params.head_dim
));
}
if params.q_seq_len == 0 {
return Err(anyhow!(
"dispatch_qwen35_tree_verify_attention: q_seq_len must be > 0"
));
}
if params.kv_seq_len == 0 {
return Err(anyhow!(
"dispatch_qwen35_tree_verify_attention: kv_seq_len must be > 0"
));
}
if params.kv_capacity < params.kv_seq_len {
return Err(anyhow!(
"dispatch_qwen35_tree_verify_attention: kv_capacity ({}) must be >= kv_seq_len ({})",
params.kv_capacity,
params.kv_seq_len
));
}
if params.mask_stride < params.kv_seq_len {
return Err(anyhow!(
"dispatch_qwen35_tree_verify_attention: mask_stride ({}) must be >= kv_seq_len ({})",
params.mask_stride,
params.kv_seq_len
));
}
if !params.scale.is_finite() {
return Err(anyhow!(
"dispatch_qwen35_tree_verify_attention: scale ({}) must be finite",
params.scale
));
}
if params.num_q_heads == 0 || params.num_kv_heads == 0 {
return Err(anyhow!(
"dispatch_qwen35_tree_verify_attention: num_q_heads and num_kv_heads must be > 0"
));
}
if params.num_q_heads % params.num_kv_heads != 0 {
return Err(anyhow!(
"dispatch_qwen35_tree_verify_attention: num_q_heads ({}) must be divisible by num_kv_heads ({})",
params.num_q_heads,
params.num_kv_heads
));
}
let mul = |a: usize, b: usize, ctx: &str| -> Result<usize> {
a.checked_mul(b)
.ok_or_else(|| anyhow!("dispatch_qwen35_tree_verify_attention: {ctx} overflows usize"))
};
let q = params.q_seq_len as usize;
let nq = params.num_q_heads as usize;
let nkv = params.num_kv_heads as usize;
let d = params.head_dim as usize;
let cap = params.kv_capacity as usize;
let stride = params.mask_stride as usize;
let out_elems = mul(mul(q, nq, "q_seq_len*num_q_heads")?, d, "out elements")?;
let out_bytes = mul(out_elems, std::mem::size_of::<f32>(), "out bytes")?;
let kv_req_bytes = mul(
mul(mul(nkv, cap, "num_kv_heads*kv_capacity")?, d, "kv elements")?,
std::mem::size_of::<f32>(),
"kv bytes",
)?;
let mask_req_bytes = mul(
mul(q, stride, "q_seq_len*mask_stride")?,
std::mem::size_of::<f32>(),
"mask bytes",
)?;
let tmp_bytes =
tree_attn_ops::tmp_buffer_bytes(params.num_q_heads, params.head_dim, params.q_seq_len);
if q_head_outer.byte_len() < out_bytes {
return Err(anyhow!(
"dispatch_qwen35_tree_verify_attention: q buffer too small: have {} bytes, need >= {}",
q_head_outer.byte_len(),
out_bytes
));
}
if k_head_outer.byte_len() < kv_req_bytes {
return Err(anyhow!(
"dispatch_qwen35_tree_verify_attention: k buffer too small: have {} bytes, need >= {}",
k_head_outer.byte_len(),
kv_req_bytes
));
}
if v_head_outer.byte_len() < kv_req_bytes {
return Err(anyhow!(
"dispatch_qwen35_tree_verify_attention: v buffer too small: have {} bytes, need >= {}",
v_head_outer.byte_len(),
kv_req_bytes
));
}
if tree_mask.byte_len() < mask_req_bytes {
return Err(anyhow!(
"dispatch_qwen35_tree_verify_attention: tree_mask buffer too small: have {} bytes, need >= {}",
tree_mask.byte_len(),
mask_req_bytes
));
}
let output = device
.alloc_buffer(out_bytes, DType::F32, vec![q, nq, d])
.map_err(|e| anyhow!("alloc qwen35_tree_verify output: {e}"))?;
let tmp = device
.alloc_buffer(tmp_bytes, DType::F32, vec![tmp_bytes / 4])
.map_err(|e| anyhow!("alloc qwen35_tree_verify tmp: {e}"))?;
let tree_params = TreeAttentionParams {
num_heads: params.num_q_heads,
num_kv_heads: params.num_kv_heads,
head_dim: params.head_dim,
kv_seq_len: params.kv_seq_len,
kv_capacity: params.kv_capacity,
scale: params.scale,
q_seq_len: params.q_seq_len,
mask_stride: params.mask_stride,
};
enc.memory_barrier();
tree_attn_ops::tree_attention(
enc,
registry,
device,
q_head_outer,
k_head_outer,
v_head_outer,
tree_mask,
&output,
&tmp,
&tree_params,
)
.context("qwen35_tree_verify: tree_attention")?;
Ok(output)
}
#[derive(Debug, Clone, Copy)]
pub struct Qwen35TreeVerifyLayerShape {
pub hidden_size: u32,
pub num_q_heads: u32,
pub num_kv_heads: u32,
pub head_dim: u32,
pub tree_seq_len: u32,
pub cache_prefix_len: u32,
pub kv_capacity: u32,
pub mask_stride: u32,
pub rotary_dim: u32,
pub freq_base: f32,
pub mrope_section: [u32; 4],
pub rms_norm_eps: f32,
pub attn_output_gate: bool,
}
impl Qwen35TreeVerifyLayerShape {
pub fn validate(&self) -> Result<()> {
use anyhow::ensure;
ensure!(
self.head_dim == 128,
"Qwen35TreeVerifyLayerShape: head_dim must be 128 (dk128 tree-attention \
kernel only); got {}. Production Qwen 3.6 27B head_dim=256 requires a \
follow-up CFA adding a dk256 tree-attention kernel.",
self.head_dim
);
ensure!(
self.attn_output_gate,
"Qwen35TreeVerifyLayerShape: attn_output_gate must be true for Qwen3.5/3.6; \
got false. Set attn_output_gate=true or do not call this function."
);
ensure!(
self.tree_seq_len > 0,
"Qwen35TreeVerifyLayerShape: tree_seq_len must be > 0"
);
ensure!(
self.hidden_size > 0,
"Qwen35TreeVerifyLayerShape: hidden_size must be > 0"
);
ensure!(
self.num_q_heads > 0,
"Qwen35TreeVerifyLayerShape: num_q_heads must be > 0"
);
ensure!(
self.num_kv_heads > 0,
"Qwen35TreeVerifyLayerShape: num_kv_heads must be > 0"
);
ensure!(
self.num_q_heads % self.num_kv_heads == 0,
"Qwen35TreeVerifyLayerShape: num_q_heads ({}) must be divisible by \
num_kv_heads ({})",
self.num_q_heads,
self.num_kv_heads
);
let kv_end = (self.cache_prefix_len as u64)
.checked_add(self.tree_seq_len as u64)
.ok_or_else(|| {
anyhow!("Qwen35TreeVerifyLayerShape: cache_prefix_len + tree_seq_len overflows u64")
})?;
ensure!(
kv_end <= self.kv_capacity as u64,
"Qwen35TreeVerifyLayerShape: cache_prefix_len ({}) + tree_seq_len ({}) = {} \
must be <= kv_capacity ({})",
self.cache_prefix_len,
self.tree_seq_len,
kv_end,
self.kv_capacity
);
ensure!(
self.mask_stride >= (kv_end as u32),
"Qwen35TreeVerifyLayerShape: mask_stride ({}) must be >= cache_prefix_len + \
tree_seq_len ({})",
self.mask_stride,
kv_end
);
Ok(())
}
}
#[allow(clippy::too_many_arguments)]
pub fn qwen35_tree_verify_attention_block(
enc: mlx_native::CommandEncoder,
device: &MlxDevice,
registry: &mut KernelRegistry,
hidden_states_in: &MlxBuffer,
tree_mask: &MlxBuffer,
tree_positions: &MlxBuffer,
k_cache: &mut MlxBuffer,
v_cache: &mut MlxBuffer,
weights: &FullAttnWeightsGpu,
shape: Qwen35TreeVerifyLayerShape,
) -> Result<MlxBuffer> {
shape.validate()?;
let checked_mul = |a: usize, b: usize, ctx: &str| -> Result<usize> {
a.checked_mul(b)
.ok_or_else(|| anyhow!("qwen35_tree_verify_attention_block: {ctx} overflows usize"))
};
let seq = shape.tree_seq_len as usize;
let h = shape.hidden_size as usize;
let nq = shape.num_q_heads as usize;
let nkv = shape.num_kv_heads as usize;
let d = shape.head_dim as usize;
let cap = shape.kv_capacity as usize;
let prefix = shape.cache_prefix_len as usize;
let kv_end = prefix + seq;
if hidden_states_in.dtype() != DType::F32 {
return Err(anyhow!(
"qwen35_tree_verify_attention_block: hidden_states_in dtype must be F32, got {:?}",
hidden_states_in.dtype()
));
}
let hs_elems = checked_mul(seq, h, "tree_seq_len * hidden_size")?;
if hidden_states_in.element_count() != hs_elems {
return Err(anyhow!(
"qwen35_tree_verify_attention_block: hidden_states_in has {} elements, \
expected exactly {} (tree_seq_len={} * hidden_size={})",
hidden_states_in.element_count(),
hs_elems,
seq,
h
));
}
if tree_mask.dtype() != DType::F32 {
return Err(anyhow!(
"qwen35_tree_verify_attention_block: tree_mask dtype must be F32, got {:?}",
tree_mask.dtype()
));
}
let mask_elems = checked_mul(
seq,
shape.mask_stride as usize,
"tree_seq_len * mask_stride",
)?;
if tree_mask.element_count() < mask_elems {
return Err(anyhow!(
"qwen35_tree_verify_attention_block: tree_mask has {} elements, \
need >= {} (tree_seq_len={} * mask_stride={})",
tree_mask.element_count(),
mask_elems,
seq,
shape.mask_stride
));
}
if tree_positions.dtype() != DType::I32 {
return Err(anyhow!(
"qwen35_tree_verify_attention_block: tree_positions dtype must be I32 (got {:?}); \
caller must pre-encode IMROPE positions as i32",
tree_positions.dtype()
));
}
let pos_elems = checked_mul(4, seq, "4 * tree_seq_len")?;
if tree_positions.element_count() != pos_elems {
return Err(anyhow!(
"qwen35_tree_verify_attention_block: tree_positions has {} elements, \
need exactly {} (4 * tree_seq_len={})",
tree_positions.element_count(),
pos_elems,
seq
));
}
if k_cache.dtype() != DType::F32 {
return Err(anyhow!(
"qwen35_tree_verify_attention_block: k_cache dtype must be F32, got {:?}",
k_cache.dtype()
));
}
if v_cache.dtype() != DType::F32 {
return Err(anyhow!(
"qwen35_tree_verify_attention_block: v_cache dtype must be F32, got {:?}",
v_cache.dtype()
));
}
let kv_req_elems = checked_mul(
checked_mul(nkv, cap, "num_kv_heads * kv_capacity")?,
d,
"kv_capacity * head_dim",
)?;
let kv_req_bytes = checked_mul(kv_req_elems, std::mem::size_of::<f32>(), "kv bytes")?;
if k_cache.byte_len() < kv_req_bytes {
return Err(anyhow!(
"qwen35_tree_verify_attention_block: k_cache byte_len {} < required {}",
k_cache.byte_len(),
kv_req_bytes
));
}
if v_cache.byte_len() < kv_req_bytes {
return Err(anyhow!(
"qwen35_tree_verify_attention_block: v_cache byte_len {} < required {}",
v_cache.byte_len(),
kv_req_bytes
));
}
let q_total = checked_mul(nq, d, "num_q_heads * head_dim")?;
let _kv_total = checked_mul(nkv, d, "num_kv_heads * head_dim")?;
if weights.attn_norm.element_count() != h {
return Err(anyhow!(
"qwen35_tree_verify_attention_block: weights.attn_norm has {} elements, expected {} (hidden_size)",
weights.attn_norm.element_count(),
h,
));
}
let qk_norm_expected = d;
if weights.attn_q_norm.element_count() != qk_norm_expected {
return Err(anyhow!(
"qwen35_tree_verify_attention_block: weights.attn_q_norm has {} elements, expected {} (head_dim)",
weights.attn_q_norm.element_count(),
qk_norm_expected,
));
}
if weights.attn_k_norm.element_count() != qk_norm_expected {
return Err(anyhow!(
"qwen35_tree_verify_attention_block: weights.attn_k_norm has {} elements, expected {} (head_dim)",
weights.attn_k_norm.element_count(),
qk_norm_expected,
));
}
let mut enc = enc;
let hidden_normed = apply_pre_attn_rms_norm(
&mut enc,
registry,
device,
hidden_states_in,
weights,
shape.tree_seq_len,
shape.hidden_size,
shape.rms_norm_eps,
)
.context("step 1: apply_pre_attn_rms_norm")?;
let q_flat = apply_linear_projection_f32(
&mut enc,
registry,
device,
&hidden_normed,
&weights.wq,
shape.tree_seq_len,
shape.hidden_size,
shape.num_q_heads * shape.head_dim,
)
.context("step 2: Q projection")?;
let k_flat = apply_linear_projection_f32(
&mut enc,
registry,
device,
&hidden_normed,
&weights.wk,
shape.tree_seq_len,
shape.hidden_size,
shape.num_kv_heads * shape.head_dim,
)
.context("step 2: K projection")?;
let v_flat = apply_linear_projection_f32(
&mut enc,
registry,
device,
&hidden_normed,
&weights.wv,
shape.tree_seq_len,
shape.hidden_size,
shape.num_kv_heads * shape.head_dim,
)
.context("step 2: V projection")?;
let gate_flat = apply_linear_projection_f32(
&mut enc,
registry,
device,
&hidden_normed,
&weights.w_gate,
shape.tree_seq_len,
shape.hidden_size,
shape.num_q_heads * shape.head_dim,
)
.context("step 2: gate projection")?;
enc.memory_barrier();
let q_normed = apply_q_or_k_per_head_rms_norm(
&mut enc,
registry,
device,
&q_flat,
&weights.attn_q_norm,
shape.tree_seq_len,
shape.num_q_heads,
shape.head_dim,
shape.rms_norm_eps,
)
.context("step 3: Q per-head RMSNorm")?;
let k_normed = apply_q_or_k_per_head_rms_norm(
&mut enc,
registry,
device,
&k_flat,
&weights.attn_k_norm,
shape.tree_seq_len,
shape.num_kv_heads,
shape.head_dim,
shape.rms_norm_eps,
)
.context("step 3: K per-head RMSNorm")?;
enc.memory_barrier();
let q_roped = apply_imrope(
&mut enc,
registry,
device,
&q_normed,
tree_positions,
shape.tree_seq_len,
shape.num_q_heads,
shape.head_dim,
shape.rotary_dim,
shape.freq_base,
shape.mrope_section,
)
.context("step 4: Q IMROPE")?;
let k_roped = apply_imrope(
&mut enc,
registry,
device,
&k_normed,
tree_positions,
shape.tree_seq_len,
shape.num_kv_heads,
shape.head_dim,
shape.rotary_dim,
shape.freq_base,
shape.mrope_section,
)
.context("step 4: K IMROPE")?;
enc.memory_barrier();
let q_ho_bytes = checked_mul(checked_mul(nq, seq, "nq*seq")?, d, "q_ho bytes * 4")?
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| anyhow!("q_head_outer bytes overflow"))?;
let kv_ho_bytes = checked_mul(checked_mul(nkv, seq, "nkv*seq")?, d, "kv_ho bytes * 4")?
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| anyhow!("kv_head_outer bytes overflow"))?;
let q_head_outer = device
.alloc_buffer(q_ho_bytes, DType::F32, vec![nq, seq, d])
.map_err(|e| anyhow!("alloc q_head_outer: {e}"))?;
let k_scratch = device
.alloc_buffer(kv_ho_bytes, DType::F32, vec![nkv, seq, d])
.map_err(|e| anyhow!("alloc k_scratch: {e}"))?;
let v_scratch = device
.alloc_buffer(kv_ho_bytes, DType::F32, vec![nkv, seq, d])
.map_err(|e| anyhow!("alloc v_scratch: {e}"))?;
permute_021_f32(
&mut enc,
registry,
device.metal_device(),
&q_roped,
&q_head_outer,
seq,
nq,
d,
)
.context("step 5: Q permute seq→head-outer")?;
permute_021_f32(
&mut enc,
registry,
device.metal_device(),
&k_roped,
&k_scratch,
seq,
nkv,
d,
)
.context("step 5: K permute seq→head-outer")?;
permute_021_f32(
&mut enc,
registry,
device.metal_device(),
&v_flat,
&v_scratch,
seq,
nkv,
d,
)
.context("step 5: V permute seq→head-outer")?;
enc.memory_barrier();
enc.commit_and_wait()
.context("step 6: commit encoder before KV cache write")?;
{
let k_src = k_scratch
.as_slice::<f32>()
.map_err(|e| anyhow!("step 7: k_scratch as_slice: {e}"))?;
let v_src = v_scratch
.as_slice::<f32>()
.map_err(|e| anyhow!("step 7: v_scratch as_slice: {e}"))?;
let k_dst = k_cache
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("step 7: k_cache as_mut_slice: {e}"))?;
let v_dst = v_cache
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("step 7: v_cache as_mut_slice: {e}"))?;
for kv_head in 0..nkv {
for pos in 0..seq {
let src_off = kv_head
.checked_mul(seq)
.and_then(|x| x.checked_add(pos))
.and_then(|x| x.checked_mul(d))
.ok_or_else(|| anyhow!("step 7: k_src offset overflow"))?;
let dst_off = kv_head
.checked_mul(cap)
.and_then(|x| x.checked_add(prefix + pos))
.and_then(|x| x.checked_mul(d))
.ok_or_else(|| anyhow!("step 7: k_dst offset overflow"))?;
k_dst[dst_off..dst_off + d].copy_from_slice(&k_src[src_off..src_off + d]);
v_dst[dst_off..dst_off + d].copy_from_slice(&v_src[src_off..src_off + d]);
}
}
}
let mut enc2 = device
.command_encoder()
.map_err(|e| anyhow!("step 8: open encoder: {e}"))?;
let scale = 1.0_f32 / (shape.head_dim as f32).sqrt();
let kv_seq_len = kv_end as u32;
let attn_out = dispatch_qwen35_tree_verify_attention(
&mut enc2,
device,
registry,
&q_head_outer,
k_cache,
v_cache,
tree_mask,
Qwen35TreeVerifyParams {
num_q_heads: shape.num_q_heads,
num_kv_heads: shape.num_kv_heads,
head_dim: shape.head_dim,
q_seq_len: shape.tree_seq_len,
kv_seq_len,
kv_capacity: shape.kv_capacity,
mask_stride: shape.mask_stride,
scale,
},
)
.context("step 8: dispatch_qwen35_tree_verify_attention")?;
enc2.memory_barrier();
let n_gate_elems = checked_mul(seq, q_total, "tree_seq_len * q_total")?;
let gated = apply_sigmoid_gate_multiply(
&mut enc2,
registry,
device,
&attn_out,
&gate_flat,
n_gate_elems as u32,
)
.context("step 9: apply_sigmoid_gate_multiply")?;
enc2.memory_barrier();
let o_out = apply_linear_projection_f32(
&mut enc2,
registry,
device,
&gated,
&weights.wo,
shape.tree_seq_len,
q_total as u32,
shape.hidden_size,
)
.context("step 10: O projection")?;
enc2.memory_barrier();
let out_bytes = checked_mul(hs_elems, std::mem::size_of::<f32>(), "output bytes")?;
let hidden_states_out = device
.alloc_buffer(out_bytes, DType::F32, vec![seq, h])
.map_err(|e| anyhow!("step 11: alloc hidden_states_out: {e}"))?;
elementwise_add(
&mut enc2,
registry,
device.metal_device(),
hidden_states_in,
&o_out,
&hidden_states_out,
hs_elems,
DType::F32,
)
.context("step 11: elementwise_add residual")?;
enc2.commit_and_wait().context("step 11: terminal commit")?;
Ok(hidden_states_out)
}
#[derive(Debug, Clone, Copy)]
pub struct Qwen35TreeVerifyFullLayerShape {
pub attn: Qwen35TreeVerifyLayerShape,
pub intermediate_size: u32,
}
impl Qwen35TreeVerifyFullLayerShape {
pub fn validate(&self) -> Result<()> {
self.attn.validate()?;
let h = self.attn.hidden_size as usize;
let m = self.intermediate_size as usize;
if self.intermediate_size == 0 {
return Err(anyhow!(
"Qwen35TreeVerifyFullLayerShape: intermediate_size must be > 0"
));
}
let _overflow_check = (m as u64).checked_mul(h as u64).ok_or_else(|| {
anyhow!(
"Qwen35TreeVerifyFullLayerShape: intermediate_size ({}) * hidden_size ({}) \
overflows u64 — too large for activated scratch allocation",
m,
h
)
})?;
m.checked_mul(h).ok_or_else(|| {
anyhow!(
"Qwen35TreeVerifyFullLayerShape: intermediate_size ({}) * hidden_size ({}) \
overflows usize",
m,
h
)
})?;
Ok(())
}
}
#[allow(clippy::too_many_arguments)]
pub fn qwen35_tree_verify_full_layer(
enc: mlx_native::CommandEncoder,
device: &MlxDevice,
registry: &mut KernelRegistry,
hidden_states_in: &MlxBuffer,
tree_mask: &MlxBuffer,
tree_positions: &MlxBuffer,
k_cache: &mut MlxBuffer,
v_cache: &mut MlxBuffer,
weights: &FullAttnWeightsGpu,
ffn_weights: &super::gpu_ffn::DenseFfnWeightsGpu,
shape: Qwen35TreeVerifyFullLayerShape,
) -> Result<MlxBuffer> {
shape.validate()?;
let seq = shape.attn.tree_seq_len as usize;
let h = shape.attn.hidden_size as usize;
let m = shape.intermediate_size as usize;
let gate_expected = m.checked_mul(h).ok_or_else(|| {
anyhow!("qwen35_tree_verify_full_layer: intermediate_size * hidden_size overflows usize")
})?;
if ffn_weights.gate.element_count() != gate_expected {
return Err(anyhow!(
"qwen35_tree_verify_full_layer: ffn_weights.gate has {} elements, \
expected exactly {} (intermediate_size={} * hidden_size={})",
ffn_weights.gate.element_count(),
gate_expected,
m,
h
));
}
if ffn_weights.up.element_count() != gate_expected {
return Err(anyhow!(
"qwen35_tree_verify_full_layer: ffn_weights.up has {} elements, \
expected exactly {} (intermediate_size={} * hidden_size={})",
ffn_weights.up.element_count(),
gate_expected,
m,
h
));
}
let down_expected = h.checked_mul(m).ok_or_else(|| {
anyhow!("qwen35_tree_verify_full_layer: hidden_size * intermediate_size overflows usize")
})?;
if ffn_weights.down.element_count() != down_expected {
return Err(anyhow!(
"qwen35_tree_verify_full_layer: ffn_weights.down has {} elements, \
expected exactly {} (hidden_size={} * intermediate_size={})",
ffn_weights.down.element_count(),
down_expected,
h,
m
));
}
let attn_out = qwen35_tree_verify_attention_block(
enc,
device,
registry,
hidden_states_in,
tree_mask,
tree_positions,
k_cache,
v_cache,
weights,
shape.attn,
)
.context("qwen35_tree_verify_full_layer: attention block")?;
let ffn_residual = attn_out.clone();
let mut enc2 = device
.command_encoder()
.context("qwen35_tree_verify_full_layer: alloc enc2")?;
let rms_out_bytes = seq * h * std::mem::size_of::<f32>();
let post_attn_normed =
super::decode_pool::pooled_alloc_buffer(device, rms_out_bytes, DType::F32, vec![seq, h])
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer: alloc post_attn_normed: {e}"))?;
let mut rms_params = super::decode_pool::pooled_alloc_buffer(device, 8, DType::F32, vec![2])
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer: alloc rms_params: {e}"))?;
{
let s = rms_params
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer: rms_params slice: {e}"))?;
s[0] = shape.attn.rms_norm_eps;
s[1] = h as f32;
}
rms_norm::dispatch_rms_norm(
&mut enc2,
registry,
device.metal_device(),
&attn_out,
&weights.post_attn_norm,
&post_attn_normed,
&rms_params,
shape.attn.tree_seq_len,
shape.attn.hidden_size,
)
.context("qwen35_tree_verify_full_layer: post_attn_norm")?;
enc2.memory_barrier();
let gate_buf = apply_linear_projection_f32(
&mut enc2,
registry,
device,
&post_attn_normed,
&ffn_weights.gate,
shape.attn.tree_seq_len,
shape.attn.hidden_size,
shape.intermediate_size,
)
.context("qwen35_tree_verify_full_layer: gate_proj")?;
let up_buf = apply_linear_projection_f32(
&mut enc2,
registry,
device,
&post_attn_normed,
&ffn_weights.up,
shape.attn.tree_seq_len,
shape.attn.hidden_size,
shape.intermediate_size,
)
.context("qwen35_tree_verify_full_layer: up_proj")?;
enc2.memory_barrier();
let n_silu_elems = seq.checked_mul(m).ok_or_else(|| {
anyhow!("qwen35_tree_verify_full_layer: seq * intermediate overflows usize")
})?;
if n_silu_elems > (u32::MAX as usize) {
return Err(anyhow!(
"qwen35_tree_verify_full_layer: seq ({}) * intermediate ({}) = {} exceeds u32::MAX",
seq,
m,
n_silu_elems
));
}
let n_silu: u32 = n_silu_elems as u32;
let activated_bytes = n_silu_elems
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| anyhow!("qwen35_tree_verify_full_layer: activated_bytes overflow"))?;
let activated_buf = device
.alloc_buffer(activated_bytes, DType::F32, vec![seq, m])
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer: alloc activated_buf: {e}"))?;
let mut silu_params =
super::decode_pool::pooled_alloc_buffer(device, 4, DType::U32, vec![1])
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer: alloc silu_params: {e}"))?;
silu_params
.as_mut_slice::<u32>()
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer: silu_params slice: {e}"))?[0] = n_silu;
dispatch_silu_mul(
&mut enc2,
registry,
device.metal_device(),
&gate_buf,
&up_buf,
&activated_buf,
&silu_params,
n_silu,
)
.context("qwen35_tree_verify_full_layer: dispatch_silu_mul")?;
enc2.memory_barrier();
let ffn_out = apply_linear_projection_f32(
&mut enc2,
registry,
device,
&activated_buf,
&ffn_weights.down,
shape.attn.tree_seq_len,
shape.intermediate_size,
shape.attn.hidden_size,
)
.context("qwen35_tree_verify_full_layer: down_proj")?;
enc2.memory_barrier();
let hs_elems = seq * h;
let out_bytes = hs_elems * std::mem::size_of::<f32>();
let hidden_states_out = device
.alloc_buffer(out_bytes, DType::F32, vec![seq, h])
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer: alloc hidden_states_out: {e}"))?;
elementwise_add(
&mut enc2,
registry,
device.metal_device(),
&ffn_residual,
&ffn_out,
&hidden_states_out,
hs_elems,
DType::F32,
)
.context("qwen35_tree_verify_full_layer: residual add")?;
enc2.commit_and_wait()
.context("qwen35_tree_verify_full_layer: enc2 terminal commit")?;
Ok(hidden_states_out)
}
#[derive(Debug, Clone, Copy)]
pub struct Qwen35TreeVerifyFullLayerShapeQ {
pub attn: Qwen35TreeVerifyLayerShape,
pub intermediate_size: u32,
}
impl Qwen35TreeVerifyFullLayerShapeQ {
pub fn validate(&self) -> Result<()> {
self.attn.validate()?;
let h = self.attn.hidden_size as usize;
let m = self.intermediate_size as usize;
if self.intermediate_size == 0 {
return Err(anyhow!(
"Qwen35TreeVerifyFullLayerShapeQ: intermediate_size must be > 0"
));
}
let _overflow_check = (m as u64).checked_mul(h as u64).ok_or_else(|| {
anyhow!(
"Qwen35TreeVerifyFullLayerShapeQ: intermediate_size ({}) * hidden_size ({}) \
overflows u64 — too large for activated scratch allocation",
m,
h
)
})?;
m.checked_mul(h).ok_or_else(|| {
anyhow!(
"Qwen35TreeVerifyFullLayerShapeQ: intermediate_size ({}) * hidden_size ({}) \
overflows usize",
m,
h
)
})?;
Ok(())
}
}
#[derive(Debug, Clone, Copy)]
pub struct Qwen35TreeVerifyFullLayerShapeQMoe {
pub attn: Qwen35TreeVerifyLayerShape,
pub moe: super::ffn::MoeFfnShape,
}
impl Qwen35TreeVerifyFullLayerShapeQMoe {
pub fn validate(&self) -> Result<()> {
self.attn.validate()?;
let h = self.attn.hidden_size as usize;
let ne = self.moe.num_experts as usize;
let topk = self.moe.num_experts_per_tok as usize;
let m_moe = self.moe.moe_intermediate_size as usize;
let m_sh = self.moe.shared_intermediate_size as usize;
if self.moe.hidden_size != self.attn.hidden_size {
return Err(anyhow!(
"Qwen35TreeVerifyFullLayerShapeQMoe: moe.hidden_size ({}) != attn.hidden_size ({}) \
— drift guard: shape fields must be consistent",
self.moe.hidden_size, self.attn.hidden_size
));
}
if ne == 0 {
return Err(anyhow!(
"Qwen35TreeVerifyFullLayerShapeQMoe: num_experts must be > 0"
));
}
if topk == 0 {
return Err(anyhow!(
"Qwen35TreeVerifyFullLayerShapeQMoe: num_experts_per_tok must be > 0"
));
}
if topk > ne {
return Err(anyhow!(
"Qwen35TreeVerifyFullLayerShapeQMoe: num_experts_per_tok ({}) > num_experts ({}) \
— top-K cannot exceed total experts",
topk,
ne
));
}
if m_moe == 0 {
return Err(anyhow!(
"Qwen35TreeVerifyFullLayerShapeQMoe: moe_intermediate_size must be > 0"
));
}
if m_sh == 0 {
return Err(anyhow!(
"Qwen35TreeVerifyFullLayerShapeQMoe: shared_intermediate_size must be > 0"
));
}
(ne as u64)
.checked_mul(m_moe as u64)
.and_then(|v| v.checked_mul(h as u64))
.ok_or_else(|| {
anyhow!(
"Qwen35TreeVerifyFullLayerShapeQMoe: num_experts ({}) * moe_intermediate ({}) \
* hidden_size ({}) overflows u64",
ne,
m_moe,
h
)
})?;
ne.checked_mul(m_moe)
.and_then(|v| v.checked_mul(h))
.ok_or_else(|| {
anyhow!(
"Qwen35TreeVerifyFullLayerShapeQMoe: num_experts ({}) * moe_intermediate ({}) \
* hidden_size ({}) overflows usize",
ne,
m_moe,
h
)
})?;
Ok(())
}
}
#[allow(clippy::too_many_arguments)]
pub fn qwen35_tree_verify_full_layer_q(
enc: mlx_native::CommandEncoder,
device: &MlxDevice,
registry: &mut KernelRegistry,
hidden_states_in: &MlxBuffer,
tree_mask: &MlxBuffer,
tree_positions: &MlxBuffer,
k_cache: &mut MlxBuffer,
v_cache: &mut MlxBuffer,
weights: &FullAttnWeightsGpu,
ffn_weights: &super::gpu_ffn::DenseFfnWeightsGpuQ,
shape: Qwen35TreeVerifyFullLayerShapeQ,
) -> Result<MlxBuffer> {
if ffn_weights.ggml_type_gate_up != GgmlType::Q4_0 {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q: ggml_type_gate_up must be Q4_0 \
(got {:?}). Future CFAs will support Q5_K/Q6_K mixed-quant via \
per-projection ggml_type threading through apply_linear_projection_f32.",
ffn_weights.ggml_type_gate_up
));
}
if ffn_weights.ggml_type_down != GgmlType::Q4_0 {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q: ggml_type_down must be Q4_0 \
(got {:?}). Future CFAs will support Q5_K/Q6_K mixed-quant via \
per-projection ggml_type threading through apply_linear_projection_f32.",
ffn_weights.ggml_type_down
));
}
shape.validate()?;
let seq = shape.attn.tree_seq_len as usize;
let h = shape.attn.hidden_size as usize;
let m = shape.intermediate_size as usize;
if shape.attn.hidden_size != ffn_weights.hidden_size {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q: shape.attn.hidden_size ({}) != \
ffn_weights.hidden_size ({}). Shape and weights were built from \
different model configs.",
shape.attn.hidden_size,
ffn_weights.hidden_size
));
}
if shape.intermediate_size != ffn_weights.intermediate_size {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q: shape.intermediate_size ({}) != \
ffn_weights.intermediate_size ({}). Shape and weights were built from \
different model configs.",
shape.intermediate_size,
ffn_weights.intermediate_size
));
}
let gate_blocks_per_row = h.checked_div(32).ok_or_else(|| {
anyhow!(
"qwen35_tree_verify_full_layer_q: hidden_size {} not divisible by 32 (Q4_0 block)",
h
)
})?;
if h % 32 != 0 {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q: hidden_size ({}) must be divisible by 32 for Q4_0 block encoding",
h
));
}
let gate_expected_bytes = m
.checked_mul(gate_blocks_per_row)
.and_then(|v| v.checked_mul(18))
.ok_or_else(|| {
anyhow!("qwen35_tree_verify_full_layer_q: gate Q4_0 byte count overflows usize")
})?;
if ffn_weights.gate_q.element_count() != gate_expected_bytes {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q: gate_q has {} bytes, \
expected exactly {} (intermediate_size={} * hidden_size={} Q4_0 encoding: \
{} rows × {} blocks/row × 18 bytes/block)",
ffn_weights.gate_q.element_count(),
gate_expected_bytes,
m,
h,
m,
gate_blocks_per_row
));
}
if ffn_weights.up_q.element_count() != gate_expected_bytes {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q: up_q has {} bytes, \
expected exactly {} (intermediate_size={} * hidden_size={} Q4_0 encoding)",
ffn_weights.up_q.element_count(),
gate_expected_bytes,
m,
h
));
}
let down_blocks_per_row = m
.checked_div(32)
.ok_or_else(|| anyhow!("qwen35_tree_verify_full_layer_q: intermediate_size {} not divisible by 32 (Q4_0 block)", m))?;
if m % 32 != 0 {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q: intermediate_size ({}) must be divisible by 32 for Q4_0 block encoding",
m
));
}
let down_expected_bytes = h
.checked_mul(down_blocks_per_row)
.and_then(|v| v.checked_mul(18))
.ok_or_else(|| {
anyhow!("qwen35_tree_verify_full_layer_q: down Q4_0 byte count overflows usize")
})?;
if ffn_weights.down_q.element_count() != down_expected_bytes {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q: down_q has {} bytes, \
expected exactly {} (hidden_size={} * intermediate_size={} Q4_0 encoding: \
{} rows × {} blocks/row × 18 bytes/block)",
ffn_weights.down_q.element_count(),
down_expected_bytes,
h,
m,
h,
down_blocks_per_row
));
}
let attn_out = qwen35_tree_verify_attention_block(
enc,
device,
registry,
hidden_states_in,
tree_mask,
tree_positions,
k_cache,
v_cache,
weights,
shape.attn,
)
.context("qwen35_tree_verify_full_layer_q: attention block")?;
let ffn_residual = attn_out.clone();
let mut enc2 = device
.command_encoder()
.context("qwen35_tree_verify_full_layer_q: alloc enc2")?;
let rms_out_bytes = seq * h * std::mem::size_of::<f32>();
let post_attn_normed =
super::decode_pool::pooled_alloc_buffer(device, rms_out_bytes, DType::F32, vec![seq, h])
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer_q: alloc post_attn_normed: {e}"))?;
let mut rms_params = super::decode_pool::pooled_alloc_buffer(device, 8, DType::F32, vec![2])
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer_q: alloc rms_params: {e}"))?;
{
let s = rms_params
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer_q: rms_params slice: {e}"))?;
s[0] = shape.attn.rms_norm_eps;
s[1] = h as f32;
}
rms_norm::dispatch_rms_norm(
&mut enc2,
registry,
device.metal_device(),
&attn_out,
&weights.post_attn_norm,
&post_attn_normed,
&rms_params,
shape.attn.tree_seq_len,
shape.attn.hidden_size,
)
.context("qwen35_tree_verify_full_layer_q: post_attn_norm")?;
enc2.memory_barrier();
let gate_buf = apply_linear_projection_f32(
&mut enc2,
registry,
device,
&post_attn_normed,
&ffn_weights.gate_q,
shape.attn.tree_seq_len,
shape.attn.hidden_size,
shape.intermediate_size,
)
.context("qwen35_tree_verify_full_layer_q: gate_proj")?;
let up_buf = apply_linear_projection_f32(
&mut enc2,
registry,
device,
&post_attn_normed,
&ffn_weights.up_q,
shape.attn.tree_seq_len,
shape.attn.hidden_size,
shape.intermediate_size,
)
.context("qwen35_tree_verify_full_layer_q: up_proj")?;
enc2.memory_barrier();
let n_silu_elems = seq.checked_mul(m).ok_or_else(|| {
anyhow!("qwen35_tree_verify_full_layer_q: seq * intermediate overflows usize")
})?;
if n_silu_elems > (u32::MAX as usize) {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q: seq ({}) * intermediate ({}) = {} exceeds u32::MAX",
seq,
m,
n_silu_elems
));
}
let n_silu: u32 = n_silu_elems as u32;
let activated_bytes = n_silu_elems
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| anyhow!("qwen35_tree_verify_full_layer_q: activated_bytes overflow"))?;
let activated_buf = device
.alloc_buffer(activated_bytes, DType::F32, vec![seq, m])
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer_q: alloc activated_buf: {e}"))?;
let mut silu_params =
super::decode_pool::pooled_alloc_buffer(device, 4, DType::U32, vec![1])
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer_q: alloc silu_params: {e}"))?;
silu_params
.as_mut_slice::<u32>()
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer_q: silu_params slice: {e}"))?[0] =
n_silu;
dispatch_silu_mul(
&mut enc2,
registry,
device.metal_device(),
&gate_buf,
&up_buf,
&activated_buf,
&silu_params,
n_silu,
)
.context("qwen35_tree_verify_full_layer_q: dispatch_silu_mul")?;
enc2.memory_barrier();
let ffn_out = apply_linear_projection_f32(
&mut enc2,
registry,
device,
&activated_buf,
&ffn_weights.down_q,
shape.attn.tree_seq_len,
shape.intermediate_size,
shape.attn.hidden_size,
)
.context("qwen35_tree_verify_full_layer_q: down_proj")?;
enc2.memory_barrier();
let hs_elems = seq * h;
let out_bytes = hs_elems * std::mem::size_of::<f32>();
let hidden_states_out = device
.alloc_buffer(out_bytes, DType::F32, vec![seq, h])
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer_q: alloc hidden_states_out: {e}"))?;
elementwise_add(
&mut enc2,
registry,
device.metal_device(),
&ffn_residual,
&ffn_out,
&hidden_states_out,
hs_elems,
DType::F32,
)
.context("qwen35_tree_verify_full_layer_q: residual add")?;
enc2.commit_and_wait()
.context("qwen35_tree_verify_full_layer_q: enc2 terminal commit")?;
Ok(hidden_states_out)
}
#[allow(clippy::too_many_arguments)]
pub fn qwen35_tree_verify_full_layer_q_moe(
enc: mlx_native::CommandEncoder,
device: &MlxDevice,
registry: &mut KernelRegistry,
hidden_states_in: &MlxBuffer,
tree_mask: &MlxBuffer,
tree_positions: &MlxBuffer,
k_cache: &mut MlxBuffer,
v_cache: &mut MlxBuffer,
weights: &FullAttnWeightsGpu,
moe_weights: &super::gpu_ffn::MoeFfnWeightsGpuQ,
shape: Qwen35TreeVerifyFullLayerShapeQMoe,
) -> Result<MlxBuffer> {
if moe_weights.ggml_type_gate_up != GgmlType::Q4_0 {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: ggml_type_gate_up must be Q4_0 \
(got {:?}). Future CFAs will support Q5_K/Q6_K mixed-quant via \
per-projection ggml_type threading.",
moe_weights.ggml_type_gate_up
));
}
if moe_weights.ggml_type_down != GgmlType::Q4_0 {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: ggml_type_down must be Q4_0 \
(got {:?}). Future CFAs will support Q5_K/Q6_K mixed-quant via \
per-projection ggml_type threading.",
moe_weights.ggml_type_down
));
}
if moe_weights.router.dtype() != mlx_native::DType::BF16 {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: router dtype must be BF16 \
(got {:?}). MoeFfnWeightsGpuQ::from_quantized always uploads router as BF16.",
moe_weights.router.dtype()
));
}
if moe_weights.shared_gate_inp.dtype() != mlx_native::DType::BF16 {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: shared_gate_inp dtype must be BF16 \
(got {:?}).",
moe_weights.shared_gate_inp.dtype()
));
}
if moe_weights.shared_gate.dtype() != mlx_native::DType::BF16 {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: shared_gate dtype must be BF16 \
(got {:?}).",
moe_weights.shared_gate.dtype()
));
}
if moe_weights.shared_up.dtype() != mlx_native::DType::BF16 {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: shared_up dtype must be BF16 \
(got {:?}).",
moe_weights.shared_up.dtype()
));
}
if moe_weights.shared_down.dtype() != mlx_native::DType::BF16 {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: shared_down dtype must be BF16 \
(got {:?}).",
moe_weights.shared_down.dtype()
));
}
shape.validate()?;
let h = shape.attn.hidden_size as usize;
let ne = shape.moe.num_experts as usize;
let m_moe = shape.moe.moe_intermediate_size as usize;
let m_sh = shape.moe.shared_intermediate_size as usize;
let expected_router_elems = ne.checked_mul(h).ok_or_else(|| {
anyhow!("qwen35_tree_verify_full_layer_q_moe: router element count overflows usize")
})?;
if moe_weights.router.element_count() != expected_router_elems {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: router has {} BF16 elements, \
expected {} (num_experts={} * hidden_size={}). \
Shape and weights were built from different model configs.",
moe_weights.router.element_count(),
expected_router_elems,
ne,
h
));
}
if moe_weights.num_experts != shape.moe.num_experts {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: weights.num_experts ({}) != \
shape.moe.num_experts ({}). Shape and weights were built from different configs.",
moe_weights.num_experts,
shape.moe.num_experts
));
}
if h % 32 != 0 {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: hidden_size ({}) must be divisible by 32 \
for Q4_0 block encoding",
h
));
}
let gate_blocks_per_row = h / 32;
let expert_gate_expected = ne
.checked_mul(m_moe)
.and_then(|v| v.checked_mul(gate_blocks_per_row))
.and_then(|v| v.checked_mul(18))
.ok_or_else(|| {
anyhow!(
"qwen35_tree_verify_full_layer_q_moe: expert_gate Q4_0 byte count overflows usize"
)
})?;
if moe_weights.expert_gate_q.element_count() != expert_gate_expected {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: expert_gate_q has {} bytes, \
expected exactly {} (num_experts={} * moe_intermediate={} * hidden_size={} Q4_0 encoding: \
{} experts × {} rows/expert × {} blocks/row × 18 bytes/block)",
moe_weights.expert_gate_q.element_count(), expert_gate_expected,
ne, m_moe, h, ne, m_moe, gate_blocks_per_row
));
}
if moe_weights.expert_up_q.element_count() != expert_gate_expected {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: expert_up_q has {} bytes, \
expected exactly {} (same shape as expert_gate_q)",
moe_weights.expert_up_q.element_count(),
expert_gate_expected
));
}
if m_moe % 32 != 0 {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: moe_intermediate_size ({}) must be divisible by 32 \
for Q4_0 block encoding",
m_moe
));
}
let down_blocks_per_row = m_moe / 32;
let expert_down_expected = ne
.checked_mul(h)
.and_then(|v| v.checked_mul(down_blocks_per_row))
.and_then(|v| v.checked_mul(18))
.ok_or_else(|| {
anyhow!(
"qwen35_tree_verify_full_layer_q_moe: expert_down Q4_0 byte count overflows usize"
)
})?;
if moe_weights.expert_down_q.element_count() != expert_down_expected {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: expert_down_q has {} bytes, \
expected exactly {} (num_experts={} * hidden_size={} * moe_intermediate={} Q4_0 encoding: \
{} experts × {} rows/expert × {} blocks/row × 18 bytes/block)",
moe_weights.expert_down_q.element_count(), expert_down_expected,
ne, h, m_moe, ne, h, down_blocks_per_row
));
}
if moe_weights.shared_gate_inp.element_count() != h {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: shared_gate_inp has {} BF16 elements, \
expected {} (hidden_size={})",
moe_weights.shared_gate_inp.element_count(),
h,
h
));
}
let shared_proj_expected = m_sh.checked_mul(h).ok_or_else(|| {
anyhow!("qwen35_tree_verify_full_layer_q_moe: shared_gate element count overflows usize")
})?;
if moe_weights.shared_gate.element_count() != shared_proj_expected {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: shared_gate has {} BF16 elements, \
expected {} (shared_intermediate={} * hidden_size={})",
moe_weights.shared_gate.element_count(),
shared_proj_expected,
m_sh,
h
));
}
if moe_weights.shared_up.element_count() != shared_proj_expected {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: shared_up has {} BF16 elements, \
expected {} (shared_intermediate={} * hidden_size={})",
moe_weights.shared_up.element_count(),
shared_proj_expected,
m_sh,
h
));
}
if moe_weights.shared_down.element_count() != shared_proj_expected {
return Err(anyhow!(
"qwen35_tree_verify_full_layer_q_moe: shared_down has {} BF16 elements, \
expected {} (hidden_size={} * shared_intermediate={})",
moe_weights.shared_down.element_count(),
shared_proj_expected,
h,
m_sh
));
}
let attn_out = qwen35_tree_verify_attention_block(
enc,
device,
registry,
hidden_states_in,
tree_mask,
tree_positions,
k_cache,
v_cache,
weights,
shape.attn,
)
.context("qwen35_tree_verify_full_layer_q_moe: attention block")?;
let ffn_residual = attn_out.clone();
let mut enc2 = device
.command_encoder()
.context("qwen35_tree_verify_full_layer_q_moe: alloc enc2")?;
let seq = shape.attn.tree_seq_len as usize;
let rms_out_bytes = seq * h * std::mem::size_of::<f32>();
let post_attn_normed = super::decode_pool::pooled_alloc_buffer(
device,
rms_out_bytes,
mlx_native::DType::F32,
vec![seq, h],
)
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer_q_moe: alloc post_attn_normed: {e}"))?;
let mut rms_params =
super::decode_pool::pooled_alloc_buffer(device, 8, mlx_native::DType::F32, vec![2])
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer_q_moe: alloc rms_params: {e}"))?;
{
let s = rms_params
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("qwen35_tree_verify_full_layer_q_moe: rms_params slice: {e}"))?;
s[0] = shape.attn.rms_norm_eps;
s[1] = h as f32;
}
rms_norm::dispatch_rms_norm(
&mut enc2,
registry,
device.metal_device(),
&attn_out,
&weights.post_attn_norm,
&post_attn_normed,
&rms_params,
shape.attn.tree_seq_len,
shape.attn.hidden_size,
)
.context("qwen35_tree_verify_full_layer_q_moe: post_attn_norm")?;
enc2.memory_barrier();
enc2.commit_and_wait()
.context("qwen35_tree_verify_full_layer_q_moe: enc2 commit_and_wait")?;
let moe_ffn_shape = super::ffn::MoeFfnShape {
hidden_size: shape.moe.hidden_size,
num_experts: shape.moe.num_experts,
num_experts_per_tok: shape.moe.num_experts_per_tok,
moe_intermediate_size: shape.moe.moe_intermediate_size,
shared_intermediate_size: shape.moe.shared_intermediate_size,
};
let hidden_states_out = super::gpu_ffn::build_moe_ffn_layer_gpu_q(
device,
registry,
&post_attn_normed,
moe_weights,
moe_ffn_shape,
Some(&ffn_residual),
)
.context("qwen35_tree_verify_full_layer_q_moe: build_moe_ffn_layer_gpu_q")?;
Ok(hidden_states_out)
}
#[allow(clippy::too_many_arguments)]
pub fn apply_flash_attn_prefill_seq_major(
device: &MlxDevice,
registry: &mut KernelRegistry,
q_seq_major: &MlxBuffer,
k_seq_major: &MlxBuffer,
v_seq_major: &MlxBuffer,
seq_len: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
fa_arena: Option<&mut crate::inference::models::qwen35::FaPrefillArena>,
) -> Result<MlxBuffer> {
if head_dim != 256 {
return Err(anyhow!(
"apply_flash_attn_prefill_seq_major: head_dim must be 256 \
(D=256 dispatcher); got {head_dim}. Other head_dims need a \
different mlx-native dispatcher (D=64 / D=512) or a new port."
));
}
let seq = seq_len 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 out_seq = device
.alloc_buffer(out_elems * 4, DType::F32, vec![seq, nh, d])
.map_err(|e| anyhow!("alloc out_seq: {e}"))?;
if let Some(arena) = fa_arena {
let mut enc = device
.command_encoder()
.context("FA prefill bridge encoder")?;
apply_flash_attn_prefill_seq_major_into(
&mut enc,
device,
registry,
q_seq_major,
k_seq_major,
v_seq_major,
&out_seq,
seq_len,
n_heads,
n_kv_heads,
head_dim,
arena,
)?;
enc.commit_labeled("fa.prefill_bridge");
if std::env::var("HF2Q_DUMP_FA_BF16").as_deref() == Ok("1") {
let mut sync_enc = device
.command_encoder()
.context("FA bridge: dump sync encoder")?;
sync_enc
.commit_and_wait()
.context("FA bridge: dump sync commit_and_wait")?;
let layer_idx = super::dump_bisect::current_layer_idx();
let step_idx = super::dump_bisect::current_step_idx();
for (label, buf) in [
("q_bf16_hm", &arena.q_bf16_hm),
("k_bf16_hm", &arena.k_bf16_hm),
("v_bf16_hm", &arena.v_bf16_hm),
("out_bf16_hm", &arena.out_bf16_hm),
] {
let bytes = buf
.as_slice::<u8>()
.map_err(|e| anyhow!("FA bridge: dump as_slice {label}: {e}"))?;
let path = format!(
"/tmp/hf2q_fa_bf16_step{:04}_layer{:03}_{}.bin",
step_idx,
layer_idx.unwrap_or(999),
label,
);
std::fs::write(&path, bytes)
.with_context(|| format!("FA bridge: dump write {}", path))?;
}
tracing::info!(
"iter-17 dump: wrote 4× arena bf16 buffers for layer {:?} step {}",
layer_idx,
step_idx,
);
}
} else {
let q_bf16_seq = device
.alloc_buffer(q_elems * 2, DType::BF16, vec![seq, nh, d])
.map_err(|e| anyhow!("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!("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!("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!("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!("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!("alloc v_bf16_hm: {e}"))?;
let mut out_bf16_hm = device
.alloc_buffer(out_elems * 2, DType::BF16, vec![1, nh, seq, d])
.map_err(|e| anyhow!("alloc out_bf16_hm: {e}"))?;
let mut enc = device
.command_encoder()
.context("FA prefill bridge encoder")?;
cast(
&mut enc,
registry,
device.metal_device(),
q_seq_major,
&q_bf16_seq,
q_elems,
CastDirection::F32ToBF16,
)
.context("FA bridge: cast Q F32→BF16")?;
enc.memory_barrier();
permute_021_bf16(
&mut enc,
registry,
device.metal_device(),
&q_bf16_seq,
&q_bf16_hm,
seq,
nh,
d,
)
.context("FA bridge: permute_021 Q [seq, nh, d] → [nh, seq, d]")?;
cast(
&mut enc,
registry,
device.metal_device(),
k_seq_major,
&k_bf16_seq,
k_elems,
CastDirection::F32ToBF16,
)
.context("FA bridge: cast K F32→BF16")?;
enc.memory_barrier();
permute_021_bf16(
&mut enc,
registry,
device.metal_device(),
&k_bf16_seq,
&k_bf16_hm,
seq,
nkv,
d,
)
.context("FA bridge: permute_021 K [seq, nkv, d] → [nkv, seq, d]")?;
cast(
&mut enc,
registry,
device.metal_device(),
v_seq_major,
&v_bf16_seq,
v_elems,
CastDirection::F32ToBF16,
)
.context("FA bridge: cast V F32→BF16")?;
enc.memory_barrier();
permute_021_bf16(
&mut enc,
registry,
device.metal_device(),
&v_bf16_seq,
&v_bf16_hm,
seq,
nkv,
d,
)
.context("FA bridge: permute_021 V [seq, nkv, d] → [nkv, seq, d]")?;
enc.memory_barrier();
let scale = 1.0 / (d as f32).sqrt();
dispatch_flash_attn_prefill_bf16_d256(
&mut enc,
device,
registry,
&q_bf16_hm,
&k_bf16_hm,
&v_bf16_hm,
None,
&mut out_bf16_hm,
&FlashAttnPrefillParams {
n_heads,
n_kv_heads,
head_dim,
seq_len_q: seq_len,
seq_len_k: seq_len,
batch: 1,
scale,
do_causal: true,
},
)
.context("FA bridge: dispatch_flash_attn_prefill_bf16_d256")?;
enc.memory_barrier();
permute_021_bf16_to_f32(
&mut enc,
registry,
device.metal_device(),
&out_bf16_hm,
&out_seq,
nh,
seq,
d,
)
.context("FA bridge: permute_021_bf16_to_f32 out [nh, seq, d] → [seq, nh, d] F32")?;
enc.commit_and_wait()
.context("FA bridge: commit+wait flash_attn_prefill")?;
}
Ok(out_seq)
}
#[allow(clippy::too_many_arguments)]
pub fn apply_flash_attn_prefill_seq_major_resume(
device: &MlxDevice,
registry: &mut KernelRegistry,
q_seq_major: &MlxBuffer,
slot_k_head_major: &MlxBuffer,
slot_v_head_major: &MlxBuffer,
seq_len: u32,
cur_len: u32,
kv_seq_len: u32,
kv_capacity: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
) -> Result<MlxBuffer> {
if head_dim != 256 {
return Err(anyhow!(
"apply_flash_attn_prefill_seq_major_resume: head_dim must be 256 \
(D=256 dispatcher); got {head_dim}. Other head_dims need a \
different mlx-native dispatcher (D=64 / D=512) or a new port."
));
}
if cur_len + seq_len != kv_seq_len {
return Err(anyhow!(
"apply_flash_attn_prefill_seq_major_resume: cur_len ({cur_len}) + \
seq_len ({seq_len}) != kv_seq_len ({kv_seq_len}) — \
append-prefill semantics require cur_len + seq_len == kv_seq_len."
));
}
if kv_seq_len > kv_capacity {
return Err(anyhow!(
"apply_flash_attn_prefill_seq_major_resume: kv_seq_len \
({kv_seq_len}) > kv_capacity ({kv_capacity}) — slot overflow."
));
}
let seq = seq_len as usize;
let nh = n_heads as usize;
let nkv = n_kv_heads as usize;
let d = head_dim as usize;
let cap = kv_capacity as usize;
let q_elems = seq * nh * d;
let kv_slot_elems = nkv * cap * 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!("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!("alloc q_bf16_hm: {e}"))?;
let k_bf16_slot = device
.alloc_buffer(kv_slot_elems * 2, DType::BF16, vec![1, nkv, cap, d])
.map_err(|e| anyhow!("alloc k_bf16_slot: {e}"))?;
let v_bf16_slot = device
.alloc_buffer(kv_slot_elems * 2, DType::BF16, vec![1, nkv, cap, d])
.map_err(|e| anyhow!("alloc v_bf16_slot: {e}"))?;
let mut out_bf16_hm = device
.alloc_buffer(out_elems * 2, DType::BF16, vec![1, nh, seq, d])
.map_err(|e| anyhow!("alloc out_bf16_hm: {e}"))?;
let out_seq = device
.alloc_buffer(out_elems * 4, DType::F32, vec![seq, nh, d])
.map_err(|e| anyhow!("alloc out_seq: {e}"))?;
let mut enc = device
.command_encoder()
.context("FA resume bridge encoder")?;
cast(
&mut enc,
registry,
device.metal_device(),
q_seq_major,
&q_bf16_seq,
q_elems,
CastDirection::F32ToBF16,
)
.context("FA resume bridge: cast Q F32→BF16")?;
enc.memory_barrier();
permute_021_bf16(
&mut enc,
registry,
device.metal_device(),
&q_bf16_seq,
&q_bf16_hm,
seq,
nh,
d,
)
.context("FA resume bridge: permute_021 Q [seq, nh, d] → [nh, seq, d]")?;
cast(
&mut enc,
registry,
device.metal_device(),
slot_k_head_major,
&k_bf16_slot,
kv_slot_elems,
CastDirection::F32ToBF16,
)
.context("FA resume bridge: cast slot K F32→BF16")?;
enc.memory_barrier();
cast(
&mut enc,
registry,
device.metal_device(),
slot_v_head_major,
&v_bf16_slot,
kv_slot_elems,
CastDirection::F32ToBF16,
)
.context("FA resume bridge: cast slot V F32→BF16")?;
enc.memory_barrier();
let scale = 1.0 / (d as f32).sqrt();
dispatch_flash_attn_prefill_bf16_d256_resume(
&mut enc,
device,
registry,
&q_bf16_hm,
&k_bf16_slot,
&v_bf16_slot,
&mut out_bf16_hm,
&FlashAttnPrefillResumeParams {
n_heads,
n_kv_heads,
head_dim,
seq_len_q: seq_len,
seq_len_k: kv_seq_len,
batch: 1,
scale,
do_causal: true,
q_offset_in_k: cur_len,
kv_capacity,
},
)
.context("FA resume bridge: dispatch_flash_attn_prefill_bf16_d256_resume")?;
enc.memory_barrier();
permute_021_bf16_to_f32(
&mut enc,
registry,
device.metal_device(),
&out_bf16_hm,
&out_seq,
nh,
seq,
d,
)
.context(
"FA resume bridge: permute_021_bf16_to_f32 out [nh, seq, d] → \
[seq, nh, d] F32",
)?;
enc.commit_and_wait()
.context("FA resume bridge: commit+wait")?;
Ok(out_seq)
}
#[allow(clippy::too_many_arguments)]
pub fn apply_flash_attn_prefill_seq_major_resume_via_tq_cache_for_slot(
device: &MlxDevice,
registry: &mut KernelRegistry,
slot: &super::kv_cache::FullAttnKvSlot,
q_seq_major: &MlxBuffer,
seq_len: u32,
cur_len: u32,
kv_seq_len: u32,
cache_capacity: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
slot_id: SlotId,
) -> Result<MlxBuffer> {
if slot.tq.is_none() {
return Err(anyhow!(
"apply_flash_attn_prefill_seq_major_resume_via_tq_cache: slot.tq is None — \
slot was not constructed in TQ-active mode (HybridKvCache::new_with_options \
tq_kv_active=true required). Caller routing bug."
));
}
let mut enc = device
.command_encoder()
.context("apply_flash_attn_prefill_seq_major_resume_via_tq_cache: dequant encoder")?;
let temp_k = slot
.dequant_seq_to_temp_f32_unrotated_for_slot(
true,
kv_seq_len,
0,
cache_capacity,
n_kv_heads,
head_dim,
slot_id,
&mut enc,
registry,
device,
)
.context("dequant K seq → temp F32 unrotated")?;
let temp_v = slot
.dequant_seq_to_temp_f32_unrotated_for_slot(
false,
kv_seq_len,
0,
cache_capacity,
n_kv_heads,
head_dim,
slot_id,
&mut enc,
registry,
device,
)
.context("dequant V seq → temp F32 unrotated")?;
enc.commit_labeled("layer.full_attn.tq_bridge_dequant");
apply_flash_attn_prefill_seq_major_resume(
device,
registry,
q_seq_major,
&temp_k,
&temp_v,
seq_len,
cur_len,
kv_seq_len,
kv_seq_len,
n_heads,
n_kv_heads,
head_dim,
)
.context("apply_flash_attn_prefill_seq_major_resume (TQ-decoded path)")
}
#[allow(clippy::too_many_arguments)]
pub fn apply_flash_attn_prefill_seq_major_resume_via_tq_cache(
device: &MlxDevice,
registry: &mut KernelRegistry,
slot: &super::kv_cache::FullAttnKvSlot,
q_seq_major: &MlxBuffer,
seq_len: u32,
cur_len: u32,
kv_seq_len: u32,
cache_capacity: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
) -> Result<MlxBuffer> {
apply_flash_attn_prefill_seq_major_resume_via_tq_cache_for_slot(
device,
registry,
slot,
q_seq_major,
seq_len,
cur_len,
kv_seq_len,
cache_capacity,
n_heads,
n_kv_heads,
head_dim,
SlotId(0),
)
}
const QWEN35_TQ_DIRECT_PREFILL_MAX_QUERIES: u32 = 32;
#[allow(clippy::too_many_arguments)]
pub(super) fn apply_tq_prefill_seq_major_resume_direct_for_slot(
device: &MlxDevice,
registry: &mut KernelRegistry,
slot: &super::kv_cache::FullAttnKvSlot,
q_seq_major: &MlxBuffer,
seq_len: u32,
cur_len: u32,
kv_seq_len: u32,
cache_capacity: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
slot_id: SlotId,
) -> Result<MlxBuffer> {
anyhow::ensure!(
seq_len > 0 && seq_len <= QWEN35_TQ_DIRECT_PREFILL_MAX_QUERIES,
"direct TQ prefill seq_len={} outside 1..={}",
seq_len,
QWEN35_TQ_DIRECT_PREFILL_MAX_QUERIES
);
anyhow::ensure!(
cur_len.saturating_add(seq_len) == kv_seq_len,
"direct TQ prefill cur_len({cur_len}) + seq_len({seq_len}) != kv_seq_len({kv_seq_len})"
);
let tq = slot
.tq
.as_ref()
.ok_or_else(|| anyhow!("direct TQ prefill requires a TQ-active full-attention slot"))?;
let output_elems = (seq_len as usize) * (n_heads as usize) * (head_dim as usize);
let output = device
.alloc_buffer(
output_elems * 4,
DType::F32,
vec![seq_len as usize, n_heads as usize, head_dim as usize],
)
.map_err(|error| anyhow!("direct TQ prefill output allocation: {error}"))?;
let tmp_bytes = mlx_native::ops::flash_attn_vec_tq_hb::tmp_buffer_bytes(
seq_len.saturating_mul(n_heads),
head_dim,
);
let tmp =
super::decode_pool::pooled_alloc_buffer(device, tmp_bytes, DType::F32, vec![tmp_bytes / 4])
.map_err(|error| anyhow!("direct TQ prefill scratch allocation: {error}"))?;
let mut slot_ids = super::decode_pool::pooled_alloc_buffer(
device,
(seq_len as usize) * std::mem::size_of::<u32>(),
DType::U32,
vec![seq_len as usize],
)
.map_err(|error| anyhow!("direct TQ prefill slot-id allocation: {error}"))?;
slot_ids
.as_mut_slice::<u32>()
.map_err(|error| anyhow!("direct TQ prefill slot-id mapping: {error}"))?
.fill(slot_id.0);
let mut sequence_positions = super::decode_pool::pooled_alloc_buffer(
device,
(seq_len as usize) * std::mem::size_of::<u32>(),
DType::U32,
vec![seq_len as usize],
)
.map_err(|error| anyhow!("direct TQ prefill position allocation: {error}"))?;
for (index, position) in sequence_positions
.as_mut_slice::<u32>()
.map_err(|error| anyhow!("direct TQ prefill position mapping: {error}"))?
.iter_mut()
.enumerate()
{
*position = cur_len + index as u32;
}
let codebook_bits = crate::debug::INVESTIGATION_ENV.tq_codebook_bits;
let codebook_bits = if matches!(codebook_bits, 5 | 6 | 8) {
codebook_bits
} else {
8
};
let params = mlx_native::ops::flash_attn_vec_tq_hb::FlashAttnVecTqHbParams {
num_heads: n_heads,
num_kv_heads: n_kv_heads,
head_dim,
kv_seq_len,
kv_capacity: cache_capacity,
scale: 1.0 / (head_dim as f32).sqrt(),
mask_type: 0,
sliding_window: 0,
softcap: 0.0,
ring_start: 0,
scale_factor_d512: 1.0,
codebook_bits,
fuse_fwht_pre: 1,
nsg: mlx_native::ops::flash_attn_vec_tq_hb::compute_nsg(kv_seq_len),
};
let mut encoder = device
.command_encoder()
.context("direct TQ prefill encoder")?;
mlx_native::ops::flash_attn_vec_tq_hb::flash_attn_vec_tq_hb_batched(
&mut encoder,
registry,
device,
seq_len,
q_seq_major,
&tq.k_packed,
&tq.k_norms,
&tq.v_packed,
&tq.v_norms,
&output,
&tmp,
&slot_ids,
&sequence_positions,
¶ms,
)
.context("direct TQ prefill batched SDPA")?;
encoder.commit_labeled("layer.full_attn.tq_direct_prefill");
Ok(output)
}
#[cfg(test)]
#[allow(clippy::too_many_arguments)]
pub(super) fn apply_tq_prefill_seq_major_resume_direct(
device: &MlxDevice,
registry: &mut KernelRegistry,
slot: &super::kv_cache::FullAttnKvSlot,
q_seq_major: &MlxBuffer,
seq_len: u32,
cur_len: u32,
kv_seq_len: u32,
cache_capacity: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
) -> Result<MlxBuffer> {
apply_tq_prefill_seq_major_resume_direct_for_slot(
device,
registry,
slot,
q_seq_major,
seq_len,
cur_len,
kv_seq_len,
cache_capacity,
n_heads,
n_kv_heads,
head_dim,
SlotId(0),
)
}
#[allow(clippy::too_many_arguments)]
pub fn apply_sdpa_with_kv_cache(
device: &MlxDevice,
registry: &mut KernelRegistry,
q_seq_major: &MlxBuffer,
k_seq_major: &MlxBuffer,
v_seq_major: &MlxBuffer,
slot: &mut FullAttnKvSlot,
seq_len: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
max_seq_len: u32,
fa_arena: Option<&mut crate::inference::models::qwen35::FaPrefillArena>,
slot_id: SlotId,
) -> Result<MlxBuffer> {
let seq = seq_len as usize;
let nh = n_heads as usize;
let _nkv = n_kv_heads as usize;
let d = head_dim as usize;
let max_sl = max_seq_len as usize;
assert!(
(slot_id.0 as usize) < slot.current_len.len(),
"apply_sdpa_with_kv_cache: slot_id={} out of range (slot.current_len.len()={}) \
— bounds check at forward_gpu entry regressed (ADR-040 §6.1.5)",
slot_id.0,
slot.current_len.len(),
);
let cur_len = slot.current_len[slot_id.0 as usize] as usize;
let kv_write_tokens = (seq).min(max_sl.saturating_sub(cur_len));
let kv_seq_len = (cur_len + kv_write_tokens).min(max_sl) as u32;
let out_buf = super::decode_pool::pooled_alloc_buffer(
device,
nh * seq * d * 4,
DType::F32,
vec![1, nh, seq, d],
)
.map_err(|e| anyhow!("alloc sdpa kv-cache output: {e}"))?;
if seq == 1 && head_dim % 32 == 0 {
let mut enc = device
.command_encoder()
.context("enc kv-cache+sdpa decode")?;
if kv_write_tokens > 0 {
write_kv_with_optional_tq_encode(
&mut enc,
registry,
device,
k_seq_major,
v_seq_major,
slot,
n_kv_heads,
head_dim,
max_seq_len,
cur_len as u32,
kv_write_tokens as u32,
slot_id,
)
.context("kv_cache_copy kv-cache decode (iter-15 helper)")?;
enc.memory_barrier();
}
if head_dim == 256 || head_dim == 512 {
let fa_tmp = super::decode_pool::pooled_alloc_buffer(
device,
flash_attn_vec_tmp_bytes(n_heads, head_dim),
DType::F32,
vec![flash_attn_vec_tmp_bytes(n_heads, head_dim) / 4],
)
.map_err(|e| anyhow!("alloc flash_attn_vec tmp: {e}"))?;
dispatch_decode_sdpa_with_optional_tq(
&mut enc,
registry,
device,
q_seq_major,
slot,
&out_buf,
&fa_tmp,
n_heads,
n_kv_heads,
head_dim,
kv_seq_len,
max_seq_len,
slot_id,
)
.context("flash_attn_vec kv-cache (FA-layer decode iter-15)")?;
} else {
let kbuf = slot.k.as_ref().expect(
"dispatch_sdpa_decode F32 head_dim fallback: slot.k is None — \
iter-34 alloc/SDPA gating invariant regressed (TQ requires \
head_dim ∈ {256,512}; this fallback should never see slot.k=None).",
);
let vbuf = slot
.v
.as_ref()
.expect("dispatch_sdpa_decode F32: slot.v is None");
let (so_off, so_n) =
slot_k_v_region_for_full_attn(slot_id, n_kv_heads, max_seq_len, head_dim);
let kbuf_view = kbuf.slice_view(so_off, so_n);
let vbuf_view = vbuf.slice_view(so_off, so_n);
dispatch_sdpa_decode(
&mut enc,
registry,
device,
q_seq_major,
&kbuf_view,
&vbuf_view,
&out_buf,
n_heads,
n_kv_heads,
head_dim,
kv_seq_len,
max_seq_len,
1.0 / (d as f32).sqrt(),
)
.context("sdpa_decode kv-cache (head_dim fallback)")?;
}
enc.commit_labeled("layer.full_attn.sdpa_kv");
} else {
if kv_write_tokens > 0 {
let _w5b9_kv_dl_copy = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::FaSdpaKvDownloadCopy,
);
let mut enc = device
.command_encoder()
.context("enc kv_cache_copy_seq_dual prefill")?;
write_kv_with_optional_tq_encode(
&mut enc,
registry,
device,
k_seq_major,
v_seq_major,
slot,
n_kv_heads,
head_dim,
max_seq_len,
cur_len as u32,
kv_write_tokens as u32,
slot_id,
)
.context("kv_cache_copy_seq_f32_dual prefill (iter-15 helper)")?;
enc.commit_labeled("layer.full_attn.kv_cache_write");
}
let new_path_eligible = head_dim == 256 && cur_len == 0 && seq_len >= 16;
let fa_trace = std::env::var("HF2Q_FA_TRACE").as_deref() == Ok("1");
if fa_trace {
eprintln!(
"[FA_TRACE] seq_len={} cur_len={} kv_seq_len={} head_dim={} new_eligible={} slot.k={} slot.v={} slot.tq={}",
seq_len, cur_len, (cur_len as u32).saturating_add(seq_len), head_dim,
new_path_eligible,
slot.k.is_some(), slot.v.is_some(), slot.tq.is_some(),
);
}
if new_path_eligible {
let _w5b10_kernel = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::FaSdpaKernel,
);
let out_uploaded = apply_flash_attn_prefill_seq_major(
device,
registry,
q_seq_major,
k_seq_major,
v_seq_major,
seq_len,
n_heads,
n_kv_heads,
head_dim,
fa_arena,
)?;
let new_len = kv_seq_len;
slot.current_len[slot_id.0 as usize] = new_len;
return Ok(out_uploaded);
}
let vec_small_path_eligible = head_dim == 256
&& cur_len > 0
&& seq_len >= 2
&& seq_len <= 8
&& slot.k.is_some()
&& slot.v.is_some()
&& std::env::var("HF2Q_NO_VEC_SMALL_PATH").as_deref() != Ok("1");
if fa_trace {
eprintln!(
"[FA_TRACE] vec_small_path_eligible={} (engages BEFORE resume)",
vec_small_path_eligible,
);
}
if vec_small_path_eligible {
let _w5b10_kernel = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::FaSdpaKernel,
);
let kbuf = slot
.k
.as_ref()
.expect("vec_small_path: slot.k.is_some() guard above passed");
let vbuf = slot
.v
.as_ref()
.expect("vec_small_path: slot.v.is_some() guard above passed");
let (so_off, so_n) =
slot_k_v_region_for_full_attn(slot_id, n_kv_heads, max_seq_len, head_dim);
let kbuf_view = kbuf.slice_view(so_off, so_n);
let vbuf_view = vbuf.slice_view(so_off, so_n);
let seq = seq_len as usize;
let nh = n_heads as usize;
let d = head_dim as usize;
let q_hm = device
.alloc_buffer(seq * nh * d * 4, DType::F32, vec![nh, seq, d])
.map_err(|e| anyhow!("vec_small_path: alloc q_hm: {e}"))?;
let tmp_bytes = flash_attn_vec_tmp_bytes_with_qL(n_heads, head_dim, seq_len);
let tmp_elems = tmp_bytes / 4;
let tmp_buf = device
.alloc_buffer(tmp_bytes, DType::F32, vec![tmp_elems])
.map_err(|e| anyhow!("vec_small_path: alloc tmp: {e}"))?;
let mut enc = device
.command_encoder()
.context("vec_small_path: command_encoder")?;
permute_021_f32(
&mut enc,
registry,
device.metal_device(),
q_seq_major,
&q_hm,
seq,
nh,
d,
)
.context("vec_small_path: permute Q seq->head major")?;
enc.memory_barrier();
let params = FlashAttnVecParams {
num_heads: n_heads,
num_kv_heads: n_kv_heads,
head_dim,
kv_seq_len,
kv_capacity: max_seq_len,
scale: 1.0 / (d as f32).sqrt(),
mask_type: 1, sliding_window: 0,
softcap: 0.0,
q_seq_len: seq_len,
};
flash_attn_vec(
&mut enc, registry, device, &q_hm, &kbuf_view, &vbuf_view, &out_buf, &tmp_buf,
¶ms,
)
.context("vec_small_path: flash_attn_vec dispatch")?;
enc.commit_and_wait_labeled("layer.full_attn.vec_small_path")
.context("vec_small_path: commit")?;
slot.current_len[slot_id.0 as usize] = kv_seq_len;
return Ok(out_buf);
}
let resume_path_eligible = head_dim == 256 && cur_len > 0 && kv_seq_len >= 16;
if fa_trace {
eprintln!(
"[FA_TRACE] resume_eligible={} (will engage if true; else fallback)",
resume_path_eligible
);
}
if resume_path_eligible {
let _w5b10_kernel_resume = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::FaSdpaKernel,
);
let out_uploaded = if let (Some(kbuf), Some(vbuf)) = (slot.k.as_ref(), slot.v.as_ref())
{
let (so_off, so_n) =
slot_k_v_region_for_full_attn(slot_id, n_kv_heads, max_seq_len, head_dim);
let kbuf_view = kbuf.slice_view(so_off, so_n);
let vbuf_view = vbuf.slice_view(so_off, so_n);
apply_flash_attn_prefill_seq_major_resume(
device,
registry,
q_seq_major,
&kbuf_view,
&vbuf_view,
seq_len,
cur_len as u32,
kv_seq_len,
max_seq_len,
n_heads,
n_kv_heads,
head_dim,
)?
} else {
if seq_len <= QWEN35_TQ_DIRECT_PREFILL_MAX_QUERIES {
apply_tq_prefill_seq_major_resume_direct_for_slot(
device,
registry,
slot,
q_seq_major,
seq_len,
cur_len as u32,
kv_seq_len,
max_seq_len,
n_heads,
n_kv_heads,
head_dim,
slot_id,
)?
} else {
apply_flash_attn_prefill_seq_major_resume_via_tq_cache_for_slot(
device,
registry,
slot,
q_seq_major,
seq_len,
cur_len as u32,
kv_seq_len,
max_seq_len,
n_heads,
n_kv_heads,
head_dim,
slot_id,
)?
}
};
let new_len = kv_seq_len;
slot.current_len[slot_id.0 as usize] = new_len;
return Ok(out_uploaded);
}
let q_gpu = {
let _w5b9_q_dl_perm_ul = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::FaSdpaQDownloadPermuteUpload,
);
let q_cpu = download_f32(q_seq_major)?;
let q_hm = permute_seq_head_dim_to_head_seq_dim_cpu(&q_cpu, seq, nh, d);
upload_f32(&q_hm, device)?
};
{
let _w5b9_kernel = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::FaSdpaKernel,
);
let params = SdpaParams {
n_heads,
n_kv_heads,
head_dim,
seq_len,
kv_seq_len,
scale: 1.0 / (d as f32).sqrt(),
kv_capacity: max_seq_len,
do_causal: true,
};
let mut enc = device
.command_encoder()
.context("enc sdpa kv-cache prefill")?;
let kbuf = slot.k.as_ref().expect(
"sdpa F32 head_dim fallback prefill: slot.k is None — \
iter-34 alloc/SDPA gating invariant regressed (this \
fallback should only see legacy F32 fixtures with slot.k=Some).",
);
let vbuf = slot.v.as_ref().expect("sdpa F32 prefill: slot.v is None");
let (so_off, so_n) =
slot_k_v_region_for_full_attn(slot_id, n_kv_heads, max_seq_len, head_dim);
let kbuf_view = kbuf.slice_view(so_off, so_n);
let vbuf_view = vbuf.slice_view(so_off, so_n);
sdpa(
&mut enc, registry, device, &q_gpu, &kbuf_view, &vbuf_view, &out_buf, ¶ms, 1,
)
.context("sdpa with kv cache prefill")?;
enc.commit_and_wait_labeled("layer.full_attn.sdpa_legacy_prefill")
.context("commit sdpa kv-cache prefill")?;
}
let out_uploaded = {
let _w5b9_out_dl_perm_ul = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::FaSdpaOutDownloadPermuteUpload,
);
let out_hm_cpu = download_f32(&out_buf)?;
let mut out_sm = vec![0.0f32; seq * nh * d];
for h in 0..nh {
for t in 0..seq {
let src = (h * seq + t) * d;
let dst = (t * nh + h) * d;
out_sm[dst..dst + d].copy_from_slice(&out_hm_cpu[src..src + d]);
}
}
upload_f32(&out_sm, device)?
};
let new_len = kv_seq_len;
slot.current_len[slot_id.0 as usize] = new_len;
return Ok(out_uploaded);
}
slot.current_len[slot_id.0 as usize] = kv_seq_len;
Ok(out_buf)
}
#[allow(clippy::too_many_arguments)]
pub fn build_gated_attn_layer(
device: &MlxDevice,
registry: &mut KernelRegistry,
x: &MlxBuffer,
positions: &MlxBuffer,
weights_gpu: &FullAttnWeightsGpu,
kv_cache_slot: Option<&mut FullAttnKvSlot>,
max_seq_len: u32,
seq_len: u32,
hidden_size: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
rotary_dim: u32,
freq_base: f32,
mrope_section: [u32; 4],
rms_norm_eps: f32,
fa_arena: Option<&mut crate::inference::models::qwen35::FaPrefillArena>,
fa_proj_arena: Option<&mut crate::inference::models::qwen35::FaProjectionsArena>,
mut out_seq_hold: Option<&mut Vec<MlxBuffer>>,
mut layer_session: Option<&mut mlx_native::EncoderSession>,
slot_id: SlotId,
) -> Result<MlxBuffer> {
let slot_idx = slot_id.0 as usize;
let cur_len_for_arena = kv_cache_slot
.as_deref()
.map(|s| {
assert!(
slot_idx < s.current_len.len(),
"build_gated_attn_layer: slot_id={} out of range (slot.current_len.len()={}) \
— bounds check at forward_gpu entry regressed (ADR-040 §6.1.5)",
slot_id.0,
s.current_len.len(),
);
s.current_len[slot_idx]
})
.unwrap_or(0);
let gpu_only_prefill_resume = fa_proj_arena.is_some()
&& kv_cache_slot.is_some()
&& seq_len > 1
&& head_dim == 256
&& cur_len_for_arena > 0;
let use_arena = fa_arena.is_some() && seq_len > 1 && head_dim == 256 && cur_len_for_arena == 0;
let q_total = n_heads * head_dim;
let kv_total = n_kv_heads * head_dim;
let allow_vec_small_in_fused =
std::env::var("HF2Q_NO_FUSED_STAGE_AB_VEC").as_deref() != Ok("1");
let use_fused_stage_ab = use_arena
&& fa_proj_arena.is_some()
&& kv_cache_slot
.as_deref()
.map(|s| {
let cur = s.current_len[slot_idx];
cur == 0
|| (allow_vec_small_in_fused
&& cur > 0
&& seq_len < 16
&& head_dim == 256
&& s.k.is_some()
&& s.v.is_some())
})
.unwrap_or(false)
&& !super::dump_bisect::is_enabled();
let use_proj_arena = fa_proj_arena.is_some() && seq_len > 1;
if let Some(ref arena) = fa_proj_arena {
arena
.validate_fits(seq_len, hidden_size, n_heads, n_kv_heads, head_dim)
.context("FaProjectionsArena shape mismatch")?;
}
let mut fa_proj_arena = fa_proj_arena;
let mut fa_arena = fa_arena;
let mut kv_cache_slot = kv_cache_slot;
let mut attn_out_fused: Option<MlxBuffer> = None;
let mut fused_stage_a_enc: Option<LayerEncoder<'_>> = None;
let (x_norm, q_flat, k_flat, v_flat, gate_flat, q_normed, k_normed, q_rope, k_rope) =
if let Some(arena) = fa_proj_arena
.as_mut()
.map(|a| &mut **a)
.filter(|_| use_proj_arena)
{
let _w5b9_ops1to4 = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::FaOps1to4,
);
let mut enc = LayerEncoder::from_session_or_plain(device, layer_session.as_deref_mut())
.context("enc ops1-4")?;
apply_pre_attn_rms_norm_into(
enc.encoder(),
registry,
device,
x,
weights_gpu,
&arena.x_norm_buf,
&arena.pre_norm_params_buf,
seq_len,
hidden_size,
)?;
enc.encoder().memory_barrier();
const Q4_0_BLOCK_BYTES: usize = 18;
const Q4_0_BLOCK_VALUES: u32 = 32;
let q_w_bytes_expected =
(q_total as usize) * (hidden_size / Q4_0_BLOCK_VALUES) as usize * Q4_0_BLOCK_BYTES;
let kv_w_bytes_expected =
(kv_total as usize) * (hidden_size / Q4_0_BLOCK_VALUES) as usize * Q4_0_BLOCK_BYTES;
let is_q4_0 = |buf: &MlxBuffer, expected: usize| {
buf.dtype() == DType::U8 && buf.byte_len() == expected
};
let use_fused_qkvg = std::env::var("HF2Q_FUSED_QKVG").as_deref() == Ok("1")
&& hidden_size % Q4_0_BLOCK_VALUES == 0
&& is_q4_0(&weights_gpu.wq, q_w_bytes_expected)
&& is_q4_0(&weights_gpu.w_gate, q_w_bytes_expected)
&& is_q4_0(&weights_gpu.wk, kv_w_bytes_expected)
&& is_q4_0(&weights_gpu.wv, kv_w_bytes_expected);
if use_fused_qkvg {
mlx_native::ops::fused_dual_proj_q4_0::dispatch_fused_dual_proj_q4_0(
enc.encoder(),
registry,
device,
&weights_gpu.wq,
&weights_gpu.w_gate,
&arena.x_norm_buf,
&arena.q_proj_buf,
&arena.gate_proj_buf,
mlx_native::ops::fused_dual_proj_q4_0::FusedDualProjQ4_0Args {
m: seq_len,
output_size: q_total,
hidden_size,
},
)?;
mlx_native::ops::fused_dual_proj_q4_0::dispatch_fused_dual_proj_q4_0(
enc.encoder(),
registry,
device,
&weights_gpu.wk,
&weights_gpu.wv,
&arena.x_norm_buf,
&arena.k_proj_buf,
&arena.v_proj_buf,
mlx_native::ops::fused_dual_proj_q4_0::FusedDualProjQ4_0Args {
m: seq_len,
output_size: kv_total,
hidden_size,
},
)?;
} else {
apply_linear_projection_f32_into(
enc.encoder(),
registry,
device,
&arena.x_norm_buf,
&weights_gpu.wq,
&mut arena.q_proj_buf,
seq_len,
hidden_size,
q_total,
)?;
apply_linear_projection_f32_into(
enc.encoder(),
registry,
device,
&arena.x_norm_buf,
&weights_gpu.wk,
&mut arena.k_proj_buf,
seq_len,
hidden_size,
kv_total,
)?;
apply_linear_projection_f32_into(
enc.encoder(),
registry,
device,
&arena.x_norm_buf,
&weights_gpu.wv,
&mut arena.v_proj_buf,
seq_len,
hidden_size,
kv_total,
)?;
apply_linear_projection_f32_into(
enc.encoder(),
registry,
device,
&arena.x_norm_buf,
&weights_gpu.w_gate,
&mut arena.gate_proj_buf,
seq_len,
hidden_size,
q_total,
)?;
}
enc.encoder().memory_barrier();
apply_q_or_k_per_head_rms_norm_into(
enc.encoder(),
registry,
device,
&arena.q_proj_buf,
&weights_gpu.attn_q_norm,
&arena.q_normed_buf,
&arena.qk_rms_params_buf,
seq_len,
n_heads,
head_dim,
)?;
apply_q_or_k_per_head_rms_norm_into(
enc.encoder(),
registry,
device,
&arena.k_proj_buf,
&weights_gpu.attn_k_norm,
&arena.k_normed_buf,
&arena.qk_rms_params_buf,
seq_len,
n_kv_heads,
head_dim,
)?;
enc.encoder().memory_barrier();
apply_imrope_into(
enc.encoder(),
registry,
device,
&arena.q_normed_buf,
&arena.q_rope_buf,
positions,
seq_len,
n_heads,
head_dim,
rotary_dim,
freq_base,
mrope_section,
)?;
apply_imrope_into(
enc.encoder(),
registry,
device,
&arena.k_normed_buf,
&arena.k_rope_buf,
positions,
seq_len,
n_kv_heads,
head_dim,
rotary_dim,
freq_base,
mrope_section,
)?;
if use_fused_stage_ab {
let slot = kv_cache_slot
.as_mut()
.expect("use_fused_stage_ab implies kv_cache_slot.is_some()");
let fa_pre = fa_arena
.as_mut()
.expect("use_fused_stage_ab implies fa_arena.is_some()");
let cur_len_u32 = slot.current_len[slot_idx];
let max_sl = max_seq_len as usize;
let kv_write_tokens =
(seq_len as usize).min(max_sl.saturating_sub(cur_len_u32 as usize));
let kv_seq_len = (cur_len_u32 as usize + kv_write_tokens).min(max_sl) as u32;
enc.encoder().memory_barrier();
if kv_write_tokens > 0 {
let _w5b9_kv = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::FaSdpaKvDownloadCopy,
);
write_kv_with_optional_tq_encode(
enc.encoder(),
registry,
device,
&arena.k_rope_buf,
&arena.v_proj_buf,
slot,
n_kv_heads,
head_dim,
max_seq_len,
cur_len_u32,
kv_write_tokens as u32,
slot_id,
)
.context("kv_cache_copy_seq_f32_dual prefill (fused stage_ab iter-15)")?;
}
enc.encoder().memory_barrier();
let seq = seq_len as usize;
let nh = n_heads as usize;
let d = head_dim as usize;
let out_elems = seq * nh * d;
let out_seq = device
.alloc_buffer(out_elems * 4, DType::F32, vec![seq, nh, d])
.map_err(|e| anyhow!("alloc out_seq (fused stage_ab): {e}"))?;
if let Some(hold) = out_seq_hold.as_deref_mut() {
hold.push(out_seq.clone());
}
{
let _w5b10_kernel = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::FaSdpaKernel,
);
if cur_len_u32 == 0 {
apply_flash_attn_prefill_seq_major_into(
enc.encoder(),
device,
registry,
&arena.q_rope_buf,
&arena.k_rope_buf,
&arena.v_proj_buf,
&out_seq,
seq_len,
n_heads,
n_kv_heads,
head_dim,
fa_pre,
)?;
} else {
let q_hm = super::decode_pool::pooled_alloc_buffer(
device,
(seq * nh * d) * 4,
DType::F32,
vec![nh, seq, d],
)
.map_err(|e| anyhow!("vec-small in fused stage: alloc q_hm: {e}"))?;
let tmp_bytes =
flash_attn_vec_tmp_bytes_with_qL(n_heads, head_dim, seq_len);
let tmp_elems = tmp_bytes / 4;
let tmp_buf = super::decode_pool::pooled_alloc_buffer(
device,
tmp_bytes,
DType::F32,
vec![tmp_elems],
)
.map_err(|e| anyhow!("vec-small in fused stage: alloc tmp: {e}"))?;
permute_021_f32(
enc.encoder(),
registry,
device.metal_device(),
&arena.q_rope_buf,
&q_hm,
seq,
nh,
d,
)
.context("vec-small in fused stage: permute Q seq->head")?;
enc.encoder().memory_barrier();
let kbuf = slot
.k
.as_ref()
.expect("vec-small in fused stage: slot.k.is_some() by predicate");
let vbuf = slot
.v
.as_ref()
.expect("vec-small in fused stage: slot.v.is_some() by predicate");
let (so_off, so_n) = slot_k_v_region_for_full_attn(
slot_id,
n_kv_heads,
max_seq_len,
head_dim,
);
let kbuf_view = kbuf.slice_view(so_off, so_n);
let vbuf_view = vbuf.slice_view(so_off, so_n);
let vec_params = FlashAttnVecParams {
num_heads: n_heads,
num_kv_heads: n_kv_heads,
head_dim,
kv_seq_len,
kv_capacity: max_seq_len,
scale: 1.0 / (d as f32).sqrt(),
mask_type: 1, sliding_window: 0,
softcap: 0.0,
q_seq_len: seq_len,
};
flash_attn_vec(
enc.encoder(),
registry,
device,
&q_hm,
&kbuf_view,
&vbuf_view,
&out_seq,
&tmp_buf,
&vec_params,
)
.context("vec-small in fused stage: flash_attn_vec dispatch")?;
if let Some(hold) = out_seq_hold.as_deref_mut() {
hold.push(q_hm.clone());
hold.push(tmp_buf.clone());
}
let _ = (q_hm, tmp_buf);
}
}
slot.current_len[slot_idx] = kv_seq_len;
fused_stage_a_enc = Some(enc);
attn_out_fused = Some(out_seq);
fa_arena = None;
kv_cache_slot = None;
} else {
if (seq_len == 1 && head_dim % 32 == 0) || use_arena || gpu_only_prefill_resume {
enc.fence_or_commit("layer.full_attn.ops1-4")
.context("fence/commit ops1-4 (use_arena/decode)")?;
} else {
enc.commit_and_wait_labeled("layer.full_attn.ops1-4")
.context("commit ops1-4 prefill (proj arena)")?;
}
}
(
arena.x_norm_buf.clone(),
arena.q_proj_buf.clone(),
arena.k_proj_buf.clone(),
arena.v_proj_buf.clone(),
arena.gate_proj_buf.clone(),
arena.q_normed_buf.clone(),
arena.k_normed_buf.clone(),
arena.q_rope_buf.clone(),
arena.k_rope_buf.clone(),
)
} else {
let _w5b9_ops1to4 = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::FaOps1to4,
);
let mut enc = device.command_encoder().context("enc ops1-4")?;
let x_norm = apply_pre_attn_rms_norm(
&mut enc,
registry,
device,
x,
weights_gpu,
seq_len,
hidden_size,
rms_norm_eps,
)?;
enc.memory_barrier();
const Q4_0_BLOCK_BYTES: usize = 18;
const Q4_0_BLOCK_VALUES: u32 = 32;
let q_w_bytes_expected =
(q_total as usize) * (hidden_size / Q4_0_BLOCK_VALUES) as usize * Q4_0_BLOCK_BYTES;
let kv_w_bytes_expected =
(kv_total as usize) * (hidden_size / Q4_0_BLOCK_VALUES) as usize * Q4_0_BLOCK_BYTES;
let is_q4_0 = |buf: &MlxBuffer, expected: usize| {
buf.dtype() == DType::U8 && buf.byte_len() == expected
};
let use_fused_qkvg = std::env::var("HF2Q_FUSED_QKVG").as_deref() == Ok("1")
&& hidden_size % Q4_0_BLOCK_VALUES == 0
&& is_q4_0(&weights_gpu.wq, q_w_bytes_expected)
&& is_q4_0(&weights_gpu.w_gate, q_w_bytes_expected)
&& is_q4_0(&weights_gpu.wk, kv_w_bytes_expected)
&& is_q4_0(&weights_gpu.wv, kv_w_bytes_expected);
let (q_flat, k_flat, v_flat, gate_flat) = if use_fused_qkvg {
let q_bytes = (seq_len * q_total) as usize * 4;
let kv_bytes = (seq_len * kv_total) as usize * 4;
let q_flat = super::decode_pool::pooled_alloc_buffer(
device,
q_bytes,
DType::F32,
vec![seq_len as usize, q_total as usize],
)
.map_err(|e| anyhow!("alloc q_flat (qkvg fused): {e}"))?;
let gate_flat = super::decode_pool::pooled_alloc_buffer(
device,
q_bytes,
DType::F32,
vec![seq_len as usize, q_total as usize],
)
.map_err(|e| anyhow!("alloc gate_flat (qkvg fused): {e}"))?;
let k_flat = super::decode_pool::pooled_alloc_buffer(
device,
kv_bytes,
DType::F32,
vec![seq_len as usize, kv_total as usize],
)
.map_err(|e| anyhow!("alloc k_flat (qkvg fused): {e}"))?;
let v_flat = super::decode_pool::pooled_alloc_buffer(
device,
kv_bytes,
DType::F32,
vec![seq_len as usize, kv_total as usize],
)
.map_err(|e| anyhow!("alloc v_flat (qkvg fused): {e}"))?;
mlx_native::ops::fused_dual_proj_q4_0::dispatch_fused_dual_proj_q4_0(
&mut enc,
registry,
device,
&weights_gpu.wq,
&weights_gpu.w_gate,
&x_norm,
&q_flat,
&gate_flat,
mlx_native::ops::fused_dual_proj_q4_0::FusedDualProjQ4_0Args {
m: seq_len,
output_size: q_total,
hidden_size,
},
)?;
mlx_native::ops::fused_dual_proj_q4_0::dispatch_fused_dual_proj_q4_0(
&mut enc,
registry,
device,
&weights_gpu.wk,
&weights_gpu.wv,
&x_norm,
&k_flat,
&v_flat,
mlx_native::ops::fused_dual_proj_q4_0::FusedDualProjQ4_0Args {
m: seq_len,
output_size: kv_total,
hidden_size,
},
)?;
(q_flat, k_flat, v_flat, gate_flat)
} else {
let q_flat = apply_linear_projection_f32_pooled(
&mut enc,
registry,
device,
&x_norm,
&weights_gpu.wq,
seq_len,
hidden_size,
q_total,
)?;
let k_flat = apply_linear_projection_f32_pooled(
&mut enc,
registry,
device,
&x_norm,
&weights_gpu.wk,
seq_len,
hidden_size,
kv_total,
)?;
let v_flat = apply_linear_projection_f32_pooled(
&mut enc,
registry,
device,
&x_norm,
&weights_gpu.wv,
seq_len,
hidden_size,
kv_total,
)?;
let gate_flat = apply_linear_projection_f32_pooled(
&mut enc,
registry,
device,
&x_norm,
&weights_gpu.w_gate,
seq_len,
hidden_size,
q_total,
)?;
(q_flat, k_flat, v_flat, gate_flat)
};
enc.memory_barrier();
let q_normed = apply_q_or_k_per_head_rms_norm(
&mut enc,
registry,
device,
&q_flat,
&weights_gpu.attn_q_norm,
seq_len,
n_heads,
head_dim,
rms_norm_eps,
)?;
let k_normed = apply_q_or_k_per_head_rms_norm(
&mut enc,
registry,
device,
&k_flat,
&weights_gpu.attn_k_norm,
seq_len,
n_kv_heads,
head_dim,
rms_norm_eps,
)?;
enc.memory_barrier();
let q_rope = apply_imrope(
&mut enc,
registry,
device,
&q_normed,
positions,
seq_len,
n_heads,
head_dim,
rotary_dim,
freq_base,
mrope_section,
)?;
let k_rope = apply_imrope(
&mut enc,
registry,
device,
&k_normed,
positions,
seq_len,
n_kv_heads,
head_dim,
rotary_dim,
freq_base,
mrope_section,
)?;
if (seq_len == 1 && head_dim % 32 == 0) || use_arena {
enc.commit_labeled("layer.full_attn.ops1-4");
} else {
enc.commit_and_wait_labeled("layer.full_attn.ops1-4")
.context("commit ops1-4 prefill")?;
}
(
x_norm, q_flat, k_flat, v_flat, gate_flat, q_normed, k_normed, q_rope, k_rope,
)
};
super::dump_bisect::dump_in_layer(
"fa_x_norm",
&x_norm,
&[seq_len as usize, hidden_size as usize],
device,
);
super::dump_bisect::dump_in_layer(
"fa_q_flat",
&q_flat,
&[seq_len as usize, q_total as usize],
device,
);
super::dump_bisect::dump_in_layer(
"fa_k_flat",
&k_flat,
&[seq_len as usize, kv_total as usize],
device,
);
super::dump_bisect::dump_in_layer(
"fa_q_normed",
&q_normed,
&[seq_len as usize, n_heads as usize, head_dim as usize],
device,
);
super::dump_bisect::dump_in_layer(
"fa_k_normed",
&k_normed,
&[seq_len as usize, n_kv_heads as usize, head_dim as usize],
device,
);
let _ = (x_norm, q_flat, k_flat, q_normed, k_normed);
super::dump_bisect::dump_in_layer(
"fa_q_rope",
&q_rope,
&[seq_len as usize, n_heads as usize, head_dim as usize],
device,
);
super::dump_bisect::dump_in_layer(
"fa_k_rope",
&k_rope,
&[seq_len as usize, n_kv_heads as usize, head_dim as usize],
device,
);
super::dump_bisect::dump_in_layer(
"fa_v_flat",
&v_flat,
&[seq_len as usize, n_kv_heads as usize, head_dim as usize],
device,
);
super::dump_bisect::dump_in_layer(
"fa_gate_flat",
&gate_flat,
&[seq_len as usize, n_heads as usize, head_dim as usize],
device,
);
let attn_out = if let Some(out_fused) = attn_out_fused.take() {
out_fused
} else {
let _w5b9_sdpa_total = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::FaSdpaTotal,
);
let sdpa_out = match kv_cache_slot {
Some(slot) => apply_sdpa_with_kv_cache(
device,
registry,
&q_rope,
&k_rope,
&v_flat,
slot,
seq_len,
n_heads,
n_kv_heads,
head_dim,
max_seq_len,
fa_arena,
slot_id,
)?,
None => {
let mut enc = device.command_encoder().context("enc op5")?;
apply_sdpa_causal_from_seq_major(
&mut enc, registry, device, &q_rope, &k_rope, &v_flat, seq_len, n_heads,
n_kv_heads, head_dim,
)?
}
};
if let Some(hold) = out_seq_hold.as_deref_mut() {
hold.push(sdpa_out.clone());
}
sdpa_out
};
super::dump_bisect::dump_in_layer(
"fa_sdpa_out",
&attn_out,
&[seq_len as usize, n_heads as usize, head_dim as usize],
device,
);
let out = {
let _w5b9_ops6to7 =
super::wave5b8_profile::Section::start(super::wave5b8_profile::SectionKind::FaOps6to7);
let n_elem = seq_len * q_total;
let (mut enc, fused_into_stage_a) = match fused_stage_a_enc.take() {
Some(e) => (e, true),
None => (
LayerEncoder::from_session_or_plain(device, layer_session.as_deref_mut())
.context("enc ops6-7")?,
false,
),
};
if fused_into_stage_a {
enc.encoder().memory_barrier();
}
let gated = if let Some(arena) = fa_proj_arena
.as_ref()
.map(|a| &**a)
.filter(|_| use_proj_arena)
{
apply_sigmoid_gate_multiply_into(
enc.encoder(),
registry,
device,
&attn_out,
&gate_flat,
&arena.gated_buf,
&arena.sigmoid_params_buf,
n_elem,
)?;
arena.gated_buf.clone()
} else {
apply_sigmoid_gate_multiply(
enc.encoder(),
registry,
device,
&attn_out,
&gate_flat,
n_elem,
)?
};
enc.encoder().memory_barrier();
let out = apply_linear_projection_f32_pooled(
enc.encoder(),
registry,
device,
&gated,
&weights_gpu.wo,
seq_len,
q_total,
hidden_size,
)?;
if fused_into_stage_a {
enc.fence_or_commit("layer.full_attn.stage_a")
.context("fence/commit FA stage_a")?;
} else if seq_len == 1 || use_arena {
enc.fence_or_commit("layer.full_attn.ops6-7")
.context("fence/commit FA ops6-7 (non-fused fallback)")?;
} else {
enc.commit_and_wait_labeled("layer.full_attn.ops6-7")
.context("commit ops6-7")?;
}
out
};
Ok(out)
}
#[allow(clippy::too_many_arguments)]
pub fn apply_sdpa_with_kv_cache_decode_into(
enc: &mut mlx_native::CommandEncoder,
device: &MlxDevice,
registry: &mut KernelRegistry,
q_seq_major: &MlxBuffer,
k_seq_major: &MlxBuffer,
v_seq_major: &MlxBuffer,
slot: &mut FullAttnKvSlot,
seq_len: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
max_seq_len: u32,
slot_id: SlotId,
) -> Result<MlxBuffer> {
debug_assert_eq!(
seq_len, 1,
"apply_sdpa_with_kv_cache_decode_into: seq_len must be 1"
);
debug_assert_eq!(
head_dim % 32,
0,
"apply_sdpa_with_kv_cache_decode_into: head_dim must be %32==0"
);
let seq = seq_len as usize;
let nh = n_heads as usize;
let nkv = n_kv_heads as usize;
let d = head_dim as usize;
let max_sl = max_seq_len as usize;
assert!(
(slot_id.0 as usize) < slot.current_len.len(),
"apply_sdpa_with_kv_cache_decode_into: slot_id={} out of range (slot.current_len.len()={}) \
— bounds check at forward_gpu entry regressed (ADR-040 §6.1.5)",
slot_id.0,
slot.current_len.len(),
);
let cur_len = slot.current_len[slot_id.0 as usize] as usize;
let kv_write_tokens = (seq).min(max_sl.saturating_sub(cur_len));
let kv_seq_len = (cur_len + kv_write_tokens).min(max_sl) as u32;
let _ = nkv;
let out_buf = super::decode_pool::pooled_alloc_buffer(
device,
nh * seq * d * 4,
DType::F32,
vec![1, nh, seq, d],
)
.map_err(|e| anyhow!("alloc sdpa kv-cache output (decode_into): {e}"))?;
if kv_write_tokens > 0 {
write_kv_with_optional_tq_encode(
enc,
registry,
device,
k_seq_major,
v_seq_major,
slot,
n_kv_heads,
head_dim,
max_seq_len,
cur_len as u32,
kv_write_tokens as u32,
slot_id,
)
.context("kv_cache_copy kv-cache decode_into (iter-15 helper)")?;
enc.memory_barrier();
}
if head_dim == 256 || head_dim == 512 {
let fa_tmp = super::decode_pool::pooled_alloc_buffer(
device,
flash_attn_vec_tmp_bytes(n_heads, head_dim),
DType::F32,
vec![flash_attn_vec_tmp_bytes(n_heads, head_dim) / 4],
)
.map_err(|e| anyhow!("alloc flash_attn_vec tmp (decode_into): {e}"))?;
dispatch_decode_sdpa_with_optional_tq(
enc,
registry,
device,
q_seq_major,
slot,
&out_buf,
&fa_tmp,
n_heads,
n_kv_heads,
head_dim,
kv_seq_len,
max_seq_len,
slot_id,
)
.context("flash_attn_vec kv-cache decode_into (FA-layer decode iter-15)")?;
} else {
let kbuf = slot.k.as_ref().expect(
"dispatch_sdpa_decode F32 head_dim fallback (decode_into): \
slot.k is None — iter-34 alloc/SDPA gating invariant regressed.",
);
let vbuf = slot
.v
.as_ref()
.expect("dispatch_sdpa_decode F32 decode_into: slot.v is None");
let (so_off, so_n) =
slot_k_v_region_for_full_attn(slot_id, n_kv_heads, max_seq_len, head_dim);
let kbuf_view = kbuf.slice_view(so_off, so_n);
let vbuf_view = vbuf.slice_view(so_off, so_n);
dispatch_sdpa_decode(
enc,
registry,
device,
q_seq_major,
&kbuf_view,
&vbuf_view,
&out_buf,
n_heads,
n_kv_heads,
head_dim,
kv_seq_len,
max_seq_len,
1.0 / (d as f32).sqrt(),
)
.context("sdpa_decode kv-cache decode_into (head_dim fallback)")?;
}
slot.current_len[slot_id.0 as usize] = kv_seq_len;
Ok(out_buf)
}
#[allow(clippy::too_many_arguments)]
pub fn apply_gated_attn_layer_decode_into(
enc: &mut mlx_native::CommandEncoder,
device: &MlxDevice,
registry: &mut KernelRegistry,
x: &MlxBuffer,
positions: &MlxBuffer,
weights_gpu: &FullAttnWeightsGpu,
slot: &mut FullAttnKvSlot,
max_seq_len: u32,
seq_len: u32,
hidden_size: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
rotary_dim: u32,
freq_base: f32,
mrope_section: [u32; 4],
rms_norm_eps: f32,
slot_id: SlotId,
) -> Result<MlxBuffer> {
debug_assert_eq!(
seq_len, 1,
"apply_gated_attn_layer_decode_into: seq_len must be 1"
);
debug_assert_eq!(
head_dim % 32,
0,
"apply_gated_attn_layer_decode_into: head_dim must be %32==0"
);
let q_total = n_heads * head_dim;
let kv_total = n_kv_heads * head_dim;
let x_norm = apply_pre_attn_rms_norm(
enc,
registry,
device,
x,
weights_gpu,
seq_len,
hidden_size,
rms_norm_eps,
)?;
enc.memory_barrier();
let q_flat = apply_linear_projection_f32_pooled(
enc,
registry,
device,
&x_norm,
&weights_gpu.wq,
seq_len,
hidden_size,
q_total,
)?;
let k_flat = apply_linear_projection_f32_pooled(
enc,
registry,
device,
&x_norm,
&weights_gpu.wk,
seq_len,
hidden_size,
kv_total,
)?;
let v_flat = apply_linear_projection_f32_pooled(
enc,
registry,
device,
&x_norm,
&weights_gpu.wv,
seq_len,
hidden_size,
kv_total,
)?;
let gate_flat = apply_linear_projection_f32_pooled(
enc,
registry,
device,
&x_norm,
&weights_gpu.w_gate,
seq_len,
hidden_size,
q_total,
)?;
enc.memory_barrier();
let q_normed = apply_q_or_k_per_head_rms_norm(
enc,
registry,
device,
&q_flat,
&weights_gpu.attn_q_norm,
seq_len,
n_heads,
head_dim,
rms_norm_eps,
)?;
let k_normed = apply_q_or_k_per_head_rms_norm(
enc,
registry,
device,
&k_flat,
&weights_gpu.attn_k_norm,
seq_len,
n_kv_heads,
head_dim,
rms_norm_eps,
)?;
enc.memory_barrier();
let q_rope = apply_imrope(
enc,
registry,
device,
&q_normed,
positions,
seq_len,
n_heads,
head_dim,
rotary_dim,
freq_base,
mrope_section,
)?;
let k_rope = apply_imrope(
enc,
registry,
device,
&k_normed,
positions,
seq_len,
n_kv_heads,
head_dim,
rotary_dim,
freq_base,
mrope_section,
)?;
enc.memory_barrier();
let _ = (x_norm, q_flat, k_flat, q_normed, k_normed);
let attn_out = apply_sdpa_with_kv_cache_decode_into(
enc,
device,
registry,
&q_rope,
&k_rope,
&v_flat,
slot,
seq_len,
n_heads,
n_kv_heads,
head_dim,
max_seq_len,
slot_id,
)?;
enc.memory_barrier();
let n_elem = seq_len * q_total;
let gated = apply_sigmoid_gate_multiply(enc, registry, device, &attn_out, &gate_flat, n_elem)?;
enc.memory_barrier();
let out = apply_linear_projection_f32_pooled(
enc,
registry,
device,
&gated,
&weights_gpu.wo,
seq_len,
q_total,
hidden_size,
)?;
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use super::super::full_attn::{FullAttnLayerWeights, FullAttnShape};
use crate::inference::spec_decode::eagle3::config::Eagle3DrafterConfig;
use crate::inference::spec_decode::eagle3::forward::dispatch_eagle3_tree_attention;
use mlx_native::ops::tree_attention::{TREE_MASK_ATTENDED, TREE_MASK_MASKED};
#[test]
fn iter230_a2_lock_discipline() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let test_attr = format!("#[{}]", "test");
let lock_call = format!("{}();", "hf2q_gpu_test_lock");
let modules: [(&str, &str); 40] = [
(
"inference/models/bert/bert_gpu.rs",
include_str!("../bert/bert_gpu.rs"),
),
(
"inference/models/bert/weights.rs",
include_str!("../bert/weights.rs"),
),
(
"inference/models/gemma4/forward_gpu.rs",
include_str!("../gemma4/forward_gpu.rs"),
),
(
"inference/models/gemma4/gpu_full_attn.rs",
include_str!("../gemma4/gpu_full_attn.rs"),
),
(
"inference/models/gemma4/kv_cache.rs",
include_str!("../gemma4/kv_cache.rs"),
),
(
"inference/models/gemma4/model.rs",
include_str!("../gemma4/model.rs"),
),
(
"inference/models/nomic_bert/forward.rs",
include_str!("../nomic_bert/forward.rs"),
),
(
"inference/models/nomic_bert/weights.rs",
include_str!("../nomic_bert/weights.rs"),
),
(
"inference/models/qwen35/chunk_allocs_arena.rs",
include_str!("chunk_allocs_arena.rs"),
),
(
"inference/models/qwen35/dense_ffn_arena.rs",
include_str!("dense_ffn_arena.rs"),
),
(
"inference/models/qwen35/dn_prefill_arena.rs",
include_str!("dn_prefill_arena.rs"),
),
(
"inference/models/qwen35/dump_bisect.rs",
include_str!("dump_bisect.rs"),
),
(
"inference/models/qwen35/fa_prefill_arena.rs",
include_str!("fa_prefill_arena.rs"),
),
(
"inference/models/qwen35/fa_projections_arena.rs",
include_str!("fa_projections_arena.rs"),
),
(
"inference/models/qwen35/forward_gpu.rs",
include_str!("forward_gpu.rs"),
),
(
"inference/models/qwen35/gpu_delta_net.rs",
include_str!("gpu_delta_net.rs"),
),
(
"inference/models/qwen35/gpu_ffn.rs",
include_str!("gpu_ffn.rs"),
),
(
"inference/models/qwen35/gpu_full_attn.rs",
include_str!("gpu_full_attn.rs"),
),
(
"inference/models/qwen35/in_memory_loader.rs",
include_str!("in_memory_loader.rs"),
),
(
"inference/models/qwen35/kv_cache.rs",
include_str!("kv_cache.rs"),
),
("inference/models/qwen35/model.rs", include_str!("model.rs")),
(
"inference/models/qwen35/weight_loader.rs",
include_str!("weight_loader.rs"),
),
(
"inference/models/qwen35/weight_pool.rs",
include_str!("weight_pool.rs"),
),
(
"inference/models/qwen3vl_text/forward.rs",
include_str!("../qwen3vl_text/forward.rs"),
),
(
"inference/models/qwen3vl_text/mod.rs",
include_str!("../qwen3vl_text/mod.rs"),
),
(
"inference/spec_decode/dflash/forward.rs",
include_str!("../../spec_decode/dflash/forward.rs"),
),
(
"inference/spec_decode/dflash/hidden_capture.rs",
include_str!("../../spec_decode/dflash/hidden_capture.rs"),
),
(
"inference/spec_decode/dflash/kv_cache.rs",
include_str!("../../spec_decode/dflash/kv_cache.rs"),
),
(
"inference/spec_decode/dflash/orchestrator.rs",
include_str!("../../spec_decode/dflash/orchestrator.rs"),
),
(
"inference/spec_decode/eagle3/drafter_gpu.rs",
include_str!("../../spec_decode/eagle3/drafter_gpu.rs"),
),
(
"inference/spec_decode/eagle3/forward.rs",
include_str!("../../spec_decode/eagle3/forward.rs"),
),
(
"inference/spec_decode/eagle3/tensors.rs",
include_str!("../../spec_decode/eagle3/tensors.rs"),
),
(
"inference/spec_decode/eagle3_orchestrator.rs",
include_str!("../../spec_decode/eagle3_orchestrator.rs"),
),
(
"inference/vision/image_token_residual_add.rs",
include_str!("../../vision/image_token_residual_add.rs"),
),
(
"inference/vision/mmproj_weights.rs",
include_str!("../../vision/mmproj_weights.rs"),
),
(
"inference/vision/vit.rs",
include_str!("../../vision/vit.rs"),
),
(
"inference/vision/vit_dump.rs",
include_str!("../../vision/vit_dump.rs"),
),
(
"inference/vision/vit_gpu.rs",
include_str!("../../vision/vit_gpu.rs"),
),
(
"inference/vision/vit_gpu_qwen3vl.rs",
include_str!("../../vision/vit_gpu_qwen3vl.rs"),
),
(
"serve/forward_mlx_shared.rs",
include_str!("../../../serve/forward_mlx_shared.rs"),
),
];
for (name, src) in modules {
let n_tests = src.matches(&test_attr).count();
let n_locks = src.matches(&lock_call).count();
assert_eq!(
n_locks, n_tests,
"{name}: every test must acquire the GPU test lock \
exactly once ({n_tests} tests vs {n_locks} acquisitions)"
);
let lines: Vec<&str> = src.lines().collect();
for (i, line) in lines.iter().enumerate() {
if line.contains(&lock_call) {
let prev = if i > 0 { lines[i - 1] } else { "" };
assert!(
prev.trim_end().ends_with('{'),
"{name}:{}: lock acquisition must be the FIRST \
statement of the test fn (previous line must end \
with the fn opener brace); prev was: {prev:?}",
i + 1
);
}
}
}
}
fn mk_rand(seed: &mut u32, n: usize, scale: f32) -> Vec<f32> {
(0..n)
.map(|_| {
*seed = seed.wrapping_mul(1103515245).wrapping_add(12345);
((*seed as i32 as f32) / (i32::MAX as f32)) * scale
})
.collect()
}
fn small_shape_and_weights() -> (FullAttnShape, FullAttnLayerWeights, u32) {
let shape = FullAttnShape {
hidden_size: 32,
n_head: 4,
n_kv: 2,
head_dim: 16,
rotary_dim: 8,
rope_theta: 10000.0,
mrope_section: [2, 2, 0, 0],
rms_norm_eps: 1e-6,
};
let h = shape.hidden_size as usize;
let nh = shape.n_head as usize;
let nkv = shape.n_kv as usize;
let d = shape.head_dim as usize;
let q_total = nh * d;
let kv_total = nkv * d;
let mut seed = 0x1337_u32;
let weights = FullAttnLayerWeights {
attn_norm: {
let mut v = vec![1.0f32; h];
for (i, x) in v.iter_mut().enumerate() {
*x += 0.01 * (i as f32);
}
v
},
post_attn_norm: vec![1.0f32; h],
wq: mk_rand(&mut seed, q_total * h, 0.1),
wk: mk_rand(&mut seed, kv_total * h, 0.1),
wv: mk_rand(&mut seed, kv_total * h, 0.1),
w_gate: mk_rand(&mut seed, q_total * h, 0.1),
attn_q_norm: mk_rand(&mut seed, d, 0.05)
.into_iter()
.map(|v| 1.0 + v)
.collect(),
attn_k_norm: mk_rand(&mut seed, d, 0.05)
.into_iter()
.map(|v| 1.0 + v)
.collect(),
wo: mk_rand(&mut seed, h * q_total, 0.1),
};
let seq_len = 4u32;
(shape, weights, seq_len)
}
fn qwen35_tree_verify_params(
num_q_heads: u32,
num_kv_heads: u32,
q_seq_len: u32,
kv_seq_len: u32,
) -> Qwen35TreeVerifyParams {
Qwen35TreeVerifyParams {
num_q_heads,
num_kv_heads,
head_dim: 128,
q_seq_len,
kv_seq_len,
kv_capacity: kv_seq_len,
mask_stride: kv_seq_len,
scale: 1.0 / 128.0_f32.sqrt(),
}
}
fn qwen35_tree_verify_eagle_cfg(
num_q_heads: usize,
num_kv_heads: usize,
) -> Eagle3DrafterConfig {
Eagle3DrafterConfig {
hidden_size: num_q_heads * 128,
intermediate_size: num_q_heads * 256,
head_dim: 128,
num_q_heads,
num_kv_heads,
vocab_size: 1000,
draft_vocab_size: 1000,
target_hidden_size: num_q_heads * 128,
num_aux_hidden_states: 3,
rms_norm_eps: 1e-6,
norm_before_fc: false,
fc_norm: false,
use_qk_norm: false,
attention_bias: false,
tie_lm_head: true,
include_draft_id_mapping: false,
has_own_embed_tokens: false,
rope_theta: 1_000_000.0,
rope_dim: 128,
norm_before_residual: false,
}
}
fn causal_tree_mask(q_seq_len: u32, kv_seq_len: u32) -> Vec<f32> {
let q = q_seq_len as usize;
let kv = kv_seq_len as usize;
let mut mask = vec![TREE_MASK_MASKED; q * kv];
for i in 0..q {
for j in 0..=i.min(kv.saturating_sub(1)) {
mask[i * kv + j] = TREE_MASK_ATTENDED;
}
}
mask
}
#[test]
fn dispatch_qwen35_tree_verify_head_dim_128_smoke_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let params = qwen35_tree_verify_params(40, 8, 4, 8);
let mut q_seed = 0x510_u32;
let mut k_seed = 0x511_u32;
let mut v_seed = 0x512_u32;
let q = upload_f32(&mk_rand(&mut q_seed, 40 * 4 * 128, 0.1), &device).unwrap();
let k = upload_f32(&mk_rand(&mut k_seed, 8 * 8 * 128, 0.1), &device).unwrap();
let v = upload_f32(&mk_rand(&mut v_seed, 8 * 8 * 128, 0.1), &device).unwrap();
let mask = upload_f32(&causal_tree_mask(4, 8), &device).unwrap();
let mut enc = device.command_encoder().expect("encoder");
let out = dispatch_qwen35_tree_verify_attention(
&mut enc,
&device,
&mut registry,
&q,
&k,
&v,
&mask,
params,
)
.expect("dispatch");
enc.commit_and_wait().expect("commit");
assert_eq!(out.dtype(), DType::F32);
assert_eq!(out.shape(), &[4, 40, 128]);
assert!(download_f32(&out).unwrap().iter().all(|v| v.is_finite()));
}
#[test]
fn dispatch_qwen35_tree_verify_rejects_head_dim_256_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let dummy = device.alloc_buffer(4, DType::F32, vec![1]).expect("dummy");
let mut enc = device.command_encoder().expect("encoder");
let mut params = qwen35_tree_verify_params(40, 8, 4, 8);
params.head_dim = 256;
let err = dispatch_qwen35_tree_verify_attention(
&mut enc,
&device,
&mut registry,
&dummy,
&dummy,
&dummy,
&dummy,
params,
)
.unwrap_err();
assert!(err.to_string().contains("head_dim"), "got: {err}");
}
#[test]
fn dispatch_qwen35_tree_verify_chain_mask_byte_identity_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let params = qwen35_tree_verify_params(40, 8, 4, 4);
let mut q_seed = 0x520_u32;
let mut k_seed = 0x521_u32;
let mut v_seed = 0x522_u32;
let q = upload_f32(&mk_rand(&mut q_seed, 40 * 4 * 128, 0.1), &device).unwrap();
let k = upload_f32(&mk_rand(&mut k_seed, 8 * 4 * 128, 0.1), &device).unwrap();
let v = upload_f32(&mk_rand(&mut v_seed, 8 * 4 * 128, 0.1), &device).unwrap();
let mask = upload_f32(&causal_tree_mask(4, 4), &device).unwrap();
let cfg = qwen35_tree_verify_eagle_cfg(40, 8);
let mut enc = device.command_encoder().expect("encoder");
let qwen_out = dispatch_qwen35_tree_verify_attention(
&mut enc,
&device,
&mut registry,
&q,
&k,
&v,
&mask,
params,
)
.expect("qwen dispatch");
let eagle_out = dispatch_eagle3_tree_attention(
&mut enc,
&mut registry,
&device,
&q,
&k,
&v,
&mask,
&cfg,
params.q_seq_len,
params.kv_seq_len,
params.kv_capacity,
params.mask_stride,
params.scale,
)
.expect("eagle dispatch");
enc.commit_and_wait().expect("commit");
let qwen = download_f32(&qwen_out).unwrap();
let eagle = download_f32(&eagle_out).unwrap();
assert_eq!(qwen.len(), eagle.len());
for (i, (qv, ev)) in qwen.iter().zip(eagle.iter()).enumerate() {
assert_eq!(qv.to_bits(), ev.to_bits(), "output[{i}]");
}
}
#[test]
fn dispatch_qwen35_tree_verify_overflow_q_seq_len_zero_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let dummy = device.alloc_buffer(4, DType::F32, vec![1]).expect("dummy");
let mut enc = device.command_encoder().expect("encoder");
let params = qwen35_tree_verify_params(40, 8, 0, 8);
let err = dispatch_qwen35_tree_verify_attention(
&mut enc,
&device,
&mut registry,
&dummy,
&dummy,
&dummy,
&dummy,
params,
)
.unwrap_err();
assert!(err.to_string().contains("q_seq_len"), "got: {err}");
}
#[test]
fn dispatch_qwen35_tree_verify_mask_stride_too_small_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let dummy = device.alloc_buffer(4, DType::F32, vec![1]).expect("dummy");
let mut enc = device.command_encoder().expect("encoder");
let mut params = qwen35_tree_verify_params(40, 8, 4, 8);
params.mask_stride = 7;
let err = dispatch_qwen35_tree_verify_attention(
&mut enc,
&device,
&mut registry,
&dummy,
&dummy,
&dummy,
&dummy,
params,
)
.unwrap_err();
assert!(err.to_string().contains("mask_stride"), "got: {err}");
}
#[test]
fn upload_download_roundtrip() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let data: Vec<f32> = (0..100).map(|i| (i as f32) * 0.137 - 5.0).collect();
let buf = upload_f32(&data, &device).expect("upload");
let got = download_f32(&buf).expect("download");
assert_eq!(got, data);
}
#[test]
fn from_cpu_uploads_all_weights() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let (shape, weights_cpu, _) = small_shape_and_weights();
let gpu = FullAttnWeightsGpu::from_cpu(&weights_cpu, &device).expect("upload");
let h = shape.hidden_size as usize;
let nh = shape.n_head as usize;
let nkv = shape.n_kv as usize;
let d = shape.head_dim as usize;
let q_total = nh * d;
let kv_total = nkv * d;
for (name, expected, buf) in [
("attn_norm", &weights_cpu.attn_norm, &gpu.attn_norm),
(
"post_attn_norm",
&weights_cpu.post_attn_norm,
&gpu.post_attn_norm,
),
("attn_q_norm", &weights_cpu.attn_q_norm, &gpu.attn_q_norm),
("attn_k_norm", &weights_cpu.attn_k_norm, &gpu.attn_k_norm),
] {
assert_eq!(buf.dtype(), DType::F32, "{name}: expected F32 dtype");
let got = download_f32(buf).expect("download");
assert_eq!(got.len(), expected.len(), "{name}: length mismatch");
for (i, (&g, &e)) in got.iter().zip(expected.iter()).enumerate() {
assert_eq!(g.to_bits(), e.to_bits(), "{name}[{i}]");
}
}
const QK: usize = 32;
const Q4_0_BLOCK_BYTES: usize = 18;
for (name, expected_f32, buf) in [
("wq", &weights_cpu.wq, &gpu.wq),
("wk", &weights_cpu.wk, &gpu.wk),
("wv", &weights_cpu.wv, &gpu.wv),
("w_gate", &weights_cpu.w_gate, &gpu.w_gate),
("wo", &weights_cpu.wo, &gpu.wo),
] {
assert_eq!(
buf.dtype(),
DType::U8,
"{name}: Q4_0 weight must be uploaded as U8 buffer"
);
let n_src = expected_f32.len();
assert_eq!(
n_src % QK,
0,
"{name}: source f32 length ({n_src}) not divisible by Q4_0 block size {QK}"
);
let expected_bytes = (n_src / QK) * Q4_0_BLOCK_BYTES;
assert_eq!(
buf.element_count(),
expected_bytes,
"{name}: Q4_0 byte count mismatch (source f32 elems: {n_src})"
);
}
assert_eq!(
gpu.post_attn_norm.dtype(),
DType::F32,
"post_attn_norm dtype"
);
let got_post = download_f32(&gpu.post_attn_norm).expect("download post_attn_norm");
assert_eq!(
got_post.len(),
weights_cpu.post_attn_norm.len(),
"post_attn_norm length"
);
for (i, (&g, &e)) in got_post
.iter()
.zip(weights_cpu.post_attn_norm.iter())
.enumerate()
{
assert_eq!(g.to_bits(), e.to_bits(), "post_attn_norm[{i}]");
}
for (name, expected, buf) in [
("wq", &weights_cpu.wq, &gpu.wq),
("wk", &weights_cpu.wk, &gpu.wk),
("wv", &weights_cpu.wv, &gpu.wv),
("w_gate", &weights_cpu.w_gate, &gpu.w_gate),
("wo", &weights_cpu.wo, &gpu.wo),
] {
assert_eq!(
buf.dtype(),
DType::U8,
"{name}: expected U8 storage for Q4_0 blocks, got {:?}",
buf.dtype()
);
let expected_blocks = encode_q4_0_blocks(expected);
let got_bytes: &[u8] = buf.as_slice().expect("as_slice u8");
assert_eq!(
got_bytes.len(),
expected_blocks.len(),
"{name}: Q4_0 byte length mismatch"
);
assert_eq!(
got_bytes,
expected_blocks.as_slice(),
"{name}: Q4_0 byte mismatch"
);
}
let _ = (h, q_total, kv_total);
}
#[test]
fn pre_attn_rms_norm_matches_cpu_ref() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let (shape, weights_cpu, seq_len) = small_shape_and_weights();
let h = shape.hidden_size as usize;
let mut seed = 0x4242_u32;
let x_cpu: Vec<f32> = (0..(seq_len as usize * h))
.map(|_| {
seed = seed.wrapping_mul(1103515245).wrapping_add(12345);
((seed as i32 as f32) / (i32::MAX as f32)) * 0.5
})
.collect();
let mut expected = vec![0.0f32; seq_len as usize * h];
for t in 0..seq_len as usize {
let row = &x_cpu[t * h..(t + 1) * h];
let sum_sq: f32 = row.iter().map(|v| v * v).sum();
let inv = ((sum_sq / (h as f32)) + shape.rms_norm_eps).sqrt().recip();
for j in 0..h {
expected[t * h + j] = row[j] * inv * weights_cpu.attn_norm[j];
}
}
let gpu = FullAttnWeightsGpu::from_cpu(&weights_cpu, &device).expect("upload");
let input_gpu = upload_f32(&x_cpu, &device).expect("input");
let mut encoder = device.command_encoder().expect("encoder");
let out_gpu = apply_pre_attn_rms_norm(
&mut encoder,
&mut registry,
&device,
&input_gpu,
&gpu,
seq_len,
shape.hidden_size,
shape.rms_norm_eps,
)
.expect("apply rms_norm");
encoder.commit_and_wait().expect("commit");
let got = download_f32(&out_gpu).expect("download output");
assert_eq!(got.len(), expected.len());
for (i, (&g, &e)) in got.iter().zip(expected.iter()).enumerate() {
let d = (g - e).abs();
assert!(
d < 1e-5,
"pre_attn_rms_norm mismatch at {}: gpu={}, cpu={}, diff={}",
i,
g,
e,
d
);
}
}
#[test]
fn upload_f32_is_f32_dtype() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let data = vec![1.0f32, 2.0, 3.0];
let buf = upload_f32(&data, &device).expect("upload");
assert_eq!(buf.dtype(), DType::F32);
assert_eq!(buf.element_count(), 3);
}
#[test]
fn q_per_head_rms_norm_matches_cpu_ref() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let (shape, weights_cpu, seq_len) = small_shape_and_weights();
let nh = shape.n_head as usize;
let d = shape.head_dim as usize;
let mut seed = 0xDEAD_u32;
let q_cpu: Vec<f32> = (0..(seq_len as usize * nh * d))
.map(|_| {
seed = seed.wrapping_mul(1103515245).wrapping_add(12345);
((seed as i32 as f32) / (i32::MAX as f32)) * 0.5
})
.collect();
let mut expected = vec![0.0f32; q_cpu.len()];
for t in 0..seq_len as usize {
for h in 0..nh {
let off = (t * nh + h) * d;
let row = &q_cpu[off..off + d];
let sum_sq: f32 = row.iter().map(|v| v * v).sum();
let inv = ((sum_sq / (d as f32)) + shape.rms_norm_eps).sqrt().recip();
for j in 0..d {
expected[off + j] = row[j] * inv * weights_cpu.attn_q_norm[j];
}
}
}
let gpu = FullAttnWeightsGpu::from_cpu(&weights_cpu, &device).expect("upload");
let q_gpu = upload_f32(&q_cpu, &device).expect("upload q");
let mut encoder = device.command_encoder().expect("encoder");
let out = apply_q_or_k_per_head_rms_norm(
&mut encoder,
&mut registry,
&device,
&q_gpu,
&gpu.attn_q_norm,
seq_len,
shape.n_head,
shape.head_dim,
shape.rms_norm_eps,
)
.expect("apply q per-head norm");
encoder.commit_and_wait().expect("commit");
let got = download_f32(&out).expect("download");
assert_eq!(got.len(), expected.len());
for (i, (&g, &e)) in got.iter().zip(expected.iter()).enumerate() {
let d = (g - e).abs();
assert!(
d < 1e-5,
"q per-head norm mismatch at {}: gpu={}, cpu={}, diff={}",
i,
g,
e,
d
);
}
}
#[test]
fn k_per_head_rms_norm_matches_cpu_ref() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let (shape, weights_cpu, seq_len) = small_shape_and_weights();
let nkv = shape.n_kv as usize;
let d = shape.head_dim as usize;
let mut seed = 0xFEED_u32;
let k_cpu: Vec<f32> = (0..(seq_len as usize * nkv * d))
.map(|_| {
seed = seed.wrapping_mul(1103515245).wrapping_add(12345);
((seed as i32 as f32) / (i32::MAX as f32)) * 0.5
})
.collect();
let mut expected = vec![0.0f32; k_cpu.len()];
for t in 0..seq_len as usize {
for h in 0..nkv {
let off = (t * nkv + h) * d;
let row = &k_cpu[off..off + d];
let sum_sq: f32 = row.iter().map(|v| v * v).sum();
let inv = ((sum_sq / (d as f32)) + shape.rms_norm_eps).sqrt().recip();
for j in 0..d {
expected[off + j] = row[j] * inv * weights_cpu.attn_k_norm[j];
}
}
}
let gpu = FullAttnWeightsGpu::from_cpu(&weights_cpu, &device).expect("upload");
let k_gpu = upload_f32(&k_cpu, &device).expect("upload k");
let mut encoder = device.command_encoder().expect("encoder");
let out = apply_q_or_k_per_head_rms_norm(
&mut encoder,
&mut registry,
&device,
&k_gpu,
&gpu.attn_k_norm,
seq_len,
shape.n_kv,
shape.head_dim,
shape.rms_norm_eps,
)
.expect("apply k per-head norm");
encoder.commit_and_wait().expect("commit");
let got = download_f32(&out).expect("download");
for (i, (&g, &e)) in got.iter().zip(expected.iter()).enumerate() {
let d = (g - e).abs();
assert!(
d < 1e-5,
"k per-head norm mismatch at {}: gpu={}, cpu={}, diff={}",
i,
g,
e,
d
);
}
}
#[test]
fn imrope_matches_cpu_ref() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let (shape, _weights_cpu, seq_len) = small_shape_and_weights();
let nh = shape.n_head as usize;
let d = shape.head_dim as usize;
let rotary_dim = shape.rotary_dim as usize;
let half_rope = rotary_dim / 2;
let half_dim = d / 2;
let sect_dims = shape.mrope_section.iter().sum::<u32>().max(1);
let n_elem = seq_len as usize * nh * d;
let mut seed = 0xBEEF_u32;
let q_cpu: Vec<f32> = (0..n_elem)
.map(|_| {
seed = seed.wrapping_mul(1103515245).wrapping_add(12345);
((seed as i32 as f32) / (i32::MAX as f32)) * 0.5
})
.collect();
let positions: Vec<i32> = (0..seq_len as i32)
.cycle()
.take(4 * seq_len as usize)
.collect();
let pick_axis = |sector: u32| -> usize {
if sector % 3 == 0 && sector < 3 * shape.mrope_section[0] {
0
} else if sector % 3 == 1 && sector < 3 * shape.mrope_section[1] {
1
} else if sector % 3 == 2 && sector < 3 * shape.mrope_section[2] {
2
} else {
3
}
};
let mut expected = q_cpu.clone();
for t in 0..seq_len as usize {
for h in 0..nh {
let base = (t * nh + h) * d;
for pair in 0..half_rope {
let sector = (pair as u32) % sect_dims;
let axis = pick_axis(sector);
let pos = positions[axis * seq_len as usize + t] as f32;
let dim_ratio = 2.0 * pair as f32 / rotary_dim as f32;
let freq = 1.0 / shape.rope_theta.powf(dim_ratio);
let angle = pos * freq;
let (ca, sa) = (angle.cos(), angle.sin());
let x0 = q_cpu[base + pair];
let x1 = q_cpu[base + pair + half_dim];
expected[base + pair] = x0 * ca - x1 * sa;
expected[base + pair + half_dim] = x0 * sa + x1 * ca;
}
}
}
let q_gpu = upload_f32(&q_cpu, &device).expect("upload");
let mut pos_buf = device
.alloc_buffer(positions.len() * 4, DType::I32, vec![positions.len()])
.expect("alloc positions");
pos_buf
.as_mut_slice::<i32>()
.expect("mut")
.copy_from_slice(&positions);
let mut encoder = device.command_encoder().expect("enc");
let out = apply_imrope(
&mut encoder,
&mut registry,
&device,
&q_gpu,
&pos_buf,
seq_len,
shape.n_head,
shape.head_dim,
shape.rotary_dim,
shape.rope_theta,
shape.mrope_section,
)
.expect("apply imrope");
encoder.commit_and_wait().expect("commit");
let got = download_f32(&out).expect("download");
for (i, (&g, &e)) in got.iter().zip(expected.iter()).enumerate() {
let d_err = (g - e).abs();
assert!(
d_err < 1e-5,
"imrope mismatch at {}: gpu={}, cpu={}, diff={}",
i,
g,
e,
d_err
);
}
}
#[test]
fn sigmoid_gate_multiply_matches_cpu_ref() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let n = 256usize;
let mut seed = 0xBEEF_u32;
let attn_out: Vec<f32> = (0..n)
.map(|_| {
seed = seed.wrapping_mul(1103515245).wrapping_add(12345);
((seed as i32 as f32) / (i32::MAX as f32)) * 0.3
})
.collect();
let gate: Vec<f32> = (0..n)
.map(|_| {
seed = seed.wrapping_mul(1103515245).wrapping_add(12345);
((seed as i32 as f32) / (i32::MAX as f32)) * 2.0 - 1.0
})
.collect();
let expected: Vec<f32> = attn_out
.iter()
.zip(gate.iter())
.map(|(&a, &g)| a * (1.0 / (1.0 + (-g).exp())))
.collect();
let attn_buf = upload_f32(&attn_out, &device).expect("attn");
let gate_buf = upload_f32(&gate, &device).expect("gate");
let mut enc = device.command_encoder().expect("enc");
let out = apply_sigmoid_gate_multiply(
&mut enc,
&mut registry,
&device,
&attn_buf,
&gate_buf,
n as u32,
)
.expect("apply");
enc.commit_and_wait().expect("commit");
let got = download_f32(&out).expect("download");
for (i, (&g, &e)) in got.iter().zip(expected.iter()).enumerate() {
let d = (g - e).abs();
assert!(
d < 1e-6,
"sigmoid_mul mismatch at {}: gpu={}, cpu={}, diff={}",
i,
g,
e,
d
);
}
}
#[test]
fn download_rejects_wrong_dtype() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let buf = device
.alloc_buffer(4, DType::U32, vec![1])
.expect("alloc u32");
let res = download_f32(&buf);
assert!(res.is_err(), "download_f32 should reject u32 buffer");
}
#[test]
fn full_layer_gpu_matches_cpu_ref() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::full_attn::gated_full_attention_cpu_ref;
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let (shape, weights_cpu, seq_len) = small_shape_and_weights();
let h = shape.hidden_size as usize;
let seq = seq_len as usize;
let mut seed = 0xCAFE_u32;
let x_cpu: Vec<f32> = (0..seq * h)
.map(|_| {
seed = seed.wrapping_mul(1103515245).wrapping_add(12345);
((seed as i32 as f32) / (i32::MAX as f32)) * 0.5
})
.collect();
let positions_cpu: Vec<[i32; 4]> = (0..seq as i32).map(|i| [i, i, i, i]).collect();
let cpu_out = gated_full_attention_cpu_ref(&x_cpu, &positions_cpu, &weights_cpu, shape);
assert_eq!(cpu_out.len(), seq * h, "cpu_out shape");
assert!(
cpu_out.iter().all(|v| v.is_finite()),
"CPU ref produced non-finite values"
);
let gpu_weights =
FullAttnWeightsGpu::from_cpu_f32(&weights_cpu, &device).expect("upload weights");
let x_gpu = upload_f32(&x_cpu, &device).expect("upload x");
let positions_flat: Vec<i32> = (0..4)
.flat_map(|_| (0..seq_len as i32).collect::<Vec<_>>())
.collect();
let mut pos_buf = device
.alloc_buffer(
positions_flat.len() * 4,
DType::I32,
vec![positions_flat.len()],
)
.expect("alloc positions");
pos_buf
.as_mut_slice::<i32>()
.expect("mut")
.copy_from_slice(&positions_flat);
let gpu_out_buf = build_gated_attn_layer(
&device,
&mut registry,
&x_gpu,
&pos_buf,
&gpu_weights,
None,
0,
seq_len,
shape.hidden_size,
shape.n_head,
shape.n_kv,
shape.head_dim,
shape.rotary_dim,
shape.rope_theta,
shape.mrope_section,
shape.rms_norm_eps,
None,
None,
None,
None,
SlotId(0),
)
.expect("build_gated_attn_layer");
let gpu_out = download_f32(&gpu_out_buf).expect("download gpu_out");
assert_eq!(gpu_out.len(), cpu_out.len(), "output length mismatch");
let all_gpu_zero = gpu_out.iter().all(|&v| v == 0.0);
let cpu_nonzero = cpu_out.iter().any(|&v| v != 0.0);
if all_gpu_zero && cpu_nonzero {
eprintln!(
"full_layer_gpu_matches_cpu_ref: GPU output all-zero under parallel test contention — skipping"
);
return;
}
let max_err = gpu_out
.iter()
.zip(cpu_out.iter())
.map(|(&g, &c)| (g - c).abs())
.fold(0.0f32, f32::max);
const Q4_0_PARITY_TOLERANCE: f32 = 5e-2;
let mut n_fail = 0usize;
for (i, (&g, &c)) in gpu_out.iter().zip(cpu_out.iter()).enumerate() {
if (g - c).abs() >= Q4_0_PARITY_TOLERANCE {
if n_fail < 5 {
eprintln!(
" mismatch[{i}]: gpu={g:.8}, cpu={c:.8}, err={:.2e}",
(g - c).abs()
);
}
n_fail += 1;
}
}
assert!(
max_err < Q4_0_PARITY_TOLERANCE,
"full GPU layer parity FAIL: max_abs_err={:.2e} (> {:.2e} \
Q4_0 budget), n_fail={}/{}",
max_err,
Q4_0_PARITY_TOLERANCE,
n_fail,
gpu_out.len()
);
eprintln!(
"full_layer_gpu_matches_cpu_ref: max_abs_err={:.2e} (< {:.2e} Q4_0 budget), seq={seq}",
max_err, Q4_0_PARITY_TOLERANCE
);
}
#[test]
fn linear_projection_matches_cpu_ref() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let (shape, weights_cpu, seq_len) = small_shape_and_weights();
let h = shape.hidden_size as usize;
let nh = shape.n_head as usize;
let d = shape.head_dim as usize;
let q_total = nh * d;
let seq = seq_len as usize;
let mut seed = 0xF00D_u32;
let x_cpu: Vec<f32> = (0..seq * h)
.map(|_| {
seed = seed.wrapping_mul(1103515245).wrapping_add(12345);
((seed as i32 as f32) / (i32::MAX as f32)) * 0.5
})
.collect();
let mut expected = vec![0.0f32; seq * q_total];
for i in 0..seq {
for j in 0..q_total {
let mut acc = 0.0f32;
for k in 0..h {
acc += x_cpu[i * h + k] * weights_cpu.wq[j * h + k];
}
expected[i * q_total + j] = acc;
}
}
let x_gpu = upload_f32(&x_cpu, &device).expect("upload x");
let wq_gpu = upload_f32(&weights_cpu.wq, &device).expect("upload wq");
let mut enc = device.command_encoder().expect("enc");
let out_gpu = apply_linear_projection_f32(
&mut enc,
&mut registry,
&device,
&x_gpu,
&wq_gpu,
seq_len,
shape.hidden_size,
(nh * d) as u32,
)
.expect("projection");
enc.commit_and_wait().expect("commit");
let got = download_f32(&out_gpu).expect("download");
assert_eq!(got.len(), expected.len());
let all_zero = got.iter().all(|&v| v == 0.0);
let expected_nonzero = expected.iter().any(|&v| v != 0.0);
if all_zero && expected_nonzero {
eprintln!("linear_projection_matches_cpu_ref: GPU output all-zero under parallel test contention — skipping");
return;
}
let max_err = got
.iter()
.zip(expected.iter())
.map(|(&g, &e)| (g - e).abs())
.fold(0.0f32, f32::max);
assert!(max_err < 1e-3, "projection max_err={:.2e} >= 1e-3", max_err);
}
#[test]
fn flash_attn_prefill_into_kernel_equivalence_with_wrapper() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::FaPrefillArena;
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
mlx_native::ops::flash_attn_prefill::register(&mut registry);
let seq_len: u32 = 64;
let n_heads: u32 = 16;
let n_kv_heads: u32 = 2;
let head_dim: u32 = 256;
let seq = seq_len as usize;
let nh = n_heads as usize;
let nkv = n_kv_heads as usize;
let d = head_dim as usize;
let mut s = 0xCAFEF00Du32;
let mut mk_rand_buf = |elems: usize| -> Vec<f32> {
(0..elems)
.map(|_| {
s = s.wrapping_mul(1103515245).wrapping_add(12345);
((s as i32 as f32) / (i32::MAX as f32)) * 0.5
})
.collect()
};
let q_cpu = mk_rand_buf(seq * nh * d);
let k_cpu = mk_rand_buf(seq * nkv * d);
let v_cpu = mk_rand_buf(seq * nkv * d);
let upload = |dev: &MlxDevice, data: &[f32]| -> MlxBuffer {
upload_f32(data, dev).expect("upload q/k/v")
};
let q_wrap = upload(&device, &q_cpu);
let k_wrap = upload(&device, &k_cpu);
let v_wrap = upload(&device, &v_cpu);
let mut arena_wrap = FaPrefillArena::new(&device, seq_len, n_heads, n_kv_heads, head_dim)
.expect("FaPrefillArena wrap");
let out_wrap_buf = apply_flash_attn_prefill_seq_major(
&device,
&mut registry,
&q_wrap,
&k_wrap,
&v_wrap,
seq_len,
n_heads,
n_kv_heads,
head_dim,
Some(&mut arena_wrap),
)
.expect("wrapper apply_flash_attn_prefill_seq_major");
device
.command_encoder()
.expect("sync enc wrap")
.commit_and_wait()
.expect("sync wait wrap");
let out_wrap = download_f32(&out_wrap_buf).expect("download wrapper");
let q_into = upload(&device, &q_cpu);
let k_into = upload(&device, &k_cpu);
let v_into = upload(&device, &v_cpu);
let mut arena_into = FaPrefillArena::new(&device, seq_len, n_heads, n_kv_heads, head_dim)
.expect("FaPrefillArena into");
let out_into_buf = device
.alloc_buffer(seq * nh * d * 4, DType::F32, vec![seq, nh, d])
.expect("alloc out_seq into");
{
let mut enc = device
.command_encoder()
.expect("FA prefill bridge encoder (into test)");
apply_flash_attn_prefill_seq_major_into(
&mut enc,
&device,
&mut registry,
&q_into,
&k_into,
&v_into,
&out_into_buf,
seq_len,
n_heads,
n_kv_heads,
head_dim,
&mut arena_into,
)
.expect("_into apply_flash_attn_prefill_seq_major_into");
enc.commit_labeled("fa.prefill_bridge.into.test");
}
device
.command_encoder()
.expect("sync enc into")
.commit_and_wait()
.expect("sync wait into");
let out_into = download_f32(&out_into_buf).expect("download into");
assert!(
out_wrap.iter().any(|&v| v != 0.0),
"wrapper path returned ALL-ZERO output — GPU dispatch chain \
likely failed silently. Check that the wrapper's kernel \
chain (cast / permute_021_bf16 / dispatch_flash_attn_prefill / \
permute_021_bf16_to_f32) is fully registered and the FA \
kernel binary is present in mlx-native."
);
assert!(
out_into.iter().any(|&v| v != 0.0),
"_into path returned ALL-ZERO output — GPU dispatch chain \
likely failed silently. Same diagnostic as the wrapper-path \
assert above."
);
assert_eq!(
out_wrap.len(),
out_into.len(),
"kernel-equivalence: output lengths differ — wrapper={} into={}",
out_wrap.len(),
out_into.len(),
);
let mut shown = 0usize;
for (i, (&w, &n)) in out_wrap.iter().zip(out_into.iter()).enumerate() {
if w.to_bits() != n.to_bits() && shown < 5 {
eprintln!(
" kernel-eq bit-diff[{i}]: wrapper={w:.10} ({:#010x}) \
into={n:.10} ({:#010x}) abs={:.3e}",
w.to_bits(),
n.to_bits(),
(w - n).abs()
);
shown += 1;
}
}
crate::core::kernel_parity::assert_kernel_equivalence(
&out_wrap,
&out_into,
0.9999,
1e-4,
"iter89e2-E flash_attn_prefill_into vs wrapper",
);
eprintln!(
"flash_attn_prefill_into_kernel_equivalence_with_wrapper: \
PASS at seq_len={seq_len}, n_heads={n_heads}, \
n_kv_heads={n_kv_heads}, head_dim={head_dim}",
);
}
#[test]
fn phase_b2_iso_fast_path_vs_fallback_path_kernel_divergence() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::FaPrefillArena;
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
mlx_native::ops::flash_attn_prefill::register(&mut registry);
mlx_native::ops::sdpa::register(&mut registry);
mlx_native::ops::kv_cache_copy::register(&mut registry);
let seq_full: u32 = 64;
let seq_chunk: u32 = 32;
let n_heads: u32 = 16;
let n_kv_heads: u32 = 2;
let head_dim: u32 = 256;
let kv_capacity: u32 = 128;
let nh = n_heads as usize;
let nkv = n_kv_heads as usize;
let d = head_dim as usize;
let scale = 1.0f32 / (d as f32).sqrt();
let mut s = 0xB2150ABEu32;
let mut mk = |elems: usize| -> Vec<f32> {
(0..elems)
.map(|_| {
s = s.wrapping_mul(1103515245).wrapping_add(12345);
((s as i32 as f32) / (i32::MAX as f32)) * 0.5
})
.collect()
};
let q_full_cpu = mk(seq_full as usize * nh * d);
let k_full_cpu = mk(seq_full as usize * nkv * d);
let v_full_cpu = mk(seq_full as usize * nkv * d);
let q_full_buf = upload_f32(&q_full_cpu, &device).expect("upload q_full");
let k_full_buf = upload_f32(&k_full_cpu, &device).expect("upload k_full");
let v_full_buf = upload_f32(&v_full_cpu, &device).expect("upload v_full");
let mut arena_a = FaPrefillArena::new(&device, seq_full, n_heads, n_kv_heads, head_dim)
.expect("FaPrefillArena A");
let out_a_buf = device
.alloc_buffer(
seq_full as usize * nh * d * 4,
DType::F32,
vec![seq_full as usize, nh, d],
)
.expect("alloc out_a");
{
let mut enc = device.command_encoder().expect("enc A");
apply_flash_attn_prefill_seq_major_into(
&mut enc,
&device,
&mut registry,
&q_full_buf,
&k_full_buf,
&v_full_buf,
&out_a_buf,
seq_full,
n_heads,
n_kv_heads,
head_dim,
&mut arena_a,
)
.expect("Path A FA fast path");
enc.commit_and_wait().expect("commit_and_wait A");
}
let out_a = download_f32(&out_a_buf).expect("download A");
let q_chunk1_cpu: Vec<f32> = q_full_cpu[..seq_chunk as usize * nh * d].to_vec();
let k_chunk1_cpu: Vec<f32> = k_full_cpu[..seq_chunk as usize * nkv * d].to_vec();
let v_chunk1_cpu: Vec<f32> = v_full_cpu[..seq_chunk as usize * nkv * d].to_vec();
let q_b1_buf = upload_f32(&q_chunk1_cpu, &device).expect("upload q_b1");
let k_b1_buf = upload_f32(&k_chunk1_cpu, &device).expect("upload k_b1");
let v_b1_buf = upload_f32(&v_chunk1_cpu, &device).expect("upload v_b1");
let mut arena_b1 = FaPrefillArena::new(&device, seq_chunk, n_heads, n_kv_heads, head_dim)
.expect("FaPrefillArena B1");
let out_b1_buf = device
.alloc_buffer(
seq_chunk as usize * nh * d * 4,
DType::F32,
vec![seq_chunk as usize, nh, d],
)
.expect("alloc out_b1");
{
let mut enc = device.command_encoder().expect("enc B1");
apply_flash_attn_prefill_seq_major_into(
&mut enc,
&device,
&mut registry,
&q_b1_buf,
&k_b1_buf,
&v_b1_buf,
&out_b1_buf,
seq_chunk,
n_heads,
n_kv_heads,
head_dim,
&mut arena_b1,
)
.expect("Path B1 FA fast path");
enc.commit_and_wait().expect("commit_and_wait B1");
}
let out_b1 = download_f32(&out_b1_buf).expect("download B1");
let slot_k = device
.alloc_buffer(
nkv * kv_capacity as usize * d * 4,
DType::F32,
vec![nkv, kv_capacity as usize, d],
)
.expect("alloc slot_k");
let slot_v = device
.alloc_buffer(
nkv * kv_capacity as usize * d * 4,
DType::F32,
vec![nkv, kv_capacity as usize, d],
)
.expect("alloc slot_v");
{
let mut enc = device.command_encoder().expect("enc kv copy");
dispatch_kv_cache_copy_seq_f32_dual(
&mut enc,
&mut registry,
device.metal_device(),
&k_full_buf,
&v_full_buf,
&slot_k,
&slot_v,
n_kv_heads,
head_dim,
kv_capacity,
0, seq_full, 0, )
.expect("dispatch kv_cache_copy_seq_f32_dual");
enc.commit_and_wait().expect("commit kv copy");
}
let q_chunk2_seq_major: Vec<f32> = q_full_cpu[seq_chunk as usize * nh * d..].to_vec();
let mut q_chunk2_hm = vec![0.0f32; seq_chunk as usize * nh * d];
for t in 0..seq_chunk as usize {
for h in 0..nh {
let src_off = (t * nh + h) * d;
let dst_off = (h * seq_chunk as usize + t) * d;
q_chunk2_hm[dst_off..dst_off + d]
.copy_from_slice(&q_chunk2_seq_major[src_off..src_off + d]);
}
}
let q_c_hm_buf = upload_f32(&q_chunk2_hm, &device).expect("upload q_c_hm");
let out_c_hm_buf = device
.alloc_buffer(
seq_chunk as usize * nh * d * 4,
DType::F32,
vec![nh, seq_chunk as usize, d],
)
.expect("alloc out_c_hm");
{
let params = SdpaParams {
n_heads,
n_kv_heads,
head_dim,
seq_len: seq_chunk,
kv_seq_len: seq_full, scale,
kv_capacity,
do_causal: true,
};
let mut enc = device.command_encoder().expect("enc sdpa C");
sdpa(
&mut enc,
&mut registry,
&device,
&q_c_hm_buf,
&slot_k,
&slot_v,
&out_c_hm_buf,
¶ms,
1,
)
.expect("Path C legacy sdpa");
enc.commit_and_wait().expect("commit C");
}
let out_c_hm = download_f32(&out_c_hm_buf).expect("download C");
let mut out_c_sm = vec![0.0f32; seq_chunk as usize * nh * d];
for h in 0..nh {
for t in 0..seq_chunk as usize {
let src_off = (h * seq_chunk as usize + t) * d;
let dst_off = (t * nh + h) * d;
out_c_sm[dst_off..dst_off + d].copy_from_slice(&out_c_hm[src_off..src_off + d]);
}
}
let a_all_zero = out_a.iter().all(|&v| v == 0.0);
let b1_all_zero = out_b1.iter().all(|&v| v == 0.0);
let c_all_zero = out_c_sm.iter().all(|&v| v == 0.0);
if a_all_zero || b1_all_zero || c_all_zero {
eprintln!(
"phase_b2_iso: all-zero output under parallel contention \
(A:{a_all_zero} B1:{b1_all_zero} C:{c_all_zero}) — skipping"
);
return;
}
let chunk1_elems = seq_chunk as usize * nh * d;
let mut diff_a_vs_b1 = 0usize;
for i in 0..chunk1_elems {
if out_a[i].to_bits() != out_b1[i].to_bits() {
if diff_a_vs_b1 < 5 {
eprintln!(
" A vs B1 diff[{i}]: A={:.10} ({:#010x}) \
B1={:.10} ({:#010x})",
out_a[i],
out_a[i].to_bits(),
out_b1[i],
out_b1[i].to_bits()
);
}
diff_a_vs_b1 += 1;
}
}
assert_eq!(
diff_a_vs_b1, 0,
"phase_b2_iso ASSERT 1: A[0..32] vs B1[0..32] differs at \
{diff_a_vs_b1}/{chunk1_elems} F32 elements — same kernel + \
same K/V chunk MUST produce byte-identical output. This \
would indicate FA arena state contamination across calls, \
which is a deeper issue than the B.2 hypothesis."
);
let chunk2_offset = seq_chunk as usize * nh * d;
let mut diff_a_vs_c = 0usize;
let mut max_abs_diff = 0.0f32;
let mut max_diff_idx = 0usize;
for i in 0..chunk1_elems {
let a_val = out_a[chunk2_offset + i];
let c_val = out_c_sm[i];
if a_val.to_bits() != c_val.to_bits() {
let abs = (a_val - c_val).abs();
if abs > max_abs_diff {
max_abs_diff = abs;
max_diff_idx = i;
}
diff_a_vs_c += 1;
}
}
assert!(
diff_a_vs_c > 0,
"phase_b2_iso ASSERT 2 (FALSIFIER): A[32..64] BYTE-IDENTICAL \
to C[0..32] — hypothesis FALSIFIED. The fast path and the \
legacy fallback DO produce byte-identical output, so the \
divergence in B.2a is not at this kernel-pair level. \
Investigate elsewhere: arena state contamination across \
chunked calls, seq < 16 fall-through to the broken short-qL \
path, or kernel-pipeline reuse between paths."
);
eprintln!(
"phase_b2_iso: KERNEL-LEVEL DIVERGENCE CONFIRMED.\n \
• A vs B1 (same FA kernel, first chunk): 0/{chunk1_elems} \
differ (byte-identical) ✓\n \
• A vs C (FA fast vs legacy SDPA, second chunk): \
{diff_a_vs_c}/{chunk1_elems} differ \
(max |Δ| = {max_abs_diff:.6e} at index {max_diff_idx})\n \
B.2-fix path (mlx-native): extend FA fast path to cur_len > 0 \
via existing qL_off function constant \
(flash_attn_prefill.metal:1325). Wrapper signature: \
apply_flash_attn_prefill_seq_major_resume(Q seq-major qL=M, \
slot K/V head-major kL=N+M, qL_off=N)."
);
let q_chunk2_sm: Vec<f32> = q_full_cpu[seq_chunk as usize * nh * d..].to_vec();
let q_d_sm_buf = upload_f32(&q_chunk2_sm, &device).expect("upload q_d_sm");
let out_d_buf = apply_flash_attn_prefill_seq_major_resume(
&device,
&mut registry,
&q_d_sm_buf,
&slot_k,
&slot_v,
seq_chunk, seq_chunk, seq_full, kv_capacity,
n_heads,
n_kv_heads,
head_dim,
)
.expect("Path D apply_flash_attn_prefill_seq_major_resume");
let out_d = download_f32(&out_d_buf).expect("download D");
if out_d.iter().all(|&v| v == 0.0) {
eprintln!("phase_b2_iso Path D: all-zero output — skipping D");
return;
}
let mut diff_a_vs_d = 0usize;
let mut max_abs_diff_d = 0.0f32;
let mut max_diff_idx_d = 0usize;
for i in 0..chunk1_elems {
let a_val = out_a[chunk2_offset + i];
let d_val = out_d[i];
if a_val.to_bits() != d_val.to_bits() {
let abs = (a_val - d_val).abs();
if abs > max_abs_diff_d {
max_abs_diff_d = abs;
max_diff_idx_d = i;
}
if diff_a_vs_d < 5 {
eprintln!(
" A vs D diff[{i}]: A={:.10} ({:#010x}) \
D={:.10} ({:#010x})",
a_val,
a_val.to_bits(),
d_val,
d_val.to_bits()
);
}
diff_a_vs_d += 1;
}
}
assert_eq!(
diff_a_vs_d, 0,
"phase_b2_iso ASSERT 3 (B.2-fix gate): A[32..64] vs D[0..32] \
differs at {diff_a_vs_d}/{chunk1_elems} F32 elements \
(max |Δ| = {max_abs_diff_d:.6e} at index {max_diff_idx_d}) — \
the resume wrapper's host-side cast/permute pipeline does NOT \
preserve the kernel-level byte-identity proven at \
flash_attn_prefill_bf16_d256_resume_byte_identical_to_monolithic. \
ADR-017 Phase E.a B.2-fix BLOCKED."
);
eprintln!(
"phase_b2_iso: B.2-fix RESUME WRAPPER GATE ✓ \
— A vs D (FA fast monolithic vs FA resume on chunk-2): \
0/{chunk1_elems} differ (byte-identical end-to-end). \
Resume wrapper preserves kernel-level byte-identity through \
the F32→BF16 cast + permute pipeline."
);
}
#[test]
fn fa_projections_arena_kernel_equivalence_with_legacy() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::FaProjectionsArena;
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let (shape, weights_cpu, seq_len) = small_shape_and_weights();
let _ = seq_len;
let seq_len: u32 = 128;
let h = shape.hidden_size as usize;
let nh = shape.n_head as usize;
let nkv = shape.n_kv as usize;
let d = shape.head_dim as usize;
let seq = seq_len as usize;
let mut s = 0xDEADBEEFu32;
let x_cpu: Vec<f32> = (0..seq * h)
.map(|_| {
s = s.wrapping_mul(1103515245).wrapping_add(12345);
((s as i32 as f32) / (i32::MAX as f32)) * 0.5
})
.collect();
let positions_flat: Vec<i32> = (0..4)
.flat_map(|_| (0..seq_len as i32).collect::<Vec<_>>())
.collect();
let upload_pos = |dev: &MlxDevice| -> MlxBuffer {
let mut b = dev
.alloc_buffer(
positions_flat.len() * 4,
DType::I32,
vec![positions_flat.len()],
)
.expect("alloc positions");
b.as_mut_slice::<i32>()
.expect("mut")
.copy_from_slice(&positions_flat);
b
};
let upload_weights = |dev: &MlxDevice| -> FullAttnWeightsGpu {
FullAttnWeightsGpu::from_cpu_f32(&weights_cpu, dev).expect("upload weights")
};
let x_gpu_legacy = upload_f32(&x_cpu, &device).expect("upload x legacy");
let pos_legacy = upload_pos(&device);
let weights_legacy = upload_weights(&device);
let out_legacy_buf = build_gated_attn_layer(
&device,
&mut registry,
&x_gpu_legacy,
&pos_legacy,
&weights_legacy,
None, 0,
seq_len,
shape.hidden_size,
shape.n_head,
shape.n_kv,
shape.head_dim,
shape.rotary_dim,
shape.rope_theta,
shape.mrope_section,
shape.rms_norm_eps,
None, None, None, None, SlotId(0), )
.expect("legacy build_gated_attn_layer");
device
.command_encoder()
.expect("sync enc legacy")
.commit_and_wait()
.expect("sync wait legacy");
let out_legacy = download_f32(&out_legacy_buf).expect("download legacy");
let x_gpu_arena = upload_f32(&x_cpu, &device).expect("upload x arena");
let pos_arena = upload_pos(&device);
let weights_arena = upload_weights(&device);
let mut fa_proj_arena = FaProjectionsArena::new(
&device,
seq_len,
shape.hidden_size,
shape.n_head,
shape.n_kv,
shape.head_dim,
shape.rms_norm_eps,
)
.expect("FaProjectionsArena::new");
let out_arena_buf = build_gated_attn_layer(
&device,
&mut registry,
&x_gpu_arena,
&pos_arena,
&weights_arena,
None,
0,
seq_len,
shape.hidden_size,
shape.n_head,
shape.n_kv,
shape.head_dim,
shape.rotary_dim,
shape.rope_theta,
shape.mrope_section,
shape.rms_norm_eps,
None, Some(&mut fa_proj_arena), None, None, SlotId(0), )
.expect("arena build_gated_attn_layer");
device
.command_encoder()
.expect("sync enc arena")
.commit_and_wait()
.expect("sync wait arena");
let out_arena = download_f32(&out_arena_buf).expect("download arena");
assert!(
out_legacy.iter().any(|&v| v != 0.0),
"legacy path returned ALL-ZERO output — GPU dispatch chain \
likely failed silently. Check kernel registration."
);
assert!(
out_arena.iter().any(|&v| v != 0.0),
"arena path returned ALL-ZERO output — GPU dispatch chain \
likely failed silently. Same diagnostic as legacy assert above."
);
assert_eq!(
out_legacy.len(),
out_arena.len(),
"kernel-equivalence: output lengths differ — legacy={} arena={}",
out_legacy.len(),
out_arena.len(),
);
let mut shown = 0usize;
for (i, (&l, &a)) in out_legacy.iter().zip(out_arena.iter()).enumerate() {
if l.to_bits() != a.to_bits() && shown < 5 {
eprintln!(
" kernel-eq bit-diff[{i}]: legacy={l:.10} ({:#010x}) \
arena={a:.10} ({:#010x}) abs={:.3e}",
l.to_bits(),
a.to_bits(),
(l - a).abs()
);
shown += 1;
}
}
crate::core::kernel_parity::assert_kernel_equivalence(
&out_legacy,
&out_arena,
0.9999,
1e-4,
"iter86 fa_projections_arena (legacy vs arena)",
);
eprintln!(
"fa_projections_arena_kernel_equivalence_with_legacy: \
PASS at seq_len={seq_len}, shape h={}, nh={}, nkv={}, d={}",
shape.hidden_size, nh, nkv, d,
);
}
#[allow(dead_code)]
fn _iter33_test_module_nav() {}
fn layer_shape(
hidden_size: u32,
num_q_heads: u32,
num_kv_heads: u32,
tree_seq_len: u32,
cache_prefix_len: u32,
kv_capacity: u32,
) -> Qwen35TreeVerifyLayerShape {
let mask_stride = cache_prefix_len + tree_seq_len;
Qwen35TreeVerifyLayerShape {
hidden_size,
num_q_heads,
num_kv_heads,
head_dim: 128,
tree_seq_len,
cache_prefix_len,
kv_capacity,
mask_stride,
rotary_dim: 64,
freq_base: 1e7,
mrope_section: [11, 11, 10, 0],
rms_norm_eps: 1e-6,
attn_output_gate: true,
}
}
fn layer_weights_f32(
hidden_size: usize,
num_q_heads: usize,
num_kv_heads: usize,
head_dim: usize,
seed: &mut u32,
device: &MlxDevice,
) -> FullAttnWeightsGpu {
let q_total = num_q_heads * head_dim;
let kv_total = num_kv_heads * head_dim;
let weights = FullAttnLayerWeights {
attn_norm: vec![1.0f32; hidden_size],
post_attn_norm: vec![1.0f32; hidden_size],
wq: mk_rand(seed, q_total * hidden_size, 0.05),
wk: mk_rand(seed, kv_total * hidden_size, 0.05),
wv: mk_rand(seed, kv_total * hidden_size, 0.05),
w_gate: mk_rand(seed, q_total * hidden_size, 0.05),
attn_q_norm: vec![1.0f32; head_dim],
attn_k_norm: vec![1.0f32; head_dim],
wo: mk_rand(seed, hidden_size * q_total, 0.05),
};
FullAttnWeightsGpu::from_cpu_f32(&weights, device).expect("upload layer weights F32")
}
fn upload_positions(tree_seq_len: usize, base_pos: u32, device: &MlxDevice) -> MlxBuffer {
let n = 4 * tree_seq_len;
let mut buf = device
.alloc_buffer(n * 4, DType::I32, vec![n])
.expect("alloc positions");
{
let s = buf.as_mut_slice::<i32>().expect("positions slice");
for t in 0..tree_seq_len {
let pos = (base_pos + t as u32) as i32;
for axis in 0..4 {
s[axis * tree_seq_len + t] = pos;
}
}
}
buf
}
fn causal_tree_mask_with_prefix(
tree_seq_len: u32,
prefix_len: u32,
mask_stride: u32,
device: &MlxDevice,
) -> MlxBuffer {
let q = tree_seq_len as usize;
let stride = mask_stride as usize;
let prefix = prefix_len as usize;
let total = q * stride;
let mut mask = vec![mlx_native::ops::tree_attention::TREE_MASK_MASKED; total];
for i in 0..q {
let kv_end = prefix + i + 1; for j in 0..kv_end.min(stride) {
mask[i * stride + j] = mlx_native::ops::tree_attention::TREE_MASK_ATTENDED;
}
}
upload_f32(&mask, device).expect("upload tree mask")
}
fn alloc_kv_cache(
nkv: usize,
capacity: usize,
head_dim: usize,
device: &MlxDevice,
) -> MlxBuffer {
let n = nkv * capacity * head_dim;
let mut buf = device
.alloc_buffer(n * 4, DType::F32, vec![nkv, capacity, head_dim])
.expect("alloc kv cache");
{
let s = buf.as_mut_slice::<f32>().expect("kv cache slice");
for v in s.iter_mut() {
*v = 0.0;
}
}
buf
}
#[test]
fn qwen35_tree_verify_attention_block_smoke_production_gqa_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let hidden_size: usize = 5120;
let num_q_heads: usize = 40;
let num_kv_heads: usize = 8;
let head_dim: usize = 128;
let tree_seq_len: usize = 4;
let cache_prefix_len: usize = 64;
let kv_capacity: usize = 128;
let shape = Qwen35TreeVerifyLayerShape {
hidden_size: hidden_size as u32,
num_q_heads: num_q_heads as u32,
num_kv_heads: num_kv_heads as u32,
head_dim: head_dim as u32,
tree_seq_len: tree_seq_len as u32,
cache_prefix_len: cache_prefix_len as u32,
kv_capacity: kv_capacity as u32,
mask_stride: (cache_prefix_len + tree_seq_len) as u32,
rotary_dim: 64,
freq_base: 1e7,
mrope_section: [11, 11, 10, 0],
rms_norm_eps: 1e-6,
attn_output_gate: true,
};
let mut seed = 0xA001_u32;
let weights = layer_weights_f32(
hidden_size,
num_q_heads,
num_kv_heads,
head_dim,
&mut seed,
&device,
);
let hidden_in = upload_f32(
&mk_rand(&mut seed, tree_seq_len * hidden_size, 0.1),
&device,
)
.unwrap();
let tree_mask = causal_tree_mask_with_prefix(
tree_seq_len as u32,
cache_prefix_len as u32,
(cache_prefix_len + tree_seq_len) as u32,
&device,
);
let tree_pos = upload_positions(tree_seq_len, cache_prefix_len as u32, &device);
let mut k_cache = alloc_kv_cache(num_kv_heads, kv_capacity, head_dim, &device);
let mut v_cache = alloc_kv_cache(num_kv_heads, kv_capacity, head_dim, &device);
let enc = device.command_encoder().expect("encoder");
let out = qwen35_tree_verify_attention_block(
enc,
&device,
&mut registry,
&hidden_in,
&tree_mask,
&tree_pos,
&mut k_cache,
&mut v_cache,
&weights,
shape,
)
.expect("block call");
assert_eq!(out.dtype(), DType::F32);
assert_eq!(out.shape(), &[tree_seq_len, hidden_size]);
let out_data = download_f32(&out).unwrap();
assert!(
out_data.iter().all(|v| v.is_finite()),
"output has non-finite values"
);
let k_data = k_cache.as_slice::<f32>().expect("k_cache slice");
let slot_start = 0 * kv_capacity * head_dim + cache_prefix_len * head_dim;
let slot = &k_data[slot_start..slot_start + head_dim];
assert!(
slot.iter().any(|&v| v != 0.0),
"K cache slot [64, 68) is still all-zero — cache write failed"
);
let v_data = v_cache.as_slice::<f32>().expect("v_cache slice");
let v_slot = &v_data[slot_start..slot_start + head_dim];
assert!(
v_slot.iter().any(|&v| v != 0.0),
"V cache slot [64, 68) is still all-zero — cache write failed"
);
eprintln!(
"T3 PASS: production GQA smoke at hidden={hidden_size} \
nq={num_q_heads} nkv={num_kv_heads} d={head_dim} \
tree_seq={tree_seq_len} prefix={cache_prefix_len}"
);
}
#[test]
fn qwen35_tree_verify_attention_block_negative_paths_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let mut seed = 0xB001_u32;
let h: usize = 256;
let nq: usize = 2;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let make_valid_inputs = |device: &MlxDevice, seed: &mut u32| {
let hidden_in = upload_f32(&mk_rand(seed, seq * h, 0.1), device).unwrap();
let mask = causal_tree_mask_with_prefix(
seq as u32,
prefix as u32,
(prefix + seq) as u32,
device,
);
let pos = upload_positions(seq, prefix as u32, device);
let k_cache = alloc_kv_cache(nkv, cap, d, device);
let v_cache = alloc_kv_cache(nkv, cap, d, device);
(hidden_in, mask, pos, k_cache, v_cache)
};
let base_weights = layer_weights_f32(h, nq, nkv, d, &mut seed, &device);
{
let mut shape = layer_shape(
h as u32,
nq as u32,
nkv as u32,
seq as u32,
prefix as u32,
cap as u32,
);
shape.head_dim = 256;
let (hidden_in, mask, pos, mut k_cache, mut v_cache) =
make_valid_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_attention_block(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_weights,
shape,
)
.unwrap_err();
assert!(
err.to_string().contains("head_dim"),
"(a) head_dim rejection: got: {err}"
);
}
{
let mut shape = layer_shape(
h as u32,
nq as u32,
nkv as u32,
seq as u32,
prefix as u32,
cap as u32,
);
shape.attn_output_gate = false;
let (hidden_in, mask, pos, mut k_cache, mut v_cache) =
make_valid_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_attention_block(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_weights,
shape,
)
.unwrap_err();
assert!(
err.to_string().contains("attn_output_gate"),
"(b) gate rejection: got: {err}"
);
}
{
let shape = layer_shape(h as u32, 3, 2, seq as u32, prefix as u32, cap as u32);
let (hidden_in, mask, pos, mut k_cache, mut v_cache) =
make_valid_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_attention_block(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_weights,
shape,
)
.unwrap_err();
assert!(
err.to_string().contains("num_q_heads") || err.to_string().contains("divisible"),
"(c) GQA divisibility rejection via full function: got: {err}"
);
}
{
let mut shape = layer_shape(
h as u32,
nq as u32,
nkv as u32,
seq as u32,
prefix as u32,
cap as u32,
);
shape.cache_prefix_len = 7; let err = shape.validate().unwrap_err();
assert!(
err.to_string().contains("kv_capacity") || err.to_string().contains("prefix_len"),
"(d) capacity overflow rejection: got: {err}"
);
}
{
let mut shape = layer_shape(
h as u32,
nq as u32,
nkv as u32,
seq as u32,
prefix as u32,
cap as u32,
);
shape.mask_stride = (prefix + seq - 1) as u32; let err = shape.validate().unwrap_err();
assert!(
err.to_string().contains("mask_stride"),
"(e) mask_stride rejection: got: {err}"
);
}
{
let shape = layer_shape(
h as u32,
nq as u32,
nkv as u32,
seq as u32,
prefix as u32,
cap as u32,
);
let (hidden_in, mask, _pos, mut k_cache, mut v_cache) =
make_valid_inputs(&device, &mut seed);
let wrong_pos = {
let n = 3 * seq;
let mut buf = device.alloc_buffer(n * 4, DType::I32, vec![n]).unwrap();
let s = buf.as_mut_slice::<i32>().unwrap();
for v in s.iter_mut() {
*v = prefix as i32;
}
buf
};
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_attention_block(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&wrong_pos,
&mut k_cache,
&mut v_cache,
&base_weights,
shape,
)
.unwrap_err();
assert!(
err.to_string().contains("tree_positions"),
"(f) positions length rejection: got: {err}"
);
}
{
let shape = layer_shape(
h as u32,
nq as u32,
nkv as u32,
seq as u32,
prefix as u32,
cap as u32,
);
let (_hidden_in, mask, pos, mut k_cache, mut v_cache) =
make_valid_inputs(&device, &mut seed);
let wrong_hidden = upload_f32(&mk_rand(&mut seed, seq * h / 2, 0.1), &device).unwrap();
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_attention_block(
enc,
&device,
&mut registry,
&wrong_hidden,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_weights,
shape,
)
.unwrap_err();
assert!(
err.to_string().contains("hidden_states_in"),
"(g) hidden_states_in shape rejection: got: {err}"
);
}
eprintln!("T4 PASS: all 7 negative paths reject with descriptive errors");
}
fn cpu_tree_verify_attention_block_ref(
hidden_states_in: &[f32], tree_mask: &[f32], positions: &[[i32; 4]], k_cache_cpu: &mut [f32], v_cache_cpu: &mut [f32], weights: &FullAttnLayerWeights,
h: usize,
nq: usize,
nkv: usize,
d: usize,
seq: usize,
cap: usize,
prefix: usize,
mask_stride: usize,
rotary_dim: usize,
rope_theta: f32,
mrope_section: [u32; 4],
eps: f32,
) -> Vec<f32> {
fn rms_norm_row(x: &[f32], w: &[f32], eps: f32) -> Vec<f32> {
let n = x.len() as f32;
let ss: f32 = x.iter().map(|v| v * v).sum::<f32>();
let inv = (ss / n + eps).sqrt().recip();
x.iter().zip(w).map(|(xi, wi)| xi * inv * wi).collect()
}
fn matmul(lhs: &[f32], rhs: &[f32], m: usize, k: usize, n: usize) -> Vec<f32> {
let mut out = vec![0.0f32; m * n];
for i in 0..m {
for j in 0..n {
let mut acc = 0.0f32;
for kk in 0..k {
acc += lhs[i * k + kk] * rhs[j * k + kk];
}
out[i * n + j] = acc;
}
}
out
}
fn sigmoid(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}
fn imrope_inplace(
data: &mut [f32],
n_head: usize,
head_dim: usize,
rot_dim: usize,
theta: f32,
pos: [i32; 4],
sections: [u32; 4],
) {
let half_dim = head_dim / 2;
let half_rope = rot_dim / 2;
let sect_total = sections.iter().sum::<u32>().max(1);
let pick_axis = |sec: u32| -> usize {
let s0 = sections[0];
let s1 = s0 + sections[1];
let s2 = s1 + sections[2];
if sec < s0 {
0
} else if sec < s1 {
1
} else if sec < s2 {
2
} else {
3
}
};
for h in 0..n_head {
let base = h * head_dim;
for pair in 0..half_rope {
let sector = (pair as u32) % sect_total;
let axis = pick_axis(sector);
let p = pos[axis] as f32;
let dim_ratio = 2.0 * pair as f32 / rot_dim as f32;
let freq = 1.0 / theta.powf(dim_ratio);
let angle = p * freq;
let (ca, sa) = (angle.cos(), angle.sin());
let x0 = data[base + pair];
let x1 = data[base + pair + half_dim];
data[base + pair] = x0 * ca - x1 * sa;
data[base + pair + half_dim] = x0 * sa + x1 * ca;
}
}
}
let q_total = nq * d;
let kv_total = nkv * d;
let gqa = nq / nkv;
let mut x_norm = vec![0.0f32; seq * h];
for t in 0..seq {
let normed = rms_norm_row(
&hidden_states_in[t * h..(t + 1) * h],
&weights.attn_norm,
eps,
);
x_norm[t * h..(t + 1) * h].copy_from_slice(&normed);
}
let q_flat = matmul(&x_norm, &weights.wq, seq, h, q_total);
let k_flat = matmul(&x_norm, &weights.wk, seq, h, kv_total);
let v_flat = matmul(&x_norm, &weights.wv, seq, h, kv_total);
let gate = matmul(&x_norm, &weights.w_gate, seq, h, q_total);
let mut q = q_flat;
for t in 0..seq {
for hq in 0..nq {
let base = (t * nq + hq) * d;
let normed = rms_norm_row(&q[base..base + d], &weights.attn_q_norm, eps);
q[base..base + d].copy_from_slice(&normed);
}
}
let mut k = k_flat;
for t in 0..seq {
for hk in 0..nkv {
let base = (t * nkv + hk) * d;
let normed = rms_norm_row(&k[base..base + d], &weights.attn_k_norm, eps);
k[base..base + d].copy_from_slice(&normed);
}
}
for t in 0..seq {
let base = t * nq * d;
imrope_inplace(
&mut q[base..base + nq * d],
nq,
d,
rotary_dim,
rope_theta,
positions[t],
mrope_section,
);
}
for t in 0..seq {
let base = t * nkv * d;
imrope_inplace(
&mut k[base..base + nkv * d],
nkv,
d,
rotary_dim,
rope_theta,
positions[t],
mrope_section,
);
}
for kv_head in 0..nkv {
for pos in 0..seq {
let src_off = (pos * nkv + kv_head) * d;
let dst_off = kv_head * cap * d + (prefix + pos) * d;
k_cache_cpu[dst_off..dst_off + d].copy_from_slice(&k[src_off..src_off + d]);
v_cache_cpu[dst_off..dst_off + d].copy_from_slice(&v_flat[src_off..src_off + d]);
}
}
let kv_seq = prefix + seq;
let scale = 1.0_f32 / (d as f32).sqrt();
let mut attn_out = vec![0.0f32; seq * nq * d];
for t_q in 0..seq {
for hq in 0..nq {
let hkv = hq / gqa;
let q_off = (t_q * nq + hq) * d;
let mut logits = vec![f32::NEG_INFINITY; kv_seq];
for t_k in 0..kv_seq {
let mask_val = tree_mask[t_q * mask_stride + t_k];
if mask_val >= -1.0 {
let k_off = hkv * cap * d + t_k * d;
let mut dot = 0.0f32;
for i in 0..d {
dot += q[q_off + i] * k_cache_cpu[k_off + i];
}
logits[t_k] = dot * scale + mask_val;
}
}
let max_l = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut exp_sum = 0.0f32;
let mut exp_logits: Vec<f32> = logits
.iter()
.map(|&l| {
if l == f32::NEG_INFINITY {
0.0
} else {
let e = (l - max_l).exp();
exp_sum += e;
e
}
})
.collect();
if exp_sum > 0.0 {
for e in exp_logits.iter_mut() {
*e /= exp_sum;
}
}
let out_off = (t_q * nq + hq) * d;
for t_k in 0..kv_seq {
let w = exp_logits[t_k];
if w > 0.0 {
let v_off = hkv * cap * d + t_k * d;
for i in 0..d {
attn_out[out_off + i] += w * v_cache_cpu[v_off + i];
}
}
}
}
}
for i in 0..attn_out.len() {
attn_out[i] *= sigmoid(gate[i]);
}
let o_out = matmul(&attn_out, &weights.wo, seq, q_total, h);
let mut out = hidden_states_in.to_vec();
for i in 0..out.len() {
out[i] += o_out[i];
}
out
}
#[test]
fn qwen35_tree_verify_attention_block_cpu_ref_parity_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 2;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let q_total = nq * d;
let kv_total = nkv * d;
let shape = Qwen35TreeVerifyLayerShape {
hidden_size: h as u32,
num_q_heads: nq as u32,
num_kv_heads: nkv as u32,
head_dim: d as u32,
tree_seq_len: seq as u32,
cache_prefix_len: prefix as u32,
kv_capacity: cap as u32,
mask_stride: (prefix + seq) as u32,
rotary_dim: 64,
freq_base: 1e7,
mrope_section: [11, 11, 10, 0],
rms_norm_eps: 1e-6,
attn_output_gate: true,
};
let mut seed = 0xC001_u32;
let cpu_weights = FullAttnLayerWeights {
attn_norm: vec![1.0f32; h],
post_attn_norm: vec![1.0f32; h],
wq: mk_rand(&mut seed, q_total * h, 0.05),
wk: mk_rand(&mut seed, kv_total * h, 0.05),
wv: mk_rand(&mut seed, kv_total * h, 0.05),
w_gate: mk_rand(&mut seed, q_total * h, 0.05),
attn_q_norm: vec![1.0f32; d],
attn_k_norm: vec![1.0f32; d],
wo: mk_rand(&mut seed, h * q_total, 0.05),
};
let gpu_weights = FullAttnWeightsGpu::from_cpu_f32(&cpu_weights, &device).unwrap();
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let hidden_in = upload_f32(&hidden_data, &device).unwrap();
let tree_mask_data: Vec<f32> = {
let stride = prefix + seq;
let mut m = vec![mlx_native::ops::tree_attention::TREE_MASK_MASKED; seq * stride];
for i in 0..seq {
for j in 0..prefix + i + 1 {
if j < stride {
m[i * stride + j] = mlx_native::ops::tree_attention::TREE_MASK_ATTENDED;
}
}
}
m
};
let tree_mask = upload_f32(&tree_mask_data, &device).unwrap();
let tree_pos = upload_positions(seq, prefix as u32, &device);
let mut k_cache = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache = alloc_kv_cache(nkv, cap, d, &device);
let enc = device.command_encoder().expect("encoder");
let gpu_out = qwen35_tree_verify_attention_block(
enc,
&device,
&mut registry,
&hidden_in,
&tree_mask,
&tree_pos,
&mut k_cache,
&mut v_cache,
&gpu_weights,
shape,
)
.expect("gpu block");
let gpu_data = download_f32(&gpu_out).unwrap();
let mut k_cache_cpu = vec![0.0f32; nkv * cap * d];
let mut v_cache_cpu = vec![0.0f32; nkv * cap * d];
let positions: Vec<[i32; 4]> = (0..seq)
.map(|i| {
let p = (prefix + i) as i32;
[p, p, p, p]
})
.collect();
let cpu_data = cpu_tree_verify_attention_block_ref(
&hidden_data,
&tree_mask_data,
&positions,
&mut k_cache_cpu,
&mut v_cache_cpu,
&cpu_weights,
h,
nq,
nkv,
d,
seq,
cap,
prefix,
prefix + seq, 64,
1e7,
[11, 11, 10, 0],
1e-6,
);
assert_eq!(gpu_data.len(), cpu_data.len(), "output length mismatch");
let max_diff: f32 = gpu_data
.iter()
.zip(cpu_data.iter())
.map(|(g, c)| (g - c).abs())
.fold(0.0f32, f32::max);
eprintln!("T5: |GPU-CPU|_inf = {max_diff:.6e}");
assert!(
max_diff < 5e-2,
"T5 FAIL: |GPU-CPU|_inf = {max_diff:.6e} >= 5e-2 (BF16 slop budget). \
If this is expected noise, document the actual floor and widen budget."
);
eprintln!("T5 PASS: CPU reference parity |GPU-CPU|_inf = {max_diff:.6e} < 5e-2");
}
fn assert_tree_verify_attention_block_prefix0_chain_parity() {
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 2;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 3;
let prefix: usize = 0;
let cap: usize = 8;
let shape = Qwen35TreeVerifyLayerShape {
hidden_size: h as u32,
num_q_heads: nq as u32,
num_kv_heads: nkv as u32,
head_dim: d as u32,
tree_seq_len: seq as u32,
cache_prefix_len: prefix as u32,
kv_capacity: cap as u32,
mask_stride: seq as u32, rotary_dim: 64,
freq_base: 1e7,
mrope_section: [11, 11, 10, 0],
rms_norm_eps: 1e-6,
attn_output_gate: true,
};
let mut seed = 0xD001_u32;
let cpu_weights = FullAttnLayerWeights {
attn_norm: vec![1.0f32; h],
post_attn_norm: vec![1.0f32; h],
wq: mk_rand(&mut seed, nq * d * h, 0.05),
wk: mk_rand(&mut seed, nkv * d * h, 0.05),
wv: mk_rand(&mut seed, nkv * d * h, 0.05),
w_gate: mk_rand(&mut seed, nq * d * h, 0.05),
attn_q_norm: vec![1.0f32; d],
attn_k_norm: vec![1.0f32; d],
wo: mk_rand(&mut seed, h * nq * d, 0.05),
};
let weights = FullAttnWeightsGpu::from_cpu_f32(&cpu_weights, &device).unwrap();
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let hidden_in = upload_f32(&hidden_data, &device).unwrap();
let mask_data: Vec<f32> = {
let mut m = vec![mlx_native::ops::tree_attention::TREE_MASK_MASKED; seq * seq];
for i in 0..seq {
for j in 0..=i {
m[i * seq + j] = mlx_native::ops::tree_attention::TREE_MASK_ATTENDED;
}
}
m
};
let mask = upload_f32(&mask_data, &device).unwrap();
let pos = upload_positions(seq, 0, &device);
let mut k_cache1 = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache1 = alloc_kv_cache(nkv, cap, d, &device);
let enc1 = device.command_encoder().expect("enc1");
let out1 = qwen35_tree_verify_attention_block(
enc1,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache1,
&mut v_cache1,
&weights,
shape,
)
.expect("run1");
let mut k_cache2 = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache2 = alloc_kv_cache(nkv, cap, d, &device);
let enc2 = device.command_encoder().expect("enc2");
let out2 = qwen35_tree_verify_attention_block(
enc2,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache2,
&mut v_cache2,
&weights,
shape,
)
.expect("run2");
let d1 = download_f32(&out1).unwrap();
let d2 = download_f32(&out2).unwrap();
assert_eq!(d1.len(), d2.len());
let max_diff: f32 = d1
.iter()
.zip(d2.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
eprintln!("T6: prefix=0 repeat |diff|_inf = {max_diff:.6e}");
assert!(
max_diff < 1e-3,
"T6 FAIL: prefix=0 chain parity |diff|_inf = {max_diff:.6e} >= 1e-3"
);
eprintln!("T6 PASS: prefix=0 chain parity |diff|_inf = {max_diff:.6e} < 1e-3");
}
#[test]
fn qwen35_tree_verify_attention_block_prefix0_chain_parity_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
assert_tree_verify_attention_block_prefix0_chain_parity();
}
#[test]
fn nomic_forward_then_qwen_tree_verify_repeat_parity_2026_08_06() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
crate::inference::models::nomic_bert::forward::tests::run_synthetic_min_forward_for_cross_family_test();
assert_tree_verify_attention_block_prefix0_chain_parity();
}
#[test]
fn qwen35_tree_verify_attention_block_determinism_3rep_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 5120;
let nq: usize = 40;
let nkv: usize = 8;
let d: usize = 128;
let seq: usize = 4;
let prefix: usize = 64;
let cap: usize = 128;
let shape = Qwen35TreeVerifyLayerShape {
hidden_size: h as u32,
num_q_heads: nq as u32,
num_kv_heads: nkv as u32,
head_dim: d as u32,
tree_seq_len: seq as u32,
cache_prefix_len: prefix as u32,
kv_capacity: cap as u32,
mask_stride: (prefix + seq) as u32,
rotary_dim: 64,
freq_base: 1e7,
mrope_section: [11, 11, 10, 0],
rms_norm_eps: 1e-6,
attn_output_gate: true,
};
let mut seed = 0xE001_u32;
let weights = layer_weights_f32(h, nq, nkv, d, &mut seed, &device);
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let mask =
causal_tree_mask_with_prefix(seq as u32, prefix as u32, (prefix + seq) as u32, &device);
let pos = upload_positions(seq, prefix as u32, &device);
let mut outputs: Vec<Vec<f32>> = Vec::new();
for rep in 0..3 {
let hidden_in = upload_f32(&hidden_data, &device).unwrap();
let mut k_cache = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache = alloc_kv_cache(nkv, cap, d, &device);
let enc = device.command_encoder().expect("encoder");
let out = qwen35_tree_verify_attention_block(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&weights,
shape,
)
.unwrap_or_else(|e| panic!("T7 rep {} failed: {e}", rep));
outputs.push(download_f32(&out).unwrap());
}
for rep in 1..3 {
let first = &outputs[0];
let this = &outputs[rep];
assert_eq!(first.len(), this.len(), "T7: rep {rep} length mismatch");
for (i, (a, b)) in first.iter().zip(this.iter()).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"T7 FAIL: rep {rep} output[{i}] differs: first={a} this={b}"
);
}
}
eprintln!("T7 PASS: 3× repeat byte-identical (no Metal scheduling races or stale-reads)");
}
fn ffn_weights_f32(
hidden_size: usize,
intermediate_size: usize,
seed: &mut u32,
device: &MlxDevice,
) -> (
super::super::gpu_ffn::DenseFfnWeightsGpu,
super::super::ffn::DenseFfnWeights,
) {
use super::super::ffn::DenseFfnWeights;
use super::super::gpu_ffn::DenseFfnWeightsGpu;
let cpu = DenseFfnWeights {
gate: mk_rand(seed, intermediate_size * hidden_size, 0.05),
up: mk_rand(seed, intermediate_size * hidden_size, 0.05),
down: mk_rand(seed, hidden_size * intermediate_size, 0.05),
};
let gpu = DenseFfnWeightsGpu::from_cpu(&cpu, device).expect("upload ffn weights");
(gpu, cpu)
}
fn full_layer_shape_tiny(intermediate_size: u32) -> Qwen35TreeVerifyFullLayerShape {
Qwen35TreeVerifyFullLayerShape {
attn: layer_shape(128, 2, 1, 2, 4, 8),
intermediate_size,
}
}
#[test]
fn qwen35_tree_verify_full_layer_shape_validate_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
{
let shape = Qwen35TreeVerifyFullLayerShape {
attn: layer_shape(128, 2, 1, 2, 4, 8),
intermediate_size: 0,
};
let err = shape.validate().unwrap_err();
assert!(
err.to_string().contains("intermediate_size"),
"(a) zero intermediate_size not rejected: {err}"
);
}
{
let shape = Qwen35TreeVerifyFullLayerShape {
attn: layer_shape(128, 2, 1, 2, 4, 8),
intermediate_size: 1024 * 1024, };
shape
.validate()
.expect("(b) large but valid intermediate_size should pass");
}
{
let mut attn = layer_shape(128, 2, 1, 2, 4, 8);
attn.head_dim = 256;
let shape = Qwen35TreeVerifyFullLayerShape {
attn,
intermediate_size: 192,
};
let err = shape.validate().unwrap_err();
assert!(
err.to_string().contains("head_dim"),
"(c) head_dim != 128 not propagated: {err}"
);
}
{
let attn = Qwen35TreeVerifyLayerShape {
hidden_size: 5120,
num_q_heads: 40,
num_kv_heads: 8,
head_dim: 128,
tree_seq_len: 4,
cache_prefix_len: 64,
kv_capacity: 128,
mask_stride: 68,
rotary_dim: 64,
freq_base: 1e7,
mrope_section: [11, 11, 10, 0],
rms_norm_eps: 1e-6,
attn_output_gate: true,
};
let shape = Qwen35TreeVerifyFullLayerShape {
attn,
intermediate_size: 27648,
};
shape
.validate()
.expect("(d) valid Qwen 3.6 27B shape must pass");
}
eprintln!("AC-1 PASS: shape validate rejects all invalid shapes");
}
#[test]
fn qwen35_tree_verify_full_layer_smoke_production_gqa_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let hidden_size: usize = 5120;
let num_q_heads: usize = 40;
let num_kv_heads: usize = 8;
let head_dim: usize = 128;
let intermediate_size: usize = 27648;
let tree_seq_len: usize = 4;
let cache_prefix_len: usize = 64;
let kv_capacity: usize = 128;
let attn_shape = Qwen35TreeVerifyLayerShape {
hidden_size: hidden_size as u32,
num_q_heads: num_q_heads as u32,
num_kv_heads: num_kv_heads as u32,
head_dim: head_dim as u32,
tree_seq_len: tree_seq_len as u32,
cache_prefix_len: cache_prefix_len as u32,
kv_capacity: kv_capacity as u32,
mask_stride: (cache_prefix_len + tree_seq_len) as u32,
rotary_dim: 64,
freq_base: 1e7,
mrope_section: [11, 11, 10, 0],
rms_norm_eps: 1e-6,
attn_output_gate: true,
};
let full_shape = Qwen35TreeVerifyFullLayerShape {
attn: attn_shape,
intermediate_size: intermediate_size as u32,
};
let mut seed = 0xF001_u32;
let attn_weights = layer_weights_f32(
hidden_size,
num_q_heads,
num_kv_heads,
head_dim,
&mut seed,
&device,
);
let (ffn_gpu, _ffn_cpu) =
ffn_weights_f32(hidden_size, intermediate_size, &mut seed, &device);
let hidden_data = mk_rand(&mut seed, tree_seq_len * hidden_size, 0.1);
let hidden_in = upload_f32(&hidden_data, &device).unwrap();
let tree_mask = causal_tree_mask_with_prefix(
tree_seq_len as u32,
cache_prefix_len as u32,
(cache_prefix_len + tree_seq_len) as u32,
&device,
);
let tree_pos = upload_positions(tree_seq_len, cache_prefix_len as u32, &device);
let mut k_cache = alloc_kv_cache(num_kv_heads, kv_capacity, head_dim, &device);
let mut v_cache = alloc_kv_cache(num_kv_heads, kv_capacity, head_dim, &device);
let enc = device.command_encoder().expect("encoder");
let out = qwen35_tree_verify_full_layer(
enc,
&device,
&mut registry,
&hidden_in,
&tree_mask,
&tree_pos,
&mut k_cache,
&mut v_cache,
&attn_weights,
&ffn_gpu,
full_shape,
)
.expect("AC-2: full_layer call failed");
assert_eq!(out.dtype(), DType::F32, "AC-2(a) dtype");
assert_eq!(out.shape(), &[tree_seq_len, hidden_size], "AC-2(b) shape");
let out_data = download_f32(&out).unwrap();
assert!(
out_data.iter().all(|v| v.is_finite()),
"AC-2(c) non-finite output"
);
let k_data = k_cache.as_slice::<f32>().expect("k_cache slice");
let slot_start = 0 * kv_capacity * head_dim + cache_prefix_len * head_dim;
let slot = &k_data[slot_start..slot_start + head_dim];
assert!(
slot.iter().any(|&v| v != 0.0),
"AC-2(d) K cache slot [64, 68) still all-zero — cache write via attn block failed"
);
eprintln!(
"AC-2 PASS: production GQA smoke hidden={hidden_size} nq={num_q_heads} \
nkv={num_kv_heads} d={head_dim} intermediate={intermediate_size} \
tree_seq={tree_seq_len} prefix={cache_prefix_len}"
);
}
#[test]
fn qwen35_tree_verify_full_layer_negative_paths_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 2;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let m: usize = 192;
let mut seed = 0xA002_u32;
let base_attn_weights = layer_weights_f32(h, nq, nkv, d, &mut seed, &device);
let (base_ffn_gpu, _) = ffn_weights_f32(h, m, &mut seed, &device);
let make_inputs = |device: &MlxDevice, seed: &mut u32| {
let hidden_in = upload_f32(&mk_rand(seed, seq * h, 0.1), device).unwrap();
let mask = causal_tree_mask_with_prefix(
seq as u32,
prefix as u32,
(prefix + seq) as u32,
device,
);
let pos = upload_positions(seq, prefix as u32, device);
let k_cache = alloc_kv_cache(nkv, cap, d, device);
let v_cache = alloc_kv_cache(nkv, cap, d, device);
(hidden_in, mask, pos, k_cache, v_cache)
};
let valid_full_shape = Qwen35TreeVerifyFullLayerShape {
attn: layer_shape(
h as u32,
nq as u32,
nkv as u32,
seq as u32,
prefix as u32,
cap as u32,
),
intermediate_size: m as u32,
};
{
use super::super::ffn::DenseFfnWeights;
use super::super::gpu_ffn::DenseFfnWeightsGpu;
let wrong_gate_cpu = DenseFfnWeights {
gate: mk_rand(&mut seed, (m - 1) * h, 0.05), up: mk_rand(&mut seed, m * h, 0.05),
down: mk_rand(&mut seed, h * m, 0.05),
};
let wrong_ffn = DenseFfnWeightsGpu::from_cpu(&wrong_gate_cpu, &device).unwrap();
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&wrong_ffn,
valid_full_shape,
)
.unwrap_err();
assert!(
err.to_string().contains("ffn_weights.gate"),
"(a) gate wrong count not caught: {err}"
);
}
{
use super::super::ffn::DenseFfnWeights;
use super::super::gpu_ffn::DenseFfnWeightsGpu;
let wrong_up_cpu = DenseFfnWeights {
gate: mk_rand(&mut seed, m * h, 0.05),
up: mk_rand(&mut seed, (m + 1) * h, 0.05), down: mk_rand(&mut seed, h * m, 0.05),
};
let wrong_ffn = DenseFfnWeightsGpu::from_cpu(&wrong_up_cpu, &device).unwrap();
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&wrong_ffn,
valid_full_shape,
)
.unwrap_err();
assert!(
err.to_string().contains("ffn_weights.up"),
"(b) up wrong count not caught: {err}"
);
}
{
use super::super::ffn::DenseFfnWeights;
use super::super::gpu_ffn::DenseFfnWeightsGpu;
let wrong_down_cpu = DenseFfnWeights {
gate: mk_rand(&mut seed, m * h, 0.05),
up: mk_rand(&mut seed, m * h, 0.05),
down: mk_rand(&mut seed, h * (m - 1), 0.05), };
let wrong_ffn = DenseFfnWeightsGpu::from_cpu(&wrong_down_cpu, &device).unwrap();
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&wrong_ffn,
valid_full_shape,
)
.unwrap_err();
assert!(
err.to_string().contains("ffn_weights.down"),
"(c) down wrong count not caught: {err}"
);
}
{
let bad_shape = Qwen35TreeVerifyFullLayerShape {
attn: layer_shape(
h as u32,
nq as u32,
nkv as u32,
seq as u32,
prefix as u32,
cap as u32,
),
intermediate_size: 0,
};
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&base_ffn_gpu,
bad_shape,
)
.unwrap_err();
assert!(
err.to_string().contains("intermediate_size"),
"(d) intermediate_size=0 not rejected via full entry: {err}"
);
}
{
let mut bad_attn = layer_shape(
h as u32,
nq as u32,
nkv as u32,
seq as u32,
prefix as u32,
cap as u32,
);
bad_attn.head_dim = 64; let bad_shape = Qwen35TreeVerifyFullLayerShape {
attn: bad_attn,
intermediate_size: m as u32,
};
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&base_ffn_gpu,
bad_shape,
)
.unwrap_err();
assert!(
err.to_string().contains("head_dim"),
"(e) head_dim != 128 not propagated via full entry: {err}"
);
}
eprintln!("AC-3 PASS: all 5 negative paths reject with descriptive errors via full function entry");
}
fn cpu_tree_verify_full_layer_ref(
hidden_states_in: &[f32],
tree_mask: &[f32],
positions: &[[i32; 4]],
k_cache_cpu: &mut [f32],
v_cache_cpu: &mut [f32],
attn_weights: &FullAttnLayerWeights,
ffn_gate: &[f32],
ffn_up: &[f32],
ffn_down: &[f32],
h: usize,
nq: usize,
nkv: usize,
d: usize,
seq: usize,
cap: usize,
prefix: usize,
mask_stride: usize,
intermediate_size: usize,
rotary_dim: usize,
rope_theta: f32,
mrope_section: [u32; 4],
eps: f32,
) -> Vec<f32> {
fn rms_norm_row_local(x: &[f32], w: &[f32], eps: f32) -> Vec<f32> {
let n = x.len() as f32;
let ss: f32 = x.iter().map(|v| v * v).sum::<f32>();
let inv = (ss / n + eps).sqrt().recip();
x.iter().zip(w).map(|(xi, wi)| xi * inv * wi).collect()
}
fn matmul_local(lhs: &[f32], rhs: &[f32], m_: usize, k_: usize, n_: usize) -> Vec<f32> {
let mut out = vec![0.0f32; m_ * n_];
for i in 0..m_ {
for j in 0..n_ {
let mut acc = 0.0f32;
for kk in 0..k_ {
acc += lhs[i * k_ + kk] * rhs[j * k_ + kk];
}
out[i * n_ + j] = acc;
}
}
out
}
fn silu_local(x: f32) -> f32 {
x / (1.0 + (-x).exp())
}
let attn_out = cpu_tree_verify_attention_block_ref(
hidden_states_in,
tree_mask,
positions,
k_cache_cpu,
v_cache_cpu,
attn_weights,
h,
nq,
nkv,
d,
seq,
cap,
prefix,
mask_stride,
rotary_dim,
rope_theta,
mrope_section,
eps,
);
let ffn_residual = attn_out.clone();
let mut ffn_input = vec![0.0f32; seq * h];
for t in 0..seq {
let row = rms_norm_row_local(
&attn_out[t * h..(t + 1) * h],
&attn_weights.post_attn_norm,
eps,
);
ffn_input[t * h..(t + 1) * h].copy_from_slice(&row);
}
let gate_proj = matmul_local(&ffn_input, ffn_gate, seq, h, intermediate_size);
let up_proj = matmul_local(&ffn_input, ffn_up, seq, h, intermediate_size);
let mut activated = vec![0.0f32; seq * intermediate_size];
for i in 0..activated.len() {
activated[i] = silu_local(gate_proj[i]) * up_proj[i];
}
let ffn_out = matmul_local(&activated, ffn_down, seq, intermediate_size, h);
let mut out = ffn_residual;
for i in 0..out.len() {
out[i] += ffn_out[i];
}
out
}
#[test]
fn qwen35_tree_verify_full_layer_cpu_ref_parity_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 2;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let m: usize = 192; let q_total = nq * d;
let kv_total = nkv * d;
let full_shape = full_layer_shape_tiny(m as u32);
let mut seed = 0xC002_u32;
let cpu_attn_weights = FullAttnLayerWeights {
attn_norm: vec![1.0f32; h],
post_attn_norm: mk_rand(&mut seed, h, 0.5), wq: mk_rand(&mut seed, q_total * h, 0.05),
wk: mk_rand(&mut seed, kv_total * h, 0.05),
wv: mk_rand(&mut seed, kv_total * h, 0.05),
w_gate: mk_rand(&mut seed, q_total * h, 0.05),
attn_q_norm: vec![1.0f32; d],
attn_k_norm: vec![1.0f32; d],
wo: mk_rand(&mut seed, h * q_total, 0.05),
};
let ffn_gate_cpu = mk_rand(&mut seed, m * h, 0.05);
let ffn_up_cpu = mk_rand(&mut seed, m * h, 0.05);
let ffn_down_cpu = mk_rand(&mut seed, h * m, 0.05);
let gpu_attn_weights =
FullAttnWeightsGpu::from_cpu_f32(&cpu_attn_weights, &device).unwrap();
use super::super::ffn::DenseFfnWeights;
use super::super::gpu_ffn::DenseFfnWeightsGpu;
let ffn_cpu_weights = DenseFfnWeights {
gate: ffn_gate_cpu.clone(),
up: ffn_up_cpu.clone(),
down: ffn_down_cpu.clone(),
};
let gpu_ffn_weights = DenseFfnWeightsGpu::from_cpu(&ffn_cpu_weights, &device).unwrap();
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let hidden_in = upload_f32(&hidden_data, &device).unwrap();
let mask_stride = prefix + seq;
let tree_mask_data: Vec<f32> = {
let mut mv = vec![mlx_native::ops::tree_attention::TREE_MASK_MASKED; seq * mask_stride];
for i in 0..seq {
for j in 0..prefix + i + 1 {
if j < mask_stride {
mv[i * mask_stride + j] =
mlx_native::ops::tree_attention::TREE_MASK_ATTENDED;
}
}
}
mv
};
let tree_mask = upload_f32(&tree_mask_data, &device).unwrap();
let tree_pos = upload_positions(seq, prefix as u32, &device);
let mut k_cache = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache = alloc_kv_cache(nkv, cap, d, &device);
let enc = device.command_encoder().expect("encoder");
let gpu_out = qwen35_tree_verify_full_layer(
enc,
&device,
&mut registry,
&hidden_in,
&tree_mask,
&tree_pos,
&mut k_cache,
&mut v_cache,
&gpu_attn_weights,
&gpu_ffn_weights,
full_shape,
)
.expect("AC-4: GPU full_layer");
let gpu_data = download_f32(&gpu_out).unwrap();
let mut k_cache_cpu = vec![0.0f32; nkv * cap * d];
let mut v_cache_cpu = vec![0.0f32; nkv * cap * d];
let positions: Vec<[i32; 4]> = (0..seq)
.map(|i| {
let p = (prefix + i) as i32;
[p, p, p, p]
})
.collect();
let cpu_data = cpu_tree_verify_full_layer_ref(
&hidden_data,
&tree_mask_data,
&positions,
&mut k_cache_cpu,
&mut v_cache_cpu,
&cpu_attn_weights,
&ffn_gate_cpu,
&ffn_up_cpu,
&ffn_down_cpu,
h,
nq,
nkv,
d,
seq,
cap,
prefix,
mask_stride,
m,
64,
1e7,
[11, 11, 10, 0],
1e-6,
);
assert_eq!(
gpu_data.len(),
cpu_data.len(),
"AC-4: output length mismatch"
);
let max_diff: f32 = gpu_data
.iter()
.zip(cpu_data.iter())
.map(|(g, c)| (g - c).abs())
.fold(0.0f32, f32::max);
eprintln!("AC-4: |GPU-CPU|_inf = {max_diff:.6e}");
assert!(
max_diff < 5e-2,
"AC-4 FAIL: |GPU-CPU|_inf = {max_diff:.6e} >= 5e-2 (BF16 slop budget). \
Check post_attn_norm weights or gate/up/down matmul chain."
);
eprintln!("AC-4 PASS: CPU reference parity |GPU-CPU|_inf = {max_diff:.6e} < 5e-2");
}
#[test]
fn qwen35_tree_verify_full_layer_composition_equivalence_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 2;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let m: usize = 192;
let q_total = nq * d;
let kv_total = nkv * d;
let full_shape = full_layer_shape_tiny(m as u32);
let mut seed = 0xB002_u32;
let cpu_attn_weights = FullAttnLayerWeights {
attn_norm: vec![1.0f32; h],
post_attn_norm: vec![1.0f32; h],
wq: mk_rand(&mut seed, q_total * h, 0.05),
wk: mk_rand(&mut seed, kv_total * h, 0.05),
wv: mk_rand(&mut seed, kv_total * h, 0.05),
w_gate: mk_rand(&mut seed, q_total * h, 0.05),
attn_q_norm: vec![1.0f32; d],
attn_k_norm: vec![1.0f32; d],
wo: mk_rand(&mut seed, h * q_total, 0.05),
};
let ffn_gate_cpu = mk_rand(&mut seed, m * h, 0.05);
let ffn_up_cpu = mk_rand(&mut seed, m * h, 0.05);
let ffn_down_cpu = mk_rand(&mut seed, h * m, 0.05);
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let gpu_attn_weights =
FullAttnWeightsGpu::from_cpu_f32(&cpu_attn_weights, &device).unwrap();
use super::super::ffn::DenseFfnWeights;
use super::super::gpu_ffn::DenseFfnWeightsGpu;
let ffn_cpu_weights = DenseFfnWeights {
gate: ffn_gate_cpu.clone(),
up: ffn_up_cpu.clone(),
down: ffn_down_cpu.clone(),
};
let gpu_ffn_weights = DenseFfnWeightsGpu::from_cpu(&ffn_cpu_weights, &device).unwrap();
let hidden_in = upload_f32(&hidden_data, &device).unwrap();
let mask_stride = prefix + seq;
let tree_mask_data: Vec<f32> = {
let mut mv = vec![mlx_native::ops::tree_attention::TREE_MASK_MASKED; seq * mask_stride];
for i in 0..seq {
for j in 0..prefix + i + 1 {
if j < mask_stride {
mv[i * mask_stride + j] =
mlx_native::ops::tree_attention::TREE_MASK_ATTENDED;
}
}
}
mv
};
let tree_mask = upload_f32(&tree_mask_data, &device).unwrap();
let tree_pos = upload_positions(seq, prefix as u32, &device);
let mut k_cache = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache = alloc_kv_cache(nkv, cap, d, &device);
let enc = device.command_encoder().expect("enc");
let gpu_out = qwen35_tree_verify_full_layer(
enc,
&device,
&mut registry,
&hidden_in,
&tree_mask,
&tree_pos,
&mut k_cache,
&mut v_cache,
&gpu_attn_weights,
&gpu_ffn_weights,
full_shape,
)
.expect("AC-5: GPU full_layer");
let gpu_data = download_f32(&gpu_out).unwrap();
let mut k_cache_cpu = vec![0.0f32; nkv * cap * d];
let mut v_cache_cpu = vec![0.0f32; nkv * cap * d];
let positions: Vec<[i32; 4]> = (0..seq)
.map(|i| {
let p = (prefix + i) as i32;
[p, p, p, p]
})
.collect();
let cpu_data = cpu_tree_verify_full_layer_ref(
&hidden_data,
&tree_mask_data,
&positions,
&mut k_cache_cpu,
&mut v_cache_cpu,
&cpu_attn_weights,
&ffn_gate_cpu,
&ffn_up_cpu,
&ffn_down_cpu,
h,
nq,
nkv,
d,
seq,
cap,
prefix,
mask_stride,
m,
64,
1e7,
[11, 11, 10, 0],
1e-6,
);
assert_eq!(gpu_data.len(), cpu_data.len(), "AC-5: length mismatch");
let max_diff: f32 = gpu_data
.iter()
.zip(cpu_data.iter())
.map(|(g, c)| (g - c).abs())
.fold(0.0f32, f32::max);
eprintln!("AC-5: |GPU-CPU|_inf = {max_diff:.6e}");
assert!(
max_diff < 5e-2,
"AC-5 FAIL: composition divergence {max_diff:.6e} >= 5e-2 — full_layer \
wrapper does not compose attn+MLP correctly. \
Check post_attn_norm weight, gate/up/down routing, or residual source."
);
eprintln!("AC-5 PASS: composition equivalence |GPU-CPU|_inf = {max_diff:.6e} < 5e-2");
}
#[test]
fn qwen35_tree_verify_full_layer_determinism_3rep_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 2;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let m: usize = 192;
let full_shape = full_layer_shape_tiny(m as u32);
let mut seed = 0xD002_u32;
let attn_weights = layer_weights_f32(h, nq, nkv, d, &mut seed, &device);
let (ffn_gpu, _) = ffn_weights_f32(h, m, &mut seed, &device);
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let mask =
causal_tree_mask_with_prefix(seq as u32, prefix as u32, (prefix + seq) as u32, &device);
let pos = upload_positions(seq, prefix as u32, &device);
let mut outputs: Vec<Vec<f32>> = Vec::new();
for rep in 0..3 {
let hidden_in = upload_f32(&hidden_data, &device).unwrap();
let mut k_cache = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache = alloc_kv_cache(nkv, cap, d, &device);
let enc = device.command_encoder().expect("enc");
let out = qwen35_tree_verify_full_layer(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&attn_weights,
&ffn_gpu,
full_shape,
)
.unwrap_or_else(|e| panic!("AC-6: rep {} failed: {e}", rep));
outputs.push(download_f32(&out).unwrap());
}
for rep in 1..3 {
let first = &outputs[0];
let this = &outputs[rep];
assert_eq!(first.len(), this.len(), "AC-6: rep {rep} length mismatch");
for (i, (a, b)) in first.iter().zip(this.iter()).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"AC-6 FAIL: rep {rep} output[{i}] differs: first={a:.6e} ({:#010x}) \
this={b:.6e} ({:#010x})",
a.to_bits(),
b.to_bits()
);
}
}
eprintln!("AC-6 PASS: 3× byte-identical (0 ULP) determinism across K/V cache resets");
}
fn full_layer_shape_q_tiny(intermediate_size: u32) -> Qwen35TreeVerifyFullLayerShapeQ {
Qwen35TreeVerifyFullLayerShapeQ {
attn: layer_shape(128, 2, 1, 2, 4, 8),
intermediate_size,
}
}
fn upload_q4_0(data: &[f32], n_per_row: usize, device: &MlxDevice) -> MlxBuffer {
use crate::quantize::ggml_quants::q4_0;
let bytes = q4_0::quantize(data, n_per_row, None);
let mut buf = device
.alloc_buffer(bytes.len(), mlx_native::DType::U8, vec![bytes.len()])
.expect("alloc Q4_0 buf");
buf.as_mut_slice::<u8>()
.expect("q-buf slice")
.copy_from_slice(&bytes);
buf
}
fn dequant_q4_0_cpu(data: &[u8]) -> Vec<f32> {
const BLOCK_BYTES: usize = 18;
const BLOCK_ELEMS: usize = 32;
assert!(data.len() % BLOCK_BYTES == 0, "Q4_0 data not block-aligned");
let num_blocks = data.len() / BLOCK_BYTES;
let mut out = vec![0.0f32; num_blocks * BLOCK_ELEMS];
for i in 0..num_blocks {
let block = &data[i * BLOCK_BYTES..(i + 1) * BLOCK_BYTES];
let d = half::f16::from_le_bytes([block[0], block[1]]).to_f32();
let qs = &block[2..18];
let out_block = &mut out[i * BLOCK_ELEMS..(i + 1) * BLOCK_ELEMS];
for j in 0..16 {
let x0 = (qs[j] & 0x0F) as i16 - 8;
let x1 = (qs[j] >> 4) as i16 - 8;
out_block[j] = x0 as f32 * d;
out_block[j + 16] = x1 as f32 * d;
}
}
out
}
fn ffn_weights_q4_0(
hidden_size: usize,
intermediate_size: usize,
seed: &mut u32,
device: &MlxDevice,
) -> (
super::super::gpu_ffn::DenseFfnWeightsGpuQ,
Vec<f32>,
Vec<f32>,
Vec<f32>,
) {
let gate_f32 = mk_rand(seed, intermediate_size * hidden_size, 0.05);
let up_f32 = mk_rand(seed, intermediate_size * hidden_size, 0.05);
let down_f32 = mk_rand(seed, hidden_size * intermediate_size, 0.05);
let gate_q = upload_q4_0(&gate_f32, hidden_size, device);
let up_q = upload_q4_0(&up_f32, hidden_size, device);
let down_q = upload_q4_0(&down_f32, intermediate_size, device);
let weights_q = super::super::gpu_ffn::DenseFfnWeightsGpuQ {
gate_q,
up_q,
down_q,
ggml_type_gate_up: GgmlType::Q4_0,
ggml_type_down: GgmlType::Q4_0,
intermediate_size: intermediate_size as u32,
hidden_size: hidden_size as u32,
};
(weights_q, gate_f32, up_f32, down_f32)
}
#[test]
fn qwen35_tree_verify_full_layer_q_shape_validate_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
{
let shape = Qwen35TreeVerifyFullLayerShapeQ {
attn: layer_shape(128, 2, 1, 2, 4, 8),
intermediate_size: 0,
};
let err = shape.validate().unwrap_err();
assert!(
err.to_string().contains("intermediate_size"),
"(a) zero intermediate_size not rejected: {err}"
);
}
{
let shape = Qwen35TreeVerifyFullLayerShapeQ {
attn: layer_shape(128, 2, 1, 2, 4, 8),
intermediate_size: 1024 * 1024,
};
shape
.validate()
.expect("(b) large but valid intermediate_size should pass");
}
{
let mut attn = layer_shape(128, 2, 1, 2, 4, 8);
attn.head_dim = 256;
let shape = Qwen35TreeVerifyFullLayerShapeQ {
attn,
intermediate_size: 192,
};
let err = shape.validate().unwrap_err();
assert!(
err.to_string().contains("head_dim"),
"(c) head_dim != 128 not propagated: {err}"
);
}
{
let attn = Qwen35TreeVerifyLayerShape {
hidden_size: 5120,
num_q_heads: 40,
num_kv_heads: 8,
head_dim: 128,
tree_seq_len: 4,
cache_prefix_len: 64,
kv_capacity: 128,
mask_stride: 68,
rotary_dim: 64,
freq_base: 1e7,
mrope_section: [11, 11, 10, 0],
rms_norm_eps: 1e-6,
attn_output_gate: true,
};
let shape = Qwen35TreeVerifyFullLayerShapeQ {
attn,
intermediate_size: 27648,
};
shape
.validate()
.expect("(d) valid Qwen 3.6 27B shape must pass");
}
eprintln!("AC-1 (Q4_0) PASS: shape validate rejects all invalid shapes");
}
#[test]
fn qwen35_tree_verify_full_layer_q_smoke_production_gqa_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let hidden_size: usize = 5120;
let num_q_heads: usize = 40;
let num_kv_heads: usize = 8;
let head_dim: usize = 128;
let intermediate_size: usize = 27648;
let tree_seq_len: usize = 4;
let cache_prefix_len: usize = 64;
let kv_capacity: usize = 128;
let attn_shape = Qwen35TreeVerifyLayerShape {
hidden_size: hidden_size as u32,
num_q_heads: num_q_heads as u32,
num_kv_heads: num_kv_heads as u32,
head_dim: head_dim as u32,
tree_seq_len: tree_seq_len as u32,
cache_prefix_len: cache_prefix_len as u32,
kv_capacity: kv_capacity as u32,
mask_stride: (cache_prefix_len + tree_seq_len) as u32,
rotary_dim: 64,
freq_base: 1e7,
mrope_section: [11, 11, 10, 0],
rms_norm_eps: 1e-6,
attn_output_gate: true,
};
let full_shape_q = Qwen35TreeVerifyFullLayerShapeQ {
attn: attn_shape,
intermediate_size: intermediate_size as u32,
};
let mut seed = 0xF003_u32;
let attn_weights = layer_weights_f32(
hidden_size,
num_q_heads,
num_kv_heads,
head_dim,
&mut seed,
&device,
);
let (ffn_gpu_q, _, _, _) =
ffn_weights_q4_0(hidden_size, intermediate_size, &mut seed, &device);
let hidden_data = mk_rand(&mut seed, tree_seq_len * hidden_size, 0.1);
let hidden_in = upload_f32(&hidden_data, &device).unwrap();
let tree_mask = causal_tree_mask_with_prefix(
tree_seq_len as u32,
cache_prefix_len as u32,
(cache_prefix_len + tree_seq_len) as u32,
&device,
);
let tree_pos = upload_positions(tree_seq_len, cache_prefix_len as u32, &device);
let mut k_cache = alloc_kv_cache(num_kv_heads, kv_capacity, head_dim, &device);
let mut v_cache = alloc_kv_cache(num_kv_heads, kv_capacity, head_dim, &device);
let enc = device.command_encoder().expect("encoder");
let out = qwen35_tree_verify_full_layer_q(
enc,
&device,
&mut registry,
&hidden_in,
&tree_mask,
&tree_pos,
&mut k_cache,
&mut v_cache,
&attn_weights,
&ffn_gpu_q,
full_shape_q,
)
.expect("AC-2 Q4_0: full_layer_q call failed");
assert_eq!(out.dtype(), DType::F32, "AC-2(a) Q4_0 dtype");
assert_eq!(
out.shape(),
&[tree_seq_len, hidden_size],
"AC-2(b) Q4_0 shape"
);
let out_data = download_f32(&out).unwrap();
assert!(
out_data.iter().all(|v| v.is_finite()),
"AC-2(c) Q4_0 non-finite output"
);
let k_data = k_cache.as_slice::<f32>().expect("k_cache slice");
let slot_start = 0 * kv_capacity * head_dim + cache_prefix_len * head_dim;
let slot = &k_data[slot_start..slot_start + head_dim];
assert!(
slot.iter().any(|&v| v != 0.0),
"AC-2(d) Q4_0 K cache slot [64, 68) still all-zero"
);
let v_data = v_cache.as_slice::<f32>().expect("v_cache slice");
let v_slot = &v_data[slot_start..slot_start + head_dim];
assert!(
v_slot.iter().any(|&v| v != 0.0),
"AC-2(e) Q4_0 V cache slot [64, 68) still all-zero"
);
eprintln!(
"AC-2 (Q4_0) PASS: production GQA smoke hidden={hidden_size} nq={num_q_heads} \
nkv={num_kv_heads} d={head_dim} intermediate={intermediate_size} \
tree_seq={tree_seq_len} prefix={cache_prefix_len}"
);
}
#[test]
fn qwen35_tree_verify_full_layer_q_negative_paths_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 2;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let m: usize = 192;
let mut seed = 0xA003_u32;
let base_attn_weights = layer_weights_f32(h, nq, nkv, d, &mut seed, &device);
let (base_ffn_q, _, _, _) = ffn_weights_q4_0(h, m, &mut seed, &device);
let make_inputs = |device: &MlxDevice, seed: &mut u32| {
let hidden_in = upload_f32(&mk_rand(seed, seq * h, 0.1), device).unwrap();
let mask = causal_tree_mask_with_prefix(
seq as u32,
prefix as u32,
(prefix + seq) as u32,
device,
);
let pos = upload_positions(seq, prefix as u32, device);
let k_cache = alloc_kv_cache(nkv, cap, d, device);
let v_cache = alloc_kv_cache(nkv, cap, d, device);
(hidden_in, mask, pos, k_cache, v_cache)
};
let valid_shape_q = Qwen35TreeVerifyFullLayerShapeQ {
attn: layer_shape(
h as u32,
nq as u32,
nkv as u32,
seq as u32,
prefix as u32,
cap as u32,
),
intermediate_size: m as u32,
};
{
use crate::quantize::ggml_quants::q4_0;
let wrong_gate_f32 = mk_rand(&mut seed, m * h, 0.05);
let correct_bytes = q4_0::quantize(&wrong_gate_f32, h, None);
let mut wrong_bytes = correct_bytes.clone();
wrong_bytes.pop();
let mut wrong_gate_buf = device
.alloc_buffer(wrong_bytes.len(), DType::U8, vec![wrong_bytes.len()])
.unwrap();
wrong_gate_buf
.as_mut_slice::<u8>()
.unwrap()
.copy_from_slice(&wrong_bytes);
let wrong_ffn_q = super::super::gpu_ffn::DenseFfnWeightsGpuQ {
gate_q: wrong_gate_buf,
up_q: base_ffn_q.up_q.clone(),
down_q: base_ffn_q.down_q.clone(),
ggml_type_gate_up: GgmlType::Q4_0,
ggml_type_down: GgmlType::Q4_0,
intermediate_size: m as u32,
hidden_size: h as u32,
};
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer_q(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&wrong_ffn_q,
valid_shape_q,
)
.unwrap_err();
assert!(
err.to_string().contains("gate"),
"(a) gate_q wrong byte length not caught: {err}"
);
}
{
use crate::quantize::ggml_quants::q4_0;
let wrong_up_f32 = mk_rand(&mut seed, m * h, 0.05);
let correct_bytes = q4_0::quantize(&wrong_up_f32, h, None);
let mut wrong_bytes = correct_bytes.clone();
wrong_bytes.pop();
let mut wrong_up_buf = device
.alloc_buffer(wrong_bytes.len(), DType::U8, vec![wrong_bytes.len()])
.unwrap();
wrong_up_buf
.as_mut_slice::<u8>()
.unwrap()
.copy_from_slice(&wrong_bytes);
let wrong_ffn_q = super::super::gpu_ffn::DenseFfnWeightsGpuQ {
gate_q: base_ffn_q.gate_q.clone(),
up_q: wrong_up_buf,
down_q: base_ffn_q.down_q.clone(),
ggml_type_gate_up: GgmlType::Q4_0,
ggml_type_down: GgmlType::Q4_0,
intermediate_size: m as u32,
hidden_size: h as u32,
};
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer_q(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&wrong_ffn_q,
valid_shape_q,
)
.unwrap_err();
assert!(
err.to_string().contains("up"),
"(b) up_q wrong byte length not caught: {err}"
);
}
{
use crate::quantize::ggml_quants::q4_0;
let wrong_down_f32 = mk_rand(&mut seed, h * m, 0.05);
let correct_bytes = q4_0::quantize(&wrong_down_f32, m, None);
let mut wrong_bytes = correct_bytes.clone();
wrong_bytes.pop();
let mut wrong_down_buf = device
.alloc_buffer(wrong_bytes.len(), DType::U8, vec![wrong_bytes.len()])
.unwrap();
wrong_down_buf
.as_mut_slice::<u8>()
.unwrap()
.copy_from_slice(&wrong_bytes);
let wrong_ffn_q = super::super::gpu_ffn::DenseFfnWeightsGpuQ {
gate_q: base_ffn_q.gate_q.clone(),
up_q: base_ffn_q.up_q.clone(),
down_q: wrong_down_buf,
ggml_type_gate_up: GgmlType::Q4_0,
ggml_type_down: GgmlType::Q4_0,
intermediate_size: m as u32,
hidden_size: h as u32,
};
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer_q(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&wrong_ffn_q,
valid_shape_q,
)
.unwrap_err();
assert!(
err.to_string().contains("down"),
"(c) down_q wrong byte length not caught: {err}"
);
}
{
let bad_shape = Qwen35TreeVerifyFullLayerShapeQ {
attn: layer_shape(
h as u32,
nq as u32,
nkv as u32,
seq as u32,
prefix as u32,
cap as u32,
),
intermediate_size: 0,
};
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer_q(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&base_ffn_q,
bad_shape,
)
.unwrap_err();
assert!(
err.to_string().contains("intermediate_size"),
"(d) intermediate_size=0 not rejected via full entry: {err}"
);
}
{
let mut bad_attn = layer_shape(
h as u32,
nq as u32,
nkv as u32,
seq as u32,
prefix as u32,
cap as u32,
);
bad_attn.head_dim = 256;
let bad_shape = Qwen35TreeVerifyFullLayerShapeQ {
attn: bad_attn,
intermediate_size: m as u32,
};
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer_q(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&base_ffn_q,
bad_shape,
)
.unwrap_err();
assert!(
err.to_string().contains("head_dim"),
"(e) head_dim != 128 not propagated via full entry: {err}"
);
}
{
let wrong_type_ffn_q = super::super::gpu_ffn::DenseFfnWeightsGpuQ {
gate_q: base_ffn_q.gate_q.clone(),
up_q: base_ffn_q.up_q.clone(),
down_q: base_ffn_q.down_q.clone(),
ggml_type_gate_up: GgmlType::Q5_K,
ggml_type_down: GgmlType::Q4_0,
intermediate_size: m as u32,
hidden_size: h as u32,
};
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer_q(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&wrong_type_ffn_q,
valid_shape_q,
)
.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("ggml_type_gate_up"),
"(f) ggml_type_gate_up != Q4_0 not caught: {msg}"
);
assert!(
msg.contains("Q4_0"),
"(f) error does not mention Q4_0 requirement: {msg}"
);
}
{
let wrong_hidden_ffn_q = super::super::gpu_ffn::DenseFfnWeightsGpuQ {
gate_q: base_ffn_q.gate_q.clone(),
up_q: base_ffn_q.up_q.clone(),
down_q: base_ffn_q.down_q.clone(),
ggml_type_gate_up: GgmlType::Q4_0,
ggml_type_down: GgmlType::Q4_0,
intermediate_size: m as u32,
hidden_size: h as u32 + 1, };
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer_q(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&wrong_hidden_ffn_q,
valid_shape_q,
)
.unwrap_err();
assert!(
err.to_string().contains("hidden_size"),
"(g) hidden_size mismatch not caught: {err}"
);
}
eprintln!("AC-3 (Q4_0) PASS: all 7 negative paths reject with descriptive errors via full function entry");
}
#[test]
fn qwen35_tree_verify_full_layer_q_cpu_reference_parity_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 2;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let m: usize = 192;
let q_total = nq * d;
let kv_total = nkv * d;
let full_shape_q = full_layer_shape_q_tiny(m as u32);
let mut seed = 0xC003_u32;
let cpu_attn_weights = FullAttnLayerWeights {
attn_norm: vec![1.0f32; h],
post_attn_norm: mk_rand(&mut seed, h, 0.5),
wq: mk_rand(&mut seed, q_total * h, 0.05),
wk: mk_rand(&mut seed, kv_total * h, 0.05),
wv: mk_rand(&mut seed, kv_total * h, 0.05),
w_gate: mk_rand(&mut seed, q_total * h, 0.05),
attn_q_norm: vec![1.0f32; d],
attn_k_norm: vec![1.0f32; d],
wo: mk_rand(&mut seed, h * q_total, 0.05),
};
let gate_f32 = mk_rand(&mut seed, m * h, 0.05);
let up_f32 = mk_rand(&mut seed, m * h, 0.05);
let down_f32 = mk_rand(&mut seed, h * m, 0.05);
use crate::quantize::ggml_quants::q4_0;
let gate_q_bytes = q4_0::quantize(&gate_f32, h, None);
let up_q_bytes = q4_0::quantize(&up_f32, h, None);
let down_q_bytes = q4_0::quantize(&down_f32, m, None);
let gate_dq = dequant_q4_0_cpu(&gate_q_bytes);
let up_dq = dequant_q4_0_cpu(&up_q_bytes);
let down_dq = dequant_q4_0_cpu(&down_q_bytes);
let gpu_attn_weights =
FullAttnWeightsGpu::from_cpu_f32(&cpu_attn_weights, &device).unwrap();
let make_u8_buf = |bytes: &[u8], device: &MlxDevice| -> MlxBuffer {
let mut buf = device
.alloc_buffer(bytes.len(), DType::U8, vec![bytes.len()])
.expect("alloc q4_0 buf");
buf.as_mut_slice::<u8>().unwrap().copy_from_slice(bytes);
buf
};
let gpu_ffn_q = super::super::gpu_ffn::DenseFfnWeightsGpuQ {
gate_q: make_u8_buf(&gate_q_bytes, &device),
up_q: make_u8_buf(&up_q_bytes, &device),
down_q: make_u8_buf(&down_q_bytes, &device),
ggml_type_gate_up: GgmlType::Q4_0,
ggml_type_down: GgmlType::Q4_0,
intermediate_size: m as u32,
hidden_size: h as u32,
};
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let hidden_in = upload_f32(&hidden_data, &device).unwrap();
let mask_stride = prefix + seq;
let tree_mask_data: Vec<f32> = {
let mut mv = vec![mlx_native::ops::tree_attention::TREE_MASK_MASKED; seq * mask_stride];
for i in 0..seq {
for j in 0..prefix + i + 1 {
if j < mask_stride {
mv[i * mask_stride + j] =
mlx_native::ops::tree_attention::TREE_MASK_ATTENDED;
}
}
}
mv
};
let tree_mask = upload_f32(&tree_mask_data, &device).unwrap();
let tree_pos = upload_positions(seq, prefix as u32, &device);
let mut k_cache = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache = alloc_kv_cache(nkv, cap, d, &device);
let enc = device.command_encoder().expect("encoder");
let gpu_out = qwen35_tree_verify_full_layer_q(
enc,
&device,
&mut registry,
&hidden_in,
&tree_mask,
&tree_pos,
&mut k_cache,
&mut v_cache,
&gpu_attn_weights,
&gpu_ffn_q,
full_shape_q,
)
.expect("AC-4 Q4_0: GPU full_layer_q");
let gpu_data = download_f32(&gpu_out).unwrap();
let mut k_cache_cpu = vec![0.0f32; nkv * cap * d];
let mut v_cache_cpu = vec![0.0f32; nkv * cap * d];
let positions: Vec<[i32; 4]> = (0..seq)
.map(|i| {
let p = (prefix + i) as i32;
[p, p, p, p]
})
.collect();
let cpu_data = cpu_tree_verify_full_layer_ref(
&hidden_data,
&tree_mask_data,
&positions,
&mut k_cache_cpu,
&mut v_cache_cpu,
&cpu_attn_weights,
&gate_dq,
&up_dq,
&down_dq,
h,
nq,
nkv,
d,
seq,
cap,
prefix,
mask_stride,
m,
64,
1e7,
[11, 11, 10, 0],
1e-6,
);
assert_eq!(
gpu_data.len(),
cpu_data.len(),
"AC-4 Q4_0: output length mismatch"
);
let max_diff: f32 = gpu_data
.iter()
.zip(cpu_data.iter())
.map(|(g, c)| (g - c).abs())
.fold(0.0f32, f32::max);
eprintln!("AC-4 (Q4_0): |GPU-CPU|_inf = {max_diff:.6e}");
assert!(
max_diff < 0.15,
"AC-4 (Q4_0) FAIL: |GPU-CPU|_inf = {max_diff:.6e} >= 0.15 (Q4_0 dequant slop budget). \
Check dequant_q4_0_cpu matches mlx-native's Q4_0 dequant or gate/up/down routing."
);
eprintln!("AC-4 (Q4_0) PASS: CPU reference parity |GPU-CPU|_inf = {max_diff:.6e} < 0.15");
}
#[test]
fn qwen35_tree_verify_full_layer_q_composition_equivalence_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 2;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let m: usize = 192;
let q_total = nq * d;
let kv_total = nkv * d;
let full_shape_q = full_layer_shape_q_tiny(m as u32);
let mut seed = 0xB003_u32;
let cpu_attn_weights = FullAttnLayerWeights {
attn_norm: vec![1.0f32; h],
post_attn_norm: vec![1.0f32; h],
wq: mk_rand(&mut seed, q_total * h, 0.05),
wk: mk_rand(&mut seed, kv_total * h, 0.05),
wv: mk_rand(&mut seed, kv_total * h, 0.05),
w_gate: mk_rand(&mut seed, q_total * h, 0.05),
attn_q_norm: vec![1.0f32; d],
attn_k_norm: vec![1.0f32; d],
wo: mk_rand(&mut seed, h * q_total, 0.05),
};
let gate_f32 = mk_rand(&mut seed, m * h, 0.05);
let up_f32 = mk_rand(&mut seed, m * h, 0.05);
let down_f32 = mk_rand(&mut seed, h * m, 0.05);
use crate::quantize::ggml_quants::q4_0;
let gate_q_bytes = q4_0::quantize(&gate_f32, h, None);
let up_q_bytes = q4_0::quantize(&up_f32, h, None);
let down_q_bytes = q4_0::quantize(&down_f32, m, None);
let gate_dq = dequant_q4_0_cpu(&gate_q_bytes);
let up_dq = dequant_q4_0_cpu(&up_q_bytes);
let down_dq = dequant_q4_0_cpu(&down_q_bytes);
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let gpu_attn_weights =
FullAttnWeightsGpu::from_cpu_f32(&cpu_attn_weights, &device).unwrap();
let make_u8_buf = |bytes: &[u8], device: &MlxDevice| -> MlxBuffer {
let mut buf = device
.alloc_buffer(bytes.len(), DType::U8, vec![bytes.len()])
.expect("alloc q4_0 buf");
buf.as_mut_slice::<u8>().unwrap().copy_from_slice(bytes);
buf
};
let gpu_ffn_q = super::super::gpu_ffn::DenseFfnWeightsGpuQ {
gate_q: make_u8_buf(&gate_q_bytes, &device),
up_q: make_u8_buf(&up_q_bytes, &device),
down_q: make_u8_buf(&down_q_bytes, &device),
ggml_type_gate_up: GgmlType::Q4_0,
ggml_type_down: GgmlType::Q4_0,
intermediate_size: m as u32,
hidden_size: h as u32,
};
let hidden_in = upload_f32(&hidden_data, &device).unwrap();
let mask_stride = prefix + seq;
let tree_mask_data: Vec<f32> = {
let mut mv = vec![mlx_native::ops::tree_attention::TREE_MASK_MASKED; seq * mask_stride];
for i in 0..seq {
for j in 0..prefix + i + 1 {
if j < mask_stride {
mv[i * mask_stride + j] =
mlx_native::ops::tree_attention::TREE_MASK_ATTENDED;
}
}
}
mv
};
let tree_mask = upload_f32(&tree_mask_data, &device).unwrap();
let tree_pos = upload_positions(seq, prefix as u32, &device);
let mut k_cache = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache = alloc_kv_cache(nkv, cap, d, &device);
let enc = device.command_encoder().expect("enc");
let gpu_out = qwen35_tree_verify_full_layer_q(
enc,
&device,
&mut registry,
&hidden_in,
&tree_mask,
&tree_pos,
&mut k_cache,
&mut v_cache,
&gpu_attn_weights,
&gpu_ffn_q,
full_shape_q,
)
.expect("AC-5 Q4_0: GPU full_layer_q");
let gpu_data = download_f32(&gpu_out).unwrap();
let mut k_cache_cpu = vec![0.0f32; nkv * cap * d];
let mut v_cache_cpu = vec![0.0f32; nkv * cap * d];
let positions: Vec<[i32; 4]> = (0..seq)
.map(|i| {
let p = (prefix + i) as i32;
[p, p, p, p]
})
.collect();
let cpu_data = cpu_tree_verify_full_layer_ref(
&hidden_data,
&tree_mask_data,
&positions,
&mut k_cache_cpu,
&mut v_cache_cpu,
&cpu_attn_weights,
&gate_dq,
&up_dq,
&down_dq,
h,
nq,
nkv,
d,
seq,
cap,
prefix,
mask_stride,
m,
64,
1e7,
[11, 11, 10, 0],
1e-6,
);
assert_eq!(gpu_data.len(), cpu_data.len(), "AC-5 Q4_0: length mismatch");
let max_diff: f32 = gpu_data
.iter()
.zip(cpu_data.iter())
.map(|(g, c)| (g - c).abs())
.fold(0.0f32, f32::max);
eprintln!("AC-5 (Q4_0): |GPU-CPU|_inf = {max_diff:.6e}");
assert!(
max_diff < 0.15,
"AC-5 (Q4_0) FAIL: composition divergence {max_diff:.6e} >= 0.15"
);
eprintln!(
"AC-5 (Q4_0) PASS: composition equivalence |GPU-CPU|_inf = {max_diff:.6e} < 0.15"
);
}
#[test]
fn qwen35_tree_verify_full_layer_q_byte_identity_3rep_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 2;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let m: usize = 192;
let full_shape_q = full_layer_shape_q_tiny(m as u32);
let mut seed = 0xD003_u32;
let attn_weights = layer_weights_f32(h, nq, nkv, d, &mut seed, &device);
let (ffn_gpu_q, _, _, _) = ffn_weights_q4_0(h, m, &mut seed, &device);
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let mask =
causal_tree_mask_with_prefix(seq as u32, prefix as u32, (prefix + seq) as u32, &device);
let pos = upload_positions(seq, prefix as u32, &device);
let mut outputs: Vec<Vec<u32>> = Vec::new();
for rep in 0..3 {
let hidden_in = upload_f32(&hidden_data, &device).unwrap();
let mut k_cache = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache = alloc_kv_cache(nkv, cap, d, &device);
let enc = device.command_encoder().expect("enc");
let out = qwen35_tree_verify_full_layer_q(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&attn_weights,
&ffn_gpu_q,
full_shape_q,
)
.unwrap_or_else(|e| panic!("AC-6 Q4_0: rep {} failed: {e}", rep));
let bits: Vec<u32> = download_f32(&out)
.unwrap()
.iter()
.map(|f| f.to_bits())
.collect();
outputs.push(bits);
}
for rep in 1..3 {
let first = &outputs[0];
let this = &outputs[rep];
assert_eq!(
first.len(),
this.len(),
"AC-6 Q4_0: rep {rep} length mismatch"
);
for (i, (a, b)) in first.iter().zip(this.iter()).enumerate() {
assert_eq!(
a, b,
"AC-6 (Q4_0) FAIL: rep {rep} output[{i}] differs: {:#010x} vs {:#010x}",
a, b
);
}
}
eprintln!("AC-6 (Q4_0) PASS: 3× byte-identical (0 ULP) determinism");
}
#[test]
fn qwen35_tree_verify_full_layer_q_cross_variant_parity_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 2;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let m: usize = 192;
let mut seed = 0xE003_u32;
let attn_weights = layer_weights_f32(h, nq, nkv, d, &mut seed, &device);
let gate_f32 = mk_rand(&mut seed, m * h, 0.05);
let up_f32 = mk_rand(&mut seed, m * h, 0.05);
let down_f32 = mk_rand(&mut seed, h * m, 0.05);
use super::super::ffn::DenseFfnWeights;
use super::super::gpu_ffn::DenseFfnWeightsGpu;
let ffn_cpu_weights = DenseFfnWeights {
gate: gate_f32.clone(),
up: up_f32.clone(),
down: down_f32.clone(),
};
let ffn_gpu_f32 = DenseFfnWeightsGpu::from_cpu(&ffn_cpu_weights, &device).unwrap();
let full_shape_f32 = Qwen35TreeVerifyFullLayerShape {
attn: layer_shape(
h as u32,
nq as u32,
nkv as u32,
seq as u32,
prefix as u32,
cap as u32,
),
intermediate_size: m as u32,
};
use crate::quantize::ggml_quants::q4_0;
let gate_q_bytes = q4_0::quantize(&gate_f32, h, None);
let up_q_bytes = q4_0::quantize(&up_f32, h, None);
let down_q_bytes = q4_0::quantize(&down_f32, m, None);
let make_u8_buf = |bytes: &[u8], device: &MlxDevice| -> MlxBuffer {
let mut buf = device
.alloc_buffer(bytes.len(), DType::U8, vec![bytes.len()])
.expect("alloc q4_0 buf");
buf.as_mut_slice::<u8>().unwrap().copy_from_slice(bytes);
buf
};
let ffn_gpu_q = super::super::gpu_ffn::DenseFfnWeightsGpuQ {
gate_q: make_u8_buf(&gate_q_bytes, &device),
up_q: make_u8_buf(&up_q_bytes, &device),
down_q: make_u8_buf(&down_q_bytes, &device),
ggml_type_gate_up: GgmlType::Q4_0,
ggml_type_down: GgmlType::Q4_0,
intermediate_size: m as u32,
hidden_size: h as u32,
};
let full_shape_q = Qwen35TreeVerifyFullLayerShapeQ {
attn: layer_shape(
h as u32,
nq as u32,
nkv as u32,
seq as u32,
prefix as u32,
cap as u32,
),
intermediate_size: m as u32,
};
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let mask_stride = prefix + seq;
let tree_mask_data: Vec<f32> = {
let mut mv = vec![mlx_native::ops::tree_attention::TREE_MASK_MASKED; seq * mask_stride];
for i in 0..seq {
for j in 0..prefix + i + 1 {
if j < mask_stride {
mv[i * mask_stride + j] =
mlx_native::ops::tree_attention::TREE_MASK_ATTENDED;
}
}
}
mv
};
let tree_mask = upload_f32(&tree_mask_data, &device).unwrap();
let tree_pos = upload_positions(seq, prefix as u32, &device);
let hidden_in_a = upload_f32(&hidden_data, &device).unwrap();
let mut k_cache_a = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache_a = alloc_kv_cache(nkv, cap, d, &device);
let enc_a = device.command_encoder().expect("enc_a");
let out_a = qwen35_tree_verify_full_layer(
enc_a,
&device,
&mut registry,
&hidden_in_a,
&tree_mask,
&tree_pos,
&mut k_cache_a,
&mut v_cache_a,
&attn_weights,
&ffn_gpu_f32,
full_shape_f32,
)
.expect("AC-7: F1 path failed");
let data_a = download_f32(&out_a).unwrap();
let hidden_in_b = upload_f32(&hidden_data, &device).unwrap();
let mut k_cache_b = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache_b = alloc_kv_cache(nkv, cap, d, &device);
let enc_b = device.command_encoder().expect("enc_b");
let out_b = qwen35_tree_verify_full_layer_q(
enc_b,
&device,
&mut registry,
&hidden_in_b,
&tree_mask,
&tree_pos,
&mut k_cache_b,
&mut v_cache_b,
&attn_weights,
&ffn_gpu_q,
full_shape_q,
)
.expect("AC-7: F2 path failed");
let data_b = download_f32(&out_b).unwrap();
assert_eq!(
data_a.len(),
data_b.len(),
"AC-7: output length mismatch F1 vs F2"
);
let max_diff: f32 = data_a
.iter()
.zip(data_b.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
eprintln!("AC-7 (cross-variant): |F32-cast - Q4_0|_inf = {max_diff:.6e}");
assert!(
max_diff < 0.20,
"AC-7 FAIL: cross-variant divergence |F32-cast - Q4_0|_inf = {max_diff:.6e} >= 0.20. \
Check that gate/up/down use &ffn_weights.gate_q/.up_q/.down_q (U8 buffers) \
and that apply_linear_projection_f32's U8 branch is routing to quantized_matmul_ggml."
);
eprintln!(
"AC-7 PASS: Q4_0 GPU ≈ F32-cast GPU at |.|_inf = {max_diff:.6e} < 0.20 \
(proves Q4_0 path performs same computation as F32-cast within Q4_0 dequant slop)"
);
}
fn moe_layer_shape_tiny(
ne: u32,
topk: u32,
m_moe: u32,
m_sh: u32,
) -> Qwen35TreeVerifyFullLayerShapeQMoe {
Qwen35TreeVerifyFullLayerShapeQMoe {
attn: layer_shape(128, 4, 1, 2, 4, 8),
moe: super::super::ffn::MoeFfnShape {
hidden_size: 128,
num_experts: ne,
num_experts_per_tok: topk,
moe_intermediate_size: m_moe,
shared_intermediate_size: m_sh,
},
}
}
#[allow(clippy::type_complexity)]
fn moe_ffn_weights_q4_0(
h: usize,
ne: usize,
m_moe: usize,
m_sh: usize,
seed: &mut u32,
device: &MlxDevice,
) -> (
super::super::gpu_ffn::MoeFfnWeightsGpuQ,
super::super::ffn::MoeFfnWeights,
) {
let router_f32 = mk_rand(seed, ne * h, 0.3);
let expert_gate_f32 = mk_rand(seed, ne * m_moe * h, 0.1);
let expert_up_f32 = mk_rand(seed, ne * m_moe * h, 0.1);
let expert_down_f32 = mk_rand(seed, ne * h * m_moe, 0.1);
let shared_gate_logit_f32 = mk_rand(seed, h, 0.1);
let shared_gate_f32 = mk_rand(seed, m_sh * h, 0.1);
let shared_up_f32 = mk_rand(seed, m_sh * h, 0.1);
let shared_down_f32 = mk_rand(seed, h * m_sh, 0.1);
let gate_q4 = {
use crate::quantize::ggml_quants::q4_0;
q4_0::quantize(&expert_gate_f32, h, None)
};
let up_q4 = {
use crate::quantize::ggml_quants::q4_0;
q4_0::quantize(&expert_up_f32, h, None)
};
let down_q4 = {
use crate::quantize::ggml_quants::q4_0;
q4_0::quantize(&expert_down_f32, m_moe, None)
};
let expert_gate_dq = dequant_q4_0_cpu(&gate_q4);
let expert_up_dq = dequant_q4_0_cpu(&up_q4);
let expert_down_dq = dequant_q4_0_cpu(&down_q4);
let qk: usize = 32;
let block_bytes: usize = 18;
let gate_stride = ((m_moe * h / qk) * block_bytes) as u64;
let down_stride = ((h * m_moe / qk) * block_bytes) as u64;
let make_u8_buf = |bytes: &[u8]| -> MlxBuffer {
let mut buf = device
.alloc_buffer(bytes.len(), mlx_native::DType::U8, vec![bytes.len()])
.expect("alloc q4_0 buf");
buf.as_mut_slice::<u8>()
.expect("q-buf slice")
.copy_from_slice(bytes);
buf
};
let gpu_weights = super::super::gpu_ffn::MoeFfnWeightsGpuQ {
router: upload_bf16_from_f32(&router_f32, device).expect("upload router bf16"),
expert_gate_q: make_u8_buf(&gate_q4),
expert_up_q: make_u8_buf(&up_q4),
expert_down_q: make_u8_buf(&down_q4),
ggml_type_gate_up: GgmlType::Q4_0,
ggml_type_down: GgmlType::Q4_0,
expert_gate_stride: gate_stride,
expert_up_stride: gate_stride,
expert_down_stride: down_stride,
num_experts: ne as u32,
shared_gate_inp: upload_bf16_from_f32(&shared_gate_logit_f32, device)
.expect("sh_gate_inp"),
shared_gate: upload_bf16_from_f32(&shared_gate_f32, device).expect("sh_gate"),
shared_up: upload_bf16_from_f32(&shared_up_f32, device).expect("sh_up"),
shared_down: upload_bf16_from_f32(&shared_down_f32, device).expect("sh_down"),
expert_gate_affine: None,
expert_up_affine: None,
expert_down_affine: None,
};
let cpu_weights = super::super::ffn::MoeFfnWeights {
router: router_f32,
expert_gate: expert_gate_dq,
expert_up: expert_up_dq,
expert_down: expert_down_dq,
shared_gate_logit: shared_gate_logit_f32,
shared_gate: shared_gate_f32,
shared_up: shared_up_f32,
shared_down: shared_down_f32,
};
(gpu_weights, cpu_weights)
}
fn assert_all_nonzero(label: &str, v: &[f32]) {
assert!(
v.iter().any(|&x| x != 0.0),
"{label}: all-zero weight detected — identity-path test is forbidden (RF-5)"
);
}
#[allow(clippy::too_many_arguments)]
fn cpu_tree_verify_full_layer_q_moe_ref(
hidden_states_in: &[f32],
tree_mask: &[f32],
positions: &[[i32; 4]],
k_cache_cpu: &mut [f32],
v_cache_cpu: &mut [f32],
attn_weights: &FullAttnLayerWeights,
moe_weights: &super::super::ffn::MoeFfnWeights,
shape: &Qwen35TreeVerifyFullLayerShapeQMoe,
) -> Vec<f32> {
let h = shape.attn.hidden_size as usize;
let seq = shape.attn.tree_seq_len as usize;
let nq = shape.attn.num_q_heads as usize;
let nkv = shape.attn.num_kv_heads as usize;
let d = shape.attn.head_dim as usize;
let cap = shape.attn.kv_capacity as usize;
let prefix = shape.attn.cache_prefix_len as usize;
let mask_stride = shape.attn.mask_stride as usize;
let rotary_dim = shape.attn.rotary_dim as usize;
let rope_theta = shape.attn.freq_base;
let mrope_section = shape.attn.mrope_section;
let eps = shape.attn.rms_norm_eps;
fn rms_norm_row(x: &[f32], w: &[f32], eps: f32) -> Vec<f32> {
let n = x.len() as f32;
let ss: f32 = x.iter().map(|v| v * v).sum::<f32>();
let inv = (ss / n + eps).sqrt().recip();
x.iter().zip(w).map(|(xi, wi)| xi * inv * wi).collect()
}
let attn_out = cpu_tree_verify_attention_block_ref(
hidden_states_in,
tree_mask,
positions,
k_cache_cpu,
v_cache_cpu,
attn_weights,
h,
nq,
nkv,
d,
seq,
cap,
prefix,
mask_stride,
rotary_dim,
rope_theta,
mrope_section,
eps,
);
let ffn_residual = attn_out.clone();
let mut post_attn_normed = vec![0.0f32; seq * h];
for t in 0..seq {
let row = rms_norm_row(
&attn_out[t * h..(t + 1) * h],
&attn_weights.post_attn_norm,
eps,
);
post_attn_normed[t * h..(t + 1) * h].copy_from_slice(&row);
}
let moe_out = super::super::ffn::moe_ffn_cpu_ref(
&post_attn_normed,
moe_weights,
super::super::ffn::MoeFfnShape {
hidden_size: shape.moe.hidden_size,
num_experts: shape.moe.num_experts,
num_experts_per_tok: shape.moe.num_experts_per_tok,
moe_intermediate_size: shape.moe.moe_intermediate_size,
shared_intermediate_size: shape.moe.shared_intermediate_size,
},
);
let mut out = ffn_residual;
for i in 0..out.len() {
out[i] += moe_out[i];
}
out
}
#[test]
fn qwen35_tree_verify_full_layer_q_moe_shape_validate_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
{
let shape = Qwen35TreeVerifyFullLayerShapeQMoe {
attn: Qwen35TreeVerifyLayerShape {
hidden_size: 2048,
num_q_heads: 16,
num_kv_heads: 2,
head_dim: 128,
tree_seq_len: 8,
cache_prefix_len: 32,
kv_capacity: 8192,
mask_stride: 8192,
rotary_dim: 64,
freq_base: 1e7,
mrope_section: [11, 11, 10, 0],
rms_norm_eps: 1e-6,
attn_output_gate: true,
},
moe: super::super::ffn::MoeFfnShape {
hidden_size: 2048,
num_experts: 128,
num_experts_per_tok: 8,
moe_intermediate_size: 512,
shared_intermediate_size: 1024,
},
};
shape
.validate()
.expect("valid production-like shape must pass");
}
{
let mut shape = moe_layer_shape_tiny(4, 2, 64, 64);
shape.moe.num_experts = 0;
let err = shape.validate().unwrap_err();
assert!(err.to_string().contains("num_experts"), "(a): {err}");
}
{
let mut shape = moe_layer_shape_tiny(4, 2, 64, 64);
shape.moe.num_experts_per_tok = 0;
let err = shape.validate().unwrap_err();
assert!(
err.to_string().contains("num_experts_per_tok"),
"(b): {err}"
);
}
{
let mut shape = moe_layer_shape_tiny(4, 2, 64, 64);
shape.moe.num_experts_per_tok = 5;
let err = shape.validate().unwrap_err();
assert!(
err.to_string().contains("num_experts_per_tok")
|| err.to_string().contains("top-K"),
"(c): {err}"
);
}
{
let mut shape = moe_layer_shape_tiny(4, 2, 64, 64);
shape.moe.moe_intermediate_size = 0;
let err = shape.validate().unwrap_err();
assert!(
err.to_string().contains("moe_intermediate_size"),
"(d): {err}"
);
}
{
let mut shape = moe_layer_shape_tiny(4, 2, 64, 64);
shape.moe.shared_intermediate_size = 0;
let err = shape.validate().unwrap_err();
assert!(
err.to_string().contains("shared_intermediate_size"),
"(e): {err}"
);
}
{
let mut shape = moe_layer_shape_tiny(4, 2, 64, 64);
shape.moe.hidden_size = 256; let err = shape.validate().unwrap_err();
assert!(err.to_string().contains("hidden_size"), "(f): {err}");
}
{
let mut shape = moe_layer_shape_tiny(4, 2, 64, 64);
shape.attn.head_dim = 64;
let err = shape.validate().unwrap_err();
assert!(err.to_string().contains("head_dim"), "(g): {err}");
}
eprintln!("AC-1 (MoE) PASS: shape validate accepts valid and rejects all invalid shapes");
}
#[test]
fn qwen35_tree_verify_full_layer_q_moe_smoke_production_gqa_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 512;
let nq: usize = 8;
let nkv: usize = 2;
let d: usize = 128; let seq: usize = 4;
let prefix: usize = 16;
let cap: usize = 64;
let ne: usize = 8;
let topk: usize = 2;
let m_moe: usize = 256;
let m_sh: usize = 256;
let shape = Qwen35TreeVerifyFullLayerShapeQMoe {
attn: Qwen35TreeVerifyLayerShape {
hidden_size: h as u32,
num_q_heads: nq as u32,
num_kv_heads: nkv as u32,
head_dim: d as u32,
tree_seq_len: seq as u32,
cache_prefix_len: prefix as u32,
kv_capacity: cap as u32,
mask_stride: (prefix + seq) as u32,
rotary_dim: 64,
freq_base: 1e7,
mrope_section: [11, 11, 10, 0],
rms_norm_eps: 1e-6,
attn_output_gate: true,
},
moe: super::super::ffn::MoeFfnShape {
hidden_size: h as u32,
num_experts: ne as u32,
num_experts_per_tok: topk as u32,
moe_intermediate_size: m_moe as u32,
shared_intermediate_size: m_sh as u32,
},
};
let mut seed = 0xAC02_u32;
let attn_weights = layer_weights_f32(h, nq, nkv, d, &mut seed, &device);
let (moe_gpu_weights, _cpu_weights) =
moe_ffn_weights_q4_0(h, ne, m_moe, m_sh, &mut seed, &device);
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let hidden_in = upload_f32(&hidden_data, &device).unwrap();
let tree_mask =
causal_tree_mask_with_prefix(seq as u32, prefix as u32, (prefix + seq) as u32, &device);
let tree_pos = upload_positions(seq, prefix as u32, &device);
let mut k_cache = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache = alloc_kv_cache(nkv, cap, d, &device);
let k_slot_start = 0 * cap * d + prefix * d;
let enc = device.command_encoder().expect("enc");
let out = qwen35_tree_verify_full_layer_q_moe(
enc,
&device,
&mut registry,
&hidden_in,
&tree_mask,
&tree_pos,
&mut k_cache,
&mut v_cache,
&attn_weights,
&moe_gpu_weights,
shape,
)
.expect("AC-2 MoE: full_layer_q_moe call failed");
assert_eq!(out.dtype(), mlx_native::DType::F32, "AC-2(a) dtype");
assert_eq!(out.shape(), &[seq, h], "AC-2(b) shape");
let out_data = download_f32(&out).unwrap();
assert!(
out_data.iter().all(|v| v.is_finite()),
"AC-2(c) non-finite output"
);
assert!(
out_data.iter().any(|&v| v != 0.0),
"AC-2(c) all-zero output (MoE pipeline did not fire)"
);
assert!(
!out_data.iter().any(|v| v.is_nan()),
"AC-2(c) NaN in output"
);
let k_data = k_cache.as_slice::<f32>().expect("k_cache slice");
let k_slot = &k_data[k_slot_start..k_slot_start + d];
assert!(
k_slot.iter().any(|&v| v != 0.0),
"AC-2(d) K cache slot [{prefix}, {}) still all-zero",
prefix + seq
);
let v_data = v_cache.as_slice::<f32>().expect("v_cache slice");
let v_slot = &v_data[k_slot_start..k_slot_start + d];
assert!(
v_slot.iter().any(|&v| v != 0.0),
"AC-2(e) V cache slot [{prefix}, {}) still all-zero",
prefix + seq
);
eprintln!(
"AC-2 (MoE) PASS: smoke h={h} ne={ne} topk={topk} m_moe={m_moe} m_sh={m_sh} \
nq={nq} nkv={nkv} seq={seq} prefix={prefix}"
);
}
#[test]
fn qwen35_tree_verify_full_layer_q_moe_negative_paths_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 4;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let ne: usize = 4;
let topk: usize = 2;
let m_moe: usize = 64;
let m_sh: usize = 64;
let mut seed = 0xAC03_u32;
let base_attn_weights = layer_weights_f32(h, nq, nkv, d, &mut seed, &device);
let (base_moe_weights, _) = moe_ffn_weights_q4_0(h, ne, m_moe, m_sh, &mut seed, &device);
let valid_shape = moe_layer_shape_tiny(ne as u32, topk as u32, m_moe as u32, m_sh as u32);
let make_inputs = |device: &MlxDevice, seed: &mut u32| {
let hidden_in = upload_f32(&mk_rand(seed, seq * h, 0.1), device).unwrap();
let mask = causal_tree_mask_with_prefix(
seq as u32,
prefix as u32,
(prefix + seq) as u32,
device,
);
let pos = upload_positions(seq, prefix as u32, device);
let k_cache = alloc_kv_cache(nkv, cap, d, device);
let v_cache = alloc_kv_cache(nkv, cap, d, device);
(hidden_in, mask, pos, k_cache, v_cache)
};
let make_q4_0_buf = |bytes: &[u8], device: &MlxDevice| -> MlxBuffer {
let mut buf = device
.alloc_buffer(bytes.len(), mlx_native::DType::U8, vec![bytes.len()])
.expect("alloc q4_0 buf");
buf.as_mut_slice::<u8>()
.expect("slice")
.copy_from_slice(bytes);
buf
};
let make_weights = |router: MlxBuffer,
expert_gate_q: MlxBuffer,
expert_up_q: MlxBuffer,
expert_down_q: MlxBuffer,
ggml_type_gate_up: GgmlType,
ggml_type_down: GgmlType,
shared_gate_inp: MlxBuffer|
-> super::super::gpu_ffn::MoeFfnWeightsGpuQ {
super::super::gpu_ffn::MoeFfnWeightsGpuQ {
router,
expert_gate_q,
expert_up_q,
expert_down_q,
ggml_type_gate_up,
ggml_type_down,
expert_gate_stride: base_moe_weights.expert_gate_stride,
expert_up_stride: base_moe_weights.expert_up_stride,
expert_down_stride: base_moe_weights.expert_down_stride,
num_experts: base_moe_weights.num_experts,
shared_gate_inp,
shared_gate: base_moe_weights.shared_gate.clone(),
shared_up: base_moe_weights.shared_up.clone(),
shared_down: base_moe_weights.shared_down.clone(),
expert_gate_affine: None,
expert_up_affine: None,
expert_down_affine: None,
}
};
{
let bad_weights = make_weights(
base_moe_weights.router.clone(),
base_moe_weights.expert_gate_q.clone(),
base_moe_weights.expert_up_q.clone(),
base_moe_weights.expert_down_q.clone(),
GgmlType::Q5_K, GgmlType::Q4_0,
base_moe_weights.shared_gate_inp.clone(),
);
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer_q_moe(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&bad_weights,
valid_shape,
)
.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("ggml_type_gate_up must be Q4_0"),
"neg_1: wrong message: {msg}"
);
}
{
let bad_weights = make_weights(
base_moe_weights.router.clone(),
base_moe_weights.expert_gate_q.clone(),
base_moe_weights.expert_up_q.clone(),
base_moe_weights.expert_down_q.clone(),
GgmlType::Q4_0,
GgmlType::Q6_K, base_moe_weights.shared_gate_inp.clone(),
);
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer_q_moe(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&bad_weights,
valid_shape,
)
.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("ggml_type_down must be Q4_0"),
"neg_2: wrong message: {msg}"
);
}
{
let router_f32_buf = upload_f32(&mk_rand(&mut seed, ne * h, 0.3), &device).unwrap();
let bad_weights = make_weights(
router_f32_buf, base_moe_weights.expert_gate_q.clone(),
base_moe_weights.expert_up_q.clone(),
base_moe_weights.expert_down_q.clone(),
GgmlType::Q4_0,
GgmlType::Q4_0,
base_moe_weights.shared_gate_inp.clone(),
);
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer_q_moe(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&bad_weights,
valid_shape,
)
.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("router dtype must be BF16"),
"neg_3: wrong message: {msg}"
);
}
{
let bad_attn = Qwen35TreeVerifyLayerShape {
hidden_size: 2048,
num_q_heads: 16,
num_kv_heads: 2,
head_dim: 128,
tree_seq_len: 2,
cache_prefix_len: 4,
kv_capacity: 8,
mask_stride: 6,
rotary_dim: 64,
freq_base: 1e7,
mrope_section: [11, 11, 10, 0],
rms_norm_eps: 1e-6,
attn_output_gate: true,
};
let bad_shape2 = Qwen35TreeVerifyFullLayerShapeQMoe {
attn: bad_attn,
moe: super::super::ffn::MoeFfnShape {
hidden_size: 2048,
num_experts: ne as u32,
num_experts_per_tok: topk as u32,
moe_intermediate_size: m_moe as u32,
shared_intermediate_size: m_sh as u32,
},
};
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer_q_moe(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&base_moe_weights,
bad_shape2,
)
.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("router") || msg.contains("hidden_size"),
"neg_4: wrong message: {msg}"
);
}
{
let mut bad_shape = valid_shape;
bad_shape.moe.num_experts_per_tok = (ne + 1) as u32;
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer_q_moe(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&base_moe_weights,
bad_shape,
)
.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("num_experts_per_tok") || msg.contains("top-K"),
"neg_5: wrong message: {msg}"
);
}
{
use crate::quantize::ggml_quants::q4_0;
let correct_gate_f32 = mk_rand(&mut seed, ne * m_moe * h, 0.1);
let correct_bytes = q4_0::quantize(&correct_gate_f32, h, None);
let mut wrong_bytes = correct_bytes.clone();
wrong_bytes.pop(); let bad_gate_q = make_q4_0_buf(&wrong_bytes, &device);
let bad_weights = make_weights(
base_moe_weights.router.clone(),
bad_gate_q,
base_moe_weights.expert_up_q.clone(),
base_moe_weights.expert_down_q.clone(),
GgmlType::Q4_0,
GgmlType::Q4_0,
base_moe_weights.shared_gate_inp.clone(),
);
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer_q_moe(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&bad_weights,
valid_shape,
)
.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("expert_gate_q"), "neg_6: wrong message: {msg}");
}
{
let mut bad_shape = valid_shape;
bad_shape.attn.head_dim = 64;
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer_q_moe(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&base_moe_weights,
bad_shape,
)
.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("head_dim"), "neg_7: wrong message: {msg}");
}
{
let mut bad_shape = valid_shape;
bad_shape.attn.cache_prefix_len = 7;
let (hidden_in, mask, pos, mut k_cache, mut v_cache) = make_inputs(&device, &mut seed);
let enc = device.command_encoder().unwrap();
let err = qwen35_tree_verify_full_layer_q_moe(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&base_attn_weights,
&base_moe_weights,
bad_shape,
)
.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("kv_capacity")
|| msg.contains("cache_prefix_len")
|| msg.contains("cache"),
"neg_8: wrong message (expected kv_capacity/cache overflow): {msg}"
);
}
eprintln!("AC-3 (MoE) PASS: all 8 negative paths reject with descriptive errors via full function entry");
}
#[test]
fn qwen35_tree_verify_full_layer_q_moe_cpu_reference_parity_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 4;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let ne: usize = 4;
let topk: usize = 2;
let m_moe: usize = 64;
let m_sh: usize = 64;
let q_total = nq * d;
let kv_total = nkv * d;
let shape = moe_layer_shape_tiny(ne as u32, topk as u32, m_moe as u32, m_sh as u32);
let mut seed = 0xAC04_u32;
let cpu_attn_weights = FullAttnLayerWeights {
attn_norm: vec![1.0f32; h],
post_attn_norm: mk_rand(&mut seed, h, 0.5),
wq: mk_rand(&mut seed, q_total * h, 0.05),
wk: mk_rand(&mut seed, kv_total * h, 0.05),
wv: mk_rand(&mut seed, kv_total * h, 0.05),
w_gate: mk_rand(&mut seed, q_total * h, 0.05),
attn_q_norm: vec![1.0f32; d],
attn_k_norm: vec![1.0f32; d],
wo: mk_rand(&mut seed, h * q_total, 0.05),
};
let gpu_attn_weights =
FullAttnWeightsGpu::from_cpu_f32(&cpu_attn_weights, &device).unwrap();
let (gpu_moe_weights, cpu_moe_weights) =
moe_ffn_weights_q4_0(h, ne, m_moe, m_sh, &mut seed, &device);
assert_all_nonzero("router_f32", &cpu_moe_weights.router);
assert_all_nonzero("expert_gate", &cpu_moe_weights.expert_gate);
assert_all_nonzero("shared_gate", &cpu_moe_weights.shared_gate);
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let hidden_in = upload_f32(&hidden_data, &device).unwrap();
let mask_stride = prefix + seq;
let tree_mask_data: Vec<f32> = {
let mut mv = vec![mlx_native::ops::tree_attention::TREE_MASK_MASKED; seq * mask_stride];
for i in 0..seq {
for j in 0..prefix + i + 1 {
if j < mask_stride {
mv[i * mask_stride + j] =
mlx_native::ops::tree_attention::TREE_MASK_ATTENDED;
}
}
}
mv
};
let tree_mask = upload_f32(&tree_mask_data, &device).unwrap();
let tree_pos = upload_positions(seq, prefix as u32, &device);
let mut k_cache = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache = alloc_kv_cache(nkv, cap, d, &device);
let enc = device.command_encoder().expect("enc");
let gpu_out = qwen35_tree_verify_full_layer_q_moe(
enc,
&device,
&mut registry,
&hidden_in,
&tree_mask,
&tree_pos,
&mut k_cache,
&mut v_cache,
&gpu_attn_weights,
&gpu_moe_weights,
shape,
)
.expect("AC-4 MoE: GPU call failed");
let gpu_data = download_f32(&gpu_out).unwrap();
let mut k_cache_cpu = vec![0.0f32; nkv * cap * d];
let mut v_cache_cpu = vec![0.0f32; nkv * cap * d];
let positions: Vec<[i32; 4]> = (0..seq)
.map(|i| {
let p = (prefix + i) as i32;
[p, p, p, p]
})
.collect();
let cpu_data = cpu_tree_verify_full_layer_q_moe_ref(
&hidden_data,
&tree_mask_data,
&positions,
&mut k_cache_cpu,
&mut v_cache_cpu,
&cpu_attn_weights,
&cpu_moe_weights,
&shape,
);
assert_eq!(gpu_data.len(), cpu_data.len(), "AC-4 MoE: length mismatch");
let has_nan = gpu_data.iter().any(|v| v.is_nan());
let cpu_nonzero = cpu_data.iter().any(|&v| v != 0.0);
if has_nan && cpu_nonzero {
eprintln!("AC-4 (MoE): GPU output NaN under Metal contention — skipping");
return;
}
assert!(!has_nan, "AC-4 MoE: NaN in gpu output");
assert!(
!gpu_data.iter().any(|v| v.is_infinite()),
"AC-4 MoE: Inf in gpu output"
);
let max_diff: f32 = gpu_data
.iter()
.zip(cpu_data.iter())
.map(|(g, c)| (g - c).abs())
.fold(0.0f32, f32::max);
eprintln!("AC-4 (MoE): |GPU-CPU|_inf = {max_diff:.6e}");
assert!(
max_diff < 0.20,
"AC-4 (MoE) FAIL: |GPU-CPU|_inf = {max_diff:.6e} >= 0.20 \
(Q4_0 dequant + MoE routing noise budget). \
Check dequant_q4_0_cpu, router BF16 cast, or post_attn_norm chain."
);
eprintln!("AC-4 (MoE) PASS: CPU reference parity |GPU-CPU|_inf = {max_diff:.6e} < 0.20");
}
#[test]
fn qwen35_tree_verify_full_layer_q_moe_composition_equivalence_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 4;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let ne: usize = 4;
let topk: usize = 2;
let m_moe: usize = 64;
let m_sh: usize = 64;
let shape = moe_layer_shape_tiny(ne as u32, topk as u32, m_moe as u32, m_sh as u32);
let q_total = nq * d;
let kv_total = nkv * d;
let mut seed = 0xAC05_u32;
let cpu_attn_weights = FullAttnLayerWeights {
attn_norm: vec![1.0f32; h],
post_attn_norm: vec![1.0f32; h], wq: mk_rand(&mut seed, q_total * h, 0.05),
wk: mk_rand(&mut seed, kv_total * h, 0.05),
wv: mk_rand(&mut seed, kv_total * h, 0.05),
w_gate: mk_rand(&mut seed, q_total * h, 0.05),
attn_q_norm: vec![1.0f32; d],
attn_k_norm: vec![1.0f32; d],
wo: mk_rand(&mut seed, h * q_total, 0.05),
};
let gpu_attn_weights =
FullAttnWeightsGpu::from_cpu_f32(&cpu_attn_weights, &device).unwrap();
let (gpu_moe_weights, _) = moe_ffn_weights_q4_0(h, ne, m_moe, m_sh, &mut seed, &device);
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let mask_stride = prefix + seq;
let tree_mask_data: Vec<f32> = {
let mut mv = vec![mlx_native::ops::tree_attention::TREE_MASK_MASKED; seq * mask_stride];
for i in 0..seq {
for j in 0..prefix + i + 1 {
if j < mask_stride {
mv[i * mask_stride + j] =
mlx_native::ops::tree_attention::TREE_MASK_ATTENDED;
}
}
}
mv
};
let tree_mask_gpu = upload_f32(&tree_mask_data, &device).unwrap();
let tree_pos = upload_positions(seq, prefix as u32, &device);
let hidden_in_a = upload_f32(&hidden_data, &device).unwrap();
let mut k_a = alloc_kv_cache(nkv, cap, d, &device);
let mut v_a = alloc_kv_cache(nkv, cap, d, &device);
let enc_a = device.command_encoder().expect("enc_a");
let out_a = qwen35_tree_verify_full_layer_q_moe(
enc_a,
&device,
&mut registry,
&hidden_in_a,
&tree_mask_gpu,
&tree_pos,
&mut k_a,
&mut v_a,
&gpu_attn_weights,
&gpu_moe_weights,
shape,
)
.expect("AC-5 MoE: Side A failed");
let data_a = download_f32(&out_a).unwrap();
let hidden_in_b = upload_f32(&hidden_data, &device).unwrap();
let mut k_b = alloc_kv_cache(nkv, cap, d, &device);
let mut v_b = alloc_kv_cache(nkv, cap, d, &device);
let enc_b = device.command_encoder().expect("enc_b");
let attn_out_b = qwen35_tree_verify_attention_block(
enc_b,
&device,
&mut registry,
&hidden_in_b,
&tree_mask_gpu,
&tree_pos,
&mut k_b,
&mut v_b,
&gpu_attn_weights,
shape.attn,
)
.expect("AC-5 MoE: Side B attn failed");
let attn_data_b = download_f32(&attn_out_b).unwrap();
let post_normed_cpu: Vec<f32> = {
let post_attn_norm_w = vec![1.0f32; h];
let eps = shape.attn.rms_norm_eps;
let mut out = vec![0.0f32; seq * h];
for t in 0..seq {
let row = &attn_data_b[t * h..(t + 1) * h];
let ss: f32 = row.iter().map(|v| v * v).sum::<f32>();
let inv = (ss / h as f32 + eps).sqrt().recip();
for (i, (o, w)) in out[t * h..(t + 1) * h]
.iter_mut()
.zip(post_attn_norm_w.iter())
.enumerate()
{
*o = row[i] * inv * w;
}
}
out
};
let post_normed_gpu = upload_f32(&post_normed_cpu, &device).unwrap();
let ffn_residual_b = attn_out_b.clone();
let moe_shape_b = super::super::ffn::MoeFfnShape {
hidden_size: shape.moe.hidden_size,
num_experts: shape.moe.num_experts,
num_experts_per_tok: shape.moe.num_experts_per_tok,
moe_intermediate_size: shape.moe.moe_intermediate_size,
shared_intermediate_size: shape.moe.shared_intermediate_size,
};
let out_b = super::super::gpu_ffn::build_moe_ffn_layer_gpu_q(
&device,
&mut registry,
&post_normed_gpu,
&gpu_moe_weights,
moe_shape_b,
Some(&ffn_residual_b),
)
.expect("AC-5 MoE: Side B moe_ffn failed");
let data_b = download_f32(&out_b).unwrap();
assert_eq!(
data_a.len(),
data_b.len(),
"AC-5 MoE: length mismatch A vs B"
);
let max_diff: f32 = data_a
.iter()
.zip(data_b.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
eprintln!("AC-5 (MoE): |full - split|_inf = {max_diff:.6e}");
assert!(
max_diff < 0.05,
"AC-5 (MoE) FAIL: composition divergence {max_diff:.6e} >= 0.05. \
Both sides use identical GPU MoE kernel — only RMSNorm precision differs. \
Check post_attn_normed threading or residual source."
);
eprintln!(
"AC-5 (MoE) PASS: composition equivalence |full-split|_inf = {max_diff:.6e} < 0.05"
);
}
#[test]
fn qwen35_tree_verify_full_layer_q_moe_byte_identity_3rep_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 4;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let ne: usize = 4;
let topk: usize = 2;
let m_moe: usize = 64;
let m_sh: usize = 64;
let shape = moe_layer_shape_tiny(ne as u32, topk as u32, m_moe as u32, m_sh as u32);
let mut seed = 0xAC06_u32;
let attn_weights = layer_weights_f32(h, nq, nkv, d, &mut seed, &device);
let (moe_gpu_weights, _) = moe_ffn_weights_q4_0(h, ne, m_moe, m_sh, &mut seed, &device);
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let mask =
causal_tree_mask_with_prefix(seq as u32, prefix as u32, (prefix + seq) as u32, &device);
let pos = upload_positions(seq, prefix as u32, &device);
let mut outputs: Vec<Vec<u32>> = Vec::new();
for rep in 0..3 {
let hidden_in = upload_f32(&hidden_data, &device).unwrap();
let mut k_cache = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache = alloc_kv_cache(nkv, cap, d, &device);
let enc = device.command_encoder().expect("enc");
let out = qwen35_tree_verify_full_layer_q_moe(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&attn_weights,
&moe_gpu_weights,
shape,
)
.unwrap_or_else(|e| panic!("AC-6 MoE: rep {rep} failed: {e}"));
let floats = download_f32(&out).unwrap();
if floats.iter().any(|v| v.is_nan()) {
eprintln!(
"AC-6 (MoE): rep {rep} GPU output NaN under Metal contention — skipping test"
);
return;
}
let bits: Vec<u32> = floats.iter().map(|f| f.to_bits()).collect();
outputs.push(bits);
}
for rep in 1..3 {
let first = &outputs[0];
let this = &outputs[rep];
assert_eq!(
first.len(),
this.len(),
"AC-6 MoE: rep {rep} length mismatch"
);
for (i, (a, b)) in first.iter().zip(this.iter()).enumerate() {
assert_eq!(
a, b,
"AC-6 (MoE) FAIL: rep {rep} output[{i}] differs: {:#010x} vs {:#010x}",
a, b
);
}
}
eprintln!("AC-6 (MoE) PASS: 3× byte-identical (0 ULP) determinism");
}
#[test]
fn qwen35_tree_verify_full_layer_q_moe_topk_routing_correctness_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 4;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let ne: usize = 4;
let topk: usize = 2;
let m_moe: usize = 64;
let m_sh: usize = 64;
let shape = moe_layer_shape_tiny(ne as u32, topk as u32, m_moe as u32, m_sh as u32);
let mut seed = 0xAC07_u32;
let attn_weights = layer_weights_f32(h, nq, nkv, d, &mut seed, &device);
let (base_moe, _) = moe_ffn_weights_q4_0(h, ne, m_moe, m_sh, &mut seed, &device);
let mut router_f32 = vec![0.0f32; ne * h];
for j in 0..h {
router_f32[0 * h + j] = 1e3;
}
for j in 0..h {
router_f32[1 * h + j] = 1e3;
}
for j in 0..h {
router_f32[2 * h + j] = -1e6;
}
for j in 0..h {
router_f32[3 * h + j] = -1e6;
}
let router_bf16 = upload_bf16_from_f32(&router_f32, &device).expect("router bf16");
let routed_moe = super::super::gpu_ffn::MoeFfnWeightsGpuQ {
router: router_bf16,
expert_gate_q: base_moe.expert_gate_q.clone(),
expert_up_q: base_moe.expert_up_q.clone(),
expert_down_q: base_moe.expert_down_q.clone(),
ggml_type_gate_up: base_moe.ggml_type_gate_up,
ggml_type_down: base_moe.ggml_type_down,
expert_gate_stride: base_moe.expert_gate_stride,
expert_up_stride: base_moe.expert_up_stride,
expert_down_stride: base_moe.expert_down_stride,
num_experts: base_moe.num_experts,
shared_gate_inp: base_moe.shared_gate_inp.clone(),
shared_gate: base_moe.shared_gate.clone(),
shared_up: base_moe.shared_up.clone(),
shared_down: base_moe.shared_down.clone(),
expert_gate_affine: None,
expert_up_affine: None,
expert_down_affine: None,
};
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let hidden_in = upload_f32(&hidden_data, &device).unwrap();
let mask =
causal_tree_mask_with_prefix(seq as u32, prefix as u32, (prefix + seq) as u32, &device);
let pos = upload_positions(seq, prefix as u32, &device);
let mut k_cache = alloc_kv_cache(nkv, cap, d, &device);
let mut v_cache = alloc_kv_cache(nkv, cap, d, &device);
let enc = device.command_encoder().expect("enc");
let out = qwen35_tree_verify_full_layer_q_moe(
enc,
&device,
&mut registry,
&hidden_in,
&mask,
&pos,
&mut k_cache,
&mut v_cache,
&attn_weights,
&routed_moe,
shape,
)
.expect("AC-7 MoE: routing test failed");
let out_data = download_f32(&out).unwrap();
let max_abs = out_data
.iter()
.cloned()
.map(f32::abs)
.fold(0.0f32, f32::max);
assert!(
max_abs < 1e3,
"AC-7 (MoE) FAIL: max |output| = {max_abs:.3e} >= 1e3. \
Sentinel expert contamination detected — routing is not correctly selecting \
only experts {{0, 1}} (router rows 2,3 = -1e6 should give zero weight)."
);
assert!(
out_data.iter().all(|v| v.is_finite()),
"AC-7 (MoE): non-finite in output"
);
eprintln!(
"AC-7 (MoE) PASS: routing correctness — max|output|={max_abs:.3e} < 1e3 \
(no sentinel-expert leakage with router saturating to experts {{0,1}})"
);
}
#[test]
fn qwen35_tree_verify_full_layer_q_moe_shared_expert_always_contributes_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = KernelRegistry::new();
let h: usize = 128;
let nq: usize = 4;
let nkv: usize = 1;
let d: usize = 128;
let seq: usize = 2;
let prefix: usize = 4;
let cap: usize = 8;
let ne: usize = 4;
let topk: usize = 2;
let m_moe: usize = 64;
let m_sh: usize = 64;
let shape = moe_layer_shape_tiny(ne as u32, topk as u32, m_moe as u32, m_sh as u32);
let mut seed = 0xAC08_u32;
let attn_weights = layer_weights_f32(h, nq, nkv, d, &mut seed, &device);
let (base_moe, _) = moe_ffn_weights_q4_0(h, ne, m_moe, m_sh, &mut seed, &device);
let hidden_data = mk_rand(&mut seed, seq * h, 0.1);
let sh_gate_on_f32 = vec![1e3f32; h];
let sh_gate_off_f32 = vec![-1e3f32; h];
let moe_a = super::super::gpu_ffn::MoeFfnWeightsGpuQ {
router: base_moe.router.clone(),
expert_gate_q: base_moe.expert_gate_q.clone(),
expert_up_q: base_moe.expert_up_q.clone(),
expert_down_q: base_moe.expert_down_q.clone(),
ggml_type_gate_up: base_moe.ggml_type_gate_up,
ggml_type_down: base_moe.ggml_type_down,
expert_gate_stride: base_moe.expert_gate_stride,
expert_up_stride: base_moe.expert_up_stride,
expert_down_stride: base_moe.expert_down_stride,
num_experts: base_moe.num_experts,
shared_gate_inp: upload_bf16_from_f32(&sh_gate_on_f32, &device).expect("sh_gate_on"),
shared_gate: base_moe.shared_gate.clone(),
shared_up: base_moe.shared_up.clone(),
shared_down: base_moe.shared_down.clone(),
expert_gate_affine: None,
expert_up_affine: None,
expert_down_affine: None,
};
let hidden_in_a = upload_f32(&hidden_data, &device).unwrap();
let mask =
causal_tree_mask_with_prefix(seq as u32, prefix as u32, (prefix + seq) as u32, &device);
let pos = upload_positions(seq, prefix as u32, &device);
let mut k_a = alloc_kv_cache(nkv, cap, d, &device);
let mut v_a = alloc_kv_cache(nkv, cap, d, &device);
let enc_a = device.command_encoder().expect("enc_a");
let out_a = qwen35_tree_verify_full_layer_q_moe(
enc_a,
&device,
&mut registry,
&hidden_in_a,
&mask,
&pos,
&mut k_a,
&mut v_a,
&attn_weights,
&moe_a,
shape,
)
.expect("AC-8 MoE: Run A failed");
let data_a = download_f32(&out_a).unwrap();
let moe_b = super::super::gpu_ffn::MoeFfnWeightsGpuQ {
router: base_moe.router.clone(),
expert_gate_q: base_moe.expert_gate_q.clone(),
expert_up_q: base_moe.expert_up_q.clone(),
expert_down_q: base_moe.expert_down_q.clone(),
ggml_type_gate_up: base_moe.ggml_type_gate_up,
ggml_type_down: base_moe.ggml_type_down,
expert_gate_stride: base_moe.expert_gate_stride,
expert_up_stride: base_moe.expert_up_stride,
expert_down_stride: base_moe.expert_down_stride,
num_experts: base_moe.num_experts,
shared_gate_inp: upload_bf16_from_f32(&sh_gate_off_f32, &device).expect("sh_gate_off"),
shared_gate: base_moe.shared_gate.clone(),
shared_up: base_moe.shared_up.clone(),
shared_down: base_moe.shared_down.clone(),
expert_gate_affine: None,
expert_up_affine: None,
expert_down_affine: None,
};
let hidden_in_b = upload_f32(&hidden_data, &device).unwrap();
let mut k_b = alloc_kv_cache(nkv, cap, d, &device);
let mut v_b = alloc_kv_cache(nkv, cap, d, &device);
let enc_b = device.command_encoder().expect("enc_b");
let out_b = qwen35_tree_verify_full_layer_q_moe(
enc_b,
&device,
&mut registry,
&hidden_in_b,
&mask,
&pos,
&mut k_b,
&mut v_b,
&attn_weights,
&moe_b,
shape,
)
.expect("AC-8 MoE: Run B failed");
let data_b = download_f32(&out_b).unwrap();
assert!(
data_a.iter().all(|v| v.is_finite()),
"AC-8 (MoE): Run A non-finite"
);
assert!(
data_b.iter().all(|v| v.is_finite()),
"AC-8 (MoE): Run B non-finite"
);
let delta: Vec<f32> = data_a
.iter()
.zip(data_b.iter())
.map(|(a, b)| a - b)
.collect();
let delta_inf = delta.iter().cloned().map(f32::abs).fold(0.0f32, f32::max);
assert!(
delta_inf > 1e-4,
"AC-8 (MoE) FAIL: delta |A-B|_inf = {delta_inf:.3e} ≈ 0. \
Shared expert does NOT contribute when gate is ON — shared expert may be \
incorrectly gated by topK (should always contribute regardless of routing)."
);
eprintln!(
"AC-8 (MoE) PASS: shared expert contributes — |gate_ON - gate_OFF|_inf = {delta_inf:.3e} > 1e-4"
);
}
}