use super::config::Eagle3DrafterConfig;
use super::kv_cache::DrafterKvCache;
use super::tensors::Eagle3DrafterTensors;
use crate::inference::models::qwen35::gpu_full_attn::{apply_imrope, apply_linear_projection_f32};
use anyhow::{anyhow, Context, Result};
use mlx_native::ops::add_bias_row_2d::{
dispatch_add_bias_row_2d_f32, register as register_add_bias_row_2d,
};
use mlx_native::ops::elementwise::elementwise_add;
use mlx_native::ops::feature_concat::{
dispatch_feature_concat_f32, register as register_feature_concat,
};
use mlx_native::ops::rms_norm::dispatch_rms_norm;
use mlx_native::ops::silu_mul::dispatch_silu_mul;
use mlx_native::ops::transpose::permute_021_f32;
use mlx_native::ops::tree_attention::{self as tree_attn_ops, TreeAttentionParams};
use mlx_native::{CommandEncoder, DType, KernelRegistry, MlxBuffer, MlxDevice};
pub fn register_eagle3_forward_kernels(registry: &mut KernelRegistry) {
register_feature_concat(registry);
register_add_bias_row_2d(registry);
}
pub fn dispatch_eagle3_fc(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
concat_hidden_gpu: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
if concat_hidden_gpu.dtype() != DType::F32 {
return Err(anyhow!(
"dispatch_eagle3_fc: concat_hidden dtype must be F32, got {:?}",
concat_hidden_gpu.dtype()
));
}
let fc_in_usize = cfg.fc_input_size();
let hidden_usize = cfg.hidden_size;
let fc_in: u32 = u32::try_from(fc_in_usize).map_err(|_| {
anyhow!(
"dispatch_eagle3_fc: fc_input_size ({}) exceeds u32::MAX",
fc_in_usize
)
})?;
let hidden: u32 = u32::try_from(hidden_usize).map_err(|_| {
anyhow!(
"dispatch_eagle3_fc: hidden_size ({}) exceeds u32::MAX",
hidden_usize
)
})?;
let expected_input_elems = (seq_len as usize).checked_mul(fc_in_usize).ok_or_else(|| {
anyhow!(
"dispatch_eagle3_fc: seq_len ({}) * fc_input_size ({}) overflows usize",
seq_len,
fc_in_usize
)
})?;
let actual_elems = concat_hidden_gpu.element_count();
if actual_elems != expected_input_elems {
return Err(anyhow!(
"dispatch_eagle3_fc: concat_hidden has {} elements, expected {} (seq_len={} * fc_input_size={})",
actual_elems, expected_input_elems, seq_len, fc_in
));
}
apply_linear_projection_f32(
encoder,
registry,
device,
concat_hidden_gpu,
&tensors.fc,
seq_len,
fc_in,
hidden,
)
.context("dispatch_eagle3_fc")
}
const RMS_NORM_DIM_F32_EXACT_MAX: u32 = 1 << 24;
fn alloc_rms_norm_params_eagle3(device: &MlxDevice, eps: f32, dim: u32) -> Result<MlxBuffer> {
if dim > RMS_NORM_DIM_F32_EXACT_MAX {
return Err(anyhow!(
"alloc_rms_norm_params_eagle3: dim {} exceeds 2^24 — `as f32` would round-to-even",
dim
));
}
let mut params = device
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow!("alloc eagle3 rms_norm params: {e}"))?;
let slice = params
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("eagle3 rms_norm params slice: {e}"))?;
slice[0] = eps;
slice[1] = dim as f32;
Ok(params)
}
fn dispatch_eagle3_rms_norm_seq_x_hidden(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
norm_weight: &MlxBuffer,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
label: &str,
) -> Result<MlxBuffer> {
if input.dtype() != DType::F32 {
return Err(anyhow!(
"dispatch_eagle3_rms_norm ({}): input dtype must be F32, got {:?}",
label,
input.dtype()
));
}
if norm_weight.dtype() != DType::F32 {
return Err(anyhow!(
"dispatch_eagle3_rms_norm ({}): norm weight dtype must be F32 (RMSNorm \
weights cast BF16→F32 at upload per ADR-030 iter-106), got {:?}",
label,
norm_weight.dtype()
));
}
if seq_len == 0 {
return Err(anyhow!(
"dispatch_eagle3_rms_norm ({}): seq_len must be > 0",
label
));
}
let hidden_usize = cfg.hidden_size;
if hidden_usize == 0 {
return Err(anyhow!(
"dispatch_eagle3_rms_norm ({}): hidden_size must be > 0",
label
));
}
let hidden: u32 = u32::try_from(hidden_usize).map_err(|_| {
anyhow!(
"dispatch_eagle3_rms_norm ({}): hidden_size ({}) exceeds u32::MAX",
label,
hidden_usize
)
})?;
if norm_weight.element_count() != hidden_usize {
return Err(anyhow!(
"dispatch_eagle3_rms_norm ({}): weight has {} elements, expected hidden_size {}",
label,
norm_weight.element_count(),
hidden_usize
));
}
let expected_elems = (seq_len as usize)
.checked_mul(hidden_usize)
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_rms_norm ({}): seq_len ({}) * hidden_size ({}) overflows usize",
label,
seq_len,
hidden_usize
)
})?;
let actual = input.element_count();
if actual != expected_elems {
return Err(anyhow!(
"dispatch_eagle3_rms_norm ({}): input has {} elements, expected {} (seq_len={} * hidden_size={})",
label, actual, expected_elems, seq_len, hidden
));
}
let out_bytes = expected_elems
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_rms_norm ({}): expected_elems ({}) * 4 overflows usize",
label,
expected_elems
)
})?;
let out = device
.alloc_buffer(out_bytes, DType::F32, vec![seq_len as usize, hidden_usize])
.map_err(|e| anyhow!("alloc {label} output: {e}"))?;
let params = alloc_rms_norm_params_eagle3(device, cfg.rms_norm_eps, hidden)?;
dispatch_rms_norm(
encoder,
registry,
device.metal_device(),
input,
norm_weight,
&out,
¶ms,
seq_len,
hidden,
)
.with_context(|| format!("dispatch_rms_norm {label}"))?;
Ok(out)
}
pub fn dispatch_eagle3_input_layernorm(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
embeds_gpu: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
dispatch_eagle3_rms_norm_seq_x_hidden(
encoder,
registry,
device,
embeds_gpu,
&tensors.input_layernorm,
cfg,
seq_len,
"input_layernorm",
)
}
pub fn dispatch_eagle3_hidden_norm(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
fc_output_gpu: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
dispatch_eagle3_rms_norm_seq_x_hidden(
encoder,
registry,
device,
fc_output_gpu,
&tensors.hidden_norm,
cfg,
seq_len,
"hidden_norm",
)
}
pub fn dispatch_eagle3_concat_2x_hidden(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
embeds_normed: &MlxBuffer,
hidden_normed: &MlxBuffer,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
if embeds_normed.dtype() != DType::F32 {
return Err(anyhow!(
"dispatch_eagle3_concat_2x_hidden: embeds_normed dtype must be F32, got {:?}",
embeds_normed.dtype()
));
}
if hidden_normed.dtype() != DType::F32 {
return Err(anyhow!(
"dispatch_eagle3_concat_2x_hidden: hidden_normed dtype must be F32, got {:?}",
hidden_normed.dtype()
));
}
if seq_len == 0 {
return Err(anyhow!(
"dispatch_eagle3_concat_2x_hidden: seq_len must be > 0"
));
}
let hidden_usize = cfg.hidden_size;
if hidden_usize == 0 {
return Err(anyhow!(
"dispatch_eagle3_concat_2x_hidden: hidden_size must be > 0"
));
}
let hidden: u32 = u32::try_from(hidden_usize).map_err(|_| {
anyhow!(
"dispatch_eagle3_concat_2x_hidden: hidden_size ({}) exceeds u32::MAX",
hidden_usize
)
})?;
let dst_stride: u32 = hidden.checked_mul(2).ok_or_else(|| {
anyhow!(
"dispatch_eagle3_concat_2x_hidden: 2 * hidden_size ({}) exceeds u32::MAX",
hidden_usize
)
})?;
let per_branch_elems = (seq_len as usize)
.checked_mul(hidden_usize)
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_concat_2x_hidden: seq_len ({}) * hidden_size ({}) overflows usize",
seq_len,
hidden_usize
)
})?;
let total_elems = per_branch_elems.checked_mul(2).ok_or_else(|| {
anyhow!(
"dispatch_eagle3_concat_2x_hidden: dst total elements (2 * {}) overflows usize",
per_branch_elems
)
})?;
if embeds_normed.element_count() != per_branch_elems {
return Err(anyhow!(
"dispatch_eagle3_concat_2x_hidden: embeds_normed has {} elements, expected {} (seq_len={} * hidden_size={})",
embeds_normed.element_count(),
per_branch_elems,
seq_len,
hidden
));
}
if hidden_normed.element_count() != per_branch_elems {
return Err(anyhow!(
"dispatch_eagle3_concat_2x_hidden: hidden_normed has {} elements, expected {} (seq_len={} * hidden_size={})",
hidden_normed.element_count(),
per_branch_elems,
seq_len,
hidden
));
}
let total_bytes = total_elems
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_concat_2x_hidden: total_elems ({}) * 4 overflows usize",
total_elems
)
})?;
let dst = device
.alloc_buffer(
total_bytes,
DType::F32,
vec![seq_len as usize, 2 * hidden_usize],
)
.map_err(|e| anyhow!("alloc concat output: {e}"))?;
dispatch_feature_concat_f32(
encoder,
registry,
device.metal_device(),
embeds_normed,
&dst,
seq_len,
hidden,
0,
dst_stride,
)
.context("dispatch_feature_concat_f32 embeds branch")?;
dispatch_feature_concat_f32(
encoder,
registry,
device.metal_device(),
hidden_normed,
&dst,
seq_len,
hidden,
hidden,
dst_stride,
)
.context("dispatch_feature_concat_f32 hidden branch")?;
Ok(dst)
}
#[allow(clippy::too_many_arguments)]
fn dispatch_eagle3_projection_with_optional_bias(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
bias: Option<&MlxBuffer>,
seq_len: u32,
in_features_usize: usize,
out_features_usize: usize,
label: &str,
) -> Result<MlxBuffer> {
if input.dtype() != DType::F32 {
return Err(anyhow!(
"dispatch_eagle3_projection ({}): input dtype must be F32, got {:?}",
label,
input.dtype()
));
}
if seq_len == 0 {
return Err(anyhow!(
"dispatch_eagle3_projection ({}): seq_len must be > 0",
label
));
}
if in_features_usize == 0 || out_features_usize == 0 {
return Err(anyhow!(
"dispatch_eagle3_projection ({}): in_features ({}) and out_features ({}) must be > 0",
label,
in_features_usize,
out_features_usize
));
}
let in_features: u32 = u32::try_from(in_features_usize).map_err(|_| {
anyhow!(
"dispatch_eagle3_projection ({}): in_features ({}) exceeds u32::MAX",
label,
in_features_usize
)
})?;
let out_features: u32 = u32::try_from(out_features_usize).map_err(|_| {
anyhow!(
"dispatch_eagle3_projection ({}): out_features ({}) exceeds u32::MAX",
label,
out_features_usize
)
})?;
let expected_in_elems = (seq_len as usize)
.checked_mul(in_features_usize)
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_projection ({}): seq_len ({}) * in_features ({}) overflows usize",
label,
seq_len,
in_features_usize
)
})?;
if input.element_count() != expected_in_elems {
return Err(anyhow!(
"dispatch_eagle3_projection ({}): input has {} elements, expected {} (seq_len={} * in_features={})",
label,
input.element_count(),
expected_in_elems,
seq_len,
in_features
));
}
if weight.dtype() != DType::BF16 {
return Err(anyhow!(
"dispatch_eagle3_projection ({}): weight dtype must be BF16, got {:?}",
label,
weight.dtype()
));
}
let expected_w_elems = out_features_usize
.checked_mul(in_features_usize)
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_projection ({}): out * in overflows usize",
label
)
})?;
if weight.element_count() != expected_w_elems {
return Err(anyhow!(
"dispatch_eagle3_projection ({}): weight has {} elements, expected {} (out * in = {} * {})",
label,
weight.element_count(),
expected_w_elems,
out_features,
in_features
));
}
if let Some(b) = bias {
if b.dtype() != DType::F32 {
return Err(anyhow!(
"dispatch_eagle3_projection ({}): bias dtype must be F32 (cast from BF16 at upload), got {:?}",
label,
b.dtype()
));
}
if b.element_count() != out_features_usize {
return Err(anyhow!(
"dispatch_eagle3_projection ({}): bias has {} elements, expected out_features {}",
label,
b.element_count(),
out_features
));
}
}
let out_elems_usize = (seq_len as usize)
.checked_mul(out_features_usize)
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_projection ({}): seq_len ({}) * out_features ({}) overflows usize",
label,
seq_len,
out_features_usize
)
})?;
if out_elems_usize > (u32::MAX as usize) {
return Err(anyhow!(
"dispatch_eagle3_projection ({}): output elements ({}) exceeds u32::MAX (downstream kernels use u32 grid)",
label,
out_elems_usize
));
}
let out = apply_linear_projection_f32(
encoder,
registry,
device,
input,
weight,
seq_len,
in_features,
out_features,
)
.with_context(|| format!("apply_linear_projection_f32 {label}"))?;
if let Some(b) = bias {
encoder.memory_barrier();
dispatch_add_bias_row_2d_f32(
encoder,
registry,
device.metal_device(),
&out,
b,
&out,
seq_len,
out_features,
)
.with_context(|| format!("dispatch_add_bias_row_2d_f32 {label}"))?;
}
Ok(out)
}
pub fn dispatch_eagle3_q_proj(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
concat_input: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
dispatch_eagle3_projection_with_optional_bias(
encoder,
registry,
device,
concat_input,
&tensors.q_proj,
tensors.q_bias.as_ref(),
seq_len,
cfg.qkv_input_width(),
cfg.q_proj_out(),
"q_proj",
)
}
pub fn dispatch_eagle3_k_proj(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
concat_input: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
dispatch_eagle3_projection_with_optional_bias(
encoder,
registry,
device,
concat_input,
&tensors.k_proj,
tensors.k_bias.as_ref(),
seq_len,
cfg.qkv_input_width(),
cfg.kv_proj_out(),
"k_proj",
)
}
pub fn dispatch_eagle3_v_proj(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
concat_input: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
dispatch_eagle3_projection_with_optional_bias(
encoder,
registry,
device,
concat_input,
&tensors.v_proj,
tensors.v_bias.as_ref(),
seq_len,
cfg.qkv_input_width(),
cfg.kv_proj_out(),
"v_proj",
)
}
#[allow(clippy::too_many_arguments)]
fn dispatch_eagle3_head_norm(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
proj: &MlxBuffer,
norm_weight: &MlxBuffer,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
num_heads: u32,
label: &str,
) -> Result<MlxBuffer> {
if proj.dtype() != DType::F32 {
return Err(anyhow!(
"dispatch_eagle3_head_norm ({}): proj dtype must be F32, got {:?}",
label,
proj.dtype()
));
}
if norm_weight.dtype() != DType::F32 {
return Err(anyhow!(
"dispatch_eagle3_head_norm ({}): norm weight dtype must be F32 (BF16→F32 at upload), got {:?}",
label,
norm_weight.dtype()
));
}
if seq_len == 0 || num_heads == 0 {
return Err(anyhow!(
"dispatch_eagle3_head_norm ({}): seq_len ({}) and num_heads ({}) must both be > 0",
label,
seq_len,
num_heads
));
}
let head_dim_usize = cfg.head_dim;
if head_dim_usize == 0 {
return Err(anyhow!(
"dispatch_eagle3_head_norm ({}): head_dim must be > 0",
label
));
}
let head_dim: u32 = u32::try_from(head_dim_usize).map_err(|_| {
anyhow!(
"dispatch_eagle3_head_norm ({}): head_dim ({}) exceeds u32::MAX",
label,
head_dim_usize
)
})?;
if head_dim > RMS_NORM_DIM_F32_EXACT_MAX {
return Err(anyhow!(
"dispatch_eagle3_head_norm ({}): head_dim ({}) exceeds 2^24 — params[1] would lose F32 precision",
label,
head_dim
));
}
if norm_weight.element_count() != head_dim_usize {
return Err(anyhow!(
"dispatch_eagle3_head_norm ({}): weight has {} elements, expected head_dim {}",
label,
norm_weight.element_count(),
head_dim
));
}
let rows_usize = (seq_len as usize)
.checked_mul(num_heads as usize)
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_head_norm ({}): seq_len ({}) * num_heads ({}) overflows usize",
label,
seq_len,
num_heads
)
})?;
if rows_usize > (u32::MAX as usize) {
return Err(anyhow!(
"dispatch_eagle3_head_norm ({}): rows ({}) exceeds u32::MAX",
label,
rows_usize
));
}
let rows: u32 = rows_usize as u32;
let expected_elems = rows_usize.checked_mul(head_dim_usize).ok_or_else(|| {
anyhow!(
"dispatch_eagle3_head_norm ({}): rows ({}) * head_dim ({}) overflows usize",
label,
rows_usize,
head_dim_usize
)
})?;
if proj.element_count() != expected_elems {
return Err(anyhow!(
"dispatch_eagle3_head_norm ({}): proj has {} elements, expected {} (rows={} * head_dim={})",
label, proj.element_count(), expected_elems, rows, head_dim
));
}
let out_bytes = expected_elems
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_head_norm ({}): expected_elems ({}) * 4 overflows usize",
label,
expected_elems
)
})?;
let out = device
.alloc_buffer(
out_bytes,
DType::F32,
vec![seq_len as usize, num_heads as usize, head_dim_usize],
)
.map_err(|e| anyhow!("alloc {label} output: {e}"))?;
let params = alloc_rms_norm_params_eagle3(device, cfg.rms_norm_eps, head_dim)?;
encoder.memory_barrier();
dispatch_rms_norm(
encoder,
registry,
device.metal_device(),
proj,
norm_weight,
&out,
¶ms,
rows,
head_dim,
)
.with_context(|| format!("dispatch_rms_norm {label}"))?;
Ok(out)
}
pub fn dispatch_eagle3_q_head_norm(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
q_proj_out: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
if !cfg.use_qk_norm {
return Err(anyhow!(
"dispatch_eagle3_q_head_norm: cfg.use_qk_norm is false — orchestrator must check the gate before calling"
));
}
let q_norm = tensors.q_norm.as_ref().ok_or_else(|| {
anyhow!(
"dispatch_eagle3_q_head_norm: q_norm absent (use_qk_norm = {})",
cfg.use_qk_norm
)
})?;
let num_q_heads: u32 = u32::try_from(cfg.num_q_heads).map_err(|_| {
anyhow!(
"dispatch_eagle3_q_head_norm: num_q_heads ({}) exceeds u32::MAX",
cfg.num_q_heads
)
})?;
dispatch_eagle3_head_norm(
encoder,
registry,
device,
q_proj_out,
q_norm,
cfg,
seq_len,
num_q_heads,
"q_norm",
)
}
pub fn dispatch_eagle3_k_head_norm(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
k_proj_out: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
if !cfg.use_qk_norm {
return Err(anyhow!(
"dispatch_eagle3_k_head_norm: cfg.use_qk_norm is false — orchestrator must check the gate before calling"
));
}
let k_norm = tensors.k_norm.as_ref().ok_or_else(|| {
anyhow!(
"dispatch_eagle3_k_head_norm: k_norm absent (use_qk_norm = {})",
cfg.use_qk_norm
)
})?;
let num_kv_heads: u32 = u32::try_from(cfg.num_kv_heads).map_err(|_| {
anyhow!(
"dispatch_eagle3_k_head_norm: num_kv_heads ({}) exceeds u32::MAX",
cfg.num_kv_heads
)
})?;
dispatch_eagle3_head_norm(
encoder,
registry,
device,
k_proj_out,
k_norm,
cfg,
seq_len,
num_kv_heads,
"k_norm",
)
}
fn build_eagle3_pos_buf(
device: &MlxDevice,
seq_len: u32,
positions_override: Option<&[u32]>,
base_pos: u32,
) -> Result<MlxBuffer> {
let l = seq_len as usize;
if let Some(p) = positions_override {
if p.len() != l {
return Err(anyhow!(
"build_eagle3_pos_buf: positions_override len {} != seq_len {}",
p.len(),
seq_len
));
}
}
let n_pos = 4 * l;
let mut buf = device
.alloc_buffer(n_pos * 4, DType::I32, vec![n_pos])
.map_err(|e| anyhow!("alloc eagle3 rope pos_buf: {e}"))?;
let slice = buf
.as_mut_slice::<i32>()
.map_err(|e| anyhow!("eagle3 rope pos_buf slice: {e}"))?;
let pos_values: Vec<i32> = if let Some(p) = positions_override {
let mut out = Vec::with_capacity(p.len());
for (i, &v) in p.iter().enumerate() {
let iv = i32::try_from(v).map_err(|_| {
anyhow!(
"build_eagle3_pos_buf: positions_override[{}] = {} exceeds i32::MAX (kernel uses signed i32 positions)",
i,
v
)
})?;
out.push(iv);
}
out
} else {
let mut out = Vec::with_capacity(l);
for i in 0..l {
let v = (base_pos as i64).checked_add(i as i64).ok_or_else(|| {
anyhow!(
"build_eagle3_pos_buf: base_pos ({}) + {} overflows i64",
base_pos,
i
)
})?;
if v > (i32::MAX as i64) {
return Err(anyhow!(
"build_eagle3_pos_buf: linear position {} (base_pos {} + offset {}) exceeds i32::MAX",
v,
base_pos,
i
));
}
out.push(v as i32);
}
out
};
for axis in 0..4 {
let dst = &mut slice[axis * l..(axis + 1) * l];
dst.copy_from_slice(&pos_values);
}
Ok(buf)
}
#[allow(clippy::too_many_arguments)]
pub fn dispatch_eagle3_rope(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
qk_in: &MlxBuffer,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
num_heads: u32,
positions_override: Option<&[u32]>,
base_pos: u32,
label: &str,
) -> Result<MlxBuffer> {
if qk_in.dtype() != DType::F32 {
return Err(anyhow!(
"dispatch_eagle3_rope ({}): qk_in dtype must be F32, got {:?}",
label,
qk_in.dtype()
));
}
if seq_len == 0 || num_heads == 0 {
return Err(anyhow!(
"dispatch_eagle3_rope ({}): seq_len ({}) and num_heads ({}) must be > 0",
label,
seq_len,
num_heads
));
}
let head_dim_usize = cfg.head_dim;
let rope_dim_usize = cfg.rope_dim;
let head_dim: u32 = u32::try_from(head_dim_usize).map_err(|_| {
anyhow!(
"dispatch_eagle3_rope ({}): head_dim ({}) exceeds u32::MAX",
label,
head_dim_usize
)
})?;
let rope_dim: u32 = u32::try_from(rope_dim_usize).map_err(|_| {
anyhow!(
"dispatch_eagle3_rope ({}): rope_dim ({}) exceeds u32::MAX",
label,
rope_dim_usize
)
})?;
let expected_elems = (seq_len as usize)
.checked_mul(num_heads as usize)
.and_then(|v| v.checked_mul(head_dim_usize))
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_rope ({}): seq * heads * head_dim overflows usize",
label
)
})?;
if expected_elems > (u32::MAX as usize) {
return Err(anyhow!(
"dispatch_eagle3_rope ({}): expected_elems ({}) exceeds u32::MAX",
label,
expected_elems
));
}
if qk_in.element_count() != expected_elems {
return Err(anyhow!(
"dispatch_eagle3_rope ({}): qk_in has {} elements, expected {} (seq={} heads={} head_dim={})",
label,
qk_in.element_count(),
expected_elems,
seq_len,
num_heads,
head_dim
));
}
let positions = build_eagle3_pos_buf(device, seq_len, positions_override, base_pos)?;
let sections = [rope_dim / 2, 0, 0, 0];
encoder.memory_barrier();
apply_imrope(
encoder,
registry,
device,
qk_in,
&positions,
seq_len,
num_heads,
head_dim,
rope_dim,
cfg.rope_theta,
sections,
)
.with_context(|| format!("apply_imrope {label}"))
}
pub const EAGLE3_TREE_MASK_ATTENDED: f32 = 0.0;
pub const EAGLE3_TREE_MASK_MASKED: f32 = -65504.0;
#[allow(clippy::too_many_arguments)]
pub fn dispatch_eagle3_tree_attention(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
q: &MlxBuffer,
k: &MlxBuffer,
v: &MlxBuffer,
tree_mask: &MlxBuffer,
cfg: &Eagle3DrafterConfig,
q_seq_len: u32,
kv_seq_len: u32,
kv_capacity: u32,
mask_stride: u32,
scale: f32,
) -> Result<MlxBuffer> {
if q_seq_len == 0 {
return Err(anyhow!(
"dispatch_eagle3_tree_attention: q_seq_len must be > 0"
));
}
if kv_seq_len == 0 {
return Err(anyhow!(
"dispatch_eagle3_tree_attention: kv_seq_len must be > 0"
));
}
if kv_capacity < kv_seq_len {
return Err(anyhow!(
"dispatch_eagle3_tree_attention: kv_capacity ({}) must be >= kv_seq_len ({})",
kv_capacity,
kv_seq_len
));
}
if mask_stride < kv_seq_len {
return Err(anyhow!(
"dispatch_eagle3_tree_attention: mask_stride ({}) must be >= kv_seq_len ({})",
mask_stride,
kv_seq_len
));
}
if !scale.is_finite() {
return Err(anyhow!(
"dispatch_eagle3_tree_attention: scale ({}) must be finite",
scale
));
}
let num_q_heads: u32 = u32::try_from(cfg.num_q_heads).map_err(|_| {
anyhow!(
"dispatch_eagle3_tree_attention: num_q_heads ({}) exceeds u32::MAX",
cfg.num_q_heads
)
})?;
let num_kv_heads: u32 = u32::try_from(cfg.num_kv_heads).map_err(|_| {
anyhow!(
"dispatch_eagle3_tree_attention: num_kv_heads ({}) exceeds u32::MAX",
cfg.num_kv_heads
)
})?;
let head_dim: u32 = u32::try_from(cfg.head_dim).map_err(|_| {
anyhow!(
"dispatch_eagle3_tree_attention: head_dim ({}) exceeds u32::MAX",
cfg.head_dim
)
})?;
let out_elems = (q_seq_len as usize)
.checked_mul(num_q_heads as usize)
.and_then(|v| v.checked_mul(cfg.head_dim))
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_tree_attention: out elements (q={} * n_q={} * hd={}) overflows usize",
q_seq_len,
num_q_heads,
cfg.head_dim
)
})?;
if out_elems > (u32::MAX as usize) {
return Err(anyhow!(
"dispatch_eagle3_tree_attention: out elements ({}) exceeds u32::MAX",
out_elems
));
}
let out_bytes = out_elems
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_tree_attention: out bytes ({} * 4) overflows usize",
out_elems
)
})?;
let output = device
.alloc_buffer(
out_bytes,
DType::F32,
vec![q_seq_len as usize, num_q_heads as usize, cfg.head_dim],
)
.map_err(|e| anyhow!("alloc tree_attention output: {e}"))?;
let tmp_bytes = tree_attn_ops::tmp_buffer_bytes(num_q_heads, head_dim, q_seq_len);
let tmp = device
.alloc_buffer(tmp_bytes, DType::F32, vec![tmp_bytes / 4])
.map_err(|e| anyhow!("alloc tree_attention tmp: {e}"))?;
let params = TreeAttentionParams {
num_heads: num_q_heads,
num_kv_heads,
head_dim,
kv_seq_len,
kv_capacity,
scale,
q_seq_len,
mask_stride,
};
encoder.memory_barrier();
tree_attn_ops::tree_attention(
encoder, registry, device, q, k, v, tree_mask, &output, &tmp, ¶ms,
)
.context("tree_attention")?;
Ok(output)
}
pub fn dispatch_eagle3_o_proj(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
attn_out: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
q_seq_len: u32,
) -> Result<MlxBuffer> {
encoder.memory_barrier();
dispatch_eagle3_projection_with_optional_bias(
encoder,
registry,
device,
attn_out,
&tensors.o_proj,
tensors.o_bias.as_ref(),
q_seq_len,
cfg.q_proj_out(), cfg.hidden_size, "o_proj",
)
}
pub fn dispatch_eagle3_residual_add(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
a: &MlxBuffer,
b: &MlxBuffer,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
if a.dtype() != DType::F32 || b.dtype() != DType::F32 {
return Err(anyhow!(
"dispatch_eagle3_residual_add: inputs must be F32, got a={:?} b={:?}",
a.dtype(),
b.dtype()
));
}
if seq_len == 0 {
return Err(anyhow!("dispatch_eagle3_residual_add: seq_len must be > 0"));
}
let hidden_usize = cfg.hidden_size;
if hidden_usize == 0 {
return Err(anyhow!(
"dispatch_eagle3_residual_add: hidden_size must be > 0"
));
}
let n_elements = (seq_len as usize)
.checked_mul(hidden_usize)
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_residual_add: seq_len ({}) * hidden_size ({}) overflows usize",
seq_len,
hidden_usize
)
})?;
if n_elements > (u32::MAX as usize) {
return Err(anyhow!(
"dispatch_eagle3_residual_add: n_elements ({}) exceeds u32::MAX",
n_elements
));
}
if a.element_count() != n_elements {
return Err(anyhow!(
"dispatch_eagle3_residual_add: a has {} elements, expected {} (seq={} * hidden={})",
a.element_count(),
n_elements,
seq_len,
hidden_usize
));
}
if b.element_count() != n_elements {
return Err(anyhow!(
"dispatch_eagle3_residual_add: b has {} elements, expected {} (seq={} * hidden={})",
b.element_count(),
n_elements,
seq_len,
hidden_usize
));
}
let out_bytes = n_elements
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| anyhow!("dispatch_eagle3_residual_add: n_elements * 4 overflows usize"))?;
let out = device
.alloc_buffer(out_bytes, DType::F32, vec![seq_len as usize, hidden_usize])
.map_err(|e| anyhow!("alloc residual_add output: {e}"))?;
encoder.memory_barrier();
elementwise_add(
encoder,
registry,
device.metal_device(),
a,
b,
&out,
n_elements,
DType::F32,
)
.map_err(|e| anyhow!("elementwise_add residual: {e}"))?;
Ok(out)
}
pub fn dispatch_eagle3_mlp(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
if input.dtype() != DType::F32 {
return Err(anyhow!(
"dispatch_eagle3_mlp: input dtype must be F32, got {:?}",
input.dtype()
));
}
if seq_len == 0 {
return Err(anyhow!("dispatch_eagle3_mlp: seq_len must be > 0"));
}
let hidden_usize = cfg.hidden_size;
let inter_usize = cfg.intermediate_size;
if hidden_usize == 0 || inter_usize == 0 {
return Err(anyhow!(
"dispatch_eagle3_mlp: hidden_size ({}) and intermediate_size ({}) must be > 0",
hidden_usize,
inter_usize
));
}
let expected_input_elems = (seq_len as usize)
.checked_mul(hidden_usize)
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_mlp: seq_len ({}) * hidden_size ({}) overflows usize",
seq_len,
hidden_usize
)
})?;
if input.element_count() != expected_input_elems {
return Err(anyhow!(
"dispatch_eagle3_mlp: input has {} elements, expected {} (seq={} * hidden={})",
input.element_count(),
expected_input_elems,
seq_len,
hidden_usize
));
}
encoder.memory_barrier();
let gate = dispatch_eagle3_projection_with_optional_bias(
encoder,
registry,
device,
input,
&tensors.mlp_gate,
None, seq_len,
hidden_usize,
inter_usize,
"mlp_gate",
)?;
let up = dispatch_eagle3_projection_with_optional_bias(
encoder,
registry,
device,
input,
&tensors.mlp_up,
None,
seq_len,
hidden_usize,
inter_usize,
"mlp_up",
)?;
let n_h_usize = (seq_len as usize).checked_mul(inter_usize).ok_or_else(|| {
anyhow!(
"dispatch_eagle3_mlp: seq_len ({}) * intermediate ({}) overflows usize",
seq_len,
inter_usize
)
})?;
if n_h_usize > (u32::MAX as usize) {
return Err(anyhow!(
"dispatch_eagle3_mlp: silu_mul element count ({}) exceeds u32::MAX",
n_h_usize
));
}
let n_h: u32 = n_h_usize as u32;
let activated_bytes = n_h_usize
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| anyhow!("dispatch_eagle3_mlp: activated bytes overflow"))?;
let activated = device
.alloc_buffer(
activated_bytes,
DType::F32,
vec![seq_len as usize, inter_usize],
)
.map_err(|e| anyhow!("alloc mlp activated: {e}"))?;
let mut silu_params = device
.alloc_buffer(4, DType::U32, vec![1])
.map_err(|e| anyhow!("alloc mlp silu_params: {e}"))?;
silu_params
.as_mut_slice::<u32>()
.map_err(|e| anyhow!("silu_params slice: {e}"))?[0] = n_h;
encoder.memory_barrier(); dispatch_silu_mul(
encoder,
registry,
device.metal_device(),
&gate,
&up,
&activated,
&silu_params,
n_h,
)
.map_err(|e| anyhow!("dispatch_silu_mul: {e}"))?;
encoder.memory_barrier();
dispatch_eagle3_projection_with_optional_bias(
encoder,
registry,
device,
&activated,
&tensors.mlp_down,
None,
seq_len,
inter_usize,
hidden_usize,
"mlp_down",
)
}
pub fn dispatch_eagle3_drafter_forward(
device: &MlxDevice,
registry: &mut KernelRegistry,
target_aux_gpu: &MlxBuffer,
embeds_gpu: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
base_pos: u32,
) -> Result<Vec<f32>> {
let mut enc = device
.command_encoder()
.map_err(|e| anyhow!("dispatch_eagle3_drafter_forward: encoder: {e}"))?;
let fc_out = dispatch_eagle3_fc(
&mut enc,
registry,
device,
target_aux_gpu,
tensors,
cfg,
seq_len,
)?;
let embeds_normed = dispatch_eagle3_input_layernorm(
&mut enc, registry, device, embeds_gpu, tensors, cfg, seq_len,
)?;
let hidden_normed =
dispatch_eagle3_hidden_norm(&mut enc, registry, device, &fc_out, tensors, cfg, seq_len)?;
let attn_residual_src: &MlxBuffer = if cfg.norm_before_residual {
&hidden_normed
} else {
&fc_out
};
let concat = dispatch_eagle3_concat_2x_hidden(
&mut enc,
registry,
device,
&embeds_normed,
&hidden_normed,
cfg,
seq_len,
)?;
let q = dispatch_eagle3_q_proj(&mut enc, registry, device, &concat, tensors, cfg, seq_len)?;
let k = dispatch_eagle3_k_proj(&mut enc, registry, device, &concat, tensors, cfg, seq_len)?;
let v = dispatch_eagle3_v_proj(&mut enc, registry, device, &concat, tensors, cfg, seq_len)?;
let (q_normed, k_normed) = if cfg.use_qk_norm {
let qn =
dispatch_eagle3_q_head_norm(&mut enc, registry, device, &q, tensors, cfg, seq_len)?;
let kn =
dispatch_eagle3_k_head_norm(&mut enc, registry, device, &k, tensors, cfg, seq_len)?;
(qn, kn)
} else {
(q, k)
};
let q_roped = dispatch_eagle3_rope(
&mut enc,
registry,
device,
&q_normed,
cfg,
seq_len,
u32::try_from(cfg.num_q_heads).map_err(|_| anyhow!("num_q_heads overflow"))?,
None,
base_pos,
"q_rope",
)?;
let k_roped = dispatch_eagle3_rope(
&mut enc,
registry,
device,
&k_normed,
cfg,
seq_len,
u32::try_from(cfg.num_kv_heads).map_err(|_| anyhow!("num_kv_heads overflow"))?,
None,
base_pos,
"k_rope",
)?;
let q_perm = dispatch_eagle3_permute_seq_to_head_outer(
&mut enc,
registry,
device,
&q_roped,
seq_len,
u32::try_from(cfg.num_q_heads).map_err(|_| anyhow!("num_q_heads overflow"))?,
cfg.head_dim,
"q_permute",
)?;
let k_perm = dispatch_eagle3_permute_seq_to_head_outer(
&mut enc,
registry,
device,
&k_roped,
seq_len,
u32::try_from(cfg.num_kv_heads).map_err(|_| anyhow!("num_kv_heads overflow"))?,
cfg.head_dim,
"k_permute",
)?;
let v_perm = dispatch_eagle3_permute_seq_to_head_outer(
&mut enc,
registry,
device,
&v,
seq_len,
u32::try_from(cfg.num_kv_heads).map_err(|_| anyhow!("num_kv_heads overflow"))?,
cfg.head_dim,
"v_permute",
)?;
let kv_seq_len = seq_len;
let kv_capacity = seq_len;
let mask_stride = kv_seq_len;
let mask_elems = (seq_len as usize) * (mask_stride as usize);
let mask_data = vec![EAGLE3_TREE_MASK_ATTENDED; mask_elems];
let mut mask_gpu = device
.alloc_buffer(
mask_data.len() * 4,
DType::F32,
vec![seq_len as usize, mask_stride as usize],
)
.map_err(|e| anyhow!("alloc mask: {e}"))?;
mask_gpu
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("mask slice: {e}"))?
.copy_from_slice(&mask_data);
let scale = 1.0f32 / (cfg.head_dim as f32).sqrt();
let attn_out = dispatch_eagle3_tree_attention(
&mut enc,
registry,
device,
&q_perm,
&k_perm,
&v_perm,
&mask_gpu,
cfg,
seq_len,
kv_seq_len,
kv_capacity,
mask_stride,
scale,
)?;
let o_out =
dispatch_eagle3_o_proj(&mut enc, registry, device, &attn_out, tensors, cfg, seq_len)?;
let attn_residual = dispatch_eagle3_residual_add(
&mut enc,
registry,
device,
&o_out,
attn_residual_src,
cfg,
seq_len,
)?;
let post_attn_normed = dispatch_eagle3_post_attention_layernorm(
&mut enc,
registry,
device,
&attn_residual,
tensors,
cfg,
seq_len,
)?;
let mlp_out = dispatch_eagle3_mlp(
&mut enc,
registry,
device,
&post_attn_normed,
tensors,
cfg,
seq_len,
)?;
let final_residual = dispatch_eagle3_residual_add(
&mut enc,
registry,
device,
&mlp_out,
&attn_residual,
cfg,
seq_len,
)?;
let final_normed = dispatch_eagle3_final_norm(
&mut enc,
registry,
device,
&final_residual,
tensors,
cfg,
seq_len,
)?;
let logits = dispatch_eagle3_lm_head(
&mut enc,
registry,
device,
&final_normed,
tensors,
cfg,
seq_len,
)?;
enc.commit_and_wait()
.map_err(|e| anyhow!("dispatch_eagle3_drafter_forward: commit: {e}"))?;
Ok(logits
.as_slice::<f32>()
.map_err(|e| anyhow!("logits slice: {e}"))?
.to_vec())
}
pub fn dispatch_eagle3_drafter_forward_with_kv_cache(
device: &MlxDevice,
registry: &mut KernelRegistry,
target_aux_gpu: &MlxBuffer,
embeds_gpu: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
base_pos: u32,
cache: &mut DrafterKvCache,
mask_override: Option<&[f32]>,
) -> Result<Vec<f32>> {
if cache.num_kv_heads != cfg.num_kv_heads {
return Err(anyhow!(
"dispatch_eagle3_drafter_forward_with_kv_cache: cache.num_kv_heads ({}) != cfg.num_kv_heads ({})",
cache.num_kv_heads,
cfg.num_kv_heads
));
}
if cache.head_dim != cfg.head_dim {
return Err(anyhow!(
"dispatch_eagle3_drafter_forward_with_kv_cache: cache.head_dim ({}) != cfg.head_dim ({})",
cache.head_dim,
cfg.head_dim
));
}
if seq_len != 1 {
return Err(anyhow!(
"dispatch_eagle3_drafter_forward_with_kv_cache: seq_len must be 1 in Step-2 (got {})",
seq_len
));
}
let s_usize = seq_len as usize;
let post_len = cache
.len()
.checked_add(s_usize)
.ok_or_else(|| anyhow!("cache len + seq_len overflows usize"))?;
if post_len > cache.capacity {
return Err(anyhow!(
"dispatch_eagle3_drafter_forward_with_kv_cache: cache would overflow (len={} + seq_len={} > capacity={})",
cache.len(),
s_usize,
cache.capacity
));
}
let mut enc1 = device
.command_encoder()
.map_err(|e| anyhow!("dispatch_eagle3_drafter_forward_with_kv_cache enc1: {e}"))?;
let fc_out = dispatch_eagle3_fc(
&mut enc1,
registry,
device,
target_aux_gpu,
tensors,
cfg,
seq_len,
)?;
let embeds_normed = dispatch_eagle3_input_layernorm(
&mut enc1, registry, device, embeds_gpu, tensors, cfg, seq_len,
)?;
let hidden_normed =
dispatch_eagle3_hidden_norm(&mut enc1, registry, device, &fc_out, tensors, cfg, seq_len)?;
let attn_residual_src_kv: &MlxBuffer = if cfg.norm_before_residual {
&hidden_normed
} else {
&fc_out
};
let concat = dispatch_eagle3_concat_2x_hidden(
&mut enc1,
registry,
device,
&embeds_normed,
&hidden_normed,
cfg,
seq_len,
)?;
let q = dispatch_eagle3_q_proj(&mut enc1, registry, device, &concat, tensors, cfg, seq_len)?;
let k = dispatch_eagle3_k_proj(&mut enc1, registry, device, &concat, tensors, cfg, seq_len)?;
let v = dispatch_eagle3_v_proj(&mut enc1, registry, device, &concat, tensors, cfg, seq_len)?;
let (q_normed, k_normed) = if cfg.use_qk_norm {
let qn =
dispatch_eagle3_q_head_norm(&mut enc1, registry, device, &q, tensors, cfg, seq_len)?;
let kn =
dispatch_eagle3_k_head_norm(&mut enc1, registry, device, &k, tensors, cfg, seq_len)?;
(qn, kn)
} else {
(q, k)
};
let q_roped = dispatch_eagle3_rope(
&mut enc1,
registry,
device,
&q_normed,
cfg,
seq_len,
u32::try_from(cfg.num_q_heads).map_err(|_| anyhow!("num_q_heads overflow"))?,
None,
base_pos,
"q_rope",
)?;
let k_roped = dispatch_eagle3_rope(
&mut enc1,
registry,
device,
&k_normed,
cfg,
seq_len,
u32::try_from(cfg.num_kv_heads).map_err(|_| anyhow!("num_kv_heads overflow"))?,
None,
base_pos,
"k_rope",
)?;
let q_perm = dispatch_eagle3_permute_seq_to_head_outer(
&mut enc1,
registry,
device,
&q_roped,
seq_len,
u32::try_from(cfg.num_q_heads).map_err(|_| anyhow!("num_q_heads overflow"))?,
cfg.head_dim,
"q_permute",
)?;
let k_perm = dispatch_eagle3_permute_seq_to_head_outer(
&mut enc1,
registry,
device,
&k_roped,
seq_len,
u32::try_from(cfg.num_kv_heads).map_err(|_| anyhow!("num_kv_heads overflow"))?,
cfg.head_dim,
"k_permute",
)?;
let v_perm = dispatch_eagle3_permute_seq_to_head_outer(
&mut enc1,
registry,
device,
&v,
seq_len,
u32::try_from(cfg.num_kv_heads).map_err(|_| anyhow!("num_kv_heads overflow"))?,
cfg.head_dim,
"v_permute",
)?;
enc1.commit_and_wait()
.map_err(|e| anyhow!("dispatch_eagle3_drafter_forward_with_kv_cache enc1 commit: {e}"))?;
let num_kv_heads = cfg.num_kv_heads;
let head_dim = cfg.head_dim;
let row_elems = num_kv_heads
.checked_mul(head_dim)
.ok_or_else(|| anyhow!("num_kv_heads * head_dim overflows usize"))?;
let k_perm_data: Vec<f32> = k_perm
.as_slice::<f32>()
.map_err(|e| anyhow!("k_perm slice: {e}"))?
.to_vec();
let v_perm_data: Vec<f32> = v_perm
.as_slice::<f32>()
.map_err(|e| anyhow!("v_perm slice: {e}"))?
.to_vec();
for p in 0..s_usize {
let mut k_row = vec![0.0_f32; row_elems];
let mut v_row = vec![0.0_f32; row_elems];
for h in 0..num_kv_heads {
let src_offset = h * s_usize * head_dim + p * head_dim;
let dst_offset = h * head_dim;
k_row[dst_offset..dst_offset + head_dim]
.copy_from_slice(&k_perm_data[src_offset..src_offset + head_dim]);
v_row[dst_offset..dst_offset + head_dim]
.copy_from_slice(&v_perm_data[src_offset..src_offset + head_dim]);
}
cache
.append(&k_row, &v_row)
.with_context(|| format!("cache append at position {} of {}", p, s_usize))?;
}
let kv_seq_len = u32::try_from(cache.len())
.map_err(|_| anyhow!("cache len {} exceeds u32::MAX", cache.len()))?;
let kv_capacity = u32::try_from(cache.capacity)
.map_err(|_| anyhow!("cache capacity {} exceeds u32::MAX", cache.capacity))?;
let mut enc2 = device
.command_encoder()
.map_err(|e| anyhow!("dispatch_eagle3_drafter_forward_with_kv_cache enc2: {e}"))?;
let mask_stride = kv_seq_len;
let mask_elems = s_usize
.checked_mul(mask_stride as usize)
.ok_or_else(|| anyhow!("mask seq_len * mask_stride overflows usize"))?;
let mask_data = if let Some(m) = mask_override {
if m.len() != mask_elems {
return Err(anyhow!(
"dispatch_eagle3_drafter_forward_with_kv_cache: mask_override has {} elements, \
expected {} (seq_len {} * (cache.len()+seq_len) {})",
m.len(),
mask_elems,
s_usize,
mask_stride
));
}
m.to_vec()
} else {
vec![EAGLE3_TREE_MASK_ATTENDED; mask_elems]
};
let mask_bytes = mask_data
.len()
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| anyhow!("mask bytes overflows usize"))?;
let mut mask_gpu = device
.alloc_buffer(mask_bytes, DType::F32, vec![s_usize, mask_stride as usize])
.map_err(|e| anyhow!("alloc mask: {e}"))?;
mask_gpu
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("mask slice: {e}"))?
.copy_from_slice(&mask_data);
let scale = 1.0f32 / (cfg.head_dim as f32).sqrt();
let attn_out = dispatch_eagle3_tree_attention(
&mut enc2,
registry,
device,
&q_perm,
&cache.k_buf,
&cache.v_buf,
&mask_gpu,
cfg,
seq_len,
kv_seq_len,
kv_capacity,
mask_stride,
scale,
)?;
let o_out = dispatch_eagle3_o_proj(
&mut enc2, registry, device, &attn_out, tensors, cfg, seq_len,
)?;
let attn_residual = dispatch_eagle3_residual_add(
&mut enc2,
registry,
device,
&o_out,
attn_residual_src_kv,
cfg,
seq_len,
)?;
let post_attn_normed = dispatch_eagle3_post_attention_layernorm(
&mut enc2,
registry,
device,
&attn_residual,
tensors,
cfg,
seq_len,
)?;
let mlp_out = dispatch_eagle3_mlp(
&mut enc2,
registry,
device,
&post_attn_normed,
tensors,
cfg,
seq_len,
)?;
let final_residual = dispatch_eagle3_residual_add(
&mut enc2,
registry,
device,
&mlp_out,
&attn_residual,
cfg,
seq_len,
)?;
let final_normed = dispatch_eagle3_final_norm(
&mut enc2,
registry,
device,
&final_residual,
tensors,
cfg,
seq_len,
)?;
let logits = dispatch_eagle3_lm_head(
&mut enc2,
registry,
device,
&final_normed,
tensors,
cfg,
seq_len,
)?;
enc2.commit_and_wait()
.map_err(|e| anyhow!("dispatch_eagle3_drafter_forward_with_kv_cache enc2 commit: {e}"))?;
Ok(logits
.as_slice::<f32>()
.map_err(|e| anyhow!("logits slice: {e}"))?
.to_vec())
}
pub fn dispatch_eagle3_post_attention_layernorm(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
attn_residual: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
encoder.memory_barrier();
dispatch_eagle3_rms_norm_seq_x_hidden(
encoder,
registry,
device,
attn_residual,
&tensors.post_attention_layernorm,
cfg,
seq_len,
"post_attention_layernorm",
)
}
pub fn dispatch_eagle3_permute_seq_to_head_outer(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
seq_outer: &MlxBuffer,
seq_len: u32,
num_heads: u32,
head_dim_usize: usize,
label: &str,
) -> Result<MlxBuffer> {
if seq_outer.dtype() != DType::F32 {
return Err(anyhow!(
"dispatch_eagle3_permute_seq_to_head_outer ({}): input dtype must be F32, got {:?}",
label,
seq_outer.dtype()
));
}
if seq_len == 0 || num_heads == 0 || head_dim_usize == 0 {
return Err(anyhow!(
"dispatch_eagle3_permute_seq_to_head_outer ({}): all dims must be > 0 (seq={}, heads={}, hd={})",
label, seq_len, num_heads, head_dim_usize
));
}
let total_elems = (seq_len as usize)
.checked_mul(num_heads as usize)
.and_then(|v| v.checked_mul(head_dim_usize))
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_permute_seq_to_head_outer ({}): seq * heads * head_dim overflows usize",
label
)
})?;
if total_elems > (u32::MAX as usize) {
return Err(anyhow!(
"dispatch_eagle3_permute_seq_to_head_outer ({}): total elements ({}) exceeds u32::MAX",
label,
total_elems
));
}
if seq_outer.element_count() != total_elems {
return Err(anyhow!(
"dispatch_eagle3_permute_seq_to_head_outer ({}): input has {} elements, expected {} (seq={} * heads={} * hd={})",
label, seq_outer.element_count(), total_elems, seq_len, num_heads, head_dim_usize
));
}
let total_bytes = total_elems
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| {
anyhow!(
"dispatch_eagle3_permute_seq_to_head_outer ({}): byte size overflows usize",
label
)
})?;
let out = device
.alloc_buffer(
total_bytes,
DType::F32,
vec![num_heads as usize, seq_len as usize, head_dim_usize],
)
.map_err(|e| anyhow!("alloc {label} permute output: {e}"))?;
encoder.memory_barrier();
permute_021_f32(
encoder,
registry,
device.metal_device(),
seq_outer,
&out,
seq_len as usize,
num_heads as usize,
head_dim_usize,
)
.map_err(|e| anyhow!("permute_021_f32 {label}: {e}"))?;
Ok(out)
}
pub fn dispatch_eagle3_final_norm(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
final_residual: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
encoder.memory_barrier();
dispatch_eagle3_rms_norm_seq_x_hidden(
encoder,
registry,
device,
final_residual,
&tensors.norm,
cfg,
seq_len,
"final_norm",
)
}
pub fn dispatch_eagle3_lm_head(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
normed_hidden: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
encoder.memory_barrier();
let (weight, out_features) = if cfg.tie_lm_head {
let emb = tensors.embed_tokens.as_ref().ok_or_else(|| {
anyhow!(
"dispatch_eagle3_lm_head: tie_lm_head=true requires embed_tokens, but has_own_embed_tokens=false (drafter shares target embeddings — caller must supply the target's embedding table)"
)
})?;
if cfg.draft_vocab_size != cfg.vocab_size {
return Err(anyhow!(
"dispatch_eagle3_lm_head: tie_lm_head=true requires draft_vocab_size ({}) == vocab_size ({}); use a separate lm_head.weight for fast-vocab-projection",
cfg.draft_vocab_size,
cfg.vocab_size,
));
}
(emb, cfg.vocab_size)
} else {
let lh = tensors.lm_head.as_ref().ok_or_else(|| {
anyhow!("dispatch_eagle3_lm_head: tie_lm_head=false requires lm_head tensor")
})?;
(lh, cfg.draft_vocab_size)
};
dispatch_eagle3_projection_with_optional_bias(
encoder,
registry,
device,
normed_hidden,
weight,
None, seq_len,
cfg.hidden_size,
out_features,
"lm_head",
)
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::unwrap_used, clippy::panic)]
mod tests {
use super::*;
use crate::inference::spec_decode::eagle3::weights::{
expected_manifest, Eagle3Weights, ExpectedTensor,
};
use mlx_native::DType;
use safetensors::tensor::{Dtype as SafeDtype, TensorView};
use std::collections::BTreeMap;
fn eagle3_test_kernel_registry() -> KernelRegistry {
let mut registry = KernelRegistry::new();
register_eagle3_forward_kernels(&mut registry);
registry
}
fn tiny_cfg() -> Eagle3DrafterConfig {
Eagle3DrafterConfig {
hidden_size: 256,
intermediate_size: 512,
head_dim: 32,
num_q_heads: 8,
num_kv_heads: 4,
vocab_size: 1000,
draft_vocab_size: 1000,
target_hidden_size: 256,
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: 32, norm_before_residual: false,
}
}
fn f32_to_bf16_bytes(values: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(values.len() * 2);
for v in values {
let bf16_bits = (v.to_bits() >> 16) as u16;
out.push((bf16_bits & 0xff) as u8);
out.push(((bf16_bits >> 8) & 0xff) as u8);
}
out
}
fn bf16_quantize_f32(v: f32) -> f32 {
let bits = v.to_bits() & 0xFFFF0000;
f32::from_bits(bits)
}
fn cpu_fc_reference(
input: &[f32], weight_bf16_q: &[f32], seq_len: usize,
in_features: usize,
out_features: usize,
) -> Vec<f32> {
let mut out = vec![0.0f32; seq_len * out_features];
for s in 0..seq_len {
for h in 0..out_features {
let mut acc = 0.0f64;
for k in 0..in_features {
acc += (input[s * in_features + k] as f64)
* (weight_bf16_q[h * in_features + k] as f64);
}
out[s * out_features + h] = acc as f32;
}
}
out
}
fn pseudo_random(seed: u64) -> f32 {
let x = seed
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
let bits = ((x >> 33) as u32) & 0x7FFFFF;
(bits as f32 / 0x7FFFFF as f32) * 2.0 - 1.0
}
fn fill_random(buf: &mut [f32], seed: u64) {
for (i, v) in buf.iter_mut().enumerate() {
*v = pseudo_random(seed.wrapping_add(i as u64));
}
}
fn build_blob_with_fc_weight(manifest: &[ExpectedTensor], fc_bytes: &[u8]) -> Vec<u8> {
let mut storage: Vec<Vec<u8>> = Vec::with_capacity(manifest.len());
for exp in manifest {
let elem_bytes = match exp.dtype {
SafeDtype::BF16 => 2,
SafeDtype::I64 => 8,
_ => panic!("unexpected dtype in test"),
};
let nelem: usize = exp.shape.iter().product();
if exp.name == "fc.weight" {
assert_eq!(fc_bytes.len(), nelem * elem_bytes);
storage.push(fc_bytes.to_vec());
} else {
storage.push(vec![0u8; nelem * elem_bytes]);
}
}
let mut tensors: BTreeMap<String, TensorView> = BTreeMap::new();
for (i, exp) in manifest.iter().enumerate() {
let view = TensorView::new(exp.dtype, exp.shape.clone(), storage[i].as_slice())
.expect("synthetic view");
tensors.insert(exp.name.clone(), view);
}
safetensors::serialize(&tensors, None::<std::collections::HashMap<String, String>>)
.expect("serialize")
}
fn upload_f32_to_gpu(device: &MlxDevice, data: &[f32], shape: Vec<usize>) -> MlxBuffer {
let bytes = data.len() * 4;
let mut buf = device
.alloc_buffer(bytes, DType::F32, shape)
.expect("alloc input");
buf.as_mut_slice::<f32>()
.expect("input slice")
.copy_from_slice(data);
buf
}
#[test]
fn adr_037_e4b2_fc_cpu_parity_seq_4_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let seq_len: u32 = 4;
let fc_in = cfg.fc_input_size();
let hidden = cfg.hidden_size;
let mut input_data = vec![0.0f32; (seq_len as usize) * fc_in];
fill_random(&mut input_data, 0xA00);
let mut weight_f32 = vec![0.0f32; hidden * fc_in];
fill_random(&mut weight_f32, 0xB00);
let weight_bf16_bytes = f32_to_bf16_bytes(&weight_f32);
let weight_bf16_q: Vec<f32> = weight_f32.iter().map(|&v| bf16_quantize_f32(v)).collect();
let cpu_out =
cpu_fc_reference(&input_data, &weight_bf16_q, seq_len as usize, fc_in, hidden);
let blob = build_blob_with_fc_weight(&manifest, &weight_bf16_bytes);
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let tensors =
Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload tensors");
let input_gpu = upload_f32_to_gpu(&device, &input_data, vec![seq_len as usize, fc_in]);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_fc(
&mut enc,
&mut registry,
&device,
&input_gpu,
&tensors,
&cfg,
seq_len,
)
.expect("dispatch_eagle3_fc");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
assert_eq!(gpu_out.len(), (seq_len as usize) * hidden, "output shape");
let mut max_diff = 0.0f32;
for (i, (&g, &c)) in gpu_out.iter().zip(cpu_out.iter()).enumerate() {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(
d < 5e-2,
"FC parity violated at idx {i}: gpu={g} cpu={c} diff={d}"
);
}
eprintln!("fc parity seq=4 max_diff={max_diff:.6e}");
}
#[test]
fn adr_037_e4b2_fc_cpu_parity_seq_1_decode_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let _seq_len: u32 = 1; let fc_in = cfg.fc_input_size();
let hidden = cfg.hidden_size;
let mut input_data = vec![0.0f32; fc_in];
fill_random(&mut input_data, 0xC00);
let mut weight_f32 = vec![0.0f32; hidden * fc_in];
fill_random(&mut weight_f32, 0xD00);
let weight_bf16_bytes = f32_to_bf16_bytes(&weight_f32);
let weight_bf16_q: Vec<f32> = weight_f32.iter().map(|&v| bf16_quantize_f32(v)).collect();
let cpu_out = cpu_fc_reference(&input_data, &weight_bf16_q, 1, fc_in, hidden);
let blob = build_blob_with_fc_weight(&manifest, &weight_bf16_bytes);
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let input_gpu = upload_f32_to_gpu(&device, &input_data, vec![1, fc_in]);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_fc(
&mut enc,
&mut registry,
&device,
&input_gpu,
&tensors,
&cfg,
1,
)
.expect("dispatch_eagle3_fc (seq=1)");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
assert_eq!(gpu_out.len(), hidden);
let mut max_diff = 0.0f32;
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(d < 5e-2, "seq=1 GEMV parity: diff={d} > 5e-2");
}
eprintln!("fc parity seq=1 (GEMV) max_diff={max_diff:.6e}");
}
#[test]
fn adr_037_e4b2_fc_rejects_wrong_input_shape_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_fc_weight(
&manifest,
&vec![0u8; cfg.hidden_size * cfg.fc_input_size() * 2],
);
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let bad_data = vec![0.0f32; 100];
let bad_input = upload_f32_to_gpu(&device, &bad_data, vec![100]);
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_fc(
&mut enc,
&mut registry,
&device,
&bad_input,
&tensors,
&cfg,
4,
)
.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("concat_hidden has"),
"expected shape error, got: {msg}"
);
}
#[test]
fn adr_037_e4b2_gate_fc_rejects_non_f32_input_dtype_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_fc_weight(
&manifest,
&vec![0u8; cfg.hidden_size * cfg.fc_input_size() * 2],
);
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let seq_len = 2_u32;
let elem_count = (seq_len as usize) * cfg.fc_input_size();
let bad_input = device
.alloc_buffer(
elem_count * 2, DType::BF16,
vec![seq_len as usize, cfg.fc_input_size()],
)
.expect("alloc bad input");
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_fc(
&mut enc,
&mut registry,
&device,
&bad_input,
&tensors,
&cfg,
seq_len,
)
.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("dtype must be F32"),
"expected F32-dtype error, got: {msg}"
);
}
fn cpu_rms_norm_f32(
input: &[f32], weight: &[f32], seq_len: usize,
dim: usize,
eps: f32,
) -> Vec<f32> {
let mut out = vec![0.0f32; seq_len * dim];
for s in 0..seq_len {
let row = &input[s * dim..(s + 1) * dim];
let mean_sq: f64 =
row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / (dim as f64);
let inv_rms = 1.0 / ((mean_sq as f32 + eps).sqrt());
for d in 0..dim {
out[s * dim + d] = row[d] * inv_rms * weight[d];
}
}
out
}
fn build_blob_with_overrides(
manifest: &[ExpectedTensor],
overrides: &std::collections::HashMap<String, Vec<u8>>,
) -> Vec<u8> {
let mut storage: Vec<Vec<u8>> = Vec::with_capacity(manifest.len());
for exp in manifest {
let elem_bytes = match exp.dtype {
SafeDtype::BF16 => 2,
SafeDtype::I64 => 8,
_ => panic!("unexpected dtype in test"),
};
let nelem: usize = exp.shape.iter().product();
if let Some(bytes) = overrides.get(&exp.name) {
assert_eq!(
bytes.len(),
nelem * elem_bytes,
"override bytes for {}",
exp.name
);
storage.push(bytes.clone());
} else {
storage.push(vec![0u8; nelem * elem_bytes]);
}
}
let mut tensors: BTreeMap<String, TensorView> = BTreeMap::new();
for (i, exp) in manifest.iter().enumerate() {
let view = TensorView::new(exp.dtype, exp.shape.clone(), storage[i].as_slice())
.expect("synthetic view");
tensors.insert(exp.name.clone(), view);
}
safetensors::serialize(&tensors, None::<std::collections::HashMap<String, String>>)
.expect("serialize")
}
#[test]
fn adr_037_e4b3_input_layernorm_cpu_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let seq_len: u32 = 4;
let hidden = cfg.hidden_size;
let mut input_data = vec![0.0f32; (seq_len as usize) * hidden];
fill_random(&mut input_data, 0xE10);
let mut weight_f32 = vec![0.0f32; hidden];
fill_random(&mut weight_f32, 0xE11);
let weight_bf16_bytes = f32_to_bf16_bytes(&weight_f32);
let weight_bf16_q: Vec<f32> = weight_f32.iter().map(|&v| bf16_quantize_f32(v)).collect();
let cpu_out = cpu_rms_norm_f32(
&input_data,
&weight_bf16_q,
seq_len as usize,
hidden,
cfg.rms_norm_eps,
);
let mut overrides = std::collections::HashMap::new();
overrides.insert(
"layers.0.input_layernorm.weight".to_string(),
weight_bf16_bytes,
);
let blob = build_blob_with_overrides(&manifest, &overrides);
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let input_gpu = upload_f32_to_gpu(&device, &input_data, vec![seq_len as usize, hidden]);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_input_layernorm(
&mut enc,
&mut registry,
&device,
&input_gpu,
&tensors,
&cfg,
seq_len,
)
.expect("dispatch_eagle3_input_layernorm");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
let mut max_diff = 0.0f32;
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(d < 1e-3, "input_layernorm parity: diff={d} > 1e-3");
}
eprintln!("input_layernorm parity max_diff={max_diff:.6e}");
}
#[test]
fn adr_037_e4b3_hidden_norm_cpu_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let seq_len: u32 = 4;
let hidden = cfg.hidden_size;
let mut input_data = vec![0.0f32; (seq_len as usize) * hidden];
fill_random(&mut input_data, 0xE20);
let mut weight_f32 = vec![0.0f32; hidden];
fill_random(&mut weight_f32, 0xE21);
let weight_bf16_bytes = f32_to_bf16_bytes(&weight_f32);
let weight_bf16_q: Vec<f32> = weight_f32.iter().map(|&v| bf16_quantize_f32(v)).collect();
let cpu_out = cpu_rms_norm_f32(
&input_data,
&weight_bf16_q,
seq_len as usize,
hidden,
cfg.rms_norm_eps,
);
let mut overrides = std::collections::HashMap::new();
overrides.insert("layers.0.hidden_norm.weight".to_string(), weight_bf16_bytes);
let blob = build_blob_with_overrides(&manifest, &overrides);
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let input_gpu = upload_f32_to_gpu(&device, &input_data, vec![seq_len as usize, hidden]);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_hidden_norm(
&mut enc,
&mut registry,
&device,
&input_gpu,
&tensors,
&cfg,
seq_len,
)
.expect("dispatch_eagle3_hidden_norm");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
let mut max_diff = 0.0f32;
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(d < 1e-3, "hidden_norm parity: diff={d} > 1e-3");
}
eprintln!("hidden_norm parity max_diff={max_diff:.6e}");
}
#[test]
fn adr_037_e4b3_concat_2x_hidden_layout_matches_vllm_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = eagle3_test_kernel_registry();
let cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let _tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let seq_len: u32 = 3;
let hidden = cfg.hidden_size;
let mut embeds_data = vec![0.0f32; (seq_len as usize) * hidden];
let mut hidden_data = vec![0.0f32; (seq_len as usize) * hidden];
for s in 0..(seq_len as usize) {
for d in 0..hidden {
embeds_data[s * hidden + d] = (s * 1000 + d + 1) as f32;
hidden_data[s * hidden + d] = -((s * 1000 + d + 1) as f32);
}
}
let embeds_gpu = upload_f32_to_gpu(&device, &embeds_data, vec![seq_len as usize, hidden]);
let hidden_gpu = upload_f32_to_gpu(&device, &hidden_data, vec![seq_len as usize, hidden]);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_concat_2x_hidden(
&mut enc,
&mut registry,
&device,
&embeds_gpu,
&hidden_gpu,
&cfg,
seq_len,
)
.expect("concat dispatch");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
assert_eq!(gpu_out.len(), (seq_len as usize) * 2 * hidden);
for s in 0..(seq_len as usize) {
for d in 0..hidden {
let left = gpu_out[s * 2 * hidden + d];
let right = gpu_out[s * 2 * hidden + hidden + d];
let expected_pos = (s * 1000 + d + 1) as f32;
assert_eq!(left, expected_pos, "embeds (s={s} d={d})");
assert_eq!(right, -expected_pos, "hidden_states (s={s} d={d})");
}
}
}
#[test]
fn adr_037_e4b3_input_layernorm_rejects_non_f32_input_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let seq_len = 2_u32;
let bf16_input = device
.alloc_buffer(
(seq_len as usize) * cfg.hidden_size * 2,
DType::BF16,
vec![seq_len as usize, cfg.hidden_size],
)
.expect("alloc bad input");
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_input_layernorm(
&mut enc,
&mut registry,
&device,
&bf16_input,
&tensors,
&cfg,
seq_len,
)
.unwrap_err();
assert!(
err.to_string().contains("dtype must be F32"),
"expected F32-dtype error, got: {err}"
);
}
#[test]
fn adr_037_e4b3_gate_input_layernorm_rejects_zero_seq_len_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let empty_input = device
.alloc_buffer(4, DType::F32, vec![1])
.expect("alloc empty");
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_input_layernorm(
&mut enc,
&mut registry,
&device,
&empty_input,
&tensors,
&cfg,
0,
)
.unwrap_err();
assert!(
err.to_string().contains("seq_len must be > 0"),
"expected seq_len-zero error, got: {err}"
);
}
#[test]
fn adr_037_e4b3_gate_concat_rejects_zero_seq_len_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let _tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let empty_a = device
.alloc_buffer(4, DType::F32, vec![1])
.expect("alloc a");
let empty_b = device
.alloc_buffer(4, DType::F32, vec![1])
.expect("alloc b");
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_concat_2x_hidden(
&mut enc,
&mut registry,
&device,
&empty_a,
&empty_b,
&cfg,
0,
)
.unwrap_err();
assert!(
err.to_string().contains("seq_len must be > 0"),
"expected seq_len-zero error, got: {err}"
);
}
fn run_q_proj_with_overrides(
device: &MlxDevice,
registry: &mut KernelRegistry,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
q_weight_f32: &[f32], q_bias_f32: Option<&[f32]>, input_data: &[f32], ) -> (Vec<f32>, Vec<f32>) {
let manifest = expected_manifest(cfg);
let mut overrides = std::collections::HashMap::new();
overrides.insert(
"layers.0.self_attn.q_proj.weight".to_string(),
f32_to_bf16_bytes(q_weight_f32),
);
if let Some(b) = q_bias_f32 {
overrides.insert(
"layers.0.self_attn.q_proj.bias".to_string(),
f32_to_bf16_bytes(b),
);
}
let blob = build_blob_with_overrides(&manifest, &overrides);
let weights = Eagle3Weights::load(&blob, cfg).expect("weights load");
let tensors = Eagle3DrafterTensors::upload(device, cfg, &weights).expect("upload");
let input_gpu = upload_f32_to_gpu(
device,
input_data,
vec![seq_len as usize, cfg.qkv_input_width()],
);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_q_proj(
&mut enc, registry, device, &input_gpu, &tensors, cfg, seq_len,
)
.expect("dispatch_eagle3_q_proj");
enc.commit_and_wait().expect("commit");
let gpu_out: Vec<f32> = out_buf.as_slice::<f32>().expect("output slice").to_vec();
let weight_bf16_q: Vec<f32> = q_weight_f32.iter().map(|&v| bf16_quantize_f32(v)).collect();
let mut cpu_out = cpu_fc_reference(
input_data,
&weight_bf16_q,
seq_len as usize,
cfg.qkv_input_width(),
cfg.q_proj_out(),
);
if let Some(b) = q_bias_f32 {
let b_q: Vec<f32> = b.iter().map(|&v| bf16_quantize_f32(v)).collect();
for s in 0..(seq_len as usize) {
for d in 0..cfg.q_proj_out() {
cpu_out[s * cfg.q_proj_out() + d] += b_q[d];
}
}
}
(gpu_out, cpu_out)
}
#[test]
fn adr_037_e4b4_q_proj_cpu_parity_no_bias_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 cfg = tiny_cfg(); let seq_len: u32 = 4;
let qkv_in = cfg.qkv_input_width();
let q_out = cfg.q_proj_out();
let mut input_data = vec![0.0f32; (seq_len as usize) * qkv_in];
fill_random(&mut input_data, 0xF40);
let mut q_weight = vec![0.0f32; q_out * qkv_in];
fill_random(&mut q_weight, 0xF41);
let (gpu_out, cpu_out) = run_q_proj_with_overrides(
&device,
&mut registry,
&cfg,
seq_len,
&q_weight,
None,
&input_data,
);
let mut max_diff = 0.0f32;
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(d < 5e-2, "q_proj parity: diff={d} > 5e-2");
}
eprintln!("q_proj parity max_diff={max_diff:.6e}");
}
fn cfg_with_bias() -> Eagle3DrafterConfig {
let mut c = tiny_cfg();
c.attention_bias = true;
c
}
#[test]
fn adr_037_e4b4_q_proj_cpu_parity_with_bias_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = eagle3_test_kernel_registry();
let cfg = cfg_with_bias();
let seq_len: u32 = 4;
let qkv_in = cfg.qkv_input_width();
let q_out = cfg.q_proj_out();
let mut input_data = vec![0.0f32; (seq_len as usize) * qkv_in];
fill_random(&mut input_data, 0xF50);
let mut q_weight = vec![0.0f32; q_out * qkv_in];
fill_random(&mut q_weight, 0xF51);
let mut q_bias = vec![0.0f32; q_out];
fill_random(&mut q_bias, 0xF52);
let (gpu_out, cpu_out) = run_q_proj_with_overrides(
&device,
&mut registry,
&cfg,
seq_len,
&q_weight,
Some(&q_bias),
&input_data,
);
let mut max_diff = 0.0f32;
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(d < 5e-2, "q_proj+bias parity: diff={d} > 5e-2");
}
eprintln!("q_proj+bias parity max_diff={max_diff:.6e}");
}
#[test]
fn adr_037_e4b4_k_proj_output_shape_matches_kv_proj_out_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let seq_len = 3_u32;
let input_data = vec![0.0f32; (seq_len as usize) * cfg.qkv_input_width()];
let input_gpu = upload_f32_to_gpu(
&device,
&input_data,
vec![seq_len as usize, cfg.qkv_input_width()],
);
let mut enc = device.command_encoder().expect("encoder");
let k_out = dispatch_eagle3_k_proj(
&mut enc,
&mut registry,
&device,
&input_gpu,
&tensors,
&cfg,
seq_len,
)
.expect("k_proj");
enc.commit_and_wait().expect("commit");
assert_eq!(k_out.dtype(), DType::F32);
assert_eq!(
k_out.element_count(),
(seq_len as usize) * cfg.kv_proj_out()
);
}
#[test]
fn adr_037_e4b4_v_proj_output_shape_matches_kv_proj_out_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let seq_len = 3_u32;
let input_data = vec![0.0f32; (seq_len as usize) * cfg.qkv_input_width()];
let input_gpu = upload_f32_to_gpu(
&device,
&input_data,
vec![seq_len as usize, cfg.qkv_input_width()],
);
let mut enc = device.command_encoder().expect("encoder");
let v_out = dispatch_eagle3_v_proj(
&mut enc,
&mut registry,
&device,
&input_gpu,
&tensors,
&cfg,
seq_len,
)
.expect("v_proj");
enc.commit_and_wait().expect("commit");
assert_eq!(v_out.dtype(), DType::F32);
assert_eq!(
v_out.element_count(),
(seq_len as usize) * cfg.kv_proj_out()
);
}
#[test]
fn adr_037_e4b4_gate_q_proj_rejects_non_f32_input_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let seq_len = 2_u32;
let bad_input = device
.alloc_buffer(
(seq_len as usize) * cfg.qkv_input_width() * 2,
DType::BF16,
vec![seq_len as usize, cfg.qkv_input_width()],
)
.expect("alloc bad");
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_q_proj(
&mut enc,
&mut registry,
&device,
&bad_input,
&tensors,
&cfg,
seq_len,
)
.unwrap_err();
assert!(err.to_string().contains("dtype must be F32"), "got: {err}");
}
#[test]
fn adr_037_e4b4_gate_q_proj_rejects_zero_seq_len_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let empty = device
.alloc_buffer(4, DType::F32, vec![1])
.expect("alloc empty");
let mut enc = device.command_encoder().expect("encoder");
let err =
dispatch_eagle3_q_proj(&mut enc, &mut registry, &device, &empty, &tensors, &cfg, 0)
.unwrap_err();
assert!(
err.to_string().contains("seq_len must be > 0"),
"got: {err}"
);
}
fn cfg_qk_norm_tiny() -> Eagle3DrafterConfig {
let mut c = tiny_cfg();
c.use_qk_norm = true;
c
}
#[test]
fn adr_037_e4b5a_q_head_norm_cpu_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 cfg = cfg_qk_norm_tiny();
let seq_len: u32 = 4;
let n_heads = cfg.num_q_heads;
let head_dim = cfg.head_dim;
let total_elems = (seq_len as usize) * n_heads * head_dim;
let mut proj_data = vec![0.0f32; total_elems];
fill_random(&mut proj_data, 0xA51);
let mut weight_f32 = vec![0.0f32; head_dim];
fill_random(&mut weight_f32, 0xA52);
let weight_bf16_q: Vec<f32> = weight_f32.iter().map(|&v| bf16_quantize_f32(v)).collect();
let cpu_out = cpu_rms_norm_f32(
&proj_data,
&weight_bf16_q,
(seq_len as usize) * n_heads, head_dim,
cfg.rms_norm_eps,
);
let manifest = expected_manifest(&cfg);
let mut overrides = std::collections::HashMap::new();
overrides.insert(
"layers.0.self_attn.q_norm.weight".to_string(),
f32_to_bf16_bytes(&weight_f32),
);
let blob = build_blob_with_overrides(&manifest, &overrides);
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let proj_gpu = upload_f32_to_gpu(
&device,
&proj_data,
vec![seq_len as usize, n_heads * head_dim],
);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_q_head_norm(
&mut enc,
&mut registry,
&device,
&proj_gpu,
&tensors,
&cfg,
seq_len,
)
.expect("q_head_norm");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
let mut max_diff = 0.0f32;
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(d < 1e-3, "q_head_norm parity: diff={d} > 1e-3");
}
eprintln!("q_head_norm parity max_diff={max_diff:.6e}");
}
#[test]
fn adr_037_e4b5a_k_head_norm_cpu_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 cfg = cfg_qk_norm_tiny();
let seq_len: u32 = 4;
let n_heads = cfg.num_kv_heads;
let head_dim = cfg.head_dim;
let total_elems = (seq_len as usize) * n_heads * head_dim;
let mut proj_data = vec![0.0f32; total_elems];
fill_random(&mut proj_data, 0xA61);
let mut weight_f32 = vec![0.0f32; head_dim];
fill_random(&mut weight_f32, 0xA62);
let weight_bf16_q: Vec<f32> = weight_f32.iter().map(|&v| bf16_quantize_f32(v)).collect();
let cpu_out = cpu_rms_norm_f32(
&proj_data,
&weight_bf16_q,
(seq_len as usize) * n_heads,
head_dim,
cfg.rms_norm_eps,
);
let manifest = expected_manifest(&cfg);
let mut overrides = std::collections::HashMap::new();
overrides.insert(
"layers.0.self_attn.k_norm.weight".to_string(),
f32_to_bf16_bytes(&weight_f32),
);
let blob = build_blob_with_overrides(&manifest, &overrides);
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let proj_gpu = upload_f32_to_gpu(
&device,
&proj_data,
vec![seq_len as usize, n_heads * head_dim],
);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_k_head_norm(
&mut enc,
&mut registry,
&device,
&proj_gpu,
&tensors,
&cfg,
seq_len,
)
.expect("k_head_norm");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
let mut max_diff = 0.0f32;
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(d < 1e-3, "k_head_norm parity: diff={d} > 1e-3");
}
eprintln!("k_head_norm parity max_diff={max_diff:.6e}");
}
#[test]
fn adr_037_e4b5a_q_head_norm_errors_when_use_qk_norm_false_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 cfg = tiny_cfg(); let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let seq_len = 2_u32;
let n = (seq_len as usize) * cfg.num_q_heads * cfg.head_dim;
let proj_data = vec![0.0f32; n];
let proj_gpu = upload_f32_to_gpu(
&device,
&proj_data,
vec![seq_len as usize, cfg.num_q_heads * cfg.head_dim],
);
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_q_head_norm(
&mut enc,
&mut registry,
&device,
&proj_gpu,
&tensors,
&cfg,
seq_len,
)
.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("use_qk_norm is false") || msg.contains("q_norm absent"),
"expected gate or absent-tensor error, got: {err}"
);
}
#[test]
fn adr_037_e4b5a_gate_q_head_norm_rejects_when_cfg_off_with_tensor_present_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 load_cfg = cfg_qk_norm_tiny();
let manifest = expected_manifest(&load_cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &load_cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &load_cfg, &weights).expect("upload");
let mut dispatch_cfg = load_cfg.clone();
dispatch_cfg.use_qk_norm = false;
let seq_len = 2_u32;
let n = (seq_len as usize) * dispatch_cfg.num_q_heads * dispatch_cfg.head_dim;
let proj_data = vec![0.0f32; n];
let proj_gpu = upload_f32_to_gpu(
&device,
&proj_data,
vec![
seq_len as usize,
dispatch_cfg.num_q_heads * dispatch_cfg.head_dim,
],
);
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_q_head_norm(
&mut enc,
&mut registry,
&device,
&proj_gpu,
&tensors,
&dispatch_cfg,
seq_len,
)
.unwrap_err();
assert!(
err.to_string().contains("cfg.use_qk_norm is false"),
"expected gate error, got: {err}"
);
}
#[test]
fn adr_037_e4b5a_gate_q_head_norm_rejects_non_f32_input_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 cfg = cfg_qk_norm_tiny();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let seq_len = 2_u32;
let n = (seq_len as usize) * cfg.num_q_heads * cfg.head_dim;
let bad = device
.alloc_buffer(
n * 2,
DType::BF16,
vec![seq_len as usize, cfg.num_q_heads * cfg.head_dim],
)
.expect("alloc bad");
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_q_head_norm(
&mut enc,
&mut registry,
&device,
&bad,
&tensors,
&cfg,
seq_len,
)
.unwrap_err();
assert!(err.to_string().contains("dtype must be F32"), "got: {err}");
}
fn cpu_neox_rope(
input: &[f32], positions: &[u32], seq_len: usize,
num_heads: usize,
head_dim: usize,
rope_dim: usize,
freq_base: f32,
) -> Vec<f32> {
let half = rope_dim / 2;
let mut out = vec![0.0f32; seq_len * num_heads * head_dim];
for s in 0..seq_len {
let pos = positions[s] as f32;
for h in 0..num_heads {
let row_base = (s * num_heads + h) * head_dim;
for d in rope_dim..head_dim {
out[row_base + d] = input[row_base + d];
}
for d in 0..half {
let inv_freq = (freq_base as f64).powf(-(2.0 * d as f64) / (rope_dim as f64));
let theta = (pos as f64) * inv_freq;
let (s_t, c_t) = (theta.sin() as f32, theta.cos() as f32);
let x = input[row_base + d];
let y = input[row_base + d + half];
out[row_base + d] = x * c_t - y * s_t;
out[row_base + d + half] = x * s_t + y * c_t;
}
}
}
out
}
#[test]
fn adr_037_e4b5b_rope_cpu_parity_linear_positions_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 cfg = tiny_cfg();
let seq_len: u32 = 4;
let num_heads = cfg.num_q_heads;
let head_dim = cfg.head_dim;
let rope_dim = cfg.rope_dim;
let total = (seq_len as usize) * num_heads * head_dim;
let mut input_data = vec![0.0f32; total];
fill_random(&mut input_data, 0xB10);
let input_gpu = upload_f32_to_gpu(
&device,
&input_data,
vec![seq_len as usize * num_heads, head_dim],
);
let base_pos: u32 = 17;
let positions: Vec<u32> = (0..seq_len).map(|i| base_pos + i).collect();
let cpu_out = cpu_neox_rope(
&input_data,
&positions,
seq_len as usize,
num_heads,
head_dim,
rope_dim,
cfg.rope_theta,
);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_rope(
&mut enc,
&mut registry,
&device,
&input_gpu,
&cfg,
seq_len,
num_heads as u32,
None, base_pos,
"rope_test_linear",
)
.expect("rope dispatch");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
assert_eq!(gpu_out.len(), total);
let mut max_diff = 0.0f32;
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(d < 1e-4, "rope linear parity: diff={d} > 1e-4");
}
eprintln!("rope linear parity max_diff={max_diff:.6e}");
}
#[test]
fn adr_037_e4b5b_rope_cpu_parity_tree_positions_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 cfg = tiny_cfg();
let seq_len: u32 = 5;
let num_heads = cfg.num_q_heads;
let head_dim = cfg.head_dim;
let rope_dim = cfg.rope_dim;
let total = (seq_len as usize) * num_heads * head_dim;
let mut input_data = vec![0.0f32; total];
fill_random(&mut input_data, 0xB20);
let input_gpu = upload_f32_to_gpu(
&device,
&input_data,
vec![seq_len as usize * num_heads, head_dim],
);
let depths: [u32; 5] = [0, 1, 1, 2, 2];
let base_pos: u32 = 42;
let positions: Vec<u32> = depths.iter().map(|d| base_pos + d).collect();
let cpu_out = cpu_neox_rope(
&input_data,
&positions,
seq_len as usize,
num_heads,
head_dim,
rope_dim,
cfg.rope_theta,
);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_rope(
&mut enc,
&mut registry,
&device,
&input_gpu,
&cfg,
seq_len,
num_heads as u32,
Some(&positions),
base_pos, "rope_test_tree",
)
.expect("rope tree dispatch");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
let mut max_diff = 0.0f32;
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(d < 1e-4, "rope tree parity: diff={d} > 1e-4");
}
eprintln!("rope tree parity max_diff={max_diff:.6e}");
}
#[test]
fn adr_037_e4b5b_rope_rejects_positions_len_mismatch_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 cfg = tiny_cfg();
let seq_len = 4_u32;
let num_heads = cfg.num_q_heads as u32;
let total = (seq_len as usize) * (num_heads as usize) * cfg.head_dim;
let input = upload_f32_to_gpu(
&device,
&vec![0.0f32; total],
vec![seq_len as usize * num_heads as usize, cfg.head_dim],
);
let bad_positions = vec![0u32, 1u32, 2u32]; let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_rope(
&mut enc,
&mut registry,
&device,
&input,
&cfg,
seq_len,
num_heads,
Some(&bad_positions),
0,
"rope_bad_pos",
)
.unwrap_err();
assert!(
err.to_string().contains("positions_override len"),
"expected positions-len error, got: {err}"
);
}
#[test]
fn adr_037_e4b5b_gate_rope_rejects_position_above_i32_max_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 cfg = tiny_cfg();
let seq_len = 2_u32;
let num_heads = cfg.num_q_heads as u32;
let total = (seq_len as usize) * (num_heads as usize) * cfg.head_dim;
let input = upload_f32_to_gpu(
&device,
&vec![0.0f32; total],
vec![seq_len as usize * num_heads as usize, cfg.head_dim],
);
let bad_positions = vec![0u32, (i32::MAX as u32) + 1];
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_rope(
&mut enc,
&mut registry,
&device,
&input,
&cfg,
seq_len,
num_heads,
Some(&bad_positions),
0,
"rope_overflow",
)
.unwrap_err();
assert!(
err.to_string().contains("exceeds i32::MAX"),
"expected i32::MAX rejection, got: {err}"
);
}
#[test]
fn adr_037_e4b5b_gate_rope_rejects_non_f32_input_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 cfg = tiny_cfg();
let seq_len = 2_u32;
let num_heads = cfg.num_q_heads as u32;
let total = (seq_len as usize) * (num_heads as usize) * cfg.head_dim;
let bad = device
.alloc_buffer(
total * 2,
DType::BF16,
vec![seq_len as usize * num_heads as usize, cfg.head_dim],
)
.expect("alloc bad");
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_rope(
&mut enc,
&mut registry,
&device,
&bad,
&cfg,
seq_len,
num_heads,
None,
0,
"rope_bad_dtype",
)
.unwrap_err();
assert!(err.to_string().contains("dtype must be F32"), "got: {err}");
}
fn cfg_for_attention_dk128() -> Eagle3DrafterConfig {
Eagle3DrafterConfig {
hidden_size: 512,
intermediate_size: 1024,
head_dim: 128,
num_q_heads: 4,
num_kv_heads: 2,
vocab_size: 1000,
draft_vocab_size: 1000,
target_hidden_size: 512,
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,
}
}
#[allow(clippy::too_many_arguments)]
fn cpu_tree_attention_reference(
q: &[f32], k: &[f32], v: &[f32], mask: &[f32], num_q_heads: usize,
num_kv_heads: usize,
head_dim: usize,
q_seq_len: usize,
kv_seq_len: usize,
kv_capacity: usize,
mask_stride: usize,
scale: f32,
) -> Vec<f32> {
let heads_per_kv = num_q_heads / num_kv_heads;
let mut out = vec![0.0f32; q_seq_len * num_q_heads * head_dim];
for h in 0..num_q_heads {
let kv_h = h / heads_per_kv;
for iq1 in 0..q_seq_len {
let q_off = h * q_seq_len * head_dim + iq1 * head_dim;
let mask_row = iq1 * mask_stride;
let mut scores = Vec::<(usize, f32)>::new();
for k_pos in 0..kv_seq_len {
if mask[mask_row + k_pos] == EAGLE3_TREE_MASK_MASKED {
continue;
}
let k_off = kv_h * kv_capacity * head_dim + k_pos * head_dim;
let mut dot = 0.0f64;
for d in 0..head_dim {
dot += q[q_off + d] as f64 * k[k_off + d] as f64;
}
scores.push((k_pos, dot as f32 * scale));
}
if scores.is_empty() {
continue;
}
let max_s = scores
.iter()
.map(|(_, s)| *s)
.fold(f32::NEG_INFINITY, f32::max);
let exp_s: Vec<f32> = scores.iter().map(|(_, s)| (*s - max_s).exp()).collect();
let sum_e: f32 = exp_s.iter().sum();
let inv = if sum_e == 0.0 { 0.0 } else { 1.0 / sum_e };
let o_off = iq1 * num_q_heads * head_dim + h * head_dim;
for d in 0..head_dim {
let mut acc = 0.0f32;
for ((k_pos, _), &es) in scores.iter().zip(exp_s.iter()) {
let weight = es * inv;
let v_off = kv_h * kv_capacity * head_dim + k_pos * head_dim;
acc += weight * v[v_off + d];
}
out[o_off + d] = acc;
}
}
}
out
}
#[test]
fn adr_037_e4b6_tree_attention_cpu_parity_dk128_fixed_square_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 cfg = cfg_for_attention_dk128();
let num_q_heads = cfg.num_q_heads;
let num_kv_heads = cfg.num_kv_heads;
let head_dim = cfg.head_dim;
let q_seq_len = 5_u32; let prefix_len = 27_usize;
let kv_seq_len = (prefix_len + q_seq_len as usize) as u32; let kv_capacity = 64_u32;
let mask_stride = kv_seq_len;
let scale = 1.0_f32 / (head_dim as f32).sqrt();
let q_elems = num_q_heads * (q_seq_len as usize) * head_dim;
let mut q_data = vec![0.0f32; q_elems];
fill_random(&mut q_data, 0xC60);
let kv_elems = num_kv_heads * (kv_capacity as usize) * head_dim;
let mut k_data = vec![0.0f32; kv_elems];
fill_random(&mut k_data, 0xC61);
let mut v_data = vec![0.0f32; kv_elems];
fill_random(&mut v_data, 0xC62);
let mask_elems = (q_seq_len as usize) * (mask_stride as usize);
let mut mask_data = vec![EAGLE3_TREE_MASK_MASKED; mask_elems];
for iq1 in 0..(q_seq_len as usize) {
let row_base = iq1 * (mask_stride as usize);
for k in 0..prefix_len {
mask_data[row_base + k] = EAGLE3_TREE_MASK_ATTENDED;
}
mask_data[row_base + prefix_len + iq1] = EAGLE3_TREE_MASK_ATTENDED;
if iq1 > 0 {
mask_data[row_base + prefix_len] = EAGLE3_TREE_MASK_ATTENDED;
}
}
let cpu_out = cpu_tree_attention_reference(
&q_data,
&k_data,
&v_data,
&mask_data,
num_q_heads,
num_kv_heads,
head_dim,
q_seq_len as usize,
kv_seq_len as usize,
kv_capacity as usize,
mask_stride as usize,
scale,
);
let q_gpu = upload_f32_to_gpu(
&device,
&q_data,
vec![num_q_heads, q_seq_len as usize, head_dim],
);
let k_gpu = upload_f32_to_gpu(
&device,
&k_data,
vec![num_kv_heads, kv_capacity as usize, head_dim],
);
let v_gpu = upload_f32_to_gpu(
&device,
&v_data,
vec![num_kv_heads, kv_capacity as usize, head_dim],
);
let mask_gpu = upload_f32_to_gpu(
&device,
&mask_data,
vec![q_seq_len as usize, mask_stride as usize],
);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_tree_attention(
&mut enc,
&mut registry,
&device,
&q_gpu,
&k_gpu,
&v_gpu,
&mask_gpu,
&cfg,
q_seq_len,
kv_seq_len,
kv_capacity,
mask_stride,
scale,
)
.expect("tree attn dispatch");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
assert_eq!(gpu_out.len(), cpu_out.len());
let mut max_diff = 0.0f32;
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(d < 5e-3, "tree_attention parity: diff={d} > 5e-3");
}
eprintln!("tree_attention dk128 fixed-square parity max_diff={max_diff:.6e}");
}
#[test]
fn adr_037_e4b6_gate_tree_attention_rejects_zero_q_seq_len_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 cfg = cfg_for_attention_dk128();
let dummy = device
.alloc_buffer(4, DType::F32, vec![1])
.expect("alloc dummy");
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_tree_attention(
&mut enc,
&mut registry,
&device,
&dummy,
&dummy,
&dummy,
&dummy,
&cfg,
0,
1,
1,
1,
1.0,
)
.unwrap_err();
assert!(
err.to_string().contains("q_seq_len must be > 0"),
"got: {err}"
);
}
#[test]
fn adr_037_e4b6_gate_tree_attention_rejects_kv_capacity_less_than_seq_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 cfg = cfg_for_attention_dk128();
let dummy = device.alloc_buffer(4, DType::F32, vec![1]).expect("alloc");
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_tree_attention(
&mut enc,
&mut registry,
&device,
&dummy,
&dummy,
&dummy,
&dummy,
&cfg,
1,
10,
5,
10,
1.0, )
.unwrap_err();
assert!(err.to_string().contains("kv_capacity"), "got: {err}");
}
#[test]
fn adr_037_e4b7_o_proj_cpu_parity_no_bias_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 cfg = tiny_cfg(); let seq_len: u32 = 4;
let in_features = cfg.q_proj_out();
let out_features = cfg.hidden_size;
let mut input_data = vec![0.0f32; (seq_len as usize) * in_features];
fill_random(&mut input_data, 0xD70);
let mut weight_f32 = vec![0.0f32; out_features * in_features];
fill_random(&mut weight_f32, 0xD71);
let weight_bf16_q: Vec<f32> = weight_f32.iter().map(|&v| bf16_quantize_f32(v)).collect();
let cpu_out = cpu_fc_reference(
&input_data,
&weight_bf16_q,
seq_len as usize,
in_features,
out_features,
);
let manifest = expected_manifest(&cfg);
let mut overrides = std::collections::HashMap::new();
overrides.insert(
"layers.0.self_attn.o_proj.weight".to_string(),
f32_to_bf16_bytes(&weight_f32),
);
let blob = build_blob_with_overrides(&manifest, &overrides);
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let input_gpu =
upload_f32_to_gpu(&device, &input_data, vec![seq_len as usize, in_features]);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_o_proj(
&mut enc,
&mut registry,
&device,
&input_gpu,
&tensors,
&cfg,
seq_len,
)
.expect("o_proj dispatch");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
let mut max_diff = 0.0f32;
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(d < 5e-2, "o_proj parity: diff={d} > 5e-2");
}
eprintln!("o_proj parity max_diff={max_diff:.6e}");
}
#[test]
fn adr_037_e4b7_residual_add_cpu_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 cfg = tiny_cfg();
let seq_len: u32 = 4;
let hidden = cfg.hidden_size;
let n = (seq_len as usize) * hidden;
let mut a_data = vec![0.0f32; n];
let mut b_data = vec![0.0f32; n];
fill_random(&mut a_data, 0xE70);
fill_random(&mut b_data, 0xE71);
let cpu_out: Vec<f32> = a_data
.iter()
.zip(b_data.iter())
.map(|(a, b)| a + b)
.collect();
let a_gpu = upload_f32_to_gpu(&device, &a_data, vec![seq_len as usize, hidden]);
let b_gpu = upload_f32_to_gpu(&device, &b_data, vec![seq_len as usize, hidden]);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_residual_add(
&mut enc,
&mut registry,
&device,
&a_gpu,
&b_gpu,
&cfg,
seq_len,
)
.expect("residual_add");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
assert_eq!(g.to_bits(), c.to_bits(), "residual_add bit-equal expected");
}
}
#[test]
fn adr_037_e4b7_gate_residual_add_rejects_non_f32_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 cfg = tiny_cfg();
let seq_len = 2_u32;
let n = (seq_len as usize) * cfg.hidden_size;
let bf16 = device
.alloc_buffer(n * 2, DType::BF16, vec![seq_len as usize, cfg.hidden_size])
.expect("alloc bf16");
let f32_buf = device
.alloc_buffer(n * 4, DType::F32, vec![seq_len as usize, cfg.hidden_size])
.expect("alloc f32");
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_residual_add(
&mut enc,
&mut registry,
&device,
&bf16,
&f32_buf,
&cfg,
seq_len,
)
.unwrap_err();
assert!(err.to_string().contains("must be F32"), "got: {err}");
}
#[test]
fn adr_037_e4b7_gate_residual_add_rejects_shape_mismatch_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 cfg = tiny_cfg();
let seq_len = 4_u32;
let n = (seq_len as usize) * cfg.hidden_size;
let good = upload_f32_to_gpu(
&device,
&vec![0.0f32; n],
vec![seq_len as usize, cfg.hidden_size],
);
let bad = upload_f32_to_gpu(&device, &vec![0.0f32; 10], vec![10]);
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_residual_add(
&mut enc,
&mut registry,
&device,
&good,
&bad,
&cfg,
seq_len,
)
.unwrap_err();
assert!(err.to_string().contains("b has"), "got: {err}");
}
fn silu_f32(x: f32) -> f32 {
x / (1.0 + (-x).exp())
}
#[allow(clippy::too_many_arguments)]
fn cpu_swiglu_mlp_reference(
input: &[f32], gate_weight_bf16_q: &[f32], up_weight_bf16_q: &[f32], down_weight_bf16_q: &[f32], seq_len: usize,
hidden: usize,
inter: usize,
) -> Vec<f32> {
let gate = cpu_fc_reference(input, gate_weight_bf16_q, seq_len, hidden, inter);
let up = cpu_fc_reference(input, up_weight_bf16_q, seq_len, hidden, inter);
let mut activated = vec![0.0f32; seq_len * inter];
for i in 0..(seq_len * inter) {
activated[i] = silu_f32(gate[i]) * up[i];
}
cpu_fc_reference(&activated, down_weight_bf16_q, seq_len, inter, hidden)
}
#[test]
fn adr_037_e4b8_mlp_cpu_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 cfg = tiny_cfg();
let seq_len: u32 = 4;
let hidden = cfg.hidden_size;
let inter = cfg.intermediate_size;
let mut input_data = vec![0.0f32; (seq_len as usize) * hidden];
fill_random(&mut input_data, 0xE80);
let mut gate_w = vec![0.0f32; inter * hidden];
fill_random(&mut gate_w, 0xE81);
let mut up_w = vec![0.0f32; inter * hidden];
fill_random(&mut up_w, 0xE82);
let mut down_w = vec![0.0f32; hidden * inter];
fill_random(&mut down_w, 0xE83);
let gate_q: Vec<f32> = gate_w.iter().map(|&v| bf16_quantize_f32(v)).collect();
let up_q: Vec<f32> = up_w.iter().map(|&v| bf16_quantize_f32(v)).collect();
let down_q: Vec<f32> = down_w.iter().map(|&v| bf16_quantize_f32(v)).collect();
let cpu_out = cpu_swiglu_mlp_reference(
&input_data,
&gate_q,
&up_q,
&down_q,
seq_len as usize,
hidden,
inter,
);
let manifest = expected_manifest(&cfg);
let mut overrides = std::collections::HashMap::new();
overrides.insert(
"layers.0.mlp.gate_proj.weight".to_string(),
f32_to_bf16_bytes(&gate_w),
);
overrides.insert(
"layers.0.mlp.up_proj.weight".to_string(),
f32_to_bf16_bytes(&up_w),
);
overrides.insert(
"layers.0.mlp.down_proj.weight".to_string(),
f32_to_bf16_bytes(&down_w),
);
let blob = build_blob_with_overrides(&manifest, &overrides);
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let input_gpu = upload_f32_to_gpu(&device, &input_data, vec![seq_len as usize, hidden]);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_mlp(
&mut enc,
&mut registry,
&device,
&input_gpu,
&tensors,
&cfg,
seq_len,
)
.expect("dispatch_eagle3_mlp");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
assert_eq!(gpu_out.len(), (seq_len as usize) * hidden);
let max_abs_cpu = cpu_out.iter().map(|v| v.abs()).fold(0.0f32, f32::max);
let rel_tol = max_abs_cpu * 1e-2;
let abs_tol = 1e-3; let tol = rel_tol.max(abs_tol);
let mut max_diff = 0.0f32;
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(
d < tol,
"mlp parity: diff={d} > tol={tol} (max_abs={max_abs_cpu})"
);
}
eprintln!(
"mlp parity max_diff={max_diff:.6e} (max_abs={max_abs_cpu:.6e}, rel={:.6e})",
max_diff / max_abs_cpu.max(1e-9)
);
}
#[test]
fn adr_037_e4b8_gate_mlp_rejects_non_f32_input_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let seq_len = 2_u32;
let bf16 = device
.alloc_buffer(
(seq_len as usize) * cfg.hidden_size * 2,
DType::BF16,
vec![seq_len as usize, cfg.hidden_size],
)
.expect("alloc bad");
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_mlp(
&mut enc,
&mut registry,
&device,
&bf16,
&tensors,
&cfg,
seq_len,
)
.unwrap_err();
assert!(err.to_string().contains("dtype must be F32"), "got: {err}");
}
#[test]
fn adr_037_e4b8_gate_mlp_rejects_wrong_input_shape_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let seq_len = 4_u32;
let bad_data = vec![0.0f32; 100];
let bad = upload_f32_to_gpu(&device, &bad_data, vec![100]);
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_mlp(
&mut enc,
&mut registry,
&device,
&bad,
&tensors,
&cfg,
seq_len,
)
.unwrap_err();
assert!(
err.to_string().contains("input has 100 elements"),
"got: {err}"
);
}
#[test]
fn adr_037_e4b8_gate_mlp_rejects_zero_seq_len_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let empty = device.alloc_buffer(4, DType::F32, vec![1]).expect("alloc");
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_mlp(&mut enc, &mut registry, &device, &empty, &tensors, &cfg, 0)
.unwrap_err();
assert!(
err.to_string().contains("seq_len must be > 0"),
"got: {err}"
);
}
#[test]
fn adr_037_e4b9_final_norm_cpu_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 cfg = tiny_cfg();
let seq_len: u32 = 4;
let hidden = cfg.hidden_size;
let mut input_data = vec![0.0f32; (seq_len as usize) * hidden];
fill_random(&mut input_data, 0xF90);
let mut weight_f32 = vec![0.0f32; hidden];
fill_random(&mut weight_f32, 0xF91);
let weight_bf16_q: Vec<f32> = weight_f32.iter().map(|&v| bf16_quantize_f32(v)).collect();
let cpu_out = cpu_rms_norm_f32(
&input_data,
&weight_bf16_q,
seq_len as usize,
hidden,
cfg.rms_norm_eps,
);
let manifest = expected_manifest(&cfg);
let mut overrides = std::collections::HashMap::new();
overrides.insert("norm.weight".to_string(), f32_to_bf16_bytes(&weight_f32));
let blob = build_blob_with_overrides(&manifest, &overrides);
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let input_gpu = upload_f32_to_gpu(&device, &input_data, vec![seq_len as usize, hidden]);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_final_norm(
&mut enc,
&mut registry,
&device,
&input_gpu,
&tensors,
&cfg,
seq_len,
)
.expect("final_norm");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
let mut max_diff = 0.0f32;
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(d < 1e-3, "final_norm parity: diff={d} > 1e-3");
}
eprintln!("final_norm parity max_diff={max_diff:.6e}");
}
fn cfg_for_lm_head_test() -> Eagle3DrafterConfig {
let mut c = tiny_cfg();
c.tie_lm_head = false;
c
}
#[test]
fn adr_037_e4b9_lm_head_cpu_parity_untied_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 cfg = cfg_for_lm_head_test();
let seq_len: u32 = 3;
let hidden = cfg.hidden_size;
let dvocab = cfg.draft_vocab_size;
let mut input_data = vec![0.0f32; (seq_len as usize) * hidden];
fill_random(&mut input_data, 0xFA0);
let mut weight_f32 = vec![0.0f32; dvocab * hidden];
fill_random(&mut weight_f32, 0xFA1);
let weight_bf16_q: Vec<f32> = weight_f32.iter().map(|&v| bf16_quantize_f32(v)).collect();
let cpu_out = cpu_fc_reference(
&input_data,
&weight_bf16_q,
seq_len as usize,
hidden,
dvocab,
);
let manifest = expected_manifest(&cfg);
let mut overrides = std::collections::HashMap::new();
overrides.insert("lm_head.weight".to_string(), f32_to_bf16_bytes(&weight_f32));
let blob = build_blob_with_overrides(&manifest, &overrides);
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let input_gpu = upload_f32_to_gpu(&device, &input_data, vec![seq_len as usize, hidden]);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_lm_head(
&mut enc,
&mut registry,
&device,
&input_gpu,
&tensors,
&cfg,
seq_len,
)
.expect("lm_head untied");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
assert_eq!(gpu_out.len(), (seq_len as usize) * dvocab);
let mut max_diff = 0.0f32;
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(d < 5e-2, "lm_head parity: diff={d} > 5e-2");
}
eprintln!("lm_head untied parity max_diff={max_diff:.6e}");
}
#[test]
fn adr_037_e4b9_gate_lm_head_tied_requires_full_vocab_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 cfg = tiny_cfg();
cfg.tie_lm_head = true;
cfg.has_own_embed_tokens = true;
cfg.draft_vocab_size = 500; let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let seq_len = 2_u32;
let input = upload_f32_to_gpu(
&device,
&vec![0.0f32; (seq_len as usize) * cfg.hidden_size],
vec![seq_len as usize, cfg.hidden_size],
);
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_lm_head(
&mut enc,
&mut registry,
&device,
&input,
&tensors,
&cfg,
seq_len,
)
.unwrap_err();
assert!(err.to_string().contains("draft_vocab_size"), "got: {err}");
}
#[test]
fn adr_037_e4b9_gate_lm_head_tied_requires_embed_tokens_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 cfg = tiny_cfg();
cfg.tie_lm_head = true;
cfg.has_own_embed_tokens = false; cfg.draft_vocab_size = cfg.vocab_size; let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let seq_len = 2_u32;
let input = upload_f32_to_gpu(
&device,
&vec![0.0f32; (seq_len as usize) * cfg.hidden_size],
vec![seq_len as usize, cfg.hidden_size],
);
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_lm_head(
&mut enc,
&mut registry,
&device,
&input,
&tensors,
&cfg,
seq_len,
)
.unwrap_err();
assert!(err.to_string().contains("embed_tokens"), "got: {err}");
}
#[test]
fn adr_037_e4b10b1_permute_sentinel_layout_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 seq = 3_u32;
let n_heads = 2_u32;
let hd = 4_usize;
let total = (seq as usize) * (n_heads as usize) * hd;
let mut input_data = vec![0.0f32; total];
for s in 0..(seq as usize) {
for h in 0..(n_heads as usize) {
for d in 0..hd {
let val = (s * 1000 + h * 100 + d) as f32;
let idx = s * (n_heads as usize) * hd + h * hd + d;
input_data[idx] = val;
}
}
}
let input_gpu = upload_f32_to_gpu(
&device,
&input_data,
vec![seq as usize, n_heads as usize, hd],
);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_permute_seq_to_head_outer(
&mut enc,
&mut registry,
&device,
&input_gpu,
seq,
n_heads,
hd,
"permute_sentinel",
)
.expect("permute dispatch");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
assert_eq!(gpu_out.len(), total);
for h in 0..(n_heads as usize) {
for s in 0..(seq as usize) {
for d in 0..hd {
let expected = (s * 1000 + h * 100 + d) as f32;
let out_idx = h * (seq as usize) * hd + s * hd + d;
assert_eq!(
gpu_out[out_idx], expected,
"head={} seq={} dim={} mismatch: got {}",
h, s, d, gpu_out[out_idx]
);
}
}
}
}
#[test]
fn adr_037_e4b10b1_permute_cpu_parity_random_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 seq = 5_u32;
let n_heads = 8_u32;
let hd = 16_usize;
let total = (seq as usize) * (n_heads as usize) * hd;
let mut input_data = vec![0.0f32; total];
fill_random(&mut input_data, 0xB10B);
let input_gpu = upload_f32_to_gpu(
&device,
&input_data,
vec![seq as usize, n_heads as usize, hd],
);
let mut cpu_out = vec![0.0f32; total];
for s in 0..(seq as usize) {
for h in 0..(n_heads as usize) {
for d in 0..hd {
let src = s * (n_heads as usize) * hd + h * hd + d;
let dst = h * (seq as usize) * hd + s * hd + d;
cpu_out[dst] = input_data[src];
}
}
}
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_permute_seq_to_head_outer(
&mut enc,
&mut registry,
&device,
&input_gpu,
seq,
n_heads,
hd,
"permute_random",
)
.expect("permute dispatch");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
for (i, (g, c)) in gpu_out.iter().zip(cpu_out.iter()).enumerate() {
assert_eq!(
g.to_bits(),
c.to_bits(),
"byte-identity violated at idx {i}: gpu={g} cpu={c}"
);
}
}
#[test]
fn adr_037_e4b10b1_gate_permute_rejects_non_f32_input_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 bf16 = device
.alloc_buffer(64, DType::BF16, vec![2, 4, 4])
.expect("alloc bad");
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_permute_seq_to_head_outer(
&mut enc,
&mut registry,
&device,
&bf16,
2,
4,
4,
"permute_bad",
)
.unwrap_err();
assert!(err.to_string().contains("dtype must be F32"), "got: {err}");
}
#[test]
fn adr_037_e4b10b1_gate_permute_rejects_wrong_element_count_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 bad_data = vec![0.0f32; 8];
let bad = upload_f32_to_gpu(&device, &bad_data, vec![8]);
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_permute_seq_to_head_outer(
&mut enc,
&mut registry,
&device,
&bad,
2,
4,
4,
"permute_wrong_count",
)
.unwrap_err();
assert!(
err.to_string().contains("input has 8 elements"),
"got: {err}"
);
}
#[test]
fn adr_037_e4b10b1_gate_permute_rejects_zero_dim_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("alloc");
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_permute_seq_to_head_outer(
&mut enc,
&mut registry,
&device,
&dummy,
0,
4,
4,
"permute_zero",
)
.unwrap_err();
assert!(
err.to_string().contains("all dims must be > 0"),
"got: {err}"
);
}
#[test]
fn adr_037_e4b10b2_post_attention_layernorm_cpu_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 cfg = tiny_cfg();
let seq_len: u32 = 4;
let hidden = cfg.hidden_size;
let mut input_data = vec![0.0f32; (seq_len as usize) * hidden];
fill_random(&mut input_data, 0xB22);
let mut weight_f32 = vec![0.0f32; hidden];
fill_random(&mut weight_f32, 0xB23);
let weight_bf16_q: Vec<f32> = weight_f32.iter().map(|&v| bf16_quantize_f32(v)).collect();
let cpu_out = cpu_rms_norm_f32(
&input_data,
&weight_bf16_q,
seq_len as usize,
hidden,
cfg.rms_norm_eps,
);
let manifest = expected_manifest(&cfg);
let mut overrides = std::collections::HashMap::new();
overrides.insert(
"layers.0.post_attention_layernorm.weight".to_string(),
f32_to_bf16_bytes(&weight_f32),
);
let blob = build_blob_with_overrides(&manifest, &overrides);
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let input_gpu = upload_f32_to_gpu(&device, &input_data, vec![seq_len as usize, hidden]);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_post_attention_layernorm(
&mut enc,
&mut registry,
&device,
&input_gpu,
&tensors,
&cfg,
seq_len,
)
.expect("post_attention_layernorm");
enc.commit_and_wait().expect("commit");
let gpu_out: &[f32] = out_buf.as_slice::<f32>().expect("output slice");
let mut max_diff = 0.0f32;
for (g, c) in gpu_out.iter().zip(cpu_out.iter()) {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
}
assert!(d < 1e-3, "post_attention_layernorm parity: diff={d}");
}
eprintln!("post_attention_layernorm parity max_diff={max_diff:.6e}");
}
fn cfg_for_full_forward_test() -> Eagle3DrafterConfig {
let mut c = cfg_for_attention_dk128();
c.tie_lm_head = false; c
}
#[test]
fn adr_037_e4b10b2_full_forward_chain_finite_and_deterministic_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = eagle3_test_kernel_registry();
let cfg = cfg_for_full_forward_test();
let seq_len: u32 = 2;
let hidden = cfg.hidden_size;
let dvocab = cfg.draft_vocab_size;
let manifest = expected_manifest(&cfg);
let mut overrides = std::collections::HashMap::new();
for tensor in &manifest {
let name_hash: u64 = tensor
.name
.bytes()
.fold(0u64, |acc, b| acc.wrapping_mul(31).wrapping_add(b as u64));
let n_elem: usize = tensor.shape.iter().product();
let elem_bytes = match tensor.dtype {
safetensors::tensor::Dtype::BF16 => 2,
safetensors::tensor::Dtype::I64 => 8,
_ => panic!("unexpected dtype in full forward test"),
};
if tensor.dtype == safetensors::tensor::Dtype::BF16 {
let mut vals = vec![0.0f32; n_elem];
for (i, v) in vals.iter_mut().enumerate() {
let seed = name_hash.wrapping_add(i as u64);
let is_norm = tensor.name.contains("norm");
*v = if is_norm {
1.0 + pseudo_random(seed) * 0.1
} else {
pseudo_random(seed) * 0.044 };
}
overrides.insert(tensor.name.clone(), f32_to_bf16_bytes(&vals));
} else {
let _ = elem_bytes;
}
}
let blob = build_blob_with_overrides(&manifest, &overrides);
let weights = Eagle3Weights::load(&blob, &cfg).expect("load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let mut target_aux = vec![0.0f32; (seq_len as usize) * cfg.fc_input_size()];
for (i, v) in target_aux.iter_mut().enumerate() {
*v = pseudo_random(0xC0FFEE + i as u64) * 0.5;
}
let mut embeds = vec![0.0f32; (seq_len as usize) * hidden];
for (i, v) in embeds.iter_mut().enumerate() {
*v = pseudo_random(0xD0FFEE + i as u64) * 0.5;
}
let target_aux_gpu = upload_f32_to_gpu(
&device,
&target_aux,
vec![seq_len as usize, cfg.fc_input_size()],
);
let embeds_gpu = upload_f32_to_gpu(&device, &embeds, vec![seq_len as usize, hidden]);
let logits_run1 = run_full_eagle3_forward(
&device,
&mut registry,
&target_aux_gpu,
&embeds_gpu,
&tensors,
&cfg,
seq_len,
);
let logits_run2 = run_full_eagle3_forward(
&device,
&mut registry,
&target_aux_gpu,
&embeds_gpu,
&tensors,
&cfg,
seq_len,
);
assert_eq!(
logits_run1.len(),
(seq_len as usize) * dvocab,
"logits shape"
);
for (i, &v) in logits_run1.iter().enumerate() {
assert!(v.is_finite(), "logits[{i}] = {v} is not finite");
}
let max_abs = logits_run1.iter().map(|v| v.abs()).fold(0.0f32, f32::max);
eprintln!(
"full forward chain at tiny-cfg: max_abs_logit = {max_abs:.6e} \
(small magnitudes expected from BF16 underflow at this synthetic shape; \
real trained weights validate signal preservation)"
);
for (i, (a, b)) in logits_run1.iter().zip(logits_run2.iter()).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"non-deterministic at idx {i}: run1={a} run2={b}"
);
}
let row0_logits = &logits_run1[..dvocab];
let top_k = crate::inference::spec_decode::eagle3::drafter::extract_top_k_from_row_logits(
row0_logits,
5,
)
.expect("top-K extraction");
assert_eq!(top_k.len(), 5);
crate::inference::spec_decode::eagle3::drafter::validate_candidates(&top_k, 5)
.expect("Phase E4a contract");
}
fn run_full_eagle3_forward(
device: &MlxDevice,
registry: &mut KernelRegistry,
target_aux_gpu: &MlxBuffer,
embeds_gpu: &MlxBuffer,
tensors: &Eagle3DrafterTensors,
cfg: &Eagle3DrafterConfig,
seq_len: u32,
) -> Vec<f32> {
let mut enc = device.command_encoder().expect("encoder");
let fc_out = dispatch_eagle3_fc(
&mut enc,
registry,
device,
target_aux_gpu,
tensors,
cfg,
seq_len,
)
.expect("fc");
let embeds_normed = dispatch_eagle3_input_layernorm(
&mut enc, registry, device, embeds_gpu, tensors, cfg, seq_len,
)
.expect("input_layernorm");
let hidden_normed =
dispatch_eagle3_hidden_norm(&mut enc, registry, device, &fc_out, tensors, cfg, seq_len)
.expect("hidden_norm");
let concat = dispatch_eagle3_concat_2x_hidden(
&mut enc,
registry,
device,
&embeds_normed,
&hidden_normed,
cfg,
seq_len,
)
.expect("concat");
let q = dispatch_eagle3_q_proj(&mut enc, registry, device, &concat, tensors, cfg, seq_len)
.expect("q_proj");
let k = dispatch_eagle3_k_proj(&mut enc, registry, device, &concat, tensors, cfg, seq_len)
.expect("k_proj");
let v = dispatch_eagle3_v_proj(&mut enc, registry, device, &concat, tensors, cfg, seq_len)
.expect("v_proj");
let (q_normed, k_normed) = if cfg.use_qk_norm {
let qn =
dispatch_eagle3_q_head_norm(&mut enc, registry, device, &q, tensors, cfg, seq_len)
.expect("q_head_norm");
let kn =
dispatch_eagle3_k_head_norm(&mut enc, registry, device, &k, tensors, cfg, seq_len)
.expect("k_head_norm");
(qn, kn)
} else {
(q, k)
};
let q_roped = dispatch_eagle3_rope(
&mut enc,
registry,
device,
&q_normed,
cfg,
seq_len,
cfg.num_q_heads as u32,
None,
0,
"q_rope",
)
.expect("q_rope");
let k_roped = dispatch_eagle3_rope(
&mut enc,
registry,
device,
&k_normed,
cfg,
seq_len,
cfg.num_kv_heads as u32,
None,
0,
"k_rope",
)
.expect("k_rope");
let q_perm = dispatch_eagle3_permute_seq_to_head_outer(
&mut enc,
registry,
device,
&q_roped,
seq_len,
cfg.num_q_heads as u32,
cfg.head_dim,
"q_permute",
)
.expect("q_permute");
let k_perm = dispatch_eagle3_permute_seq_to_head_outer(
&mut enc,
registry,
device,
&k_roped,
seq_len,
cfg.num_kv_heads as u32,
cfg.head_dim,
"k_permute",
)
.expect("k_permute");
let v_perm = dispatch_eagle3_permute_seq_to_head_outer(
&mut enc,
registry,
device,
&v,
seq_len,
cfg.num_kv_heads as u32,
cfg.head_dim,
"v_permute",
)
.expect("v_permute");
let kv_seq_len = seq_len;
let kv_capacity = seq_len;
let mask_stride = kv_seq_len;
let mask_elems = (seq_len as usize) * (mask_stride as usize);
let mask_data = vec![EAGLE3_TREE_MASK_ATTENDED; mask_elems];
let mask_gpu = upload_f32_to_gpu(
device,
&mask_data,
vec![seq_len as usize, mask_stride as usize],
);
let scale = 1.0f32 / (cfg.head_dim as f32).sqrt();
let attn_out = dispatch_eagle3_tree_attention(
&mut enc,
registry,
device,
&q_perm,
&k_perm,
&v_perm,
&mask_gpu,
cfg,
seq_len,
kv_seq_len,
kv_capacity,
mask_stride,
scale,
)
.expect("tree_attention");
let o_out =
dispatch_eagle3_o_proj(&mut enc, registry, device, &attn_out, tensors, cfg, seq_len)
.expect("o_proj");
let attn_residual = dispatch_eagle3_residual_add(
&mut enc,
registry,
device,
&o_out,
&hidden_normed,
cfg,
seq_len,
)
.expect("attn_residual_add");
let post_attn_normed = dispatch_eagle3_post_attention_layernorm(
&mut enc,
registry,
device,
&attn_residual,
tensors,
cfg,
seq_len,
)
.expect("post_attention_layernorm");
let mlp_out = dispatch_eagle3_mlp(
&mut enc,
registry,
device,
&post_attn_normed,
tensors,
cfg,
seq_len,
)
.expect("mlp");
let final_residual = dispatch_eagle3_residual_add(
&mut enc,
registry,
device,
&mlp_out,
&attn_residual,
cfg,
seq_len,
)
.expect("final_residual_add");
let final_normed = dispatch_eagle3_final_norm(
&mut enc,
registry,
device,
&final_residual,
tensors,
cfg,
seq_len,
)
.expect("final_norm");
let logits = dispatch_eagle3_lm_head(
&mut enc,
registry,
device,
&final_normed,
tensors,
cfg,
seq_len,
)
.expect("lm_head");
enc.commit_and_wait().expect("commit");
logits.as_slice::<f32>().expect("logits slice").to_vec()
}
#[test]
fn adr_037_e4b3_concat_rejects_wrong_input_elements_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_overrides(&manifest, &std::collections::HashMap::new());
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let _tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let seq_len = 2_u32;
let good_data = vec![0.0f32; (seq_len as usize) * cfg.hidden_size];
let good_gpu =
upload_f32_to_gpu(&device, &good_data, vec![seq_len as usize, cfg.hidden_size]);
let bad_data = vec![0.0f32; 10];
let bad_gpu = upload_f32_to_gpu(&device, &bad_data, vec![10]);
let mut enc = device.command_encoder().expect("encoder");
let err = dispatch_eagle3_concat_2x_hidden(
&mut enc,
&mut registry,
&device,
&good_gpu,
&bad_gpu,
&cfg,
seq_len,
)
.unwrap_err();
assert!(
err.to_string().contains("hidden_normed has"),
"expected element-count error, got: {err}"
);
}
fn e5b_test_setup(
device: &MlxDevice,
seq_len: u32,
) -> Option<(
Eagle3DrafterConfig,
Eagle3DrafterTensors,
MlxBuffer,
MlxBuffer,
)> {
let cfg = cfg_for_full_forward_test();
let manifest = expected_manifest(&cfg);
let mut overrides = std::collections::HashMap::new();
for tensor in &manifest {
let name_hash: u64 = tensor
.name
.bytes()
.fold(0u64, |acc, b| acc.wrapping_mul(31).wrapping_add(b as u64));
let n_elem: usize = tensor.shape.iter().product();
if tensor.dtype == safetensors::tensor::Dtype::BF16 {
let mut vals = vec![0.0f32; n_elem];
for (i, v) in vals.iter_mut().enumerate() {
let seed = name_hash.wrapping_add(i as u64);
let is_norm = tensor.name.contains("norm");
*v = if is_norm {
1.0 + pseudo_random(seed) * 0.1
} else {
pseudo_random(seed) * 0.044
};
}
overrides.insert(tensor.name.clone(), f32_to_bf16_bytes(&vals));
}
}
let blob = build_blob_with_overrides(&manifest, &overrides);
let weights = Eagle3Weights::load(&blob, &cfg).ok()?;
let tensors = Eagle3DrafterTensors::upload(device, &cfg, &weights).ok()?;
let mut target_aux = vec![0.0f32; (seq_len as usize) * cfg.fc_input_size()];
for (i, v) in target_aux.iter_mut().enumerate() {
*v = pseudo_random(0xC0FFEE + i as u64) * 0.5;
}
let mut embeds = vec![0.0f32; (seq_len as usize) * cfg.hidden_size];
for (i, v) in embeds.iter_mut().enumerate() {
*v = pseudo_random(0xD0FFEE + i as u64) * 0.5;
}
let target_aux_gpu = upload_f32_to_gpu(
device,
&target_aux,
vec![seq_len as usize, cfg.fc_input_size()],
);
let embeds_gpu =
upload_f32_to_gpu(device, &embeds, vec![seq_len as usize, cfg.hidden_size]);
Some((cfg, tensors, target_aux_gpu, embeds_gpu))
}
#[test]
fn adr_037_e5b_step2_smoke_empty_cache_appends_one_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = eagle3_test_kernel_registry();
let (cfg, tensors, target_aux_gpu, embeds_gpu) = match e5b_test_setup(&device, 1) {
Some(t) => t,
None => return,
};
let mut cache =
DrafterKvCache::new(&device, cfg.num_kv_heads, 1, cfg.head_dim).expect("alloc cache");
assert_eq!(cache.len(), 0);
let logits = dispatch_eagle3_drafter_forward_with_kv_cache(
&device,
&mut registry,
&target_aux_gpu,
&embeds_gpu,
&tensors,
&cfg,
1,
0,
&mut cache,
None,
)
.expect("forward with cache");
assert_eq!(cache.len(), 1);
assert_eq!(logits.len(), cfg.draft_vocab_size);
for (i, &v) in logits.iter().enumerate() {
assert!(v.is_finite(), "logits[{i}] = {v} not finite");
}
}
#[test]
fn adr_037_e5b_step2_rejects_wrong_cache_num_kv_heads_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 (cfg, tensors, target_aux_gpu, embeds_gpu) = match e5b_test_setup(&device, 1) {
Some(t) => t,
None => return,
};
let mut cache = DrafterKvCache::new(&device, cfg.num_kv_heads + 1, 1, cfg.head_dim)
.expect("alloc cache");
let err = dispatch_eagle3_drafter_forward_with_kv_cache(
&device,
&mut registry,
&target_aux_gpu,
&embeds_gpu,
&tensors,
&cfg,
1,
0,
&mut cache,
None,
)
.unwrap_err();
assert!(
err.to_string().contains("num_kv_heads"),
"expected num_kv_heads mismatch error, got: {err}"
);
}
#[test]
fn adr_037_e5b_step2_rejects_wrong_cache_head_dim_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 (cfg, tensors, target_aux_gpu, embeds_gpu) = match e5b_test_setup(&device, 1) {
Some(t) => t,
None => return,
};
let mut cache = DrafterKvCache::new(&device, cfg.num_kv_heads, 1, cfg.head_dim + 1)
.expect("alloc cache");
let err = dispatch_eagle3_drafter_forward_with_kv_cache(
&device,
&mut registry,
&target_aux_gpu,
&embeds_gpu,
&tensors,
&cfg,
1,
0,
&mut cache,
None,
)
.unwrap_err();
assert!(
err.to_string().contains("head_dim"),
"expected head_dim mismatch error, got: {err}"
);
}
#[test]
fn adr_037_e5b_step2_rejects_seq_len_not_1_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 (cfg, tensors, target_aux_gpu, embeds_gpu) = match e5b_test_setup(&device, 2) {
Some(t) => t,
None => return,
};
let mut cache =
DrafterKvCache::new(&device, cfg.num_kv_heads, 4, cfg.head_dim).expect("alloc cache");
let err = dispatch_eagle3_drafter_forward_with_kv_cache(
&device,
&mut registry,
&target_aux_gpu,
&embeds_gpu,
&tensors,
&cfg,
2,
0,
&mut cache,
None,
)
.unwrap_err();
assert!(
err.to_string().contains("seq_len must be 1"),
"expected seq_len reject, got: {err}"
);
}
#[test]
fn adr_037_e5b_step2_rejects_cache_overflow_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = eagle3_test_kernel_registry();
let (cfg, tensors, target_aux_gpu, embeds_gpu) = match e5b_test_setup(&device, 1) {
Some(t) => t,
None => return,
};
let mut cache =
DrafterKvCache::new(&device, cfg.num_kv_heads, 1, cfg.head_dim).expect("alloc cache");
let _ = dispatch_eagle3_drafter_forward_with_kv_cache(
&device,
&mut registry,
&target_aux_gpu,
&embeds_gpu,
&tensors,
&cfg,
1,
0,
&mut cache,
None,
)
.expect("first call");
assert_eq!(cache.len(), 1);
let err = dispatch_eagle3_drafter_forward_with_kv_cache(
&device,
&mut registry,
&target_aux_gpu,
&embeds_gpu,
&tensors,
&cfg,
1,
1,
&mut cache,
None,
)
.unwrap_err();
assert!(
err.to_string().contains("would overflow"),
"expected overflow, got: {err}"
);
}
#[test]
fn adr_037_e5b_step2_equivalence_with_unbatched_at_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 = eagle3_test_kernel_registry();
let (cfg, tensors, target_aux_gpu, embeds_gpu) = match e5b_test_setup(&device, 1) {
Some(t) => t,
None => return,
};
let logits_ref = dispatch_eagle3_drafter_forward(
&device,
&mut registry,
&target_aux_gpu,
&embeds_gpu,
&tensors,
&cfg,
1,
0,
)
.expect("unbatched");
let mut cache =
DrafterKvCache::new(&device, cfg.num_kv_heads, 1, cfg.head_dim).expect("cache");
let logits_cache = dispatch_eagle3_drafter_forward_with_kv_cache(
&device,
&mut registry,
&target_aux_gpu,
&embeds_gpu,
&tensors,
&cfg,
1,
0,
&mut cache,
None,
)
.expect("with cache");
assert_eq!(logits_ref.len(), logits_cache.len());
for (i, (a, b)) in logits_ref.iter().zip(logits_cache.iter()).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"logits drift at idx {}: ref={} cache={}",
i,
a,
b,
);
}
assert_eq!(cache.len(), 1);
}
#[test]
fn adr_037_e5b_step2_incremental_two_calls_grows_cache_2026_05_22() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(_) => return,
};
let mut registry = eagle3_test_kernel_registry();
let (cfg, tensors, target_aux_gpu, embeds_gpu) = match e5b_test_setup(&device, 1) {
Some(t) => t,
None => return,
};
let mut cache =
DrafterKvCache::new(&device, cfg.num_kv_heads, 2, cfg.head_dim).expect("cache");
let logits1 = dispatch_eagle3_drafter_forward_with_kv_cache(
&device,
&mut registry,
&target_aux_gpu,
&embeds_gpu,
&tensors,
&cfg,
1,
0,
&mut cache,
None,
)
.expect("first");
assert_eq!(cache.len(), 1);
let logits2 = dispatch_eagle3_drafter_forward_with_kv_cache(
&device,
&mut registry,
&target_aux_gpu,
&embeds_gpu,
&tensors,
&cfg,
1,
1,
&mut cache,
None,
)
.expect("second");
assert_eq!(cache.len(), 2);
for (i, &v) in logits1.iter().enumerate() {
assert!(v.is_finite(), "first logits[{i}] = {v} not finite");
}
for (i, &v) in logits2.iter().enumerate() {
assert!(v.is_finite(), "second logits[{i}] = {v} not finite");
}
}
#[test]
fn adr_037_e4b2_fc_output_shape_correct_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 cfg = tiny_cfg();
let manifest = expected_manifest(&cfg);
let blob = build_blob_with_fc_weight(
&manifest,
&vec![0u8; cfg.hidden_size * cfg.fc_input_size() * 2],
);
let weights = Eagle3Weights::load(&blob, &cfg).expect("weights load");
let tensors = Eagle3DrafterTensors::upload(&device, &cfg, &weights).expect("upload");
let seq_len = 8_u32;
let input_data = vec![0.0f32; (seq_len as usize) * cfg.fc_input_size()];
let input_gpu = upload_f32_to_gpu(
&device,
&input_data,
vec![seq_len as usize, cfg.fc_input_size()],
);
let mut enc = device.command_encoder().expect("encoder");
let out_buf = dispatch_eagle3_fc(
&mut enc,
&mut registry,
&device,
&input_gpu,
&tensors,
&cfg,
seq_len,
)
.expect("dispatch");
enc.commit_and_wait().expect("commit");
assert_eq!(out_buf.dtype(), DType::F32);
assert_eq!(
out_buf.element_count(),
(seq_len as usize) * cfg.hidden_size
);
}
}