use super::config::DFlashConfig;
use super::tensors::DFlashLayerTensors;
use crate::inference::models::qwen35::gpu_full_attn::{
apply_imrope, apply_linear_projection_f32, apply_sdpa_causal_from_seq_major,
};
use anyhow::{anyhow, Context, Result};
use mlx_native::ops::elementwise::elementwise_add;
use mlx_native::ops::rms_norm::dispatch_rms_norm;
use mlx_native::ops::sdpa::{sdpa, SdpaParams};
use mlx_native::ops::silu_mul::dispatch_silu_mul;
use mlx_native::ops::softcap::dispatch_softcap;
use mlx_native::ops::transpose::permute_021_f32;
use mlx_native::{CommandEncoder, DType, KernelRegistry, MlxBuffer, MlxDevice};
fn alloc_rms_norm_params(device: &MlxDevice, eps: f32, dim: u32) -> Result<MlxBuffer> {
let mut params = device
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow!("alloc rms_norm params: {e}"))?;
let slice = params
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("rms_norm params slice: {e}"))?;
slice[0] = eps;
slice[1] = dim as f32;
Ok(params)
}
pub fn dispatch_dflash_input_layernorm(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
h_input: &MlxBuffer,
layer: &DFlashLayerTensors,
cfg: &DFlashConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
let hidden = cfg.hidden_size as u32;
let element_count = (seq_len as usize) * (hidden as usize);
if h_input.element_count() != element_count {
return Err(anyhow!(
"dflash input_layernorm: h_input element count {} != L({}) * hidden({})",
h_input.element_count(),
seq_len,
hidden
));
}
let normed = device
.alloc_buffer(
element_count * 4,
DType::F32,
vec![seq_len as usize, hidden as usize],
)
.map_err(|e| anyhow!("alloc input_layernorm output: {e}"))?;
let params = alloc_rms_norm_params(device, cfg.rms_norm_eps, hidden)?;
dispatch_rms_norm(
encoder,
registry,
device.metal_device(),
h_input,
&layer.input_layernorm,
&normed,
¶ms,
seq_len,
hidden,
)
.context("dispatch_rms_norm input_layernorm")?;
Ok(normed)
}
pub fn dispatch_dflash_q_proj(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
layer: &DFlashLayerTensors,
cfg: &DFlashConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
let hidden = cfg.hidden_size as u32;
let q_out_dim = (cfg.num_attention_heads * cfg.head_dim) as u32;
apply_linear_projection_f32(
encoder,
registry,
device,
input,
&layer.q_proj,
seq_len,
hidden,
q_out_dim,
)
.context("dispatch_dflash_q_proj")
}
pub fn dispatch_dflash_k_proj(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
layer: &DFlashLayerTensors,
cfg: &DFlashConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
let hidden = cfg.hidden_size as u32;
let kv_out_dim = (cfg.num_key_value_heads * cfg.head_dim) as u32;
apply_linear_projection_f32(
encoder,
registry,
device,
input,
&layer.k_proj,
seq_len,
hidden,
kv_out_dim,
)
.context("dispatch_dflash_k_proj")
}
pub fn dispatch_dflash_head_norm(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
proj: &MlxBuffer,
norm_weight: &MlxBuffer,
cfg: &DFlashConfig,
seq_len: u32,
num_heads: u32,
) -> Result<MlxBuffer> {
let head_dim = cfg.head_dim as u32;
let rows = seq_len * num_heads;
let expected_elem = (rows as usize) * (head_dim as usize);
if proj.element_count() != expected_elem {
return Err(anyhow!(
"dflash head_norm: proj element count {} != rows({}) * head_dim({})",
proj.element_count(),
rows,
head_dim
));
}
let normed = device
.alloc_buffer(
expected_elem * 4,
DType::F32,
vec![seq_len as usize, num_heads as usize, head_dim as usize],
)
.map_err(|e| anyhow!("alloc head_norm output: {e}"))?;
let params = alloc_rms_norm_params(device, cfg.rms_norm_eps, head_dim)?;
dispatch_rms_norm(
encoder,
registry,
device.metal_device(),
proj,
norm_weight,
&normed,
¶ms,
rows,
head_dim,
)
.context("dispatch_rms_norm head_norm")?;
Ok(normed)
}
pub fn dispatch_dflash_v_proj(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
layer: &DFlashLayerTensors,
cfg: &DFlashConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
let hidden = cfg.hidden_size as u32;
let kv_out_dim = (cfg.num_key_value_heads * cfg.head_dim) as u32;
apply_linear_projection_f32(
encoder,
registry,
device,
input,
&layer.v_proj,
seq_len,
hidden,
kv_out_dim,
)
.context("dispatch_dflash_v_proj")
}
pub fn dispatch_dflash_mlp(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
layer: &DFlashLayerTensors,
cfg: &DFlashConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
let hidden = cfg.hidden_size as u32;
let inter = cfg.intermediate_size as u32;
let gate = apply_linear_projection_f32(
encoder,
registry,
device,
input,
&layer.mlp_gate,
seq_len,
hidden,
inter,
)
.context("dispatch_dflash_mlp: gate_proj")?;
let up = apply_linear_projection_f32(
encoder,
registry,
device,
input,
&layer.mlp_up,
seq_len,
hidden,
inter,
)
.context("dispatch_dflash_mlp: up_proj")?;
encoder.memory_barrier();
let n_h = seq_len * inter;
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;
let activated = device
.alloc_buffer(
(n_h as usize) * 4,
DType::F32,
vec![seq_len as usize, inter as usize],
)
.map_err(|e| anyhow!("alloc mlp activated: {e}"))?;
dispatch_silu_mul(
encoder,
registry,
device.metal_device(),
&gate,
&up,
&activated,
&silu_params,
n_h,
)
.context("dispatch_dflash_mlp: silu_mul")?;
encoder.memory_barrier();
apply_linear_projection_f32(
encoder,
registry,
device,
&activated,
&layer.mlp_down,
seq_len,
inter,
hidden,
)
.context("dispatch_dflash_mlp: down_proj")
}
fn build_dflash_pos_buf(device: &MlxDevice, seq_len: u32, offset: u32) -> Result<MlxBuffer> {
let n_pos = 4 * (seq_len as usize);
let mut buf = device
.alloc_buffer(n_pos * 4, DType::I32, vec![n_pos])
.map_err(|e| anyhow!("alloc rope pos_buf: {e}"))?;
let slice = buf
.as_mut_slice::<i32>()
.map_err(|e| anyhow!("rope pos_buf slice: {e}"))?;
let l = seq_len as usize;
let base = offset as i32;
for axis in 0..4 {
let dst = &mut slice[axis * l..(axis + 1) * l];
for (i, v) in dst.iter_mut().enumerate() {
*v = base + (i as i32);
}
}
Ok(buf)
}
pub fn dispatch_dflash_rope(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
qk_in: &MlxBuffer,
cfg: &DFlashConfig,
seq_len: u32,
num_heads: u32,
offset: u32,
) -> Result<MlxBuffer> {
let head_dim = cfg.head_dim as u32;
let rope_dim = head_dim; let freq_base = cfg.rope_theta;
let sections = [head_dim / 2, 0, 0, 0];
let positions = build_dflash_pos_buf(device, seq_len, offset)?;
apply_imrope(
encoder, registry, device, qk_in, &positions, seq_len, num_heads, head_dim, rope_dim,
freq_base, sections,
)
.context("dispatch_dflash_rope")
}
pub fn dispatch_dflash_fc(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
target_hidden_concat: &MlxBuffer,
model: &super::tensors::DFlashModelTensors,
cfg: &DFlashConfig,
ctx_seq_len: u32,
) -> Result<MlxBuffer> {
let fc_in = cfg.fc_input_dim() as u32;
let hidden = cfg.hidden_size as u32;
apply_linear_projection_f32(
encoder,
registry,
device,
target_hidden_concat,
&model.fc,
ctx_seq_len,
fc_in,
hidden,
)
.context("dispatch_dflash_fc")
}
pub fn dispatch_dflash_hidden_norm(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
fc_out: &MlxBuffer,
model: &super::tensors::DFlashModelTensors,
cfg: &DFlashConfig,
ctx_seq_len: u32,
) -> Result<MlxBuffer> {
let hidden = cfg.hidden_size as u32;
let element_count = (ctx_seq_len as usize) * (hidden as usize);
if fc_out.element_count() != element_count {
return Err(anyhow!(
"dflash hidden_norm: fc_out element count {} != S({}) * hidden({})",
fc_out.element_count(),
ctx_seq_len,
hidden
));
}
let normed = device
.alloc_buffer(
element_count * 4,
DType::F32,
vec![ctx_seq_len as usize, hidden as usize],
)
.map_err(|e| anyhow!("alloc hidden_norm output: {e}"))?;
let params = alloc_rms_norm_params(device, cfg.rms_norm_eps, hidden)?;
dispatch_rms_norm(
encoder,
registry,
device.metal_device(),
fc_out,
&model.hidden_norm,
&normed,
¶ms,
ctx_seq_len,
hidden,
)
.context("dispatch_rms_norm hidden_norm")?;
Ok(normed)
}
pub fn dispatch_dflash_final_norm(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
h: &MlxBuffer,
model: &super::tensors::DFlashModelTensors,
cfg: &DFlashConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
let hidden = cfg.hidden_size as u32;
let element_count = (seq_len as usize) * (hidden as usize);
if h.element_count() != element_count {
return Err(anyhow!(
"dflash final_norm: h element count {} != L({}) * hidden({})",
h.element_count(),
seq_len,
hidden
));
}
let normed = device
.alloc_buffer(
element_count * 4,
DType::F32,
vec![seq_len as usize, hidden as usize],
)
.map_err(|e| anyhow!("alloc final_norm output: {e}"))?;
let params = alloc_rms_norm_params(device, cfg.rms_norm_eps, hidden)?;
dispatch_rms_norm(
encoder,
registry,
device.metal_device(),
h,
&model.final_norm,
&normed,
¶ms,
seq_len,
hidden,
)
.context("dispatch_rms_norm final_norm")?;
Ok(normed)
}
pub fn dispatch_dflash_softcap(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
logits: &MlxBuffer,
cfg: &DFlashConfig,
) -> Result<Option<MlxBuffer>> {
let cap = match cfg.final_logit_softcapping {
Some(c) => c,
None => return Ok(None),
};
let n = logits.element_count();
let shape: Vec<usize> = logits.shape().to_vec();
let mut out = device
.alloc_buffer(n * 4, DType::F32, shape)
.map_err(|e| anyhow!("alloc softcap output: {e}"))?;
let mut params = device
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow!("alloc softcap params: {e}"))?;
{
let s = params
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("softcap params slice: {e}"))?;
s[0] = cap;
s[1] = f32::from_bits(n as u32);
}
dispatch_softcap(
encoder,
registry,
device.metal_device(),
logits,
&out,
¶ms,
cap,
)
.context("dispatch_softcap")?;
let _ = registry; Ok(Some({
let _ = &mut out; out
}))
}
pub fn dispatch_dflash_residual_add(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
a: &MlxBuffer,
b: &MlxBuffer,
) -> Result<MlxBuffer> {
let n = a.element_count();
if b.element_count() != n {
return Err(anyhow!(
"dflash residual_add: length mismatch a={} b={}",
n,
b.element_count()
));
}
let shape: Vec<usize> = a.shape().to_vec();
let out = device
.alloc_buffer(n * 4, DType::F32, shape)
.map_err(|e| anyhow!("alloc residual_add output: {e}"))?;
elementwise_add(
encoder,
registry,
device.metal_device(),
a,
b,
&out,
n,
DType::F32,
)
.map_err(|e| anyhow!("elementwise_add: {e}"))?;
Ok(out)
}
pub fn dispatch_dflash_sdpa_self_attn(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
q_roped: &MlxBuffer,
k_roped: &MlxBuffer,
v: &MlxBuffer,
cfg: &DFlashConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
let n_heads = cfg.num_attention_heads as u32;
let n_kv_heads = cfg.num_key_value_heads as u32;
let head_dim = cfg.head_dim as u32;
apply_sdpa_causal_from_seq_major(
encoder, registry, device, q_roped, k_roped, v, seq_len, n_heads, n_kv_heads, head_dim,
)
.context("dispatch_dflash_sdpa_self_attn")
}
pub fn dispatch_dflash_decoder_layer_attention(
registry: &mut KernelRegistry,
device: &MlxDevice,
h: &MlxBuffer,
h_ctx: &MlxBuffer,
layer_weights: &DFlashLayerTensors,
cache_layer: &mut super::kv_cache::DFlashLayerKvCache,
cfg: &DFlashConfig,
layer_idx: usize,
block_size: u32,
ctx_chunk_size: u32,
) -> Result<MlxBuffer> {
let do_causal = matches!(
cfg.layer_types[layer_idx],
super::config::LayerType::SlidingAttention
);
let prior_offset = cache_layer.seq_len;
let n_q = cfg.num_attention_heads as u32;
let n_kv = cfg.num_key_value_heads as u32;
let head_dim = cfg.head_dim as u32;
let l = block_size;
let s = ctx_chunk_size;
let mut enc = device
.command_encoder()
.context("decoder_layer_attention: open encoder A")?;
let normed_h =
dispatch_dflash_input_layernorm(&mut enc, registry, device, h, layer_weights, cfg, l)
.context("layer attn: input_layernorm")?;
enc.memory_barrier();
let q = dispatch_dflash_q_proj(&mut enc, registry, device, &normed_h, layer_weights, cfg, l)
.context("layer attn: q_proj")?;
let k_ctx = dispatch_dflash_k_proj(&mut enc, registry, device, h_ctx, layer_weights, cfg, s)
.context("layer attn: k_ctx_proj")?;
let v_ctx = dispatch_dflash_v_proj(&mut enc, registry, device, h_ctx, layer_weights, cfg, s)
.context("layer attn: v_ctx_proj")?;
let k_prop =
dispatch_dflash_k_proj(&mut enc, registry, device, &normed_h, layer_weights, cfg, l)
.context("layer attn: k_prop_proj")?;
let v_prop =
dispatch_dflash_v_proj(&mut enc, registry, device, &normed_h, layer_weights, cfg, l)
.context("layer attn: v_prop_proj")?;
enc.memory_barrier();
let q_normed = dispatch_dflash_head_norm(
&mut enc,
registry,
device,
&q,
&layer_weights.q_norm,
cfg,
l,
n_q,
)
.context("layer attn: q_norm")?;
let k_ctx_normed = dispatch_dflash_head_norm(
&mut enc,
registry,
device,
&k_ctx,
&layer_weights.k_norm,
cfg,
s,
n_kv,
)
.context("layer attn: k_ctx_norm")?;
let k_prop_normed = dispatch_dflash_head_norm(
&mut enc,
registry,
device,
&k_prop,
&layer_weights.k_norm,
cfg,
l,
n_kv,
)
.context("layer attn: k_prop_norm")?;
enc.memory_barrier();
let q_roped = dispatch_dflash_rope(
&mut enc,
registry,
device,
&q_normed,
cfg,
l,
n_q,
prior_offset + s,
)
.context("layer attn: q rope")?;
let k_ctx_roped = dispatch_dflash_rope(
&mut enc,
registry,
device,
&k_ctx_normed,
cfg,
s,
n_kv,
prior_offset,
)
.context("layer attn: k_ctx rope")?;
let k_prop_roped = dispatch_dflash_rope(
&mut enc,
registry,
device,
&k_prop_normed,
cfg,
l,
n_kv,
prior_offset + s,
)
.context("layer attn: k_prop rope")?;
enc.memory_barrier();
cache_layer
.append_seq_major_kv_gpu(
&mut enc,
registry,
device.metal_device(),
&k_ctx_roped,
&v_ctx,
s,
n_kv,
head_dim,
)
.context("layer attn: append_seq_major_kv_gpu")?;
debug_assert_eq!(cache_layer.seq_len, prior_offset + s);
enc.memory_barrier();
cache_layer
.write_slack_kv_gpu(
&mut enc,
registry,
device.metal_device(),
&k_prop_roped,
&v_prop,
l,
n_kv,
head_dim,
)
.context("layer attn: write_slack_kv_gpu")?;
debug_assert_eq!(
cache_layer.seq_len,
prior_offset + s,
"slack must not advance seq_len"
);
enc.commit_labeled("dflash.decoder_layer_attention.phase_a_kv_writes");
let kv_seq_len = cache_layer.seq_len + l;
let mut sdpa_enc = device
.command_encoder()
.context("decoder_layer_attention: open encoder for sdpa+o_proj")?;
let attn_out = dispatch_dflash_sdpa_cross_length(
&mut sdpa_enc,
registry,
device,
&q_roped,
cache_layer,
cfg,
l,
kv_seq_len,
do_causal,
)
.context("layer attn: sdpa cross-length")?;
let attn_proj = dispatch_dflash_o_proj(
&mut sdpa_enc,
registry,
device,
&attn_out,
layer_weights,
cfg,
l,
)
.context("layer attn: o_proj (fused)")?;
sdpa_enc.commit_labeled("dflash.decoder_layer_attention.sdpa_oproj_fused");
Ok(attn_proj)
}
pub fn dispatch_dflash_model_forward(
registry: &mut KernelRegistry,
device: &MlxDevice,
h: &MlxBuffer,
target_hidden_concat: &MlxBuffer,
model: &super::tensors::DFlashModelTensors,
cache: &mut super::kv_cache::DFlashKvCache,
cfg: &DFlashConfig,
block_size: u32,
ctx_chunk_size: u32,
) -> Result<MlxBuffer> {
let mut enc = device
.command_encoder()
.context("model_forward: open encoder for fc + hidden_norm")?;
let fc_out = dispatch_dflash_fc(
&mut enc,
registry,
device,
target_hidden_concat,
model,
cfg,
ctx_chunk_size,
)
.context("model_forward: fc")?;
enc.memory_barrier();
let h_ctx = dispatch_dflash_hidden_norm(
&mut enc,
registry,
device,
&fc_out,
model,
cfg,
ctx_chunk_size,
)
.context("model_forward: hidden_norm")?;
enc.commit_labeled("dflash.model_forward.prelude_fc_hidden_norm");
let mut h_curr_owned: Option<MlxBuffer> = None;
for layer_idx in 0..cfg.num_hidden_layers {
let h_in: &MlxBuffer = match h_curr_owned.as_ref() {
Some(b) => b,
None => h,
};
let h_out = dispatch_dflash_decoder_layer(
registry,
device,
h_in,
&h_ctx,
&model.layers[layer_idx],
&mut cache.layers[layer_idx],
cfg,
layer_idx,
block_size,
ctx_chunk_size,
)
.with_context(|| format!("model_forward: layer {layer_idx}"))?;
h_curr_owned = Some(h_out);
}
let h_after_layers = h_curr_owned.expect("at least 1 decoder layer in drafter");
let mut enc = device
.command_encoder()
.context("model_forward: open encoder for final_norm")?;
let h_final = dispatch_dflash_final_norm(
&mut enc,
registry,
device,
&h_after_layers,
model,
cfg,
block_size,
)
.context("model_forward: final_norm")?;
enc.commit_and_wait()
.context("model_forward: commit epilogue")?;
Ok(h_final)
}
pub fn dispatch_dflash_decoder_layer(
registry: &mut KernelRegistry,
device: &MlxDevice,
h: &MlxBuffer,
h_ctx: &MlxBuffer,
layer_weights: &DFlashLayerTensors,
cache_layer: &mut super::kv_cache::DFlashLayerKvCache,
cfg: &DFlashConfig,
layer_idx: usize,
block_size: u32,
ctx_chunk_size: u32,
) -> Result<MlxBuffer> {
let attn_proj = dispatch_dflash_decoder_layer_attention(
registry,
device,
h,
h_ctx,
layer_weights,
cache_layer,
cfg,
layer_idx,
block_size,
ctx_chunk_size,
)
.context("decoder_layer: attention sub-block")?;
let mut enc = device
.command_encoder()
.context("decoder_layer: open encoder for residual + MLP")?;
let h_after_attn = dispatch_dflash_residual_add(&mut enc, registry, device, h, &attn_proj)
.context("decoder_layer: residual 1")?;
enc.memory_barrier();
let post_normed = dispatch_dflash_post_attention_layernorm(
&mut enc,
registry,
device,
&h_after_attn,
layer_weights,
cfg,
block_size,
)
.context("decoder_layer: post_attention_layernorm")?;
enc.memory_barrier();
let mlp_out = dispatch_dflash_mlp(
&mut enc,
registry,
device,
&post_normed,
layer_weights,
cfg,
block_size,
)
.context("decoder_layer: mlp")?;
enc.memory_barrier();
let h_out = dispatch_dflash_residual_add(&mut enc, registry, device, &h_after_attn, &mlp_out)
.context("decoder_layer: residual 2")?;
enc.commit_labeled("dflash.decoder_layer.residual_mlp");
Ok(h_out)
}
pub fn dispatch_dflash_sdpa_cross_length(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
q_seq_major: &MlxBuffer,
cache_layer: &super::kv_cache::DFlashLayerKvCache,
cfg: &DFlashConfig,
q_seq_len: u32,
kv_seq_len: u32,
do_causal: bool,
) -> Result<MlxBuffer> {
let n_heads = cfg.num_attention_heads as u32;
let n_kv_heads = cfg.num_key_value_heads as u32;
let head_dim = cfg.head_dim as u32;
let scale = 1.0f32 / (head_dim as f32).sqrt();
if kv_seq_len > cache_layer.capacity {
return Err(anyhow!(
"sdpa cross-length: kv_seq_len {} > cache.capacity {}",
kv_seq_len,
cache_layer.capacity
));
}
let l = q_seq_len as usize;
let nh = n_heads as usize;
let d = head_dim as usize;
let out_elem = nh * l * d;
let q_hm = device
.alloc_buffer(out_elem * 4, DType::F32, vec![nh, l, d])
.map_err(|e| anyhow!("alloc q_hm: {e}"))?;
let out_hm = device
.alloc_buffer(out_elem * 4, DType::F32, vec![1, nh, l, d])
.map_err(|e| anyhow!("alloc sdpa cross-length output_hm: {e}"))?;
let out_sm = device
.alloc_buffer(out_elem * 4, DType::F32, vec![l, nh, d])
.map_err(|e| anyhow!("alloc sdpa cross-length output_sm: {e}"))?;
permute_021_f32(
encoder,
registry,
device.metal_device(),
q_seq_major,
&q_hm,
l,
nh,
d,
)
.map_err(|e| anyhow!("permute_021_f32 q seq->hm: {e}"))?;
encoder.memory_barrier();
let params = SdpaParams {
n_heads,
n_kv_heads,
head_dim,
seq_len: q_seq_len,
kv_seq_len,
scale,
kv_capacity: cache_layer.capacity,
do_causal,
};
sdpa(
encoder,
registry,
device,
&q_hm,
&cache_layer.keys,
&cache_layer.values,
&out_hm,
¶ms,
1,
)
.map_err(|e| anyhow!("sdpa cross-length: {e}"))?;
encoder.memory_barrier();
permute_021_f32(
encoder,
registry,
device.metal_device(),
&out_hm,
&out_sm,
nh,
l,
d,
)
.map_err(|e| anyhow!("permute_021_f32 out hm->seq: {e}"))?;
encoder.memory_barrier();
Ok(out_sm)
}
pub fn dispatch_dflash_o_proj(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
attn_out: &MlxBuffer,
layer: &DFlashLayerTensors,
cfg: &DFlashConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
let hidden = cfg.hidden_size as u32;
let q_in_dim = (cfg.num_attention_heads * cfg.head_dim) as u32;
apply_linear_projection_f32(
encoder,
registry,
device,
attn_out,
&layer.o_proj,
seq_len,
q_in_dim,
hidden,
)
.context("dispatch_dflash_o_proj")
}
pub fn dispatch_dflash_post_attention_layernorm(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
h: &MlxBuffer,
layer: &DFlashLayerTensors,
cfg: &DFlashConfig,
seq_len: u32,
) -> Result<MlxBuffer> {
let hidden = cfg.hidden_size as u32;
let element_count = (seq_len as usize) * (hidden as usize);
if h.element_count() != element_count {
return Err(anyhow!(
"dflash post_attention_layernorm: h element count {} != L({}) * hidden({})",
h.element_count(),
seq_len,
hidden
));
}
let normed = device
.alloc_buffer(
element_count * 4,
DType::F32,
vec![seq_len as usize, hidden as usize],
)
.map_err(|e| anyhow!("alloc post_attention_layernorm output: {e}"))?;
let params = alloc_rms_norm_params(device, cfg.rms_norm_eps, hidden)?;
dispatch_rms_norm(
encoder,
registry,
device.metal_device(),
h,
&layer.post_attention_layernorm,
&normed,
¶ms,
seq_len,
hidden,
)
.context("dispatch_rms_norm post_attention_layernorm")?;
Ok(normed)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::inference::spec_decode::dflash::{
config::DFlashConfig,
tensors::DFlashModelTensors,
weights::{DFlashWeights, DFlashWeightsFile},
};
fn gemma4_26b_a4b_dflash_config() -> DFlashConfig {
DFlashConfig::from_json_str(
crate::inference::spec_decode::dflash::config::tests::GEMMA4_26B_A4B_DFLASH_CONFIG,
)
.expect("test fixture must parse")
}
#[test]
#[ignore = "requires Metal device + drafter HF cache"]
fn smoke_input_norm_and_q_proj() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = gemma4_26b_a4b_dflash_config();
let device = MlxDevice::new().expect("Metal device available on M5 Max");
let mut registry = KernelRegistry::new();
let home = std::env::var("HOME").expect("HOME set");
let path = format!(
"{home}/.cache/huggingface/hub/models--z-lab--gemma-4-26B-A4B-it-DFlash/snapshots/77d4202772dfe50b2396ec7bac9cfffc7b9e7057/model.safetensors"
);
let file = DFlashWeightsFile::open(&path).expect("file open");
let weights = DFlashWeights::load(file.bytes(), &cfg).expect("validated load");
let tensors = DFlashModelTensors::upload(&device, &cfg, &weights).expect("GPU upload");
let block_size = 8u32; let hidden = cfg.hidden_size as u32;
let elem = (block_size as usize) * (hidden as usize);
let mut h_input = device
.alloc_buffer(
elem * 4,
DType::F32,
vec![block_size as usize, hidden as usize],
)
.expect("alloc h_input");
{
let slice = h_input.as_mut_slice::<f32>().expect("h_input slice");
for v in slice.iter_mut() {
*v = 1.0;
}
}
let mut encoder = device.command_encoder().expect("encoder");
let layer = &tensors.layers[0];
let normed = dispatch_dflash_input_layernorm(
&mut encoder,
&mut registry,
&device,
&h_input,
layer,
&cfg,
block_size,
)
.expect("input_layernorm dispatch");
encoder.memory_barrier(); let q = dispatch_dflash_q_proj(
&mut encoder,
&mut registry,
&device,
&normed,
layer,
&cfg,
block_size,
)
.expect("q_proj dispatch");
let k = dispatch_dflash_k_proj(
&mut encoder,
&mut registry,
&device,
&normed,
layer,
&cfg,
block_size,
)
.expect("k_proj dispatch");
let v = dispatch_dflash_v_proj(
&mut encoder,
&mut registry,
&device,
&normed,
layer,
&cfg,
block_size,
)
.expect("v_proj dispatch");
encoder.memory_barrier();
let q_normed = dispatch_dflash_head_norm(
&mut encoder,
&mut registry,
&device,
&q,
&layer.q_norm,
&cfg,
block_size,
cfg.num_attention_heads as u32,
)
.expect("q_norm dispatch");
let k_normed = dispatch_dflash_head_norm(
&mut encoder,
&mut registry,
&device,
&k,
&layer.k_norm,
&cfg,
block_size,
cfg.num_key_value_heads as u32,
)
.expect("k_norm dispatch");
encoder.memory_barrier(); let q_roped = dispatch_dflash_rope(
&mut encoder,
&mut registry,
&device,
&q_normed,
&cfg,
block_size,
cfg.num_attention_heads as u32,
0,
)
.expect("q rope dispatch");
let k_roped = dispatch_dflash_rope(
&mut encoder,
&mut registry,
&device,
&k_normed,
&cfg,
block_size,
cfg.num_key_value_heads as u32,
0,
)
.expect("k rope dispatch");
encoder.memory_barrier(); let attn_out = dispatch_dflash_sdpa_self_attn(
&mut encoder,
&mut registry,
&device,
&q_roped,
&k_roped,
&v,
&cfg,
block_size,
)
.expect("sdpa dispatch");
let mut encoder = device.command_encoder().expect("encoder2");
let h_out = dispatch_dflash_o_proj(
&mut encoder,
&mut registry,
&device,
&attn_out,
layer,
&cfg,
block_size,
)
.expect("o_proj dispatch");
encoder.commit_and_wait().expect("commit2");
let q_dim = (cfg.num_attention_heads * cfg.head_dim) as usize;
let kv_dim = (cfg.num_key_value_heads * cfg.head_dim) as usize;
let l = block_size as usize;
assert_eq!(q.element_count(), l * q_dim);
assert_eq!(k.element_count(), l * kv_dim);
assert_eq!(v.element_count(), l * kv_dim);
assert_eq!(q_normed.element_count(), l * q_dim);
assert_eq!(k_normed.element_count(), l * kv_dim);
assert_eq!(q_roped.element_count(), l * q_dim);
assert_eq!(k_roped.element_count(), l * kv_dim);
assert_eq!(attn_out.element_count(), l * q_dim);
let h_dim = cfg.hidden_size;
assert_eq!(h_out.element_count(), l * h_dim);
for (name, buf, dim) in [
("Q", &q, q_dim),
("K", &k, kv_dim),
("V", &v, kv_dim),
("Q_normed", &q_normed, q_dim),
("K_normed", &k_normed, kv_dim),
("Q_roped", &q_roped, q_dim),
("K_roped", &k_roped, kv_dim),
("Attn_out", &attn_out, q_dim),
("H_out", &h_out, h_dim),
] {
let host: &[f32] = buf.as_slice::<f32>().expect("host slice");
assert_eq!(host.len(), l * dim, "{name} length");
let n_finite = host.iter().filter(|v| v.is_finite()).count();
let n_zero = host.iter().filter(|v| **v == 0.0).count();
assert_eq!(
n_finite,
host.len(),
"{name}: all values must be finite (got {n_finite}/{}; n_zero={n_zero})",
host.len()
);
assert!(
n_zero < host.len() / 2,
"{name} suspiciously sparse (n_zero={n_zero}/{})",
host.len()
);
}
}
#[test]
#[ignore = "requires Metal device + drafter HF cache"]
fn smoke_model_forward_with_cache() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::kv_cache::DFlashKvCache;
let cfg = gemma4_26b_a4b_dflash_config();
let device = MlxDevice::new().expect("Metal device available on M5 Max");
let mut registry = KernelRegistry::new();
let home = std::env::var("HOME").expect("HOME set");
let path = format!(
"{home}/.cache/huggingface/hub/models--z-lab--gemma-4-26B-A4B-it-DFlash/snapshots/77d4202772dfe50b2396ec7bac9cfffc7b9e7057/model.safetensors"
);
let file = DFlashWeightsFile::open(&path).expect("file open");
let weights = DFlashWeights::load(file.bytes(), &cfg).expect("validated load");
let tensors = DFlashModelTensors::upload(&device, &cfg, &weights).expect("GPU upload");
let mut cache = DFlashKvCache::new(&device, &cfg, 128).expect("cache");
let block_size = 8u32;
let ctx_chunk = 4u32;
let hidden = cfg.hidden_size as u32;
let fc_in = cfg.fc_input_dim() as u32;
let h_elem = (block_size as usize) * (hidden as usize);
let mut h = device
.alloc_buffer(
h_elem * 4,
DType::F32,
vec![block_size as usize, hidden as usize],
)
.expect("alloc h");
{
let s = h.as_mut_slice::<f32>().expect("h slice");
for v in s.iter_mut() {
*v = 1.0;
}
}
let thc_elem = (ctx_chunk as usize) * (fc_in as usize);
let mut target_hidden = device
.alloc_buffer(
thc_elem * 4,
DType::F32,
vec![ctx_chunk as usize, fc_in as usize],
)
.expect("alloc target_hidden");
{
let s = target_hidden
.as_mut_slice::<f32>()
.expect("target_hidden slice");
for (i, v) in s.iter_mut().enumerate() {
*v = 0.1 + ((i % 17) as f32) / 170.0;
}
}
let initial_seq_lens: Vec<u32> = cache.layers.iter().map(|l| l.seq_len).collect();
let h_final = dispatch_dflash_model_forward(
&mut registry,
&device,
&h,
&target_hidden,
&tensors,
&mut cache,
&cfg,
block_size,
ctx_chunk,
)
.expect("model forward");
assert_eq!(
h_final.element_count(),
(block_size as usize) * (hidden as usize)
);
let host: &[f32] = h_final.as_slice::<f32>().expect("h_final slice");
let n_finite = host.iter().filter(|v| v.is_finite()).count();
let n_zero = host.iter().filter(|v| **v == 0.0).count();
assert_eq!(
n_finite,
host.len(),
"model forward output must be all finite (got {n_finite}/{}; n_zero={n_zero})",
host.len()
);
assert!(
n_zero < host.len() / 2,
"model forward output suspiciously sparse (n_zero={n_zero}/{})",
host.len()
);
for (i, l) in cache.layers.iter().enumerate() {
assert_eq!(
l.seq_len,
initial_seq_lens[i] + ctx_chunk,
"layer {i} cache should advance by ctx_chunk_size"
);
}
}
#[test]
#[ignore = "requires Metal device + drafter HF cache"]
fn smoke_decoder_layer_full_with_cache() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::kv_cache::DFlashKvCache;
let cfg = gemma4_26b_a4b_dflash_config();
let device = MlxDevice::new().expect("Metal device available on M5 Max");
let mut registry = KernelRegistry::new();
let home = std::env::var("HOME").expect("HOME set");
let path = format!(
"{home}/.cache/huggingface/hub/models--z-lab--gemma-4-26B-A4B-it-DFlash/snapshots/77d4202772dfe50b2396ec7bac9cfffc7b9e7057/model.safetensors"
);
let file = DFlashWeightsFile::open(&path).expect("file open");
let weights = DFlashWeights::load(file.bytes(), &cfg).expect("validated load");
let tensors = DFlashModelTensors::upload(&device, &cfg, &weights).expect("GPU upload");
let mut cache = DFlashKvCache::new(&device, &cfg, 64).expect("cache");
let block_size = 8u32;
let ctx_chunk = 4u32;
let hidden = cfg.hidden_size as u32;
let h_elem = (block_size as usize) * (hidden as usize);
let mut h = device
.alloc_buffer(
h_elem * 4,
DType::F32,
vec![block_size as usize, hidden as usize],
)
.expect("alloc h");
{
let s = h.as_mut_slice::<f32>().expect("h slice");
for v in s.iter_mut() {
*v = 1.0;
}
}
let hctx_elem = (ctx_chunk as usize) * (hidden as usize);
let mut h_ctx = device
.alloc_buffer(
hctx_elem * 4,
DType::F32,
vec![ctx_chunk as usize, hidden as usize],
)
.expect("alloc h_ctx");
{
let s = h_ctx.as_mut_slice::<f32>().expect("h_ctx slice");
for (i, v) in s.iter_mut().enumerate() {
*v = 0.5 + ((i % 31) as f32) / 31.0;
}
}
let layer_idx = 4usize; let initial_seq_len = cache.layers[layer_idx].seq_len;
let h_out = dispatch_dflash_decoder_layer(
&mut registry,
&device,
&h,
&h_ctx,
&tensors.layers[layer_idx],
&mut cache.layers[layer_idx],
&cfg,
layer_idx,
block_size,
ctx_chunk,
)
.expect("decoder layer forward");
assert_eq!(
h_out.element_count(),
(block_size as usize) * (hidden as usize)
);
let host: &[f32] = h_out.as_slice::<f32>().expect("h_out slice");
let n_finite = host.iter().filter(|v| v.is_finite()).count();
let n_zero = host.iter().filter(|v| **v == 0.0).count();
assert_eq!(
n_finite,
host.len(),
"decoder layer forward output must be all finite (got {n_finite}/{}; n_zero={n_zero})",
host.len()
);
assert!(
n_zero < host.len() / 2,
"decoder layer forward output suspiciously sparse (n_zero={n_zero}/{})",
host.len()
);
assert_eq!(
cache.layers[layer_idx].seq_len,
initial_seq_len + ctx_chunk,
"cache should advance by ctx_chunk only after full layer forward"
);
}
#[test]
#[ignore = "requires Metal device + drafter HF cache"]
fn smoke_decoder_layer_attention_with_cache() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::kv_cache::DFlashKvCache;
let cfg = gemma4_26b_a4b_dflash_config();
let device = MlxDevice::new().expect("Metal device available on M5 Max");
let mut registry = KernelRegistry::new();
let home = std::env::var("HOME").expect("HOME set");
let path = format!(
"{home}/.cache/huggingface/hub/models--z-lab--gemma-4-26B-A4B-it-DFlash/snapshots/77d4202772dfe50b2396ec7bac9cfffc7b9e7057/model.safetensors"
);
let file = DFlashWeightsFile::open(&path).expect("file open");
let weights = DFlashWeights::load(file.bytes(), &cfg).expect("validated load");
let tensors = DFlashModelTensors::upload(&device, &cfg, &weights).expect("GPU upload");
let mut cache = DFlashKvCache::new(&device, &cfg, 64).expect("cache");
let block_size = 8u32;
let ctx_chunk = 4u32;
let hidden = cfg.hidden_size as u32;
let h_elem = (block_size as usize) * (hidden as usize);
let mut h = device
.alloc_buffer(
h_elem * 4,
DType::F32,
vec![block_size as usize, hidden as usize],
)
.expect("alloc h");
{
let s = h.as_mut_slice::<f32>().expect("h slice");
for v in s.iter_mut() {
*v = 1.0;
}
}
let hctx_elem = (ctx_chunk as usize) * (hidden as usize);
let mut h_ctx = device
.alloc_buffer(
hctx_elem * 4,
DType::F32,
vec![ctx_chunk as usize, hidden as usize],
)
.expect("alloc h_ctx");
{
let s = h_ctx.as_mut_slice::<f32>().expect("h_ctx slice");
for (i, v) in s.iter_mut().enumerate() {
*v = 0.5 + ((i % 31) as f32) / 31.0;
}
}
let layer_idx = 4usize;
let initial_seq_len = cache.layers[layer_idx].seq_len;
let attn_out = dispatch_dflash_decoder_layer_attention(
&mut registry,
&device,
&h,
&h_ctx,
&tensors.layers[layer_idx],
&mut cache.layers[layer_idx],
&cfg,
layer_idx,
block_size,
ctx_chunk,
)
.expect("decoder layer attention");
assert_eq!(
attn_out.element_count(),
(block_size as usize) * (hidden as usize)
);
let host: &[f32] = attn_out.as_slice::<f32>().expect("attn_out slice");
let n_finite = host.iter().filter(|v| v.is_finite()).count();
let n_zero = host.iter().filter(|v| **v == 0.0).count();
assert_eq!(
n_finite,
host.len(),
"decoder layer attention output must be finite (got {n_finite}/{}; n_zero={n_zero})",
host.len()
);
assert!(
n_zero < host.len() / 2,
"decoder layer attention output suspiciously sparse (n_zero={n_zero}/{})",
host.len()
);
assert_eq!(
cache.layers[layer_idx].seq_len,
initial_seq_len + ctx_chunk,
"cache should advance by ctx_chunk_size only; prop lives in slack"
);
}
#[test]
#[ignore = "requires Metal device + drafter HF cache"]
fn smoke_sdpa_cross_length() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::kv_cache::DFlashKvCache;
let cfg = gemma4_26b_a4b_dflash_config();
let device = MlxDevice::new().expect("Metal device available on M5 Max");
let mut registry = KernelRegistry::new();
let mut cache = DFlashKvCache::new(&device, &cfg, 64).expect("cache");
let layer_idx = 4usize; let n_kv = cfg.num_key_value_heads as u32;
let d = cfg.head_dim as u32;
let ctx_len = 12u32;
let total = (ctx_len as usize) * (n_kv as usize) * (d as usize);
let mut k_input = vec![0.0f32; total];
let mut v_input = vec![0.0f32; total];
for (i, v) in k_input.iter_mut().enumerate() {
*v = (i % 17) as f32 / 17.0; }
for (i, v) in v_input.iter_mut().enumerate() {
*v = (i % 23) as f32 / 23.0;
}
cache.layers[layer_idx]
.append_seq_major_kv(&k_input, &v_input, ctx_len, n_kv, d)
.expect("populate cache");
assert_eq!(cache.layers[layer_idx].seq_len, ctx_len);
let q_seq_len = 8u32;
let n_q = cfg.num_attention_heads as u32;
let q_elem = (q_seq_len as usize) * (n_q as usize) * (d as usize);
let mut q = device
.alloc_buffer(
q_elem * 4,
DType::F32,
vec![q_seq_len as usize, n_q as usize, d as usize],
)
.expect("alloc q");
{
let s = q.as_mut_slice::<f32>().expect("q slice");
for (i, v) in s.iter_mut().enumerate() {
*v = ((i % 11) as f32) * 0.01;
}
}
let mut encoder = device.command_encoder().expect("encoder");
let out = dispatch_dflash_sdpa_cross_length(
&mut encoder,
&mut registry,
&device,
&q,
&cache.layers[layer_idx],
&cfg,
q_seq_len,
ctx_len,
true,
)
.expect("sdpa cross-length");
let expected_elem = (q_seq_len as usize) * (n_q as usize) * (d as usize);
assert_eq!(out.element_count(), expected_elem);
let host: &[f32] = out.as_slice::<f32>().expect("out slice");
let n_finite = host.iter().filter(|v| v.is_finite()).count();
assert_eq!(
n_finite,
host.len(),
"cross-length SDPA output must be all finite (got {n_finite}/{})",
host.len()
);
}
#[test]
#[ignore = "requires Metal device + drafter HF cache"]
fn smoke_model_level_globals() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = gemma4_26b_a4b_dflash_config();
let device = MlxDevice::new().expect("Metal device available on M5 Max");
let mut registry = KernelRegistry::new();
let home = std::env::var("HOME").expect("HOME set");
let path = format!(
"{home}/.cache/huggingface/hub/models--z-lab--gemma-4-26B-A4B-it-DFlash/snapshots/77d4202772dfe50b2396ec7bac9cfffc7b9e7057/model.safetensors"
);
let file = DFlashWeightsFile::open(&path).expect("file open");
let weights = DFlashWeights::load(file.bytes(), &cfg).expect("validated load");
let tensors = DFlashModelTensors::upload(&device, &cfg, &weights).expect("GPU upload");
let ctx_seq_len = 4u32;
let fc_in = cfg.fc_input_dim() as u32;
let target_hidden_elem = (ctx_seq_len as usize) * (fc_in as usize);
let mut target_hidden = device
.alloc_buffer(
target_hidden_elem * 4,
DType::F32,
vec![ctx_seq_len as usize, fc_in as usize],
)
.expect("alloc target_hidden");
{
let s = target_hidden
.as_mut_slice::<f32>()
.expect("target_hidden slice");
for v in s.iter_mut() {
*v = 1.0;
}
}
let mut encoder = device.command_encoder().expect("encoder1");
let fc_out = dispatch_dflash_fc(
&mut encoder,
&mut registry,
&device,
&target_hidden,
&tensors,
&cfg,
ctx_seq_len,
)
.expect("fc dispatch");
encoder.memory_barrier();
let h_ctx = dispatch_dflash_hidden_norm(
&mut encoder,
&mut registry,
&device,
&fc_out,
&tensors,
&cfg,
ctx_seq_len,
)
.expect("hidden_norm dispatch");
let block_size = 8u32;
let hidden = cfg.hidden_size as u32;
let h_elem = (block_size as usize) * (hidden as usize);
let mut h = device
.alloc_buffer(
h_elem * 4,
DType::F32,
vec![block_size as usize, hidden as usize],
)
.expect("alloc h");
{
let s = h.as_mut_slice::<f32>().expect("h slice");
for v in s.iter_mut() {
*v = 1.0;
}
}
let h_final = dispatch_dflash_final_norm(
&mut encoder,
&mut registry,
&device,
&h,
&tensors,
&cfg,
block_size,
)
.expect("final_norm dispatch");
let logits_n = (block_size as usize) * 256; let mut logits = device
.alloc_buffer(logits_n * 4, DType::F32, vec![block_size as usize, 256])
.expect("alloc logits");
{
let s = logits.as_mut_slice::<f32>().expect("logits slice");
for (i, v) in s.iter_mut().enumerate() {
*v = if i % 17 == 0 { 100.0 } else { 5.0 };
}
}
let capped = dispatch_dflash_softcap(&mut encoder, &mut registry, &device, &logits, &cfg)
.expect("softcap dispatch")
.expect("softcap should be Some (cfg has final_logit_softcapping=30.0)");
encoder.commit_and_wait().expect("commit");
let h_dim = cfg.hidden_size;
assert_eq!(fc_out.element_count(), (ctx_seq_len as usize) * h_dim);
assert_eq!(h_ctx.element_count(), (ctx_seq_len as usize) * h_dim);
assert_eq!(h_final.element_count(), (block_size as usize) * h_dim);
assert_eq!(capped.element_count(), logits_n);
for (name, buf) in [
("fc_out", &fc_out),
("h_ctx", &h_ctx),
("h_final", &h_final),
("softcap_out", &capped),
] {
let host: &[f32] = buf.as_slice::<f32>().expect("host slice");
let n_finite = host.iter().filter(|v| v.is_finite()).count();
assert_eq!(
n_finite,
host.len(),
"{name}: all values must be finite (got {n_finite}/{})",
host.len()
);
}
let cap = cfg.final_logit_softcapping.unwrap();
let s_host: &[f32] = capped.as_slice::<f32>().expect("capped slice");
let max_abs = s_host.iter().map(|v| v.abs()).fold(0.0f32, f32::max);
assert!(
max_abs < cap,
"softcap output |v|.max = {max_abs} >= cap {cap}; softcap kernel broken"
);
let near_cap = s_host.iter().filter(|v| **v > cap - 0.5).count();
assert!(
near_cap > 0,
"expected some outputs near cap; got max_abs={max_abs}"
);
}
#[test]
#[ignore = "requires Metal device + drafter HF cache"]
fn smoke_decoder_layer_self_attn() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = gemma4_26b_a4b_dflash_config();
let device = MlxDevice::new().expect("Metal device available on M5 Max");
let mut registry = KernelRegistry::new();
let home = std::env::var("HOME").expect("HOME set");
let path = format!(
"{home}/.cache/huggingface/hub/models--z-lab--gemma-4-26B-A4B-it-DFlash/snapshots/77d4202772dfe50b2396ec7bac9cfffc7b9e7057/model.safetensors"
);
let file = DFlashWeightsFile::open(&path).expect("file open");
let weights = DFlashWeights::load(file.bytes(), &cfg).expect("validated load");
let tensors = DFlashModelTensors::upload(&device, &cfg, &weights).expect("GPU upload");
let block_size = 8u32;
let hidden = cfg.hidden_size as u32;
let elem = (block_size as usize) * (hidden as usize);
let mut h0 = device
.alloc_buffer(
elem * 4,
DType::F32,
vec![block_size as usize, hidden as usize],
)
.expect("alloc h0");
{
let s = h0.as_mut_slice::<f32>().expect("h0 slice");
for v in s.iter_mut() {
*v = 1.0;
}
}
let layer = &tensors.layers[0];
let l = block_size as usize;
let q_dim = (cfg.num_attention_heads * cfg.head_dim) as usize;
let mut encoder = device.command_encoder().expect("encoder1");
let h0_normed = dispatch_dflash_input_layernorm(
&mut encoder,
&mut registry,
&device,
&h0,
layer,
&cfg,
block_size,
)
.expect("input_norm");
encoder.memory_barrier();
let q = dispatch_dflash_q_proj(
&mut encoder,
&mut registry,
&device,
&h0_normed,
layer,
&cfg,
block_size,
)
.expect("q_proj");
let k = dispatch_dflash_k_proj(
&mut encoder,
&mut registry,
&device,
&h0_normed,
layer,
&cfg,
block_size,
)
.expect("k_proj");
let v = dispatch_dflash_v_proj(
&mut encoder,
&mut registry,
&device,
&h0_normed,
layer,
&cfg,
block_size,
)
.expect("v_proj");
encoder.memory_barrier();
let q_normed = dispatch_dflash_head_norm(
&mut encoder,
&mut registry,
&device,
&q,
&layer.q_norm,
&cfg,
block_size,
cfg.num_attention_heads as u32,
)
.expect("q_norm");
let k_normed = dispatch_dflash_head_norm(
&mut encoder,
&mut registry,
&device,
&k,
&layer.k_norm,
&cfg,
block_size,
cfg.num_key_value_heads as u32,
)
.expect("k_norm");
encoder.memory_barrier();
let q_roped = dispatch_dflash_rope(
&mut encoder,
&mut registry,
&device,
&q_normed,
&cfg,
block_size,
cfg.num_attention_heads as u32,
0,
)
.expect("q rope");
let k_roped = dispatch_dflash_rope(
&mut encoder,
&mut registry,
&device,
&k_normed,
&cfg,
block_size,
cfg.num_key_value_heads as u32,
0,
)
.expect("k rope");
encoder.memory_barrier();
let attn_out = dispatch_dflash_sdpa_self_attn(
&mut encoder,
&mut registry,
&device,
&q_roped,
&k_roped,
&v,
&cfg,
block_size,
)
.expect("sdpa");
let mut encoder = device.command_encoder().expect("encoder2");
let h_after_attn_proj = dispatch_dflash_o_proj(
&mut encoder,
&mut registry,
&device,
&attn_out,
layer,
&cfg,
block_size,
)
.expect("o_proj");
encoder.memory_barrier();
let h_after_attn = dispatch_dflash_residual_add(
&mut encoder,
&mut registry,
&device,
&h0,
&h_after_attn_proj,
)
.expect("residual1");
encoder.memory_barrier();
let post_normed = dispatch_dflash_post_attention_layernorm(
&mut encoder,
&mut registry,
&device,
&h_after_attn,
layer,
&cfg,
block_size,
)
.expect("post_norm");
encoder.memory_barrier();
let mlp_out = dispatch_dflash_mlp(
&mut encoder,
&mut registry,
&device,
&post_normed,
layer,
&cfg,
block_size,
)
.expect("mlp");
encoder.memory_barrier();
let h_out = dispatch_dflash_residual_add(
&mut encoder,
&mut registry,
&device,
&h_after_attn,
&mlp_out,
)
.expect("residual2");
encoder.commit_and_wait().expect("final commit");
let h_dim = cfg.hidden_size;
assert_eq!(h_out.element_count(), l * h_dim);
assert_eq!(h_after_attn.element_count(), l * h_dim);
let host_final: &[f32] = h_out.as_slice::<f32>().expect("h_out slice");
let n_finite = host_final.iter().filter(|v| v.is_finite()).count();
let n_zero = host_final.iter().filter(|v| **v == 0.0).count();
assert_eq!(
n_finite,
host_final.len(),
"decoder layer output must be all finite (got {n_finite}/{}; n_zero={n_zero})",
host_final.len()
);
assert!(
n_zero < host_final.len() / 2,
"decoder layer output suspiciously sparse (n_zero={n_zero}/{})",
host_final.len()
);
let _ = q_dim; }
#[test]
#[ignore = "requires Metal device + drafter HF cache"]
fn smoke_o_proj() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = gemma4_26b_a4b_dflash_config();
let device = MlxDevice::new().expect("Metal device available on M5 Max");
let mut registry = KernelRegistry::new();
let home = std::env::var("HOME").expect("HOME set");
let path = format!(
"{home}/.cache/huggingface/hub/models--z-lab--gemma-4-26B-A4B-it-DFlash/snapshots/77d4202772dfe50b2396ec7bac9cfffc7b9e7057/model.safetensors"
);
let file = DFlashWeightsFile::open(&path).expect("file open");
let weights = DFlashWeights::load(file.bytes(), &cfg).expect("validated load");
let tensors = DFlashModelTensors::upload(&device, &cfg, &weights).expect("GPU upload");
let block_size = 8u32;
let q_dim = (cfg.num_attention_heads * cfg.head_dim) as u32; let elem = (block_size as usize) * (q_dim as usize);
let mut attn_out = device
.alloc_buffer(
elem * 4,
DType::F32,
vec![block_size as usize, q_dim as usize],
)
.expect("alloc attn_out");
{
let s = attn_out.as_mut_slice::<f32>().expect("attn_out slice");
for v in s.iter_mut() {
*v = 1.0;
}
}
let mut encoder = device.command_encoder().expect("encoder");
let layer = &tensors.layers[0];
let h_out = dispatch_dflash_o_proj(
&mut encoder,
&mut registry,
&device,
&attn_out,
layer,
&cfg,
block_size,
)
.expect("o_proj dispatch");
encoder.commit_and_wait().expect("commit");
let h_dim = cfg.hidden_size;
assert_eq!(h_out.element_count(), (block_size as usize) * h_dim);
let host: &[f32] = h_out.as_slice::<f32>().expect("h_out slice");
let n_finite = host.iter().filter(|v| v.is_finite()).count();
let n_zero = host.iter().filter(|v| **v == 0.0).count();
assert_eq!(
n_finite,
host.len(),
"O projection output must be all finite (got {n_finite}/{}; n_zero={n_zero})",
host.len()
);
assert!(
n_zero < host.len() / 2,
"O projection output suspiciously sparse (n_zero={n_zero}/{})",
host.len()
);
}
#[test]
#[ignore = "requires Metal device + drafter HF cache"]
fn smoke_post_norm_and_mlp() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = gemma4_26b_a4b_dflash_config();
let device = MlxDevice::new().expect("Metal device available on M5 Max");
let mut registry = KernelRegistry::new();
let home = std::env::var("HOME").expect("HOME set");
let path = format!(
"{home}/.cache/huggingface/hub/models--z-lab--gemma-4-26B-A4B-it-DFlash/snapshots/77d4202772dfe50b2396ec7bac9cfffc7b9e7057/model.safetensors"
);
let file = DFlashWeightsFile::open(&path).expect("file open");
let weights = DFlashWeights::load(file.bytes(), &cfg).expect("validated load");
let tensors = DFlashModelTensors::upload(&device, &cfg, &weights).expect("GPU upload");
let block_size = 8u32;
let hidden = cfg.hidden_size as u32;
let elem = (block_size as usize) * (hidden as usize);
let mut h = device
.alloc_buffer(
elem * 4,
DType::F32,
vec![block_size as usize, hidden as usize],
)
.expect("alloc h");
{
let slice = h.as_mut_slice::<f32>().expect("h slice");
for v in slice.iter_mut() {
*v = 1.0;
}
}
let mut encoder = device.command_encoder().expect("encoder");
let layer = &tensors.layers[0];
let post_normed = dispatch_dflash_post_attention_layernorm(
&mut encoder,
&mut registry,
&device,
&h,
layer,
&cfg,
block_size,
)
.expect("post_norm");
encoder.memory_barrier();
let mlp_out = dispatch_dflash_mlp(
&mut encoder,
&mut registry,
&device,
&post_normed,
layer,
&cfg,
block_size,
)
.expect("mlp");
encoder.commit_and_wait().expect("commit");
let h_dim = cfg.hidden_size;
assert_eq!(mlp_out.element_count(), (block_size as usize) * h_dim);
let host: &[f32] = mlp_out.as_slice::<f32>().expect("mlp_out host slice");
let n_finite = host.iter().filter(|v| v.is_finite()).count();
let n_zero = host.iter().filter(|v| **v == 0.0).count();
assert_eq!(
n_finite,
host.len(),
"MLP output must be all finite (got {n_finite}/{}; n_zero={n_zero})",
host.len()
);
assert!(
n_zero < host.len() / 2,
"MLP output suspiciously sparse (n_zero={n_zero}/{})",
host.len()
);
}
}