use anyhow::{anyhow, Context, Result};
use mlx_native::ops::chunk_gated_delta_rule::{
dispatch_chunk_gated_delta_rule_fwd, dispatch_chunk_gated_delta_rule_fwd_with_arena,
ChunkGatedDeltaRuleParams, ChunkInternalArena, FIXED_BT,
};
use mlx_native::ops::compute_g_beta::dispatch_compute_g_beta;
use mlx_native::ops::dense_mm_bf16::{dense_matmul_bf16_f32_tensor, DenseMmBf16F32Params};
use mlx_native::ops::elementwise::{cast, scalar_mul_f32, CastDirection};
use mlx_native::ops::gated_delta_net::{
build_gated_delta_net_params, dispatch_gated_delta_net, GatedDeltaNetParams,
};
use mlx_native::ops::gated_delta_net_decode::dispatch_gated_delta_net_decode;
use mlx_native::ops::l2_norm::dispatch_l2_norm;
use mlx_native::ops::qkv_split::{dispatch_qkv_split_f32, QkvSplitParams};
use mlx_native::ops::quantized_matmul_ggml::{
quantized_matmul_ggml, GgmlQuantizedMatmulParams, GgmlType,
};
use mlx_native::ops::repeat_tiled::{dispatch_repeat_tiled_f32, RepeatTiledParams};
use mlx_native::ops::rms_norm;
use mlx_native::ops::ssm_conv::{dispatch_ssm_conv, dispatch_ssm_conv_with_capture, SsmConvParams};
use mlx_native::ops::ssm_norm_gate::dispatch_ssm_norm_gate;
use mlx_native::{DType, KernelRegistry, MlxBuffer, MlxDevice};
use super::delta_net::DeltaNetLayerWeights;
use super::encoder_stage::LayerEncoder;
use super::gpu_full_attn::{download_f32, upload_f32, upload_f32_weight, upload_q4_0_from_f32};
use crate::debug::INVESTIGATION_ENV;
use crate::serve::multi_seq_kv::SlotId;
pub const CHUNK_THRESHOLD: u32 = 64;
fn chunk_path_eligible(seq_len: u32, d_k: u32) -> bool {
INVESTIGATION_ENV.chunk_scan_prefill
&& seq_len > CHUNK_THRESHOLD
&& seq_len % mlx_native::ops::chunk_gated_delta_rule::FIXED_BT == 0
&& (d_k == mlx_native::ops::chunk_gated_delta_rule::MAX_K || d_k == mlx_native::ops::chunk_gated_delta_rule_bank_split::BANK_SPLIT_K)
}
#[inline]
fn slot_recurrent_region(slot_id: SlotId, d_k: u32, d_v: u32, n_v_heads: u32) -> (u64, usize) {
let n_elements = (d_k as usize) * (d_v as usize) * (n_v_heads as usize);
let byte_offset = (slot_id.0 as u64)
.checked_mul(n_elements as u64)
.and_then(|e| e.checked_mul(std::mem::size_of::<f32>() as u64))
.expect("slot recurrent byte offset overflow (slot * D_k * D_v * n_v_heads * 4)");
(byte_offset, n_elements)
}
#[inline]
fn slot_conv_state_region(slot_id: SlotId, conv_channels: u32, k_minus_one: u32) -> (u64, usize) {
let n_elements = (conv_channels as usize) * (k_minus_one as usize);
let byte_offset = (slot_id.0 as u64)
.checked_mul(n_elements as u64)
.and_then(|e| e.checked_mul(std::mem::size_of::<f32>() as u64))
.expect("slot conv_state byte offset overflow (slot * channels * (K-1) * 4)");
(byte_offset, n_elements)
}
#[inline]
fn slot_capture_states_region(
slot_id: SlotId,
d_k: u32,
d_v: u32,
n_v_heads: u32,
n_tokens_max: u32,
) -> (u64, usize) {
let per_seq_elems =
(d_k as usize) * (d_v as usize) * (n_v_heads as usize) * (n_tokens_max as usize);
let byte_offset = (slot_id.0 as u64)
.checked_mul(per_seq_elems as u64)
.and_then(|e| e.checked_mul(std::mem::size_of::<f32>() as u64))
.expect("slot capture_states byte offset overflow");
(byte_offset, per_seq_elems)
}
#[inline]
fn slot_conv_capture_region(
slot_id: SlotId,
conv_channels: u32,
k_minus_one: u32,
n_tokens_max: u32,
) -> (u64, usize) {
let per_seq_elems = (n_tokens_max as usize) * (k_minus_one as usize) * (conv_channels as usize);
let byte_offset = (slot_id.0 as u64)
.checked_mul(per_seq_elems as u64)
.and_then(|e| e.checked_mul(std::mem::size_of::<f32>() as u64))
.expect("slot conv_capture byte offset overflow");
(byte_offset, per_seq_elems)
}
struct LaPingPongSlotView {
conv_state_in: MlxBuffer,
conv_state_out: MlxBuffer,
recurrent_in: MlxBuffer,
recurrent_out: MlxBuffer,
}
#[inline]
fn narrow_la_ping_pong_to_slot(
conv_state_in: &MlxBuffer,
conv_state_out: &MlxBuffer,
recurrent_in: &MlxBuffer,
recurrent_out: &MlxBuffer,
slot_id: SlotId,
d_k: u32,
d_v: u32,
n_v_heads: u32,
conv_channels: u32,
k_minus_one: u32,
) -> LaPingPongSlotView {
let (rec_off, rec_n) = slot_recurrent_region(slot_id, d_k, d_v, n_v_heads);
let (conv_off, conv_n) = slot_conv_state_region(slot_id, conv_channels, k_minus_one);
LaPingPongSlotView {
conv_state_in: conv_state_in.slice_view(conv_off, conv_n),
conv_state_out: conv_state_out.slice_view(conv_off, conv_n),
recurrent_in: recurrent_in.slice_view(rec_off, rec_n),
recurrent_out: recurrent_out.slice_view(rec_off, rec_n),
}
}
#[inline]
fn narrow_capture_states_to_slot(
capture: &MlxBuffer,
slot_id: SlotId,
d_k: u32,
d_v: u32,
n_v_heads: u32,
n_seqs_alloc: u32,
) -> MlxBuffer {
let state_elems = (d_k as usize) * (d_v as usize) * (n_v_heads as usize);
let total = capture.element_count();
let per_seq_elems = total / (n_seqs_alloc as usize).max(1);
debug_assert_eq!(
per_seq_elems * (n_seqs_alloc as usize),
total,
"capture_states total ({}) must divide evenly by n_seqs ({})",
total,
n_seqs_alloc,
);
let n_tokens_max = if state_elems == 0 {
0
} else {
per_seq_elems / state_elems
};
let (off, n) = slot_capture_states_region(slot_id, d_k, d_v, n_v_heads, n_tokens_max as u32);
capture.slice_view(off, n)
}
#[inline]
fn narrow_conv_capture_to_slot(
conv_capture: &MlxBuffer,
slot_id: SlotId,
conv_channels: u32,
k_minus_one: u32,
n_seqs_alloc: u32,
) -> MlxBuffer {
let total = conv_capture.element_count();
let per_seq_elems = total / (n_seqs_alloc as usize).max(1);
debug_assert_eq!(
per_seq_elems * (n_seqs_alloc as usize),
total,
"conv_capture_states total ({}) must divide evenly by n_seqs ({})",
total,
n_seqs_alloc,
);
let per_token_elems = (conv_channels as usize) * (k_minus_one as usize);
let n_tokens_max = if per_token_elems == 0 {
0
} else {
per_seq_elems / per_token_elems
};
let (off, n) =
slot_conv_capture_region(slot_id, conv_channels, k_minus_one, n_tokens_max as u32);
conv_capture.slice_view(off, n)
}
const FORWARD_DISPATCH_N_SEQS: u32 = 1;
#[inline]
fn qkv_channels_for(n_k_heads: u32, n_v_heads: u32, d_k: u32, d_v: u32) -> u32 {
2 * n_k_heads * d_k + n_v_heads * d_v
}
pub struct DeltaNetWeightsGpu {
pub attn_norm: MlxBuffer,
pub post_attn_norm: MlxBuffer,
pub attn_qkv: MlxBuffer,
pub attn_gate: MlxBuffer,
pub ssm_conv1d: MlxBuffer,
pub ssm_alpha: MlxBuffer,
pub ssm_dt_bias: MlxBuffer,
pub ssm_dt_bias_cpu: Vec<f32>,
pub ssm_beta: MlxBuffer,
pub ssm_a: MlxBuffer,
pub ssm_a_cpu: Vec<f32>,
pub ssm_norm: MlxBuffer,
pub ssm_norm_cpu: Vec<f32>,
pub ssm_out: MlxBuffer,
}
impl DeltaNetWeightsGpu {
pub fn from_cpu(
weights: &DeltaNetLayerWeights,
device: &MlxDevice,
k_width: usize,
qkv_channels: usize,
) -> Result<Self> {
let conv1d_t =
transpose_k_channels_to_channels_k(&weights.ssm_conv1d, k_width, qkv_channels);
Ok(Self {
attn_norm: upload_f32_weight(&weights.attn_norm, device)?,
post_attn_norm: upload_f32_weight(&weights.post_attn_norm, device)?,
ssm_conv1d: upload_f32_weight(&conv1d_t, device)?,
ssm_dt_bias: upload_f32_weight(&weights.ssm_dt_bias, device)?,
ssm_dt_bias_cpu: weights.ssm_dt_bias.clone(),
ssm_a: upload_f32_weight(&weights.ssm_a, device)?,
ssm_a_cpu: weights.ssm_a.clone(),
ssm_norm: upload_f32_weight(&weights.ssm_norm, device)?,
ssm_norm_cpu: weights.ssm_norm.clone(),
attn_qkv: upload_q4_0_from_f32(&weights.attn_qkv, device)?,
attn_gate: upload_q4_0_from_f32(&weights.attn_gate, device)?,
ssm_alpha: upload_q4_0_from_f32(&weights.ssm_alpha, device)?,
ssm_beta: upload_q4_0_from_f32(&weights.ssm_beta, device)?,
ssm_out: upload_q4_0_from_f32(&weights.ssm_out, device)?,
})
}
#[cfg(test)]
pub fn from_cpu_f32(
weights: &DeltaNetLayerWeights,
device: &MlxDevice,
k_width: usize,
qkv_channels: usize,
) -> Result<Self> {
let conv1d_t =
transpose_k_channels_to_channels_k(&weights.ssm_conv1d, k_width, qkv_channels);
Ok(Self {
attn_norm: upload_f32(&weights.attn_norm, device)?,
post_attn_norm: upload_f32(&weights.post_attn_norm, device)?,
ssm_conv1d: upload_f32(&conv1d_t, device)?,
ssm_dt_bias: upload_f32(&weights.ssm_dt_bias, device)?,
ssm_dt_bias_cpu: weights.ssm_dt_bias.clone(),
ssm_a: upload_f32(&weights.ssm_a, device)?,
ssm_a_cpu: weights.ssm_a.clone(),
ssm_norm: upload_f32(&weights.ssm_norm, device)?,
ssm_norm_cpu: weights.ssm_norm.clone(),
attn_qkv: upload_f32(&weights.attn_qkv, device)?,
attn_gate: upload_f32(&weights.attn_gate, device)?,
ssm_alpha: upload_f32(&weights.ssm_alpha, device)?,
ssm_beta: upload_f32(&weights.ssm_beta, device)?,
ssm_out: upload_f32(&weights.ssm_out, device)?,
})
}
}
fn transpose_k_channels_to_channels_k(src: &[f32], k: usize, channels: usize) -> Vec<f32> {
let mut dst = vec![0.0f32; k * channels];
for ki in 0..k {
for c in 0..channels {
dst[c * k + ki] = src[ki * channels + c];
}
}
dst
}
fn transpose_state_km1_c_to_c_km1(src: &[f32], km1: usize, channels: usize) -> Vec<f32> {
let mut dst = vec![0.0f32; km1 * channels];
for i in 0..km1 {
for c in 0..channels {
dst[c * km1 + i] = src[i * channels + c];
}
}
dst
}
fn transpose_state_c_km1_to_km1_c(src: &[f32], km1: usize, channels: usize) -> Vec<f32> {
let mut dst = vec![0.0f32; km1 * channels];
for c in 0..channels {
for i in 0..km1 {
dst[i * channels + c] = src[c * km1 + i];
}
}
dst
}
pub fn apply_pre_norm(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
norm_weight: &MlxBuffer,
seq_len: u32,
hidden_size: u32,
eps: f32,
) -> Result<MlxBuffer> {
let out = super::decode_pool::pooled_alloc_buffer(
device,
(seq_len * hidden_size) as usize * 4,
DType::F32,
vec![seq_len as usize, hidden_size as usize],
)
.map_err(|e| anyhow!("alloc pre_norm out: {e}"))?;
let mut params = super::decode_pool::pooled_alloc_buffer(device, 8, DType::F32, vec![2])
.map_err(|e| anyhow!("alloc params: {e}"))?;
{
let s = params.as_mut_slice::<f32>().map_err(|e| anyhow!("{e}"))?;
s[0] = eps;
s[1] = hidden_size as f32;
}
rms_norm::dispatch_rms_norm(
encoder,
registry,
device.metal_device(),
input,
norm_weight,
&out,
¶ms,
seq_len,
hidden_size,
)
.context("dispatch_rms_norm pre_norm")?;
Ok(out)
}
pub fn apply_proj(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
seq_len: u32,
in_features: u32,
out_features: u32,
) -> Result<MlxBuffer> {
let out_bytes = (seq_len * out_features) as usize * 4;
let mut dst = super::decode_pool::pooled_alloc_buffer(
device,
out_bytes,
DType::F32,
vec![seq_len as usize, out_features as usize],
)
.map_err(|e| anyhow!("alloc proj out: {e}"))?;
match weight.dtype() {
DType::U8 => {
let params = GgmlQuantizedMatmulParams {
m: seq_len,
n: out_features,
k: in_features,
ggml_type: GgmlType::Q4_0,
};
quantized_matmul_ggml(encoder, registry, device, input, weight, &mut dst, ¶ms)
.context("quantized_matmul_ggml Q4_0")?;
}
DType::BF16 => {
let params = DenseMmBf16F32Params {
m: seq_len,
n: out_features,
k: in_features,
src0_batch: 1,
src1_batch: 1,
};
dense_matmul_bf16_f32_tensor(
encoder, registry, device, weight, input, &mut dst, ¶ms,
)
.context("dense_matmul_bf16_f32_tensor proj")?;
}
DType::F32 => {
let n_w = (out_features * in_features) as usize;
let weight_bf16 = super::decode_pool::pooled_alloc_buffer(
device,
n_w * 2,
DType::BF16,
vec![out_features as usize, in_features as usize],
)
.map_err(|e| anyhow!("alloc weight_bf16 (pooled): {e}"))?;
cast(
encoder,
registry,
device.metal_device(),
weight,
&weight_bf16,
n_w,
CastDirection::F32ToBF16,
)
.context("cast weight F32→BF16")?;
encoder.memory_barrier();
let params = DenseMmBf16F32Params {
m: seq_len,
n: out_features,
k: in_features,
src0_batch: 1,
src1_batch: 1,
};
dense_matmul_bf16_f32_tensor(
encoder,
registry,
device,
&weight_bf16,
input,
&mut dst,
¶ms,
)
.context("dense_matmul_bf16_f32_tensor proj (F32 legacy)")?;
}
other => {
return Err(anyhow!("apply_proj: unsupported weight dtype {:?}", other));
}
}
Ok(dst)
}
pub fn prepare_ssm_conv_buffers(
device: &MlxDevice,
old_conv_state_km1_c: &[f32], seq_len: u32,
qkv_channels: u32,
k_width: u32,
) -> Result<(MlxBuffer, MlxBuffer, MlxBuffer, MlxBuffer, SsmConvParams)> {
let km1 = (k_width - 1) as usize;
let channels = qkv_channels as usize;
let old_state_ck = transpose_state_km1_c_to_c_km1(old_conv_state_km1_c, km1, channels);
let old_state_buf = upload_f32(&old_state_ck, device)?;
let s_elems = km1 * channels;
let new_state_buf = device
.alloc_buffer(s_elems * 4, DType::F32, vec![channels, km1])
.map_err(|e| anyhow!("alloc new_state_buf: {e}"))?;
let y = device
.alloc_buffer(
(seq_len * qkv_channels) as usize * 4,
DType::F32,
vec![seq_len as usize, qkv_channels as usize],
)
.map_err(|e| anyhow!("alloc ssm_conv y: {e}"))?;
let mut params_buf = device
.alloc_buffer(4 * 4, DType::U32, vec![4])
.map_err(|e| anyhow!("alloc ssm_conv params: {e}"))?;
{
let s = params_buf
.as_mut_slice::<u32>()
.map_err(|e| anyhow!("{e}"))?;
s[0] = qkv_channels;
s[1] = seq_len;
s[2] = FORWARD_DISPATCH_N_SEQS;
s[3] = k_width;
}
let conv_params = SsmConvParams {
channels: qkv_channels,
n_tokens: seq_len,
n_seqs: FORWARD_DISPATCH_N_SEQS,
k_width,
};
Ok((old_state_buf, new_state_buf, y, params_buf, conv_params))
}
pub fn extract_new_conv_state(
new_state_buf: &MlxBuffer,
km1: usize,
channels: usize,
) -> Result<Vec<f32>> {
let new_state_ck = download_f32(new_state_buf)?;
Ok(transpose_state_c_km1_to_km1_c(&new_state_ck, km1, channels))
}
pub fn apply_ssm_conv(
device: &MlxDevice,
registry: &mut KernelRegistry,
qkv_seq_major: &MlxBuffer,
conv_kernel_transposed: &MlxBuffer, old_conv_state_km1_c: &[f32], seq_len: u32,
qkv_channels: u32,
k_width: u32,
) -> Result<(MlxBuffer, Vec<f32>)> {
let km1 = (k_width - 1) as usize;
let channels = qkv_channels as usize;
let (old_state_buf, new_state_buf, y, params_buf, conv_params) =
prepare_ssm_conv_buffers(device, old_conv_state_km1_c, seq_len, qkv_channels, k_width)?;
let mut enc = device.command_encoder().context("enc ssm_conv")?;
dispatch_ssm_conv(
&mut enc,
registry,
device.metal_device(),
qkv_seq_major,
conv_kernel_transposed,
&old_state_buf,
&new_state_buf,
&y,
¶ms_buf,
conv_params,
)
.context("dispatch_ssm_conv")?;
enc.commit_and_wait().context("commit ssm_conv")?;
let new_state_km1_c = extract_new_conv_state(&new_state_buf, km1, channels)?;
Ok((y, new_state_km1_c))
}
pub fn apply_l2_norm_per_head(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
seq_len: u32,
n_heads: u32,
head_dim: u32,
eps: f32,
) -> Result<MlxBuffer> {
let rows = seq_len * n_heads;
let dim = head_dim;
let out = super::decode_pool::pooled_alloc_buffer(
device,
(rows * dim) as usize * 4,
DType::F32,
vec![rows as usize, dim as usize],
)
.map_err(|e| anyhow!("alloc l2_norm out: {e}"))?;
let mut params_buf = super::decode_pool::pooled_alloc_buffer(device, 8, DType::F32, vec![2])
.map_err(|e| anyhow!("alloc l2_norm params: {e}"))?;
{
let s = params_buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("{e}"))?;
s[0] = eps;
s[1] = dim as f32;
}
dispatch_l2_norm(
encoder,
registry,
device.metal_device(),
input,
&out,
¶ms_buf,
rows,
dim,
)
.context("dispatch_l2_norm")?;
Ok(out)
}
#[allow(clippy::too_many_arguments)]
pub fn apply_pre_norm_into(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
norm_weight: &MlxBuffer,
out: &MlxBuffer,
params_buf: &mut MlxBuffer,
seq_len: u32,
hidden_size: u32,
eps: f32,
) -> Result<()> {
{
let s = params_buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("apply_pre_norm_into: params as_mut_slice: {e}"))?;
s[0] = eps;
s[1] = hidden_size as f32;
}
rms_norm::dispatch_rms_norm(
encoder,
registry,
device.metal_device(),
input,
norm_weight,
out,
params_buf,
seq_len,
hidden_size,
)
.context("dispatch_rms_norm pre_norm (arena)")?;
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn apply_proj_into(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
dst: &mut MlxBuffer,
seq_len: u32,
in_features: u32,
out_features: u32,
) -> Result<()> {
match weight.dtype() {
DType::U8 => {
let params = GgmlQuantizedMatmulParams {
m: seq_len,
n: out_features,
k: in_features,
ggml_type: GgmlType::Q4_0,
};
quantized_matmul_ggml(encoder, registry, device, input, weight, dst, ¶ms)
.context("quantized_matmul_ggml Q4_0 (arena)")?;
}
DType::BF16 => {
let params = DenseMmBf16F32Params {
m: seq_len,
n: out_features,
k: in_features,
src0_batch: 1,
src1_batch: 1,
};
dense_matmul_bf16_f32_tensor(encoder, registry, device, weight, input, dst, ¶ms)
.context("dense_matmul_bf16_f32_tensor proj (arena)")?;
}
DType::F32 => {
let n_w = (out_features * in_features) as usize;
let weight_bf16 = super::decode_pool::pooled_alloc_buffer(
device,
n_w * 2,
DType::BF16,
vec![out_features as usize, in_features as usize],
)
.map_err(|e| anyhow!("alloc weight_bf16 (pooled, arena F32 fallback): {e}"))?;
cast(
encoder,
registry,
device.metal_device(),
weight,
&weight_bf16,
n_w,
CastDirection::F32ToBF16,
)
.context("cast weight F32→BF16 (arena F32 fallback)")?;
encoder.memory_barrier();
let params = DenseMmBf16F32Params {
m: seq_len,
n: out_features,
k: in_features,
src0_batch: 1,
src1_batch: 1,
};
dense_matmul_bf16_f32_tensor(
encoder,
registry,
device,
&weight_bf16,
input,
dst,
¶ms,
)
.context("dense_matmul_bf16_f32_tensor proj (arena F32 fallback)")?;
}
other => {
return Err(anyhow!(
"apply_proj_into: unsupported weight dtype {:?}",
other
));
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn apply_l2_norm_per_head_into(
encoder: &mut mlx_native::CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
out: &MlxBuffer,
params_buf: &mut MlxBuffer,
seq_len: u32,
n_heads: u32,
head_dim: u32,
eps: f32,
) -> Result<()> {
let rows = seq_len * n_heads;
let dim = head_dim;
{
let s = params_buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("apply_l2_norm_per_head_into: params as_mut_slice: {e}"))?;
s[0] = eps;
s[1] = dim as f32;
}
dispatch_l2_norm(
encoder,
registry,
device.metal_device(),
input,
out,
params_buf,
rows,
dim,
)
.context("dispatch_l2_norm (arena)")?;
Ok(())
}
pub fn compute_g_and_beta_cpu(
alpha_logit_cpu: &[f32], beta_logit_cpu: &[f32], dt_bias: &[f32], ssm_a: &[f32], seq_len: usize,
nv: usize,
) -> (Vec<f32>, Vec<f32>) {
let mut g = vec![0.0f32; seq_len * nv];
let mut beta = vec![0.0f32; seq_len * nv];
for t in 0..seq_len {
for vh in 0..nv {
let a_logit = alpha_logit_cpu[t * nv + vh] + dt_bias[vh];
g[t * nv + vh] = softplus_f32(a_logit) * (-ssm_a[vh]);
beta[t * nv + vh] = sigmoid_f32(beta_logit_cpu[t * nv + vh]);
}
}
(g, beta)
}
#[inline(always)]
fn softplus_f32(x: f32) -> f32 {
if x > 20.0 {
x
} else if x < -20.0 {
0.0
} else {
(1.0 + x.exp()).ln()
}
}
#[inline(always)]
fn sigmoid_f32(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}
#[allow(clippy::too_many_arguments)]
pub fn apply_gated_delta_net(
device: &MlxDevice,
registry: &mut KernelRegistry,
q: &MlxBuffer,
k: &MlxBuffer,
v: &MlxBuffer,
g_buf: &MlxBuffer,
beta_buf: &MlxBuffer,
state_in: &MlxBuffer,
seq_len: u32,
n_k_heads: u32,
n_v_heads: u32,
d_k: u32,
d_v: u32,
) -> Result<(MlxBuffer, MlxBuffer)> {
let n_seqs = FORWARD_DISPATCH_N_SEQS;
let out_elems = (n_v_heads * seq_len * d_v) as usize;
let output_buf = device
.alloc_buffer(out_elems * 4, DType::F32, vec![out_elems])
.map_err(|e| anyhow!("alloc gdn output: {e}"))?;
let state_elems = (d_k * d_v * n_v_heads) as usize; let state_out_buf = device
.alloc_buffer(state_elems * 4, DType::F32, vec![state_elems])
.map_err(|e| anyhow!("alloc gdn state_out: {e}"))?;
let params = GatedDeltaNetParams {
d_k,
d_v,
n_k_heads,
n_v_heads,
n_tokens: seq_len,
n_seqs,
};
let params_buf = build_gated_delta_net_params(device, params)
.map_err(|e| anyhow!("build gdn params: {e}"))?;
let mut enc = device.command_encoder().context("enc gdn")?;
dispatch_gated_delta_net(
&mut enc,
registry,
device.metal_device(),
q,
k,
v,
g_buf,
beta_buf,
state_in,
&output_buf,
&state_out_buf,
¶ms_buf,
params,
)
.context("dispatch_gated_delta_net")?;
enc.commit_and_wait().context("commit gdn")?;
Ok((output_buf, state_out_buf))
}
#[allow(clippy::too_many_arguments)]
pub fn apply_gated_delta_net_chunk(
device: &MlxDevice,
registry: &mut KernelRegistry,
q: &MlxBuffer,
k: &MlxBuffer,
v: &MlxBuffer,
g_buf: &MlxBuffer,
beta_buf: &MlxBuffer,
state_in: &MlxBuffer,
output_buf: &MlxBuffer,
final_state: &MlxBuffer,
seq_len: u32,
n_k_heads: u32,
n_v_heads: u32,
d_k: u32,
d_v: u32,
use_qk_l2norm: bool,
) -> Result<()> {
if use_qk_l2norm {
return Err(anyhow!(
"apply_gated_delta_net_chunk: use_qk_l2norm=true is reserved for a \
future iter that defers l2-norm to the chunk dispatch. Iter 5 \
callers must pre-apply l2-norm (matching apply_gated_delta_net) \
and pass use_qk_l2norm=false."
));
}
if seq_len == 0 || seq_len % FIXED_BT != 0 {
return Err(anyhow!(
"apply_gated_delta_net_chunk: seq_len ({}) must be a positive multiple \
of FIXED_BT ({})",
seq_len,
FIXED_BT
));
}
let n_seqs = FORWARD_DISPATCH_N_SEQS;
let q_elems_exp = (seq_len * n_v_heads * d_k) as usize; let v_elems = (seq_len * n_v_heads * d_v) as usize; let g_elems = (seq_len * n_v_heads) as usize; let state_elems = (d_k * d_v * n_v_heads) as usize; let out_elems_bf16 = v_elems; let out_elems_f32 = (n_v_heads * seq_len * d_v) as usize;
let _w5b8_expand = crate::inference::models::qwen35::wave5b8_profile::Section::start(
crate::inference::models::qwen35::wave5b8_profile::SectionKind::ChunkGqaExpand,
);
let q_expanded = device
.alloc_buffer(q_elems_exp * 4, DType::F32, vec![q_elems_exp])
.map_err(|e| anyhow!("alloc q_expanded f32: {e}"))?;
let k_expanded = device
.alloc_buffer(q_elems_exp * 4, DType::F32, vec![q_elems_exp])
.map_err(|e| anyhow!("alloc k_expanded f32: {e}"))?;
drop(_w5b8_expand);
let _w5b8_allocs = crate::inference::models::qwen35::wave5b8_profile::Section::start(
crate::inference::models::qwen35::wave5b8_profile::SectionKind::ChunkAllocs,
);
let q_bf16 = device
.alloc_buffer(q_elems_exp * 2, DType::BF16, vec![q_elems_exp])
.map_err(|e| anyhow!("alloc q_bf16: {e}"))?;
let k_bf16 = device
.alloc_buffer(q_elems_exp * 2, DType::BF16, vec![q_elems_exp])
.map_err(|e| anyhow!("alloc k_bf16: {e}"))?;
let v_bf16 = device
.alloc_buffer(v_elems * 2, DType::BF16, vec![v_elems])
.map_err(|e| anyhow!("alloc v_bf16: {e}"))?;
let g_log_decay = device
.alloc_buffer(g_elems * 4, DType::F32, vec![g_elems])
.map_err(|e| anyhow!("alloc g_log_decay: {e}"))?;
let o_bf16 = device
.alloc_buffer(out_elems_bf16 * 2, DType::BF16, vec![out_elems_bf16])
.map_err(|e| anyhow!("alloc o_bf16: {e}"))?;
let final_state_bytes_needed = state_elems * std::mem::size_of::<f32>();
if final_state.byte_len() < final_state_bytes_needed {
return Err(anyhow!(
"apply_gated_delta_net_chunk: final_state byte_len {} < required {} (state_elems={})",
final_state.byte_len(),
final_state_bytes_needed,
state_elems
));
}
let output_bytes_needed = out_elems_f32 * std::mem::size_of::<f32>();
if output_buf.byte_len() < output_bytes_needed {
return Err(anyhow!(
"apply_gated_delta_net_chunk: output_buf byte_len {} < required {} (out_elems_f32={})",
output_buf.byte_len(),
output_bytes_needed,
out_elems_f32
));
}
let p = ChunkGatedDeltaRuleParams {
b: n_seqs,
t: seq_len,
hg: n_v_heads,
h: n_v_heads,
k: d_k,
v: d_v,
bt: FIXED_BT,
scale: 1.0_f32,
use_qk_l2norm: false,
};
drop(_w5b8_allocs);
let _w5b8_encbuild = crate::inference::models::qwen35::wave5b8_profile::Section::start(
crate::inference::models::qwen35::wave5b8_profile::SectionKind::ChunkEncBuild,
);
let mut enc = device
.command_encoder()
.context("enc apply_gated_delta_net_chunk")?;
let rt_params = RepeatTiledParams {
seq: seq_len,
hg: n_k_heads,
h: n_v_heads,
k: d_k,
};
dispatch_repeat_tiled_f32(
&mut enc,
registry,
device.metal_device(),
q,
&q_expanded,
&rt_params,
)
.map_err(|e| anyhow!("dispatch_repeat_tiled_f32 q (W-5b.20): {e}"))?;
dispatch_repeat_tiled_f32(
&mut enc,
registry,
device.metal_device(),
k,
&k_expanded,
&rt_params,
)
.map_err(|e| anyhow!("dispatch_repeat_tiled_f32 k (W-5b.20): {e}"))?;
enc.memory_barrier();
cast(
&mut enc,
registry,
device.metal_device(),
&q_expanded,
&q_bf16,
q_elems_exp,
CastDirection::F32ToBF16,
)
.context("cast q_expanded F32→BF16")?;
cast(
&mut enc,
registry,
device.metal_device(),
&k_expanded,
&k_bf16,
q_elems_exp,
CastDirection::F32ToBF16,
)
.context("cast k_expanded F32→BF16")?;
cast(
&mut enc,
registry,
device.metal_device(),
v,
&v_bf16,
v_elems,
CastDirection::F32ToBF16,
)
.context("cast v F32→BF16")?;
scalar_mul_f32(
&mut enc,
registry,
device.metal_device(),
g_buf,
&g_log_decay,
g_elems,
-1.0_f32,
)
.context("scalar_mul_f32 g_log_decay = -g")?;
enc.memory_barrier();
if p.k == mlx_native::ops::chunk_gated_delta_rule_bank_split::BANK_SPLIT_K {
mlx_native::ops::chunk_gated_delta_rule_bank_split::dispatch_chunk_gated_delta_rule_fwd_k256_bank_split(
&mut enc,
registry,
device,
&q_bf16,
&k_bf16,
&v_bf16,
&g_log_decay,
beta_buf,
state_in,
&o_bf16,
final_state,
p,
)
.map_err(|e| anyhow!("dispatch_chunk_gated_delta_rule_fwd_k256: {e}"))?;
} else {
dispatch_chunk_gated_delta_rule_fwd(
&mut enc,
registry,
device,
&q_bf16,
&k_bf16,
&v_bf16,
&g_log_decay,
beta_buf,
state_in,
&o_bf16,
final_state,
p,
)
.map_err(|e| anyhow!("dispatch_chunk_gated_delta_rule_fwd: {e}"))?;
}
enc.memory_barrier();
cast(
&mut enc,
registry,
device.metal_device(),
&o_bf16,
output_buf,
out_elems_bf16,
CastDirection::BF16ToF32,
)
.context("cast output BF16→F32")?;
drop(_w5b8_encbuild);
{
let _w5b8_commit = crate::inference::models::qwen35::wave5b8_profile::Section::start(
crate::inference::models::qwen35::wave5b8_profile::SectionKind::ChunkCommitWait,
);
enc.commit_and_wait_labeled("layer.gdn.chunk_attn")
.context("commit_and_wait apply_gated_delta_net_chunk")?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn apply_gated_delta_net_chunk_with_arena(
device: &MlxDevice,
registry: &mut KernelRegistry,
q: &MlxBuffer,
k: &MlxBuffer,
v: &MlxBuffer,
g_buf: &MlxBuffer,
beta_buf: &MlxBuffer,
state_in: &MlxBuffer,
output_buf: &MlxBuffer,
final_state: &MlxBuffer,
arena: &mut super::ChunkAllocsArena,
chunk_internal_arena: Option<&mut ChunkInternalArena>,
seq_len: u32,
n_k_heads: u32,
n_v_heads: u32,
d_k: u32,
d_v: u32,
use_qk_l2norm: bool,
) -> Result<()> {
if use_qk_l2norm {
return Err(anyhow!(
"apply_gated_delta_net_chunk_with_arena: use_qk_l2norm=true is reserved \
for a future iter that defers l2-norm to the chunk dispatch. iter78 \
callers must pre-apply l2-norm (matching apply_gated_delta_net) and \
pass use_qk_l2norm=false."
));
}
if seq_len == 0 || seq_len % FIXED_BT != 0 {
return Err(anyhow!(
"apply_gated_delta_net_chunk_with_arena: seq_len ({}) must be a positive \
multiple of FIXED_BT ({})",
seq_len,
FIXED_BT
));
}
arena
.validate_fits(seq_len, n_v_heads, d_k, d_v)
.context("ChunkAllocsArena shape mismatch")?;
let n_seqs = FORWARD_DISPATCH_N_SEQS;
let q_elems_exp = (seq_len * n_v_heads * d_k) as usize;
let v_elems = (seq_len * n_v_heads * d_v) as usize;
let g_elems = (seq_len * n_v_heads) as usize;
let state_elems = (d_k * d_v * n_v_heads) as usize;
let out_elems_bf16 = v_elems;
let out_elems_f32 = (n_v_heads * seq_len * d_v) as usize;
let _w5b8_expand = crate::inference::models::qwen35::wave5b8_profile::Section::start(
crate::inference::models::qwen35::wave5b8_profile::SectionKind::ChunkGqaExpand,
);
drop(_w5b8_expand);
let _w5b8_allocs = crate::inference::models::qwen35::wave5b8_profile::Section::start(
crate::inference::models::qwen35::wave5b8_profile::SectionKind::ChunkAllocs,
);
drop(_w5b8_allocs);
let final_state_bytes_needed = state_elems * std::mem::size_of::<f32>();
if final_state.byte_len() < final_state_bytes_needed {
return Err(anyhow!(
"apply_gated_delta_net_chunk_with_arena: final_state byte_len {} < \
required {} (state_elems={})",
final_state.byte_len(),
final_state_bytes_needed,
state_elems
));
}
let output_bytes_needed = out_elems_f32 * std::mem::size_of::<f32>();
if output_buf.byte_len() < output_bytes_needed {
return Err(anyhow!(
"apply_gated_delta_net_chunk_with_arena: output_buf byte_len {} < \
required {} (out_elems_f32={})",
output_buf.byte_len(),
output_bytes_needed,
out_elems_f32
));
}
let p = ChunkGatedDeltaRuleParams {
b: n_seqs,
t: seq_len,
hg: n_v_heads,
h: n_v_heads,
k: d_k,
v: d_v,
bt: FIXED_BT,
scale: 1.0_f32,
use_qk_l2norm: false,
};
let _w5b8_encbuild = crate::inference::models::qwen35::wave5b8_profile::Section::start(
crate::inference::models::qwen35::wave5b8_profile::SectionKind::ChunkEncBuild,
);
let mut enc = device
.command_encoder()
.context("enc apply_gated_delta_net_chunk_with_arena")?;
let rt_params = RepeatTiledParams {
seq: seq_len,
hg: n_k_heads,
h: n_v_heads,
k: d_k,
};
dispatch_repeat_tiled_f32(
&mut enc,
registry,
device.metal_device(),
q,
&arena.q_expanded_buf,
&rt_params,
)
.map_err(|e| anyhow!("dispatch_repeat_tiled_f32 q (arena): {e}"))?;
dispatch_repeat_tiled_f32(
&mut enc,
registry,
device.metal_device(),
k,
&arena.k_expanded_buf,
&rt_params,
)
.map_err(|e| anyhow!("dispatch_repeat_tiled_f32 k (arena): {e}"))?;
enc.memory_barrier();
cast(
&mut enc,
registry,
device.metal_device(),
&arena.q_expanded_buf,
&arena.q_bf16_buf,
q_elems_exp,
CastDirection::F32ToBF16,
)
.context("cast q_expanded F32→BF16 (arena)")?;
cast(
&mut enc,
registry,
device.metal_device(),
&arena.k_expanded_buf,
&arena.k_bf16_buf,
q_elems_exp,
CastDirection::F32ToBF16,
)
.context("cast k_expanded F32→BF16 (arena)")?;
cast(
&mut enc,
registry,
device.metal_device(),
v,
&arena.v_bf16_buf,
v_elems,
CastDirection::F32ToBF16,
)
.context("cast v F32→BF16 (arena)")?;
scalar_mul_f32(
&mut enc,
registry,
device.metal_device(),
g_buf,
&arena.g_log_decay_buf,
g_elems,
-1.0_f32,
)
.context("scalar_mul_f32 g_log_decay = -g (arena)")?;
enc.memory_barrier();
if p.k == mlx_native::ops::chunk_gated_delta_rule_bank_split::BANK_SPLIT_K {
mlx_native::ops::chunk_gated_delta_rule_bank_split::dispatch_chunk_gated_delta_rule_fwd_k256_bank_split(
&mut enc,
registry,
device,
&arena.q_bf16_buf,
&arena.k_bf16_buf,
&arena.v_bf16_buf,
&arena.g_log_decay_buf,
beta_buf,
state_in,
&arena.o_bf16_buf,
final_state,
p,
)
.map_err(|e| anyhow!("dispatch_chunk_gated_delta_rule_fwd_k256 (arena fallback): {e}"))?;
} else if let Some(ci_arena) = chunk_internal_arena {
dispatch_chunk_gated_delta_rule_fwd_with_arena(
&mut enc,
registry,
device,
&arena.q_bf16_buf,
&arena.k_bf16_buf,
&arena.v_bf16_buf,
&arena.g_log_decay_buf,
beta_buf,
state_in,
&arena.o_bf16_buf,
final_state,
ci_arena,
p,
)
.map_err(|e| anyhow!("dispatch_chunk_gated_delta_rule_fwd_with_arena: {e}"))?;
} else {
dispatch_chunk_gated_delta_rule_fwd(
&mut enc,
registry,
device,
&arena.q_bf16_buf,
&arena.k_bf16_buf,
&arena.v_bf16_buf,
&arena.g_log_decay_buf,
beta_buf,
state_in,
&arena.o_bf16_buf,
final_state,
p,
)
.map_err(|e| anyhow!("dispatch_chunk_gated_delta_rule_fwd (arena): {e}"))?;
}
enc.memory_barrier();
cast(
&mut enc,
registry,
device.metal_device(),
&arena.o_bf16_buf,
output_buf,
out_elems_bf16,
CastDirection::BF16ToF32,
)
.context("cast output BF16→F32 (arena)")?;
drop(_w5b8_encbuild);
{
let _w5b8_commit = crate::inference::models::qwen35::wave5b8_profile::Section::start(
crate::inference::models::qwen35::wave5b8_profile::SectionKind::ChunkCommitWait,
);
enc.commit_and_wait_labeled("layer.gdn.chunk_attn")
.context("commit_and_wait apply_gated_delta_net_chunk_with_arena")?;
}
Ok(())
}
pub fn apply_ssm_norm_and_gate(
_encoder: &mut mlx_native::CommandEncoder,
_registry: &mut KernelRegistry,
device: &MlxDevice,
attn_out: &MlxBuffer,
z_flat: &MlxBuffer,
ssm_norm_w_cpu: &[f32], seq_len: u32,
z_channels: u32, eps: f32,
) -> Result<MlxBuffer> {
let attn_out_cpu = download_f32(attn_out).context("download attn_out ssm_norm")?;
let z_cpu = download_f32(z_flat).context("download z ssm_norm")?;
let n_total = (seq_len * z_channels) as usize;
let dv = ssm_norm_w_cpu.len(); let nv = if dv > 0 { z_channels as usize / dv } else { 1 };
let mut gated = vec![0.0f32; n_total];
let seq = seq_len as usize;
for t in 0..seq {
for vh in 0..nv {
let head_off = t * nv * dv + vh * dv;
let head_row = &attn_out_cpu[head_off..head_off + dv];
let sum_sq: f32 = head_row.iter().map(|v| v * v).sum();
let inv = ((sum_sq / (dv as f32)) + eps).sqrt().recip();
for d in 0..dv {
let normed_val = head_row[d] * inv * ssm_norm_w_cpu[d];
let z_val = z_cpu[head_off + d];
let z_silu = z_val / (1.0 + (-z_val).exp());
gated[head_off + d] = normed_val * z_silu;
}
}
}
upload_f32(&gated, device).context("upload ssm_norm_gated")
}
#[allow(clippy::too_many_arguments)]
#[allow(clippy::too_many_arguments)]
pub fn build_delta_net_layer(
device: &MlxDevice,
registry: &mut KernelRegistry,
x: &MlxBuffer,
weights: &DeltaNetWeightsGpu,
conv_state_in: &MlxBuffer, conv_state_out: &MlxBuffer, state_in: &MlxBuffer, state_out: &MlxBuffer, seq_len: u32,
hidden_size: u32,
n_k_heads: u32,
n_v_heads: u32,
d_k: u32,
d_v: u32,
k_width: u32,
rms_norm_eps: f32,
state_capture: Option<&MlxBuffer>,
conv_state_capture: Option<&MlxBuffer>,
slot_id: SlotId,
) -> Result<MlxBuffer> {
let per_seq_rec_elems = (d_k as usize) * (d_v as usize) * (n_v_heads as usize);
let n_seqs_alloc: u32 = if per_seq_rec_elems == 0 {
1
} else {
let total = state_in.element_count();
if total % per_seq_rec_elems != 0 {
return Err(anyhow!(
"build_delta_net_layer: state_in element_count {} not divisible by \
per-seq recurrent elems {} (D_k*D_v*n_v_heads={}). Shape drift \
between caller and kernel-native layout — ADR-040 §6.1.40 lift",
total,
per_seq_rec_elems,
per_seq_rec_elems,
));
}
(total / per_seq_rec_elems) as u32
};
if slot_id.0 >= n_seqs_alloc {
return Err(anyhow!(
"build_delta_net_layer: slot_id={} out of range (state_in n_seqs={}). \
ADR-040 Phase A2b-cont per-slot routing contract.",
slot_id.0,
n_seqs_alloc,
));
}
let la_view = narrow_la_ping_pong_to_slot(
conv_state_in,
conv_state_out,
state_in,
state_out,
slot_id,
d_k,
d_v,
n_v_heads,
qkv_channels_for(n_k_heads, n_v_heads, d_k, d_v),
k_width - 1,
);
let conv_state_in = &la_view.conv_state_in;
let conv_state_out = &la_view.conv_state_out;
let state_in = &la_view.recurrent_in;
let state_out = &la_view.recurrent_out;
let state_capture_view = state_capture
.map(|b| narrow_capture_states_to_slot(b, slot_id, d_k, d_v, n_v_heads, n_seqs_alloc));
let conv_state_capture_view = conv_state_capture.map(|b| {
narrow_conv_capture_to_slot(
b,
slot_id,
qkv_channels_for(n_k_heads, n_v_heads, d_k, d_v),
k_width - 1,
n_seqs_alloc,
)
});
let state_capture: Option<&MlxBuffer> = state_capture_view.as_ref();
let conv_state_capture: Option<&MlxBuffer> = conv_state_capture_view.as_ref();
let qkv_channels = 2 * n_k_heads * d_k + n_v_heads * d_v;
let z_channels = n_v_heads * d_v;
let q_span = n_k_heads * d_k;
let k_span = n_k_heads * d_k;
let _km1 = (k_width - 1) as usize;
let _channels = qkv_channels as usize;
let seq = seq_len as usize;
let _nk = n_k_heads as usize;
let nv = n_v_heads as usize;
let dk = d_k as usize;
let dv = d_v as usize;
let qkv_ch = qkv_channels as usize;
let q_sp = q_span as usize;
let k_sp = k_span as usize;
let n_q_elems = (seq_len * n_k_heads * d_k) as usize;
let q_scale_val = 1.0_f32 / (dk as f32).sqrt();
let g_n = (seq_len * n_v_heads) as usize;
let rows_op8 = seq_len * n_v_heads;
let gated_elems = (rows_op8 * d_v) as usize;
let n_seqs = FORWARD_DISPATCH_N_SEQS;
let out_elems = (n_v_heads * seq_len * d_v) as usize;
let qkv_conv = super::decode_pool::pooled_alloc_buffer(
device,
(seq_len * qkv_channels) as usize * 4,
DType::F32,
vec![seq as usize, qkv_ch],
)
.map_err(|e| anyhow!("alloc qkv_conv: {e}"))?;
let mut ssm_params_buf =
super::decode_pool::pooled_alloc_buffer(device, 4 * 4, DType::U32, vec![4])
.map_err(|e| anyhow!("alloc ssm params: {e}"))?;
{
let s = ssm_params_buf
.as_mut_slice::<u32>()
.map_err(|e| anyhow!("{e}"))?;
s[0] = qkv_channels;
s[1] = seq_len;
s[2] = FORWARD_DISPATCH_N_SEQS;
s[3] = k_width;
}
let ssm_conv_params = SsmConvParams {
channels: qkv_channels,
n_tokens: seq_len,
n_seqs: FORWARD_DISPATCH_N_SEQS,
k_width,
};
let q_scaled =
super::decode_pool::pooled_alloc_buffer(device, n_q_elems * 4, DType::F32, vec![n_q_elems])
.map_err(|e| anyhow!("alloc q_scaled: {e}"))?;
let g_buf = super::decode_pool::pooled_alloc_buffer(device, g_n * 4, DType::F32, vec![g_n])
.map_err(|e| anyhow!("alloc g_buf: {e}"))?;
let beta_buf = super::decode_pool::pooled_alloc_buffer(device, g_n * 4, DType::F32, vec![g_n])
.map_err(|e| anyhow!("alloc beta_buf: {e}"))?;
let mut g_params_buf = super::decode_pool::pooled_alloc_buffer(device, 8, DType::U32, vec![2])
.map_err(|e| anyhow!("alloc g_params: {e}"))?;
{
let s = g_params_buf
.as_mut_slice::<u32>()
.map_err(|e| anyhow!("{e}"))?;
s[0] = n_v_heads;
s[1] = seq_len;
}
let gated_buf = super::decode_pool::pooled_alloc_buffer(
device,
gated_elems * 4,
DType::F32,
vec![gated_elems],
)
.map_err(|e| anyhow!("alloc op8 gated: {e}"))?;
let mut op8_params = super::decode_pool::pooled_alloc_buffer(device, 8, DType::F32, vec![2])
.map_err(|e| anyhow!("alloc op8 params (pooled): {e}"))?;
{
let s = op8_params
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("{e}"))?;
s[0] = rms_norm_eps;
s[1] = d_v as f32;
}
let attn_out_buf =
super::decode_pool::pooled_alloc_buffer(device, out_elems * 4, DType::F32, vec![out_elems])
.map_err(|e| anyhow!("alloc gdn output: {e}"))?;
let gdn_params = GatedDeltaNetParams {
d_k,
d_v,
n_k_heads,
n_v_heads,
n_tokens: seq_len,
n_seqs,
};
let mut gdn_params_buf =
super::decode_pool::pooled_alloc_buffer(device, 9 * 4, DType::U32, vec![9])
.map_err(|e| anyhow!("alloc gdn params (pooled): {e}"))?;
{
let s = gdn_params_buf
.as_mut_slice::<u32>()
.map_err(|e| anyhow!("{e}"))?;
s[0] = gdn_params.d_k;
s[1] = gdn_params.d_v;
s[2] = gdn_params.n_k_heads;
s[3] = gdn_params.n_v_heads;
s[4] = gdn_params.n_tokens;
s[5] = gdn_params.n_seqs;
s[6] = 0;
s[7] = 0;
s[8] = 0; }
let output = if seq == 1 {
let q_gpu = qkv_conv.slice_view(0, q_sp);
let k_gpu = qkv_conv.slice_view((q_sp * 4) as u64, k_sp);
let v_gpu = qkv_conv.slice_view(((q_sp + k_sp) * 4) as u64, nv * dv);
let mut enc = device.command_encoder().context("enc ops1-9 decode")?;
let x_norm = apply_pre_norm(
&mut enc,
registry,
device,
x,
&weights.attn_norm,
seq_len,
hidden_size,
rms_norm_eps,
)?;
enc.memory_barrier();
let qkv_raw = apply_proj(
&mut enc,
registry,
device,
&x_norm,
&weights.attn_qkv,
seq_len,
hidden_size,
qkv_channels,
)?;
let z = apply_proj(
&mut enc,
registry,
device,
&x_norm,
&weights.attn_gate,
seq_len,
hidden_size,
z_channels,
)?;
enc.memory_barrier();
if let Some(conv_capture_buf) = conv_state_capture {
dispatch_ssm_conv_with_capture(
&mut enc,
registry,
device.metal_device(),
&qkv_raw,
&weights.ssm_conv1d,
conv_state_in,
&qkv_conv,
conv_capture_buf,
&ssm_params_buf,
ssm_conv_params,
)
.context("dispatch_ssm_conv_with_capture ops3 decode (K=N spec)")?;
} else {
dispatch_ssm_conv(
&mut enc,
registry,
device.metal_device(),
&qkv_raw,
&weights.ssm_conv1d,
conv_state_in,
conv_state_out,
&qkv_conv,
&ssm_params_buf,
ssm_conv_params,
)
.context("dispatch_ssm_conv ops3")?;
}
enc.memory_barrier();
let q_l2 = apply_l2_norm_per_head(
&mut enc,
registry,
device,
&q_gpu,
seq_len,
n_k_heads,
d_k,
rms_norm_eps,
)?;
let k_normed = apply_l2_norm_per_head(
&mut enc,
registry,
device,
&k_gpu,
seq_len,
n_k_heads,
d_k,
rms_norm_eps,
)?;
let alpha_logit_buf = apply_proj(
&mut enc,
registry,
device,
&x_norm,
&weights.ssm_alpha,
seq_len,
hidden_size,
n_v_heads,
)?;
let beta_logit_buf = apply_proj(
&mut enc,
registry,
device,
&x_norm,
&weights.ssm_beta,
seq_len,
hidden_size,
n_v_heads,
)?;
enc.memory_barrier();
scalar_mul_f32(
&mut enc,
registry,
device.metal_device(),
&q_l2,
&q_scaled,
n_q_elems,
q_scale_val,
)
.context("scalar_mul_f32 q_scale")?;
dispatch_compute_g_beta(
&mut enc,
registry,
device.metal_device(),
&alpha_logit_buf,
&beta_logit_buf,
&weights.ssm_dt_bias,
&weights.ssm_a,
&g_buf,
&beta_buf,
&g_params_buf,
seq_len,
n_v_heads,
)
.context("dispatch_compute_g_beta")?;
enc.memory_barrier();
let nsg_compatible =
d_k % 32 == 0 && d_k / 32 <= mlx_native::ops::gated_delta_net_decode::MAX_NSG;
if nsg_compatible {
if let Some(capture_buf) = state_capture {
use mlx_native::ops::gated_delta_net_decode::dispatch_gated_delta_net_decode_with_capture;
dispatch_gated_delta_net_decode_with_capture(
&mut enc,
registry,
device.metal_device(),
&q_scaled,
&k_normed,
&v_gpu,
&g_buf,
&beta_buf,
state_in,
&attn_out_buf,
state_out,
&gdn_params_buf,
capture_buf,
gdn_params,
)
.context(
"dispatch_gated_delta_net_decode_with_capture (build_delta_net_layer K=N spec)",
)?;
} else {
dispatch_gated_delta_net_decode(
&mut enc,
registry,
device.metal_device(),
&q_scaled,
&k_normed,
&v_gpu,
&g_buf,
&beta_buf,
state_in,
&attn_out_buf,
state_out,
&gdn_params_buf,
gdn_params,
)
.context("dispatch_gated_delta_net_decode (build_delta_net_layer decode)")?;
}
} else {
if state_capture.is_some() {
return Err(anyhow!(
"build_delta_net_layer: state_capture is Some but D_k={} \
is not NSG-compatible (only D_k % 32 == 0 with D_k/32 \
<= MAX_NSG supports the capture kernel). Either disable \
capture for this shape or port the capture path to the \
legacy 128-thread kernel.",
d_k
));
}
dispatch_gated_delta_net(
&mut enc,
registry,
device.metal_device(),
&q_scaled,
&k_normed,
&v_gpu,
&g_buf,
&beta_buf,
state_in,
&attn_out_buf,
state_out,
&gdn_params_buf,
gdn_params,
)
.context("dispatch_gated_delta_net (build_delta_net_layer decode, non-NSG shape)")?;
}
enc.memory_barrier();
dispatch_ssm_norm_gate(
&mut enc,
registry,
device.metal_device(),
&attn_out_buf,
&weights.ssm_norm,
&z,
&gated_buf,
&op8_params,
rows_op8,
d_v,
)
.context("dispatch_ssm_norm_gate")?;
enc.memory_barrier();
let output = apply_proj(
&mut enc,
registry,
device,
&gated_buf,
&weights.ssm_out,
seq_len,
z_channels,
hidden_size,
)?;
enc.commit_labeled("layer.delta_net.ops1-9");
output
} else {
let (x_norm, qkv_conv_out, z) = {
let _w5b8 = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::LayerOps1to3,
);
let mut enc = device.command_encoder().context("enc ops1-3 prefill")?;
let x_norm = apply_pre_norm(
&mut enc,
registry,
device,
x,
&weights.attn_norm,
seq_len,
hidden_size,
rms_norm_eps,
)?;
enc.memory_barrier();
let qkv_raw = apply_proj(
&mut enc,
registry,
device,
&x_norm,
&weights.attn_qkv,
seq_len,
hidden_size,
qkv_channels,
)?;
let z = apply_proj(
&mut enc,
registry,
device,
&x_norm,
&weights.attn_gate,
seq_len,
hidden_size,
z_channels,
)?;
enc.memory_barrier();
if let Some(conv_capture_buf) = conv_state_capture {
dispatch_ssm_conv_with_capture(
&mut enc,
registry,
device.metal_device(),
&qkv_raw,
&weights.ssm_conv1d,
conv_state_in,
&qkv_conv,
conv_capture_buf,
&ssm_params_buf,
ssm_conv_params,
)
.context("dispatch_ssm_conv_with_capture ops3 prefill (K=N spec)")?;
} else {
dispatch_ssm_conv(
&mut enc,
registry,
device.metal_device(),
&qkv_raw,
&weights.ssm_conv1d,
conv_state_in,
conv_state_out,
&qkv_conv,
&ssm_params_buf,
ssm_conv_params,
)
.context("dispatch_ssm_conv ops3 prefill")?;
}
enc.commit_labeled("layer.gdn.ops1-3");
(x_norm, qkv_conv, z)
};
let v_sp = nv * dv;
let (q_gpu, k_gpu, v_gpu) = {
let _w5b8 = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::LayerQkvDeinterleave,
);
let _w5b17 = super::wave5b8_profile::Section::start_w5b17(
super::wave5b8_profile::SectionKind::DnQkvGpuSplit,
);
let q_gpu = super::decode_pool::pooled_alloc_buffer(
device,
seq * q_sp * 4,
DType::F32,
vec![seq, q_sp],
)
.map_err(|e| anyhow!("alloc q_gpu (W-5b.18): {e}"))?;
let k_gpu = super::decode_pool::pooled_alloc_buffer(
device,
seq * k_sp * 4,
DType::F32,
vec![seq, k_sp],
)
.map_err(|e| anyhow!("alloc k_gpu (W-5b.18): {e}"))?;
let v_gpu = super::decode_pool::pooled_alloc_buffer(
device,
seq * v_sp * 4,
DType::F32,
vec![seq, v_sp],
)
.map_err(|e| anyhow!("alloc v_gpu (W-5b.18): {e}"))?;
let params = QkvSplitParams {
seq: seq_len,
q_sp: q_sp as u32,
k_sp: k_sp as u32,
v_sp: v_sp as u32,
};
let mut enc = device
.command_encoder()
.context("enc qkv_split (W-5b.18) prefill")?;
dispatch_qkv_split_f32(
&mut enc,
registry,
device.metal_device(),
&qkv_conv_out,
&q_gpu,
&k_gpu,
&v_gpu,
¶ms,
)
.map_err(|e| anyhow!("dispatch_qkv_split_f32 (W-5b.18): {e}"))?;
enc.commit_labeled("layer.gdn.qkv_split");
(q_gpu, k_gpu, v_gpu)
};
let chunk_route = chunk_path_eligible(seq_len, d_k);
{
use std::sync::atomic::{AtomicBool, Ordering};
static CHUNK_LOGGED: AtomicBool = AtomicBool::new(false);
static AUTOREG_LOGGED: AtomicBool = AtomicBool::new(false);
if chunk_route {
if !CHUNK_LOGGED.swap(true, Ordering::Relaxed) {
tracing::info!(
target: "hf2q::wave5b3",
seq_len, d_k,
"chunk-pipeline ENGAGED for prefill (HF2Q_CHUNK_SCAN_PREFILL=1)"
);
}
} else if INVESTIGATION_ENV.chunk_scan_prefill
&& !AUTOREG_LOGGED.swap(true, Ordering::Relaxed)
{
tracing::info!(
target: "hf2q::wave5b3",
seq_len, d_k,
"chunk-pipeline gate set but predicate FAILED — autoregressive path running (seq_len%64={}, d_k={})",
seq_len % 64, d_k
);
}
}
let output = if chunk_route {
let k_normed_buf = {
let _w5b8 = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::LayerChunkPrep,
);
let mut enc = device.command_encoder().context("enc chunk-prep prefill")?;
let q_l2 = apply_l2_norm_per_head(
&mut enc,
registry,
device,
&q_gpu,
seq_len,
n_k_heads,
d_k,
rms_norm_eps,
)?;
let k_normed = apply_l2_norm_per_head(
&mut enc,
registry,
device,
&k_gpu,
seq_len,
n_k_heads,
d_k,
rms_norm_eps,
)?;
let alpha_logit_buf = apply_proj(
&mut enc,
registry,
device,
&x_norm,
&weights.ssm_alpha,
seq_len,
hidden_size,
n_v_heads,
)?;
let beta_logit_buf = apply_proj(
&mut enc,
registry,
device,
&x_norm,
&weights.ssm_beta,
seq_len,
hidden_size,
n_v_heads,
)?;
enc.memory_barrier();
scalar_mul_f32(
&mut enc,
registry,
device.metal_device(),
&q_l2,
&q_scaled,
n_q_elems,
q_scale_val,
)
.context("scalar_mul_f32 q_scale chunk prefill")?;
dispatch_compute_g_beta(
&mut enc,
registry,
device.metal_device(),
&alpha_logit_buf,
&beta_logit_buf,
&weights.ssm_dt_bias,
&weights.ssm_a,
&g_buf,
&beta_buf,
&g_params_buf,
seq_len,
n_v_heads,
)
.context("dispatch_compute_g_beta chunk prefill")?;
enc.commit_and_wait().context("commit chunk-prep prefill")?;
k_normed
};
{
let _w5b8 = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::LayerChunkCall,
);
apply_gated_delta_net_chunk(
device,
registry,
&q_scaled,
&k_normed_buf,
&v_gpu,
&g_buf,
&beta_buf,
state_in,
&attn_out_buf,
state_out,
seq_len,
n_k_heads,
n_v_heads,
d_k,
d_v,
false,
)
.context("apply_gated_delta_net_chunk prefill")?;
}
let _w5b8_ops89 = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::LayerChunkOps8to9,
);
let mut enc = device
.command_encoder()
.context("enc chunk ops8-9 prefill")?;
dispatch_ssm_norm_gate(
&mut enc,
registry,
device.metal_device(),
&attn_out_buf,
&weights.ssm_norm,
&z,
&gated_buf,
&op8_params,
rows_op8,
d_v,
)
.context("dispatch_ssm_norm_gate chunk prefill")?;
enc.memory_barrier();
let output = apply_proj(
&mut enc,
registry,
device,
&gated_buf,
&weights.ssm_out,
seq_len,
z_channels,
hidden_size,
)?;
enc.commit_and_wait_labeled("layer.gdn.ops8-9")
.context("commit chunk ops8-9 prefill")?;
output
} else {
let _w5b8 = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::LayerAutoregOps5to9,
);
let mut enc = device.command_encoder().context("enc ops5-9 prefill")?;
let q_l2 = apply_l2_norm_per_head(
&mut enc,
registry,
device,
&q_gpu,
seq_len,
n_k_heads,
d_k,
rms_norm_eps,
)?;
let k_normed = apply_l2_norm_per_head(
&mut enc,
registry,
device,
&k_gpu,
seq_len,
n_k_heads,
d_k,
rms_norm_eps,
)?;
let alpha_logit_buf = apply_proj(
&mut enc,
registry,
device,
&x_norm,
&weights.ssm_alpha,
seq_len,
hidden_size,
n_v_heads,
)?;
let beta_logit_buf = apply_proj(
&mut enc,
registry,
device,
&x_norm,
&weights.ssm_beta,
seq_len,
hidden_size,
n_v_heads,
)?;
enc.memory_barrier();
scalar_mul_f32(
&mut enc,
registry,
device.metal_device(),
&q_l2,
&q_scaled,
n_q_elems,
q_scale_val,
)
.context("scalar_mul_f32 q_scale prefill")?;
dispatch_compute_g_beta(
&mut enc,
registry,
device.metal_device(),
&alpha_logit_buf,
&beta_logit_buf,
&weights.ssm_dt_bias,
&weights.ssm_a,
&g_buf,
&beta_buf,
&g_params_buf,
seq_len,
n_v_heads,
)
.context("dispatch_compute_g_beta prefill")?;
enc.memory_barrier();
let decode_kernel_eligible = d_k % 32 == 0 && d_k <= 128;
if decode_kernel_eligible {
if let Some(capture_buf) = state_capture {
use mlx_native::ops::gated_delta_net_decode::dispatch_gated_delta_net_decode_with_capture;
dispatch_gated_delta_net_decode_with_capture(
&mut enc,
registry,
device.metal_device(),
&q_scaled,
&k_normed,
&v_gpu,
&g_buf,
&beta_buf,
state_in,
&attn_out_buf,
state_out,
&gdn_params_buf,
capture_buf,
gdn_params,
)
.context("dispatch_gated_delta_net_decode_with_capture prefill (K=N spec)")?;
} else {
dispatch_gated_delta_net_decode(
&mut enc,
registry,
device.metal_device(),
&q_scaled,
&k_normed,
&v_gpu,
&g_buf,
&beta_buf,
state_in,
&attn_out_buf,
state_out,
&gdn_params_buf,
gdn_params,
)
.context("dispatch_gated_delta_net_decode prefill")?;
}
} else {
if state_capture.is_some() {
return Err(anyhow!(
"build_delta_net_layer prefill: state_capture is Some but \
D_k={} not decode-kernel-eligible (D_k % 32 == 0 && D_k <= 128); \
capture path requires NSG-compatible shape.",
d_k
));
}
dispatch_gated_delta_net(
&mut enc,
registry,
device.metal_device(),
&q_scaled,
&k_normed,
&v_gpu,
&g_buf,
&beta_buf,
state_in,
&attn_out_buf,
state_out,
&gdn_params_buf,
gdn_params,
)
.context("dispatch_gated_delta_net prefill (fallback)")?;
}
enc.memory_barrier();
dispatch_ssm_norm_gate(
&mut enc,
registry,
device.metal_device(),
&attn_out_buf,
&weights.ssm_norm,
&z,
&gated_buf,
&op8_params,
rows_op8,
d_v,
)
.context("dispatch_ssm_norm_gate prefill")?;
enc.memory_barrier();
let output = apply_proj(
&mut enc,
registry,
device,
&gated_buf,
&weights.ssm_out,
seq_len,
z_channels,
hidden_size,
)?;
enc.commit_labeled("layer.gdn.ops5-9");
output
};
output
};
Ok(output)
}
#[allow(clippy::too_many_arguments)]
pub fn build_delta_net_layer_with_arena(
device: &MlxDevice,
registry: &mut KernelRegistry,
x: &MlxBuffer,
weights: &DeltaNetWeightsGpu,
conv_state_in: &MlxBuffer,
conv_state_out: &MlxBuffer,
state_in: &MlxBuffer,
state_out: &MlxBuffer,
arena: &mut super::DnPrefillArena,
chunk_allocs_arena: Option<&mut super::ChunkAllocsArena>,
chunk_internal_arena: Option<&mut ChunkInternalArena>,
seq_len: u32,
hidden_size: u32,
n_k_heads: u32,
n_v_heads: u32,
d_k: u32,
d_v: u32,
k_width: u32,
rms_norm_eps: f32,
layer_session: Option<&mut mlx_native::EncoderSession>,
slot_id: SlotId,
) -> Result<MlxBuffer> {
debug_assert!(
seq_len > 1,
"build_delta_net_layer_with_arena: seq_len must be > 1 (decode uses pooled path)"
);
arena
.validate_fits(seq_len, hidden_size, n_k_heads, n_v_heads, d_k, d_v)
.context("DnPrefillArena shape mismatch")?;
let qkv_channels = 2 * n_k_heads * d_k + n_v_heads * d_v;
let z_channels = n_v_heads * d_v;
let q_span = n_k_heads * d_k;
let k_span = n_k_heads * d_k;
let per_seq_rec_elems = (d_k as usize) * (d_v as usize) * (n_v_heads as usize);
let n_seqs_alloc: u32 = if per_seq_rec_elems == 0 {
1
} else {
let total = state_in.element_count();
if total % per_seq_rec_elems != 0 {
return Err(anyhow!(
"build_delta_net_layer_with_arena: state_in element_count {} not \
divisible by per-seq recurrent elems {}. ADR-040 §6.1.40 lift.",
total,
per_seq_rec_elems,
));
}
(total / per_seq_rec_elems) as u32
};
if slot_id.0 >= n_seqs_alloc {
return Err(anyhow!(
"build_delta_net_layer_with_arena: slot_id={} out of range \
(state_in n_seqs={}). ADR-040 Phase A2b-cont per-slot routing contract.",
slot_id.0,
n_seqs_alloc,
));
}
let la_view = narrow_la_ping_pong_to_slot(
conv_state_in,
conv_state_out,
state_in,
state_out,
slot_id,
d_k,
d_v,
n_v_heads,
qkv_channels,
k_width - 1,
);
let conv_state_in = &la_view.conv_state_in;
let conv_state_out = &la_view.conv_state_out;
let state_in = &la_view.recurrent_in;
let state_out = &la_view.recurrent_out;
let _seq = seq_len as usize;
let nv = n_v_heads as usize;
let dk = d_k as usize;
let dv = d_v as usize;
let _qkv_ch = qkv_channels as usize;
let q_sp = q_span as usize;
let k_sp = k_span as usize;
let n_q_elems = (seq_len * n_k_heads * d_k) as usize;
let q_scale_val = 1.0_f32 / (dk as f32).sqrt();
let rows_op8 = seq_len * n_v_heads;
let n_seqs = FORWARD_DISPATCH_N_SEQS;
{
let s = arena
.ssm_params_buf
.as_mut_slice::<u32>()
.map_err(|e| anyhow!("DnPrefillArena ssm_params: {e}"))?;
s[0] = qkv_channels;
s[1] = seq_len;
s[2] = FORWARD_DISPATCH_N_SEQS;
s[3] = k_width;
}
{
let s = arena
.g_params_buf
.as_mut_slice::<u32>()
.map_err(|e| anyhow!("DnPrefillArena g_params: {e}"))?;
s[0] = n_v_heads;
s[1] = seq_len;
}
{
let s = arena
.op8_params_buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("DnPrefillArena op8_params: {e}"))?;
s[0] = rms_norm_eps;
s[1] = d_v as f32;
}
let gdn_params = GatedDeltaNetParams {
d_k,
d_v,
n_k_heads,
n_v_heads,
n_tokens: seq_len,
n_seqs,
};
{
let s = arena
.gdn_params_buf
.as_mut_slice::<u32>()
.map_err(|e| anyhow!("DnPrefillArena gdn_params: {e}"))?;
s[0] = gdn_params.d_k;
s[1] = gdn_params.d_v;
s[2] = gdn_params.n_k_heads;
s[3] = gdn_params.n_v_heads;
s[4] = gdn_params.n_tokens;
s[5] = gdn_params.n_seqs;
s[6] = 0;
s[7] = 0;
s[8] = 0; }
let ssm_conv_params = SsmConvParams {
channels: qkv_channels,
n_tokens: seq_len,
n_seqs: FORWARD_DISPATCH_N_SEQS,
k_width,
};
let v_sp = nv * dv;
{
let _w5b8 = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::LayerOps1to3,
);
let _w5b17 = super::wave5b8_profile::Section::start_w5b17(
super::wave5b8_profile::SectionKind::DnQkvGpuSplit,
);
let mut enc = LayerEncoder::from_session_or_plain(device, layer_session)
.context("enc gdn stage_a prefill (arena)")?;
apply_pre_norm_into(
enc.encoder(),
registry,
device,
x,
&weights.attn_norm,
&arena.x_norm_buf,
&mut arena.pre_norm_params_buf,
seq_len,
hidden_size,
rms_norm_eps,
)?;
enc.encoder().memory_barrier();
apply_proj_into(
enc.encoder(),
registry,
device,
&arena.x_norm_buf,
&weights.attn_qkv,
&mut arena.qkv_raw_buf,
seq_len,
hidden_size,
qkv_channels,
)?;
apply_proj_into(
enc.encoder(),
registry,
device,
&arena.x_norm_buf,
&weights.attn_gate,
&mut arena.z_buf,
seq_len,
hidden_size,
z_channels,
)?;
enc.encoder().memory_barrier();
dispatch_ssm_conv(
enc.encoder(),
registry,
device.metal_device(),
&arena.qkv_raw_buf,
&weights.ssm_conv1d,
conv_state_in,
conv_state_out,
&arena.qkv_conv_buf,
&arena.ssm_params_buf,
ssm_conv_params,
)
.context("dispatch_ssm_conv ops3 (arena)")?;
enc.encoder().memory_barrier();
let qkv_split_params = QkvSplitParams {
seq: seq_len,
q_sp: q_sp as u32,
k_sp: k_sp as u32,
v_sp: v_sp as u32,
};
dispatch_qkv_split_f32(
enc.encoder(),
registry,
device.metal_device(),
&arena.qkv_conv_buf,
&arena.q_split_buf,
&arena.k_split_buf,
&arena.v_split_buf,
&qkv_split_params,
)
.map_err(|e| anyhow!("dispatch_qkv_split_f32 (arena): {e}"))?;
enc.fence_or_commit("layer.gdn.stage_a")
.context("fence/commit DN stage_a (arena)")?;
}
let _ = (k_sp, q_sp, v_sp);
let chunk_route = chunk_path_eligible(seq_len, d_k);
{
use std::sync::atomic::{AtomicBool, Ordering};
static CHUNK_LOGGED: AtomicBool = AtomicBool::new(false);
static AUTOREG_LOGGED: AtomicBool = AtomicBool::new(false);
if chunk_route {
if !CHUNK_LOGGED.swap(true, Ordering::Relaxed) {
tracing::info!(
target: "hf2q::wave5b3",
seq_len, d_k,
"chunk-pipeline ENGAGED for prefill (HF2Q_CHUNK_SCAN_PREFILL=1) [arena]"
);
}
} else if INVESTIGATION_ENV.chunk_scan_prefill
&& !AUTOREG_LOGGED.swap(true, Ordering::Relaxed)
{
tracing::info!(
target: "hf2q::wave5b3",
seq_len, d_k,
"chunk-pipeline gate set but predicate FAILED — autoregressive path running (seq_len%64={}, d_k={}) [arena]",
seq_len % 64, d_k
);
}
}
let output = if chunk_route {
{
let _w5b8 = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::LayerChunkPrep,
);
let mut enc = device
.command_encoder()
.context("enc chunk-prep prefill (arena)")?;
apply_l2_norm_per_head_into(
&mut enc,
registry,
device,
&arena.q_split_buf,
&arena.q_l2_buf,
&mut arena.l2_params_q_buf,
seq_len,
n_k_heads,
d_k,
rms_norm_eps,
)?;
apply_l2_norm_per_head_into(
&mut enc,
registry,
device,
&arena.k_split_buf,
&arena.k_normed_buf,
&mut arena.l2_params_k_buf,
seq_len,
n_k_heads,
d_k,
rms_norm_eps,
)?;
apply_proj_into(
&mut enc,
registry,
device,
&arena.x_norm_buf,
&weights.ssm_alpha,
&mut arena.alpha_logit_buf,
seq_len,
hidden_size,
n_v_heads,
)?;
apply_proj_into(
&mut enc,
registry,
device,
&arena.x_norm_buf,
&weights.ssm_beta,
&mut arena.beta_logit_buf,
seq_len,
hidden_size,
n_v_heads,
)?;
enc.memory_barrier();
scalar_mul_f32(
&mut enc,
registry,
device.metal_device(),
&arena.q_l2_buf,
&arena.q_scaled_buf,
n_q_elems,
q_scale_val,
)
.context("scalar_mul_f32 q_scale chunk prefill (arena)")?;
dispatch_compute_g_beta(
&mut enc,
registry,
device.metal_device(),
&arena.alpha_logit_buf,
&arena.beta_logit_buf,
&weights.ssm_dt_bias,
&weights.ssm_a,
&arena.g_buf,
&arena.beta_buf,
&arena.g_params_buf,
seq_len,
n_v_heads,
)
.context("dispatch_compute_g_beta chunk prefill (arena)")?;
enc.commit_and_wait()
.context("commit chunk-prep prefill (arena)")?;
}
{
let _w5b8 = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::LayerChunkCall,
);
if let Some(chunk_arena) = chunk_allocs_arena {
apply_gated_delta_net_chunk_with_arena(
device,
registry,
&arena.q_scaled_buf,
&arena.k_normed_buf,
&arena.v_split_buf,
&arena.g_buf,
&arena.beta_buf,
state_in,
&arena.attn_out_buf,
state_out,
chunk_arena,
chunk_internal_arena,
seq_len,
n_k_heads,
n_v_heads,
d_k,
d_v,
false,
)
.context("apply_gated_delta_net_chunk_with_arena prefill (arena)")?;
} else {
apply_gated_delta_net_chunk(
device,
registry,
&arena.q_scaled_buf,
&arena.k_normed_buf,
&arena.v_split_buf,
&arena.g_buf,
&arena.beta_buf,
state_in,
&arena.attn_out_buf,
state_out,
seq_len,
n_k_heads,
n_v_heads,
d_k,
d_v,
false,
)
.context("apply_gated_delta_net_chunk prefill (arena)")?;
}
}
let _w5b8_ops89 = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::LayerChunkOps8to9,
);
let mut enc = device
.command_encoder()
.context("enc chunk ops8-9 prefill (arena)")?;
dispatch_ssm_norm_gate(
&mut enc,
registry,
device.metal_device(),
&arena.attn_out_buf,
&weights.ssm_norm,
&arena.z_buf,
&arena.gated_buf,
&arena.op8_params_buf,
rows_op8,
d_v,
)
.context("dispatch_ssm_norm_gate chunk prefill (arena)")?;
enc.memory_barrier();
let output = apply_proj(
&mut enc,
registry,
device,
&arena.gated_buf,
&weights.ssm_out,
seq_len,
z_channels,
hidden_size,
)?;
enc.commit_and_wait_labeled("layer.gdn.ops8-9")
.context("commit chunk ops8-9 prefill (arena)")?;
output
} else {
let _w5b8 = super::wave5b8_profile::Section::start(
super::wave5b8_profile::SectionKind::LayerAutoregOps5to9,
);
let mut enc = device
.command_encoder()
.context("enc ops5-9 prefill (arena)")?;
apply_l2_norm_per_head_into(
&mut enc,
registry,
device,
&arena.q_split_buf,
&arena.q_l2_buf,
&mut arena.l2_params_q_buf,
seq_len,
n_k_heads,
d_k,
rms_norm_eps,
)?;
apply_l2_norm_per_head_into(
&mut enc,
registry,
device,
&arena.k_split_buf,
&arena.k_normed_buf,
&mut arena.l2_params_k_buf,
seq_len,
n_k_heads,
d_k,
rms_norm_eps,
)?;
apply_proj_into(
&mut enc,
registry,
device,
&arena.x_norm_buf,
&weights.ssm_alpha,
&mut arena.alpha_logit_buf,
seq_len,
hidden_size,
n_v_heads,
)?;
apply_proj_into(
&mut enc,
registry,
device,
&arena.x_norm_buf,
&weights.ssm_beta,
&mut arena.beta_logit_buf,
seq_len,
hidden_size,
n_v_heads,
)?;
enc.memory_barrier();
scalar_mul_f32(
&mut enc,
registry,
device.metal_device(),
&arena.q_l2_buf,
&arena.q_scaled_buf,
n_q_elems,
q_scale_val,
)
.context("scalar_mul_f32 q_scale prefill (arena)")?;
dispatch_compute_g_beta(
&mut enc,
registry,
device.metal_device(),
&arena.alpha_logit_buf,
&arena.beta_logit_buf,
&weights.ssm_dt_bias,
&weights.ssm_a,
&arena.g_buf,
&arena.beta_buf,
&arena.g_params_buf,
seq_len,
n_v_heads,
)
.context("dispatch_compute_g_beta prefill (arena)")?;
enc.memory_barrier();
let decode_kernel_eligible = d_k % 32 == 0 && d_k <= 128;
if decode_kernel_eligible {
dispatch_gated_delta_net_decode(
&mut enc,
registry,
device.metal_device(),
&arena.q_scaled_buf,
&arena.k_normed_buf,
&arena.v_split_buf,
&arena.g_buf,
&arena.beta_buf,
state_in,
&arena.attn_out_buf,
state_out,
&arena.gdn_params_buf,
gdn_params,
)
.context("dispatch_gated_delta_net_decode prefill (arena)")?;
} else {
dispatch_gated_delta_net(
&mut enc,
registry,
device.metal_device(),
&arena.q_scaled_buf,
&arena.k_normed_buf,
&arena.v_split_buf,
&arena.g_buf,
&arena.beta_buf,
state_in,
&arena.attn_out_buf,
state_out,
&arena.gdn_params_buf,
gdn_params,
)
.context("dispatch_gated_delta_net prefill (arena fallback)")?;
}
enc.memory_barrier();
dispatch_ssm_norm_gate(
&mut enc,
registry,
device.metal_device(),
&arena.attn_out_buf,
&weights.ssm_norm,
&arena.z_buf,
&arena.gated_buf,
&arena.op8_params_buf,
rows_op8,
d_v,
)
.context("dispatch_ssm_norm_gate prefill (arena)")?;
enc.memory_barrier();
let output = apply_proj(
&mut enc,
registry,
device,
&arena.gated_buf,
&weights.ssm_out,
seq_len,
z_channels,
hidden_size,
)?;
enc.commit_labeled("layer.gdn.ops5-9");
output
};
Ok(output)
}
#[allow(clippy::too_many_arguments)]
pub fn build_delta_net_layer_decode_into(
enc: &mut mlx_native::CommandEncoder,
device: &MlxDevice,
registry: &mut KernelRegistry,
x: &MlxBuffer,
weights: &DeltaNetWeightsGpu,
conv_state_in: &MlxBuffer,
conv_state_out: &MlxBuffer,
state_in: &MlxBuffer,
state_out: &MlxBuffer,
seq_len: u32,
hidden_size: u32,
n_k_heads: u32,
n_v_heads: u32,
d_k: u32,
d_v: u32,
k_width: u32,
rms_norm_eps: f32,
slot_id: SlotId,
) -> Result<MlxBuffer> {
debug_assert_eq!(
seq_len, 1,
"build_delta_net_layer_decode_into: seq_len must be 1"
);
let qkv_channels = 2 * n_k_heads * d_k + n_v_heads * d_v;
let z_channels = n_v_heads * d_v;
let q_span = n_k_heads * d_k;
let k_span = n_k_heads * d_k;
let nv = n_v_heads as usize;
let dk = d_k as usize;
let dv = d_v as usize;
let qkv_ch = qkv_channels as usize;
let q_sp = q_span as usize;
let k_sp = k_span as usize;
let per_seq_rec_elems = (d_k as usize) * (d_v as usize) * (n_v_heads as usize);
let n_seqs_alloc: u32 = if per_seq_rec_elems == 0 {
1
} else {
let total = state_in.element_count();
if total % per_seq_rec_elems != 0 {
return Err(anyhow!(
"build_delta_net_layer_decode_into: state_in element_count {} not \
divisible by per-seq recurrent elems {}. ADR-040 §6.1.40 lift.",
total,
per_seq_rec_elems,
));
}
(total / per_seq_rec_elems) as u32
};
if slot_id.0 >= n_seqs_alloc {
return Err(anyhow!(
"build_delta_net_layer_decode_into: slot_id={} out of range \
(state_in n_seqs={}). ADR-040 Phase A2b-cont per-slot routing contract.",
slot_id.0,
n_seqs_alloc,
));
}
let la_view = narrow_la_ping_pong_to_slot(
conv_state_in,
conv_state_out,
state_in,
state_out,
slot_id,
d_k,
d_v,
n_v_heads,
qkv_channels,
k_width - 1,
);
let conv_state_in = &la_view.conv_state_in;
let conv_state_out = &la_view.conv_state_out;
let state_in = &la_view.recurrent_in;
let state_out = &la_view.recurrent_out;
let n_q_elems = (seq_len * n_k_heads * d_k) as usize;
let q_scale_val = 1.0_f32 / (dk as f32).sqrt();
let g_n = (seq_len * n_v_heads) as usize;
let rows_op8 = seq_len * n_v_heads;
let gated_elems = (rows_op8 * d_v) as usize;
let n_seqs = FORWARD_DISPATCH_N_SEQS;
let out_elems = (n_v_heads * seq_len * d_v) as usize;
let qkv_conv = super::decode_pool::pooled_alloc_buffer(
device,
(seq_len * qkv_channels) as usize * 4,
DType::F32,
vec![seq_len as usize, qkv_ch],
)
.map_err(|e| anyhow!("alloc qkv_conv (decode_into): {e}"))?;
let mut ssm_params_buf =
super::decode_pool::pooled_alloc_buffer(device, 4 * 4, DType::U32, vec![4])
.map_err(|e| anyhow!("alloc ssm params (decode_into): {e}"))?;
{
let s = ssm_params_buf
.as_mut_slice::<u32>()
.map_err(|e| anyhow!("{e}"))?;
s[0] = qkv_channels;
s[1] = seq_len;
s[2] = FORWARD_DISPATCH_N_SEQS;
s[3] = k_width;
}
let ssm_conv_params = SsmConvParams {
channels: qkv_channels,
n_tokens: seq_len,
n_seqs: FORWARD_DISPATCH_N_SEQS,
k_width,
};
let q_scaled =
super::decode_pool::pooled_alloc_buffer(device, n_q_elems * 4, DType::F32, vec![n_q_elems])
.map_err(|e| anyhow!("alloc q_scaled (decode_into): {e}"))?;
let g_buf = super::decode_pool::pooled_alloc_buffer(device, g_n * 4, DType::F32, vec![g_n])
.map_err(|e| anyhow!("alloc g_buf (decode_into): {e}"))?;
let beta_buf = super::decode_pool::pooled_alloc_buffer(device, g_n * 4, DType::F32, vec![g_n])
.map_err(|e| anyhow!("alloc beta_buf (decode_into): {e}"))?;
let mut g_params_buf = super::decode_pool::pooled_alloc_buffer(device, 8, DType::U32, vec![2])
.map_err(|e| anyhow!("alloc g_params (decode_into): {e}"))?;
{
let s = g_params_buf
.as_mut_slice::<u32>()
.map_err(|e| anyhow!("{e}"))?;
s[0] = n_v_heads;
s[1] = seq_len;
}
let gated_buf = super::decode_pool::pooled_alloc_buffer(
device,
gated_elems * 4,
DType::F32,
vec![gated_elems],
)
.map_err(|e| anyhow!("alloc op8 gated (decode_into): {e}"))?;
let mut op8_params = super::decode_pool::pooled_alloc_buffer(device, 8, DType::F32, vec![2])
.map_err(|e| anyhow!("alloc op8 params (decode_into, pooled): {e}"))?;
{
let s = op8_params
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("{e}"))?;
s[0] = rms_norm_eps;
s[1] = d_v as f32;
}
let attn_out_buf =
super::decode_pool::pooled_alloc_buffer(device, out_elems * 4, DType::F32, vec![out_elems])
.map_err(|e| anyhow!("alloc gdn output (decode_into): {e}"))?;
let gdn_params = GatedDeltaNetParams {
d_k,
d_v,
n_k_heads,
n_v_heads,
n_tokens: seq_len,
n_seqs,
};
let mut gdn_params_buf =
super::decode_pool::pooled_alloc_buffer(device, 9 * 4, DType::U32, vec![9])
.map_err(|e| anyhow!("alloc gdn params (decode_into, pooled): {e}"))?;
{
let s = gdn_params_buf
.as_mut_slice::<u32>()
.map_err(|e| anyhow!("{e}"))?;
s[0] = gdn_params.d_k;
s[1] = gdn_params.d_v;
s[2] = gdn_params.n_k_heads;
s[3] = gdn_params.n_v_heads;
s[4] = gdn_params.n_tokens;
s[5] = gdn_params.n_seqs;
s[6] = 0;
s[7] = 0;
s[8] = 0; }
let q_gpu = qkv_conv.slice_view(0, q_sp);
let k_gpu = qkv_conv.slice_view((q_sp * 4) as u64, k_sp);
let v_gpu = qkv_conv.slice_view(((q_sp + k_sp) * 4) as u64, nv * dv);
let x_norm = apply_pre_norm(
enc,
registry,
device,
x,
&weights.attn_norm,
seq_len,
hidden_size,
rms_norm_eps,
)?;
enc.memory_barrier();
let qkv_raw = apply_proj(
enc,
registry,
device,
&x_norm,
&weights.attn_qkv,
seq_len,
hidden_size,
qkv_channels,
)?;
let z = apply_proj(
enc,
registry,
device,
&x_norm,
&weights.attn_gate,
seq_len,
hidden_size,
z_channels,
)?;
enc.memory_barrier();
dispatch_ssm_conv(
enc,
registry,
device.metal_device(),
&qkv_raw,
&weights.ssm_conv1d,
conv_state_in,
conv_state_out,
&qkv_conv,
&ssm_params_buf,
ssm_conv_params,
)
.context("dispatch_ssm_conv ops3 (decode_into)")?;
enc.memory_barrier();
let q_l2 = apply_l2_norm_per_head(
enc,
registry,
device,
&q_gpu,
seq_len,
n_k_heads,
d_k,
rms_norm_eps,
)?;
let k_normed = apply_l2_norm_per_head(
enc,
registry,
device,
&k_gpu,
seq_len,
n_k_heads,
d_k,
rms_norm_eps,
)?;
let alpha_logit_buf = apply_proj(
enc,
registry,
device,
&x_norm,
&weights.ssm_alpha,
seq_len,
hidden_size,
n_v_heads,
)?;
let beta_logit_buf = apply_proj(
enc,
registry,
device,
&x_norm,
&weights.ssm_beta,
seq_len,
hidden_size,
n_v_heads,
)?;
enc.memory_barrier();
scalar_mul_f32(
enc,
registry,
device.metal_device(),
&q_l2,
&q_scaled,
n_q_elems,
q_scale_val,
)
.context("scalar_mul_f32 q_scale (decode_into)")?;
dispatch_compute_g_beta(
enc,
registry,
device.metal_device(),
&alpha_logit_buf,
&beta_logit_buf,
&weights.ssm_dt_bias,
&weights.ssm_a,
&g_buf,
&beta_buf,
&g_params_buf,
seq_len,
n_v_heads,
)
.context("dispatch_compute_g_beta (decode_into)")?;
enc.memory_barrier();
dispatch_gated_delta_net_decode(
enc,
registry,
device.metal_device(),
&q_scaled,
&k_normed,
&v_gpu,
&g_buf,
&beta_buf,
state_in,
&attn_out_buf,
state_out,
&gdn_params_buf,
gdn_params,
)
.context("dispatch_gated_delta_net_decode (decode_into)")?;
enc.memory_barrier();
dispatch_ssm_norm_gate(
enc,
registry,
device.metal_device(),
&attn_out_buf,
&weights.ssm_norm,
&z,
&gated_buf,
&op8_params,
rows_op8,
d_v,
)
.context("dispatch_ssm_norm_gate (decode_into)")?;
enc.memory_barrier();
let output = apply_proj(
enc,
registry,
device,
&gated_buf,
&weights.ssm_out,
seq_len,
z_channels,
hidden_size,
)?;
Ok(output)
}
#[cfg(test)]
mod tests {
use super::super::delta_net::{
delta_net_layer_cpu_ref, DeltaNetLayerShape, DeltaNetLayerWeights,
};
use super::*;
fn mk_rand(seed: &mut u32, n: usize, scale: f32) -> Vec<f32> {
(0..n)
.map(|_| {
*seed = seed.wrapping_mul(1103515245).wrapping_add(12345);
((*seed as i32 as f32) / (i32::MAX as f32)) * scale
})
.collect()
}
fn small_shape() -> DeltaNetLayerShape {
DeltaNetLayerShape {
hidden_size: 32,
n_k_heads: 2,
n_v_heads: 4,
d_k: 8,
d_v: 8,
conv_kernel: 4,
rms_norm_eps: 1e-6,
}
}
fn synthetic_weights(shape: DeltaNetLayerShape, seed_init: u32) -> DeltaNetLayerWeights {
let h = shape.hidden_size as usize;
let _nk = shape.n_k_heads as usize;
let nv = shape.n_v_heads as usize;
let _dk = shape.d_k as usize;
let dv = shape.d_v as usize;
let k_width = shape.conv_kernel as usize;
let qkv_channels = shape.qkv_channels() as usize;
let z_channels = nv * dv;
let mut seed = seed_init;
DeltaNetLayerWeights {
attn_norm: {
let mut v = vec![1.0f32; h];
for (i, x) in v.iter_mut().enumerate() {
*x += 0.01 * (i as f32);
}
v
},
post_attn_norm: vec![1.0f32; h],
attn_qkv: mk_rand(&mut seed, qkv_channels * h, 0.1),
attn_gate: mk_rand(&mut seed, z_channels * h, 0.1),
ssm_conv1d: mk_rand(&mut seed, k_width * qkv_channels, 0.1),
ssm_alpha: mk_rand(&mut seed, nv * h, 0.1),
ssm_dt_bias: mk_rand(&mut seed, nv, 0.05),
ssm_beta: mk_rand(&mut seed, nv * h, 0.1),
ssm_a: mk_rand(&mut seed, nv, 0.1),
ssm_norm: {
let mut v = vec![1.0f32; dv];
for (i, x) in v.iter_mut().enumerate() {
*x += 0.01 * (i as f32);
}
v
},
ssm_out: mk_rand(&mut seed, h * z_channels, 0.1),
}
}
#[test]
fn full_delta_net_layer_gpu_matches_cpu_ref() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let shape = small_shape();
let weights_cpu = synthetic_weights(shape, 0x9ABC);
let seq_len = 4u32;
let h = shape.hidden_size as usize;
let seq = seq_len as usize;
let km1 = (shape.conv_kernel - 1) as usize;
let qkv_channels = shape.qkv_channels() as usize;
let state_size = (shape.d_k * shape.d_v * shape.n_v_heads) as usize;
let x_cpu: Vec<f32> = (0..seq * h).map(|i| 0.01 * (i as f32) - 0.5).collect();
let state_in = vec![0.0f32; state_size];
let conv_state = vec![0.0f32; km1 * qkv_channels];
let (cpu_out, _, _) =
delta_net_layer_cpu_ref(&x_cpu, &weights_cpu, shape, &state_in, &conv_state);
assert!(cpu_out.iter().all(|v| v.is_finite()), "CPU ref non-finite");
assert_eq!(cpu_out.len(), seq * h);
let gpu_weights = DeltaNetWeightsGpu::from_cpu_f32(
&weights_cpu,
&device,
shape.conv_kernel as usize,
qkv_channels,
)
.expect("from_cpu_f32");
let x_gpu = upload_f32(&x_cpu, &device).expect("upload x");
let state_in_gpu = upload_f32(&state_in, &device).expect("upload state_in");
let state_out_gpu = upload_f32(&state_in, &device).expect("upload state_out scratch");
let conv_state_in_gpu = upload_f32(&conv_state, &device).expect("upload conv_state_in");
let conv_state_out_gpu = upload_f32(&conv_state, &device).expect("alloc conv_state_out");
let gpu_out_buf = build_delta_net_layer(
&device,
&mut registry,
&x_gpu,
&gpu_weights,
&conv_state_in_gpu,
&conv_state_out_gpu,
&state_in_gpu,
&state_out_gpu,
seq_len,
shape.hidden_size,
shape.n_k_heads,
shape.n_v_heads,
shape.d_k,
shape.d_v,
shape.conv_kernel,
shape.rms_norm_eps,
None, None, SlotId(0), )
.expect("build_delta_net_layer");
device
.command_encoder()
.expect("sync enc full_delta_net_layer")
.commit_and_wait()
.expect("sync wait full_delta_net_layer");
let gpu_out = download_f32(&gpu_out_buf).expect("download gpu_out");
assert_eq!(gpu_out.len(), cpu_out.len(), "output length mismatch");
assert!(
gpu_out.iter().any(|&v| v != 0.0),
"full_delta_net_layer_gpu_matches_cpu_ref: GPU output all-zero — \
dispatch chain likely failed silently. cpu_first_8={:?}",
&cpu_out[..8.min(cpu_out.len())]
);
let max_err = gpu_out
.iter()
.zip(cpu_out.iter())
.map(|(&g, &c)| (g - c).abs())
.fold(0.0f32, f32::max);
const Q4_0_PARITY_TOLERANCE: f32 = 5e-2;
let mut n_fail = 0usize;
for (i, (&g, &c)) in gpu_out.iter().zip(cpu_out.iter()).enumerate() {
if (g - c).abs() >= Q4_0_PARITY_TOLERANCE {
if n_fail < 5 {
eprintln!(
" mismatch[{i}]: gpu={g:.8}, cpu={c:.8}, err={:.2e}",
(g - c).abs()
);
}
n_fail += 1;
}
}
assert!(
max_err < Q4_0_PARITY_TOLERANCE,
"DeltaNet GPU parity FAIL: max_abs_err={:.2e} (> {:.2e} \
Q4_0 budget), n_fail={}/{}",
max_err,
Q4_0_PARITY_TOLERANCE,
n_fail,
gpu_out.len()
);
eprintln!(
"full_delta_net_layer_gpu_matches_cpu_ref: max_abs_err={:.2e} (< {:.2e} Q4_0 budget), seq={seq}",
max_err, Q4_0_PARITY_TOLERANCE
);
}
#[test]
fn delta_net_layer_seq2_mid_stream_state() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let shape = small_shape();
let weights_cpu = synthetic_weights(shape, 0x9ABC);
let seq_len = 2u32; let h = shape.hidden_size as usize;
let seq = seq_len as usize;
let km1 = (shape.conv_kernel - 1) as usize;
let qkv_channels = shape.qkv_channels() as usize;
let state_size = (shape.d_k * shape.d_v * shape.n_v_heads) as usize;
let x_cpu: Vec<f32> = (0..seq * h).map(|i| 0.02 * (i as f32) - 0.7).collect();
let state_in: Vec<f32> = (0..state_size)
.map(|i| 0.013 * ((i as f32) * 0.7).sin() - 0.005)
.collect();
let conv_state: Vec<f32> = (0..km1 * qkv_channels)
.map(|i| 0.011 * ((i as f32) * 0.5).cos())
.collect();
let (cpu_out, _, _) =
delta_net_layer_cpu_ref(&x_cpu, &weights_cpu, shape, &state_in, &conv_state);
assert!(cpu_out.iter().all(|v| v.is_finite()), "CPU ref non-finite");
assert_eq!(cpu_out.len(), seq * h);
let gpu_weights = DeltaNetWeightsGpu::from_cpu_f32(
&weights_cpu,
&device,
shape.conv_kernel as usize,
qkv_channels,
)
.expect("from_cpu_f32");
let x_gpu = upload_f32(&x_cpu, &device).expect("upload x");
let state_in_gpu = upload_f32(&state_in, &device).expect("upload state_in");
let state_out_gpu = upload_f32(&state_in, &device).expect("upload state_out scratch");
let conv_state_in_gpu = upload_f32(&conv_state, &device).expect("upload conv_state_in");
let conv_state_out_gpu = upload_f32(&conv_state, &device).expect("alloc conv_state_out");
let gpu_out_buf = build_delta_net_layer(
&device,
&mut registry,
&x_gpu,
&gpu_weights,
&conv_state_in_gpu,
&conv_state_out_gpu,
&state_in_gpu,
&state_out_gpu,
seq_len,
shape.hidden_size,
shape.n_k_heads,
shape.n_v_heads,
shape.d_k,
shape.d_v,
shape.conv_kernel,
shape.rms_norm_eps,
None, None, SlotId(0), )
.expect("build_delta_net_layer");
device
.command_encoder()
.expect("sync enc")
.commit_and_wait()
.expect("sync wait");
let gpu_out = download_f32(&gpu_out_buf).expect("download gpu_out");
assert_eq!(gpu_out.len(), cpu_out.len(), "output length mismatch");
assert!(
gpu_out.iter().any(|&v| v != 0.0),
"delta_net_layer_seq2_mid_stream_state: GPU output all-zero — \
dispatch chain likely failed silently. cpu_first_8={:?}",
&cpu_out[..8.min(cpu_out.len())]
);
const Q4_0_PARITY_TOLERANCE: f32 = 5e-2;
let mut n_fail = 0usize;
for (i, (&g, &c)) in gpu_out.iter().zip(cpu_out.iter()).enumerate() {
if (g - c).abs() >= Q4_0_PARITY_TOLERANCE {
if n_fail < 5 {
eprintln!(
" [seq2_mid_stream] mismatch[{i}]: gpu={g:.6}, cpu={c:.6}, err={:.2e}",
(g - c).abs()
);
}
n_fail += 1;
}
}
let max_err = gpu_out
.iter()
.zip(cpu_out.iter())
.map(|(&g, &c)| (g - c).abs())
.fold(0.0f32, f32::max);
eprintln!(
"delta_net_layer_seq2_mid_stream_state: max_abs_err={:.2e} \
(< {:.2e} Q4_0 budget? {}), n_fail={}/{}, seq={seq_len}",
max_err,
Q4_0_PARITY_TOLERANCE,
max_err < Q4_0_PARITY_TOLERANCE,
n_fail,
gpu_out.len()
);
assert!(
max_err < Q4_0_PARITY_TOLERANCE,
"DeltaNet seq2 mid-stream non-zero-state PARITY FAIL: \
max_abs_err={:.2e} (> {:.2e}), n_fail={}/{}. \
This is the iter-178 'cur_len bug' — bug is inside \
build_delta_net_layer's prefill path setup at seq_len>1 + \
non-zero state_in. Per iter-273 test plan, bisect within \
op1 (norm), op2 (qkv proj), op3 (ssm_conv), op5+ chain.",
max_err,
Q4_0_PARITY_TOLERANCE,
n_fail,
gpu_out.len()
);
}
#[test]
fn delta_net_layer_seq1_plus_seq1_eq_seq2_at_same_initial_state() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let shape = small_shape();
let weights_cpu = synthetic_weights(shape, 0xBEEF);
let h = shape.hidden_size as usize;
let km1 = (shape.conv_kernel - 1) as usize;
let qkv_channels = shape.qkv_channels() as usize;
let state_size = (shape.d_k * shape.d_v * shape.n_v_heads) as usize;
let x_cpu: Vec<f32> = (0..2 * h).map(|i| 0.013 * (i as f32) - 0.4).collect();
let state_in0: Vec<f32> = (0..state_size)
.map(|i| 0.011 * ((i as f32) * 0.6).cos() - 0.003)
.collect();
let conv_state0: Vec<f32> = (0..km1 * qkv_channels)
.map(|i| 0.009 * ((i as f32) * 0.4).sin())
.collect();
let (cpu_out_a1, cpu_state_a1, cpu_conv_a1) =
delta_net_layer_cpu_ref(&x_cpu[0..h], &weights_cpu, shape, &state_in0, &conv_state0);
let (cpu_out_a2, cpu_state_a2, cpu_conv_a2) = delta_net_layer_cpu_ref(
&x_cpu[h..2 * h],
&weights_cpu,
shape,
&cpu_state_a1,
&cpu_conv_a1,
);
let mut cpu_seq_concat = cpu_out_a1.clone();
cpu_seq_concat.extend_from_slice(&cpu_out_a2);
let (cpu_out_b, cpu_state_b, cpu_conv_b) =
delta_net_layer_cpu_ref(&x_cpu, &weights_cpu, shape, &state_in0, &conv_state0);
let cpu_max_diff = cpu_seq_concat
.iter()
.zip(cpu_out_b.iter())
.map(|(&a, &b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(
cpu_max_diff < 1e-5,
"CPU ref self-consistency FAIL: seq=1+1 != seq=2 (max_abs_err={:.2e}). \
Indicates a bug in delta_net_layer_cpu_ref's recurrence math, \
NOT in the GPU layer.",
cpu_max_diff
);
let cpu_state_max_diff = cpu_state_a2
.iter()
.zip(cpu_state_b.iter())
.map(|(&a, &b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(
cpu_state_max_diff < 1e-5,
"CPU ref state self-consistency FAIL: state(seq=1+1) != state(seq=2). \
max_abs_err={:.2e}",
cpu_state_max_diff
);
let _ = (cpu_conv_a2, cpu_conv_b);
let gpu_weights = DeltaNetWeightsGpu::from_cpu_f32(
&weights_cpu,
&device,
shape.conv_kernel as usize,
qkv_channels,
)
.expect("from_cpu_f32");
let x_gpu = upload_f32(&x_cpu, &device).expect("upload x");
let state_in_gpu = upload_f32(&state_in0, &device).expect("upload state_in");
let state_out_gpu = upload_f32(&state_in0, &device).expect("alloc state_out");
let conv_in_gpu = upload_f32(&conv_state0, &device).expect("upload conv");
let conv_out_gpu = upload_f32(&conv_state0, &device).expect("alloc conv_out");
let gpu_out_b_buf = build_delta_net_layer(
&device,
&mut registry,
&x_gpu,
&gpu_weights,
&conv_in_gpu,
&conv_out_gpu,
&state_in_gpu,
&state_out_gpu,
2,
shape.hidden_size,
shape.n_k_heads,
shape.n_v_heads,
shape.d_k,
shape.d_v,
shape.conv_kernel,
shape.rms_norm_eps,
None, None, SlotId(0), )
.expect("build_delta_net_layer seq=2");
device
.command_encoder()
.expect("sync")
.commit_and_wait()
.expect("wait");
let gpu_out_b = download_f32(&gpu_out_b_buf).expect("download");
const Q4_0_TOL: f32 = 5e-2;
let max_err_b = gpu_out_b
.iter()
.zip(cpu_out_b.iter())
.map(|(&g, &c)| (g - c).abs())
.fold(0.0f32, f32::max);
assert!(
max_err_b < Q4_0_TOL,
"GPU seq=2 vs CPU seq=2 PARITY FAIL: max_err={:.2e} > {:.2e}. \
This indicates DeltaNet GPU+seq=2 is BROKEN at the layer level. \
(Should match iter-273 result; if iter-273 passed and this fails, \
investigate weight upload differences.)",
max_err_b,
Q4_0_TOL
);
eprintln!(
"delta_net_layer_seq1_plus_seq1_eq_seq2: \
cpu_self_consistency_max_diff={:.2e}, \
gpu_seq2_vs_cpu_max_err={:.2e}. \
CPU recurrence proves seq=1+1 ≡ seq=2 EXACTLY. \
GPU seq=2 matches CPU seq=2 within Q4_0 budget. \
⇒ Layer satisfies the K1 implicit invariant. \
Bug must be HIGHER (qwen35 forward orchestration).",
cpu_max_diff, max_err_b
);
}
#[test]
fn kernel_w_transpose_roundtrip() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let k = 4usize;
let ch = 6usize;
let mut seed = 0xDEAD_u32;
let orig: Vec<f32> = (0..k * ch)
.map(|_| {
seed = seed.wrapping_mul(1103515245).wrapping_add(12345);
(seed as i32 as f32) / (i32::MAX as f32)
})
.collect();
let transposed = transpose_k_channels_to_channels_k(&orig, k, ch);
let roundtrip: Vec<f32> = {
let mut dst = vec![0.0f32; k * ch];
for c in 0..ch {
for ki in 0..k {
dst[ki * ch + c] = transposed[c * k + ki];
}
}
dst
};
for (a, b) in orig.iter().zip(roundtrip.iter()) {
assert_eq!(a.to_bits(), b.to_bits(), "transpose roundtrip mismatch");
}
}
#[test]
fn conv_state_transpose_roundtrip() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let km1 = 3usize;
let ch = 8usize;
let mut seed = 0xFACE_u32;
let orig: Vec<f32> = (0..km1 * ch)
.map(|_| {
seed = seed.wrapping_mul(1103515245).wrapping_add(12345);
(seed as i32 as f32) / (i32::MAX as f32)
})
.collect();
let to_kernel = transpose_state_km1_c_to_c_km1(&orig, km1, ch);
let back = transpose_state_c_km1_to_km1_c(&to_kernel, km1, ch);
for (a, b) in orig.iter().zip(back.iter()) {
assert_eq!(a.to_bits(), b.to_bits(), "conv_state roundtrip mismatch");
}
}
#[test]
fn upload_download_f32_roundtrip() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let data: Vec<f32> = (0..64).map(|i| i as f32 * 0.1).collect();
let buf = upload_f32(&data, &device).expect("upload");
let got = download_f32(&buf).expect("download");
assert_eq!(got, data);
}
#[test]
fn gpu_state_propagation_chunked_vs_monolithic() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(e) => {
eprintln!("skipping: no Metal device: {e}");
return;
}
};
let mut registry = KernelRegistry::new();
let shape = small_shape();
let weights_cpu = synthetic_weights(shape, 0x1234);
let h = shape.hidden_size as usize;
let km1 = (shape.conv_kernel - 1) as usize;
let qkv_channels = shape.qkv_channels() as usize;
let state_size = (shape.d_k * shape.d_v * shape.n_v_heads) as usize;
let x_full: Vec<f32> = (0..2 * h).map(|i| 0.02 * i as f32 - 0.5).collect();
let state_zeros = vec![0.0f32; state_size];
let conv_zeros = vec![0.0f32; km1 * qkv_channels];
let gpu_weights = DeltaNetWeightsGpu::from_cpu_f32(
&weights_cpu,
&device,
shape.conv_kernel as usize,
qkv_channels,
)
.expect("from_cpu_f32");
let flush = |device: &MlxDevice| {
let mut enc = device.command_encoder().expect("flush enc");
enc.commit_and_wait().expect("flush commit_and_wait");
};
let x_full_gpu = upload_f32(&x_full, &device).expect("upload");
let state_zeros_gpu = upload_f32(&state_zeros, &device).expect("upload state_zeros mono");
let state_scratch_mono = upload_f32(&state_zeros, &device).expect("state scratch mono");
let conv_zeros_gpu_mono = upload_f32(&conv_zeros, &device).expect("upload conv mono in");
let conv_scratch_mono = upload_f32(&conv_zeros, &device).expect("alloc conv mono out");
let mono_buf = build_delta_net_layer(
&device,
&mut registry,
&x_full_gpu,
&gpu_weights,
&conv_zeros_gpu_mono,
&conv_scratch_mono,
&state_zeros_gpu,
&state_scratch_mono,
2,
shape.hidden_size,
shape.n_k_heads,
shape.n_v_heads,
shape.d_k,
shape.d_v,
shape.conv_kernel,
shape.rms_norm_eps,
None, None, SlotId(0), )
.expect("mono");
flush(&device);
let mono_out = download_f32(&mono_buf).expect("dl mono");
let x_t0 = x_full[0..h].to_vec();
let x_t1 = x_full[h..2 * h].to_vec();
let x_t0_gpu = upload_f32(&x_t0, &device).expect("upload t0");
let state_t0_in = upload_f32(&state_zeros, &device).expect("upload state t0 in");
let state_t0_out = upload_f32(&state_zeros, &device).expect("alloc state t0 out");
let conv_t0_in = upload_f32(&conv_zeros, &device).expect("upload conv t0 in");
let conv_t0_out = upload_f32(&conv_zeros, &device).expect("alloc conv t0 out");
let t0_buf = build_delta_net_layer(
&device,
&mut registry,
&x_t0_gpu,
&gpu_weights,
&conv_t0_in,
&conv_t0_out,
&state_t0_in,
&state_t0_out,
1,
shape.hidden_size,
shape.n_k_heads,
shape.n_v_heads,
shape.d_k,
shape.d_v,
shape.conv_kernel,
shape.rms_norm_eps,
None, None, SlotId(0), )
.expect("chunk t0");
flush(&device);
let t0_out = download_f32(&t0_buf).expect("dl t0");
let x_t1_gpu = upload_f32(&x_t1, &device).expect("upload t1");
let state_t1_out = upload_f32(&state_zeros, &device).expect("alloc state t1 out");
let conv_t1_out = upload_f32(&conv_zeros, &device).expect("alloc conv t1 out");
let t1_buf = build_delta_net_layer(
&device,
&mut registry,
&x_t1_gpu,
&gpu_weights,
&conv_t0_out,
&conv_t1_out,
&state_t0_out,
&state_t1_out,
1,
shape.hidden_size,
shape.n_k_heads,
shape.n_v_heads,
shape.d_k,
shape.d_v,
shape.conv_kernel,
shape.rms_norm_eps,
None, None, SlotId(0), )
.expect("chunk t1");
flush(&device);
let t1_out = download_f32(&t1_buf).expect("dl t1");
for i in 0..h {
let diff = (mono_out[i] - t0_out[i]).abs();
assert!(
diff < 1e-3,
"t0 mismatch[{i}]: mono={:.6}, chunk={:.6}, diff={:.2e}",
mono_out[i],
t0_out[i],
diff
);
}
for i in 0..h {
let diff = (mono_out[h + i] - t1_out[i]).abs();
assert!(
diff < 1e-3,
"t1 mismatch[{i}]: mono={:.6}, chunk={:.6}, diff={:.2e}",
mono_out[h + i],
t1_out[i],
diff
);
}
eprintln!("gpu_state_propagation_chunked_vs_monolithic: PASS");
}
#[test]
fn chunk_path_first_token_matches_autoregressive_at_seq128() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let seq_len: u32 = 128;
let n_k_heads: u32 = 2;
let n_v_heads: u32 = 4;
let d_k: u32 = 128;
let d_v: u32 = 128;
let n_q_elems = (seq_len * n_k_heads * d_k) as usize;
let n_v_elems = (seq_len * n_v_heads * d_v) as usize;
let n_g_elems = (seq_len * n_v_heads) as usize;
let n_state = (d_k * d_v * n_v_heads) as usize;
let n_out = (n_v_heads * seq_len * d_v) as usize;
let mut seed: u32 = 0xDEAD;
let q_cpu: Vec<f32> = mk_rand(&mut seed, n_q_elems, 0.05);
let k_cpu: Vec<f32> = mk_rand(&mut seed, n_q_elems, 0.05);
let v_cpu: Vec<f32> = mk_rand(&mut seed, n_v_elems, 0.1);
let g_cpu: Vec<f32> = (0..n_g_elems)
.map(|i| 0.001 + 0.0001 * (i as f32 % 7.0))
.collect();
let beta_cpu: Vec<f32> = (0..n_g_elems)
.map(|i| 0.5 + 0.01 * ((i as f32 % 11.0) - 5.0))
.collect();
let state_zeros = vec![0.0f32; n_state];
let q_auto = upload_f32(&q_cpu, &device).expect("upload q auto");
let k_auto = upload_f32(&k_cpu, &device).expect("upload k auto");
let v_auto = upload_f32(&v_cpu, &device).expect("upload v auto");
let g_auto = upload_f32(&g_cpu, &device).expect("upload g auto");
let beta_auto = upload_f32(&beta_cpu, &device).expect("upload beta auto");
let state_auto = upload_f32(&state_zeros, &device).expect("upload state auto");
let q_chunk = upload_f32(&q_cpu, &device).expect("upload q chunk");
let k_chunk = upload_f32(&k_cpu, &device).expect("upload k chunk");
let v_chunk = upload_f32(&v_cpu, &device).expect("upload v chunk");
let g_chunk = upload_f32(&g_cpu, &device).expect("upload g chunk");
let beta_chunk = upload_f32(&beta_cpu, &device).expect("upload beta chunk");
let state_chunk = upload_f32(&state_zeros, &device).expect("upload state chunk");
let (auto_out, _auto_state) = apply_gated_delta_net(
&device,
&mut registry,
&q_auto,
&k_auto,
&v_auto,
&g_auto,
&beta_auto,
&state_auto,
seq_len,
n_k_heads,
n_v_heads,
d_k,
d_v,
)
.expect("apply_gated_delta_net (autoregressive)");
let auto_out_cpu = download_f32(&auto_out).expect("download auto out");
assert_eq!(auto_out_cpu.len(), n_out, "auto output length");
let chunk_out = device
.alloc_buffer(n_out * 4, DType::F32, vec![n_out])
.expect("alloc chunk_out");
let chunk_state = device
.alloc_buffer(n_state * 4, DType::F32, vec![n_state])
.expect("alloc chunk_state");
apply_gated_delta_net_chunk(
&device,
&mut registry,
&q_chunk,
&k_chunk,
&v_chunk,
&g_chunk,
&beta_chunk,
&state_chunk,
&chunk_out,
&chunk_state,
seq_len,
n_k_heads,
n_v_heads,
d_k,
d_v,
false,
)
.expect("apply_gated_delta_net_chunk");
device
.command_encoder()
.expect("sync enc")
.commit_and_wait()
.expect("sync wait");
let chunk_out_cpu = download_f32(&chunk_out).expect("download chunk out");
assert_eq!(chunk_out_cpu.len(), n_out, "chunk output length");
let auto_all_zero = auto_out_cpu.iter().all(|&v| v == 0.0);
let chunk_all_zero = chunk_out_cpu.iter().all(|&v| v == 0.0);
assert!(
!auto_all_zero,
"autoregressive path returned all-zero — GPU dispatch chain \
likely failed silently. chunk_first_8={:?}",
&chunk_out_cpu[..8.min(chunk_out_cpu.len())]
);
assert!(
!chunk_all_zero,
"chunk path returned all-zero — GPU dispatch chain likely \
failed silently. auto_first_8={:?}",
&auto_out_cpu[..8.min(auto_out_cpu.len())]
);
const FULL_BUFFER_TOL: f32 = 5.0e-2;
let mut max_diff: f32 = 0.0;
let mut argmax_t: u32 = 0;
let mut argmax_h: u32 = 0;
let mut argmax_v: u32 = 0;
for t in 0..seq_len {
let token_base = (t * n_v_heads * d_v) as usize;
for h in 0..n_v_heads {
let head_base = token_base + (h * d_v) as usize;
for v_idx in 0..d_v {
let i = head_base + v_idx as usize;
let diff = (auto_out_cpu[i] - chunk_out_cpu[i]).abs();
if diff > max_diff {
max_diff = diff;
argmax_t = t;
argmax_h = h;
argmax_v = v_idx;
}
}
}
}
const FIRST_TOKEN_TOL: f32 = FULL_BUFFER_TOL;
let _ = argmax_t;
let argmax_idx =
(argmax_t * n_v_heads * d_v) as usize + (argmax_h * d_v) as usize + argmax_v as usize;
eprintln!(
"chunk_path_full_buffer_matches_autoregressive_at_seq128: \
max_diff={:.4e} at (t={}, h={}, v={}), tol={:.0e}, \
total_elements_compared={}",
max_diff, argmax_t, argmax_h, argmax_v, FIRST_TOKEN_TOL, n_out
);
assert!(
max_diff < FIRST_TOKEN_TOL,
"full-buffer output diverges between autoregressive and chunk paths: \
max_diff={:.4e} at (t={}, h={}, v={}), tol={:.0e}. \
auto[t={},h={},v={}]={:.6}, chunk[t={},h={},v={}]={:.6}",
max_diff,
argmax_t,
argmax_h,
argmax_v,
FIRST_TOKEN_TOL,
argmax_t,
argmax_h,
argmax_v,
auto_out_cpu[argmax_idx],
argmax_t,
argmax_h,
argmax_v,
chunk_out_cpu[argmax_idx],
);
}
#[test]
fn chunk_arena_byte_exact_parity_at_seq128() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let seq_len: u32 = 128;
let n_k_heads: u32 = 2;
let n_v_heads: u32 = 4;
let d_k: u32 = 128;
let d_v: u32 = 128;
let n_q_elems = (seq_len * n_k_heads * d_k) as usize;
let n_v_elems = (seq_len * n_v_heads * d_v) as usize;
let n_g_elems = (seq_len * n_v_heads) as usize;
let n_state = (d_k * d_v * n_v_heads) as usize;
let n_out = (n_v_heads * seq_len * d_v) as usize;
let mut seed: u32 = 0xCAFE;
let q_cpu: Vec<f32> = mk_rand(&mut seed, n_q_elems, 0.05);
let k_cpu: Vec<f32> = mk_rand(&mut seed, n_q_elems, 0.05);
let v_cpu: Vec<f32> = mk_rand(&mut seed, n_v_elems, 0.1);
let g_cpu: Vec<f32> = (0..n_g_elems)
.map(|i| 0.001 + 0.0001 * (i as f32 % 7.0))
.collect();
let beta_cpu: Vec<f32> = (0..n_g_elems)
.map(|i| 0.5 + 0.01 * ((i as f32 % 11.0) - 5.0))
.collect();
let state_zeros = vec![0.0f32; n_state];
let q_na = upload_f32(&q_cpu, &device).expect("upload q na");
let k_na = upload_f32(&k_cpu, &device).expect("upload k na");
let v_na = upload_f32(&v_cpu, &device).expect("upload v na");
let g_na = upload_f32(&g_cpu, &device).expect("upload g na");
let beta_na = upload_f32(&beta_cpu, &device).expect("upload beta na");
let state_na = upload_f32(&state_zeros, &device).expect("upload state na");
let out_na = device
.alloc_buffer(n_out * 4, DType::F32, vec![n_out])
.expect("alloc out_na");
let final_state_na = device
.alloc_buffer(n_state * 4, DType::F32, vec![n_state])
.expect("alloc final_state_na");
apply_gated_delta_net_chunk(
&device,
&mut registry,
&q_na,
&k_na,
&v_na,
&g_na,
&beta_na,
&state_na,
&out_na,
&final_state_na,
seq_len,
n_k_heads,
n_v_heads,
d_k,
d_v,
false,
)
.expect("apply_gated_delta_net_chunk (na)");
device
.command_encoder()
.expect("sync enc na")
.commit_and_wait()
.expect("sync wait na");
let out_na_cpu = download_f32(&out_na).expect("download out_na");
let state_na_cpu = download_f32(&final_state_na).expect("download state_na");
let q_a = upload_f32(&q_cpu, &device).expect("upload q a");
let k_a = upload_f32(&k_cpu, &device).expect("upload k a");
let v_a = upload_f32(&v_cpu, &device).expect("upload v a");
let g_a = upload_f32(&g_cpu, &device).expect("upload g a");
let beta_a = upload_f32(&beta_cpu, &device).expect("upload beta a");
let state_a = upload_f32(&state_zeros, &device).expect("upload state a");
let out_a = device
.alloc_buffer(n_out * 4, DType::F32, vec![n_out])
.expect("alloc out_a");
let final_state_a = device
.alloc_buffer(n_state * 4, DType::F32, vec![n_state])
.expect("alloc final_state_a");
let mut arena = crate::inference::models::qwen35::ChunkAllocsArena::new(
&device, seq_len, n_v_heads, d_k, d_v,
)
.expect("ChunkAllocsArena::new");
apply_gated_delta_net_chunk_with_arena(
&device,
&mut registry,
&q_a,
&k_a,
&v_a,
&g_a,
&beta_a,
&state_a,
&out_a,
&final_state_a,
&mut arena,
None,
seq_len,
n_k_heads,
n_v_heads,
d_k,
d_v,
false,
)
.expect("apply_gated_delta_net_chunk_with_arena (a)");
device
.command_encoder()
.expect("sync enc a")
.commit_and_wait()
.expect("sync wait a");
let out_a_cpu = download_f32(&out_a).expect("download out_a");
let state_a_cpu = download_f32(&final_state_a).expect("download state_a");
let na_zero = out_na_cpu.iter().all(|&v| v == 0.0);
let a_zero = out_a_cpu.iter().all(|&v| v == 0.0);
if na_zero && a_zero {
eprintln!(
"chunk_arena_byte_exact_parity_at_seq128: \
BOTH paths returned all-zero — likely parallel-contention flake; \
re-run in isolation with --test-threads=1 to confirm."
);
return;
}
assert!(!na_zero, "no-arena path returned all-zero");
assert!(!a_zero, "arena path returned all-zero");
assert_eq!(
out_na_cpu.len(),
out_a_cpu.len(),
"output buffer length mismatch"
);
let mut diffs = 0usize;
let mut max_abs: f32 = 0.0;
for (i, (&n, &a)) in out_na_cpu.iter().zip(out_a_cpu.iter()).enumerate() {
if n.to_bits() != a.to_bits() {
diffs += 1;
let d = (n - a).abs();
if d > max_abs {
max_abs = d;
}
if diffs <= 4 {
eprintln!(
"chunk_arena diff[{}] na={:.6e} bits={:#x} \
vs a={:.6e} bits={:#x} diff={:.3e}",
i,
n,
n.to_bits(),
a,
a.to_bits(),
d
);
}
}
}
assert_eq!(
diffs, 0,
"chunk arena byte-parity FAILED: {diffs} elements differ \
(max abs diff {max_abs:.3e}); arena path must produce \
bit-identical output to non-arena path"
);
assert_eq!(
state_na_cpu.len(),
state_a_cpu.len(),
"final_state length mismatch"
);
for (i, (&n, &a)) in state_na_cpu.iter().zip(state_a_cpu.iter()).enumerate() {
assert_eq!(
n.to_bits(),
a.to_bits(),
"chunk arena final_state byte-parity FAILED at index {i}: na={n:.6e} a={a:.6e}"
);
}
}
#[test]
fn chunk_internal_arena_kernel_equivalence_at_seq128() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let seq_len: u32 = 128;
let n_k_heads: u32 = 2;
let n_v_heads: u32 = 4;
let d_k: u32 = 128;
let d_v: u32 = 128;
let n_q_elems = (seq_len * n_k_heads * d_k) as usize;
let n_v_elems = (seq_len * n_v_heads * d_v) as usize;
let n_g_elems = (seq_len * n_v_heads) as usize;
let n_state = (d_k * d_v * n_v_heads) as usize;
let n_out = (n_v_heads * seq_len * d_v) as usize;
let mut seed: u32 = 0xBAB3;
let q_cpu: Vec<f32> = mk_rand(&mut seed, n_q_elems, 0.05);
let k_cpu: Vec<f32> = mk_rand(&mut seed, n_q_elems, 0.05);
let v_cpu: Vec<f32> = mk_rand(&mut seed, n_v_elems, 0.1);
let g_cpu: Vec<f32> = (0..n_g_elems)
.map(|i| 0.001 + 0.0001 * (i as f32 % 7.0))
.collect();
let beta_cpu: Vec<f32> = (0..n_g_elems)
.map(|i| 0.5 + 0.01 * ((i as f32 % 11.0) - 5.0))
.collect();
let state_zeros = vec![0.0f32; n_state];
let q_na = upload_f32(&q_cpu, &device).expect("upload q na");
let k_na = upload_f32(&k_cpu, &device).expect("upload k na");
let v_na = upload_f32(&v_cpu, &device).expect("upload v na");
let g_na = upload_f32(&g_cpu, &device).expect("upload g na");
let beta_na = upload_f32(&beta_cpu, &device).expect("upload beta na");
let state_na = upload_f32(&state_zeros, &device).expect("upload state na");
let out_na = device
.alloc_buffer(n_out * 4, DType::F32, vec![n_out])
.expect("alloc out_na");
let final_state_na = device
.alloc_buffer(n_state * 4, DType::F32, vec![n_state])
.expect("alloc final_state_na");
let mut allocs_arena_na = crate::inference::models::qwen35::ChunkAllocsArena::new(
&device, seq_len, n_v_heads, d_k, d_v,
)
.expect("ChunkAllocsArena::new (na)");
apply_gated_delta_net_chunk_with_arena(
&device,
&mut registry,
&q_na,
&k_na,
&v_na,
&g_na,
&beta_na,
&state_na,
&out_na,
&final_state_na,
&mut allocs_arena_na,
None, seq_len,
n_k_heads,
n_v_heads,
d_k,
d_v,
false,
)
.expect("apply_gated_delta_net_chunk_with_arena (na)");
device
.command_encoder()
.expect("sync enc na")
.commit_and_wait()
.expect("sync wait na");
let out_na_cpu = download_f32(&out_na).expect("download out_na");
let state_na_cpu = download_f32(&final_state_na).expect("download state_na");
let q_ia = upload_f32(&q_cpu, &device).expect("upload q ia");
let k_ia = upload_f32(&k_cpu, &device).expect("upload k ia");
let v_ia = upload_f32(&v_cpu, &device).expect("upload v ia");
let g_ia = upload_f32(&g_cpu, &device).expect("upload g ia");
let beta_ia = upload_f32(&beta_cpu, &device).expect("upload beta ia");
let state_ia = upload_f32(&state_zeros, &device).expect("upload state ia");
let out_ia = device
.alloc_buffer(n_out * 4, DType::F32, vec![n_out])
.expect("alloc out_ia");
let final_state_ia = device
.alloc_buffer(n_state * 4, DType::F32, vec![n_state])
.expect("alloc final_state_ia");
let mut allocs_arena_ia = crate::inference::models::qwen35::ChunkAllocsArena::new(
&device, seq_len, n_v_heads, d_k, d_v,
)
.expect("ChunkAllocsArena::new (ia)");
let mut internal_arena = mlx_native::ops::chunk_gated_delta_rule::ChunkInternalArena::new(
&device, 1, seq_len, n_v_heads,
n_v_heads, d_k, d_v, 64,
)
.expect("ChunkInternalArena::new");
apply_gated_delta_net_chunk_with_arena(
&device,
&mut registry,
&q_ia,
&k_ia,
&v_ia,
&g_ia,
&beta_ia,
&state_ia,
&out_ia,
&final_state_ia,
&mut allocs_arena_ia,
Some(&mut internal_arena),
seq_len,
n_k_heads,
n_v_heads,
d_k,
d_v,
false,
)
.expect("apply_gated_delta_net_chunk_with_arena (ia)");
device
.command_encoder()
.expect("sync enc ia")
.commit_and_wait()
.expect("sync wait ia");
let out_ia_cpu = download_f32(&out_ia).expect("download out_ia");
let state_ia_cpu = download_f32(&final_state_ia).expect("download state_ia");
assert!(
out_na_cpu.iter().any(|&v| v != 0.0),
"no-internal-arena path returned ALL-ZERO output — GPU \
dispatch chain likely failed silently. Check that the chunk \
pipeline kernels are registered."
);
assert!(
out_ia_cpu.iter().any(|&v| v != 0.0),
"internal-arena path returned ALL-ZERO output — GPU dispatch \
chain likely failed silently. Same diagnostic as the no-arena \
assert above."
);
let mut shown = 0usize;
for (i, (&n, &a)) in out_na_cpu.iter().zip(out_ia_cpu.iter()).enumerate() {
if n.to_bits() != a.to_bits() && shown < 4 {
eprintln!(
"internal-arena bit-diff[{}] na={:.6e} bits={:#x} \
vs ia={:.6e} bits={:#x} abs={:.3e}",
i,
n,
n.to_bits(),
a,
a.to_bits(),
(n - a).abs()
);
shown += 1;
}
}
crate::core::kernel_parity::assert_kernel_equivalence(
&out_na_cpu,
&out_ia_cpu,
0.9999,
1e-4,
"iter83 chunk_internal_arena (out)",
);
crate::core::kernel_parity::assert_kernel_equivalence(
&state_na_cpu,
&state_ia_cpu,
0.9999,
1e-4,
"iter83 chunk_internal_arena (final_state)",
);
}
#[test]
fn chunk_path_rejects_non_multiple_of_bt() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let dummy_f32 = upload_f32(&[0.0f32; 1], &device).expect("upload dummy");
let err = apply_gated_delta_net_chunk(
&device,
&mut registry,
&dummy_f32,
&dummy_f32,
&dummy_f32,
&dummy_f32,
&dummy_f32,
&dummy_f32,
&dummy_f32, &dummy_f32, 65,
1,
1,
8,
8,
false,
)
.expect_err("seq_len=65 must be rejected (not a multiple of BT=64)");
let msg = err.to_string();
assert!(
msg.contains("65") && msg.contains("64"),
"expected error to cite seq_len=65 and FIXED_BT=64, got: {msg}"
);
}
#[test]
fn chunk_path_rejects_qk_l2norm_in_kernel() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let dummy_f32 = upload_f32(&[0.0f32; 1], &device).expect("upload dummy");
let err = apply_gated_delta_net_chunk(
&device,
&mut registry,
&dummy_f32,
&dummy_f32,
&dummy_f32,
&dummy_f32,
&dummy_f32,
&dummy_f32,
&dummy_f32, &dummy_f32, 128,
1,
1,
8,
8,
true,
)
.expect_err("use_qk_l2norm=true must be rejected at iter 5");
let msg = err.to_string();
assert!(
msg.contains("use_qk_l2norm"),
"expected error to cite use_qk_l2norm, got: {msg}"
);
}
fn cpu_ref_recurrence(
q: &[f32],
k: &[f32],
v: &[f32],
g: &[f32],
beta: &[f32],
state_in: &[f32],
seq_len: u32,
n_k_heads: u32,
n_v_heads: u32,
d_k: u32,
d_v: u32,
) -> Vec<f32> {
use mlx_native::ops::gated_delta_net::cpu_reference_f32 as gdn_cpu_ref;
let p = GatedDeltaNetParams {
d_k,
d_v,
n_k_heads,
n_v_heads,
n_tokens: seq_len,
n_seqs: FORWARD_DISPATCH_N_SEQS,
};
let (out, _state) = gdn_cpu_ref(q, k, v, g, beta, state_in, p);
out
}
fn cpu_ref_recurrence_pre_decay(
q: &[f32],
k: &[f32],
v: &[f32],
g: &[f32],
beta: &[f32],
state_in: &[f32],
seq_len: u32,
n_k_heads: u32,
n_v_heads: u32,
d_k: u32,
d_v: u32,
) -> Vec<f32> {
let d_k = d_k as usize;
let d_v = d_v as usize;
let nh_k = n_k_heads as usize;
let nh_v = n_v_heads as usize;
let n_t = seq_len as usize;
let kq_token_stride = nh_k * d_k;
let v_token_stride = nh_v * d_v;
let scalar_stride = nh_v;
let state_head_stride = d_v * d_k;
let mut output = vec![0.0f32; n_t * v_token_stride];
let mut state = state_in.to_vec();
for vh in 0..nh_v {
let kh = vh % nh_k;
for t in 0..n_t {
let kq_base = t * kq_token_stride + kh * d_k;
let v_base = t * v_token_stride + vh * d_v;
let sc_idx = t * scalar_stride + vh;
let beta_val = beta[sc_idx];
let g_val = g[sc_idx];
let alpha = (-g_val).exp();
let state_base = vh * state_head_stride;
let mut delta = vec![0.0f32; d_v];
for i in 0..d_v {
let mut sk = 0.0f32;
for j in 0..d_k {
sk += state[state_base + i * d_k + j] * k[kq_base + j];
}
delta[i] = v[v_base + i] - sk;
}
for i in 0..d_v {
let beta_delta = beta_val * delta[i];
for j in 0..d_k {
let idx = state_base + i * d_k + j;
state[idx] = alpha * state[idx] + beta_delta * k[kq_base + j];
}
}
for i in 0..d_v {
let mut acc = 0.0f32;
for j in 0..d_k {
acc += state[state_base + i * d_k + j] * q[kq_base + j];
}
output[v_base + i] = acc;
}
}
}
output
}
fn run_seqlen_with_heads(
seq_len: u32,
n_k_heads: u32,
n_v_heads: u32,
) -> (f32, f32, f32, f32, f32) {
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let d_k: u32 = 128;
let d_v: u32 = 128;
let n_q_elems = (seq_len * n_k_heads * d_k) as usize;
let n_v_elems = (seq_len * n_v_heads * d_v) as usize;
let n_g_elems = (seq_len * n_v_heads) as usize;
let n_state = (d_k * d_v * n_v_heads) as usize;
let n_out = (n_v_heads * seq_len * d_v) as usize;
let mut seed: u32 = 0xDEAD;
let q_cpu: Vec<f32> = mk_rand(&mut seed, n_q_elems, 0.05);
let k_cpu: Vec<f32> = mk_rand(&mut seed, n_q_elems, 0.05);
let v_cpu: Vec<f32> = mk_rand(&mut seed, n_v_elems, 0.1);
let g_cpu: Vec<f32> = (0..n_g_elems)
.map(|i| 0.001 + 0.0001 * (i as f32 % 7.0))
.collect();
let beta_cpu: Vec<f32> = (0..n_g_elems)
.map(|i| 0.5 + 0.01 * ((i as f32 % 11.0) - 5.0))
.collect();
let state_zeros = vec![0.0f32; n_state];
let q_a = upload_f32(&q_cpu, &device).expect("upload");
let k_a = upload_f32(&k_cpu, &device).expect("upload");
let v_a = upload_f32(&v_cpu, &device).expect("upload");
let g_a = upload_f32(&g_cpu, &device).expect("upload");
let b_a = upload_f32(&beta_cpu, &device).expect("upload");
let s_a = upload_f32(&state_zeros, &device).expect("upload");
let (auto_out, _) = apply_gated_delta_net(
&device,
&mut registry,
&q_a,
&k_a,
&v_a,
&g_a,
&b_a,
&s_a,
seq_len,
n_k_heads,
n_v_heads,
d_k,
d_v,
)
.expect("auto");
let auto_cpu = download_f32(&auto_out).expect("dl auto");
assert_eq!(auto_cpu.len(), n_out);
let q_c = upload_f32(&q_cpu, &device).expect("upload");
let k_c = upload_f32(&k_cpu, &device).expect("upload");
let v_c = upload_f32(&v_cpu, &device).expect("upload");
let g_c = upload_f32(&g_cpu, &device).expect("upload");
let b_c = upload_f32(&beta_cpu, &device).expect("upload");
let s_c = upload_f32(&state_zeros, &device).expect("upload");
let chunk_out = device
.alloc_buffer(n_out * 4, DType::F32, vec![n_out])
.expect("alloc chunk_out");
let chunk_state = device
.alloc_buffer(n_state * 4, DType::F32, vec![n_state])
.expect("alloc chunk_state");
apply_gated_delta_net_chunk(
&device,
&mut registry,
&q_c,
&k_c,
&v_c,
&g_c,
&b_c,
&s_c,
&chunk_out,
&chunk_state,
seq_len,
n_k_heads,
n_v_heads,
d_k,
d_v,
false,
)
.expect("chunk");
let _ = &chunk_state; device
.command_encoder()
.expect("sync enc")
.commit_and_wait()
.expect("sync wait");
let chunk_cpu = download_f32(&chunk_out).expect("dl chunk");
let cpu_post = cpu_ref_recurrence(
&q_cpu,
&k_cpu,
&v_cpu,
&g_cpu,
&beta_cpu,
&state_zeros,
seq_len,
n_k_heads,
n_v_heads,
d_k,
d_v,
);
let cpu_pre = cpu_ref_recurrence_pre_decay(
&q_cpu,
&k_cpu,
&v_cpu,
&g_cpu,
&beta_cpu,
&state_zeros,
seq_len,
n_k_heads,
n_v_heads,
d_k,
d_v,
);
let mut max_ac: f32 = 0.0;
let mut max_a_cpost: f32 = 0.0;
let mut max_c_cpost: f32 = 0.0;
let mut max_a_cpre: f32 = 0.0;
let mut max_c_cpre: f32 = 0.0;
let mut max_abs_auto: f32 = 0.0;
let mut max_abs_chunk: f32 = 0.0;
let mut sum_sq_diff: f64 = 0.0;
let mut sum_sq_auto: f64 = 0.0;
let mut argmax_idx: usize = 0;
for i in 0..n_out {
let a = auto_cpu[i];
let c = chunk_cpu[i];
let cpost = cpu_post[i];
let cpre = cpu_pre[i];
let d_ac = (a - c).abs();
if d_ac > max_ac {
max_ac = d_ac;
argmax_idx = i;
}
max_a_cpost = max_a_cpost.max((a - cpost).abs());
max_c_cpost = max_c_cpost.max((c - cpost).abs());
max_a_cpre = max_a_cpre.max((a - cpre).abs());
max_c_cpre = max_c_cpre.max((c - cpre).abs());
max_abs_auto = max_abs_auto.max(a.abs());
max_abs_chunk = max_abs_chunk.max(c.abs());
sum_sq_diff += (a as f64 - c as f64) * (a as f64 - c as f64);
sum_sq_auto += (a as f64) * (a as f64);
}
let rms_diff = (sum_sq_diff / n_out as f64).sqrt() as f32;
let rms_auto = (sum_sq_auto / n_out as f64).sqrt() as f32;
let argmax_t = argmax_idx / (n_v_heads as usize * d_v as usize);
let argmax_h = (argmax_idx / d_v as usize) % n_v_heads as usize;
let argmax_v = argmax_idx % d_v as usize;
eprintln!(
"[w5b4] seq_len={:5} auto_vs_chunk={:.4e} \
auto_vs_cpu_post={:.4e} chunk_vs_cpu_post={:.4e} \
auto_vs_cpu_pre={:.4e} chunk_vs_cpu_pre={:.4e} \
max|auto|={:.4e} max|chunk|={:.4e} \
rel_drift={:.4e} rms_drift={:.4e} rms_auto={:.4e} \
argmax(t={},h={},v={}) auto_at={:.6} chunk_at={:.6}",
seq_len,
max_ac,
max_a_cpost,
max_c_cpost,
max_a_cpre,
max_c_cpre,
max_abs_auto,
max_abs_chunk,
max_c_cpost / max_abs_auto.max(1e-9),
rms_diff,
rms_auto,
argmax_t,
argmax_h,
argmax_v,
auto_cpu[argmax_idx],
chunk_cpu[argmax_idx]
);
(max_ac, max_a_cpost, max_c_cpost, max_a_cpre, max_c_cpre)
}
#[test]
fn chunk_vs_autoreg_divergence_scaling() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if std::env::var("HF2Q_W5B4_DIVERGENCE").as_deref() != Ok("1") {
eprintln!(
"chunk_vs_autoreg_divergence_scaling: skipped \
(set HF2Q_W5B4_DIVERGENCE=1 to enable)"
);
return;
}
eprintln!("---- GQA case: n_k=2, n_v=4 (group_ratio=2; tiled vs block differ) ----");
for &n in &[128u32, 256, 512, 1024, 2048] {
run_seqlen_with_heads(n, 2, 4);
}
eprintln!("---- No-GQA case: n_k=2, n_v=2 (group_ratio=1; tiled == block) ----");
for &n in &[128u32, 256, 512, 1024, 2048] {
run_seqlen_with_heads(n, 2, 2);
}
}
#[test]
fn dn_stage_a_byte_exact_parity_with_pre_phase3a() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
let shape = DeltaNetLayerShape {
hidden_size: 64,
n_k_heads: 2,
n_v_heads: 4,
d_k: 32,
d_v: 32,
conv_kernel: 4,
rms_norm_eps: 1e-6,
};
let seq_len: u32 = 16;
let _h = shape.hidden_size as usize;
let qkv_channels = shape.qkv_channels();
let nv = shape.n_v_heads as usize;
let dv = shape.d_v as usize;
let state_size = nv * shape.d_k as usize * dv;
let km1 = (shape.conv_kernel - 1) as usize;
let weights_cpu = synthetic_weights(shape, 0xDEADBEEF);
let gpu_weights = DeltaNetWeightsGpu::from_cpu_f32(
&weights_cpu,
&device,
shape.conv_kernel as usize,
qkv_channels as usize,
)
.expect("from_cpu_f32");
let x_cpu: Vec<f32> = (0..seq_len as usize * shape.hidden_size as usize)
.map(|i| 0.01 * (i as f32) - 0.5)
.collect();
let state_in = vec![0.0f32; state_size];
let conv_state = vec![0.0f32; km1 * qkv_channels as usize];
let x_prod = upload_f32(&x_cpu, &device).expect("upload x prod");
let state_in_prod = upload_f32(&state_in, &device).expect("upload state_in prod");
let state_out_prod = upload_f32(&state_in, &device).expect("upload state_out prod");
let conv_in_prod = upload_f32(&conv_state, &device).expect("upload conv_in prod");
let conv_out_prod = upload_f32(&conv_state, &device).expect("upload conv_out prod");
let mut arena_prod = super::super::DnPrefillArena::new(
&device,
seq_len,
shape.hidden_size,
shape.n_k_heads,
shape.n_v_heads,
shape.d_k,
shape.d_v,
)
.expect("DnPrefillArena prod");
let prod_out_buf = build_delta_net_layer_with_arena(
&device,
&mut registry,
&x_prod,
&gpu_weights,
&conv_in_prod,
&conv_out_prod,
&state_in_prod,
&state_out_prod,
&mut arena_prod,
None, None, seq_len,
shape.hidden_size,
shape.n_k_heads,
shape.n_v_heads,
shape.d_k,
shape.d_v,
shape.conv_kernel,
shape.rms_norm_eps,
None,
SlotId(0), )
.expect("production consolidated Stage-A path");
{
let mut flush_enc = device.command_encoder().expect("flush prod");
flush_enc
.commit_and_wait()
.expect("flush commit_and_wait prod");
}
let prod_out = download_f32(&prod_out_buf).expect("download prod");
let x_ref = upload_f32(&x_cpu, &device).expect("upload x ref");
let state_in_ref = upload_f32(&state_in, &device).expect("upload state_in ref");
let state_out_ref = upload_f32(&state_in, &device).expect("upload state_out ref");
let conv_in_ref = upload_f32(&conv_state, &device).expect("upload conv_in ref");
let conv_out_ref = upload_f32(&conv_state, &device).expect("upload conv_out ref");
let mut arena_ref = super::super::DnPrefillArena::new(
&device,
seq_len,
shape.hidden_size,
shape.n_k_heads,
shape.n_v_heads,
shape.d_k,
shape.d_v,
)
.expect("DnPrefillArena ref");
let n_q_elems = (seq_len * shape.n_k_heads * shape.d_k) as usize;
let q_scale_val = 1.0_f32 / (shape.d_k as f32).sqrt();
let rows_op8 = seq_len * shape.n_v_heads;
let z_channels = shape.n_v_heads * shape.d_v;
let q_sp = (shape.n_k_heads * shape.d_k) as usize;
let k_sp = (shape.n_k_heads * shape.d_k) as usize;
let v_sp = nv * dv;
{
let s = arena_ref
.ssm_params_buf
.as_mut_slice::<u32>()
.expect("ssm_params");
s[0] = qkv_channels;
s[1] = seq_len;
s[2] = FORWARD_DISPATCH_N_SEQS;
s[3] = shape.conv_kernel;
}
{
let s = arena_ref
.g_params_buf
.as_mut_slice::<u32>()
.expect("g_params");
s[0] = shape.n_v_heads;
s[1] = seq_len;
}
{
let s = arena_ref
.op8_params_buf
.as_mut_slice::<f32>()
.expect("op8_params");
s[0] = shape.rms_norm_eps;
s[1] = shape.d_v as f32;
}
let gdn_params = GatedDeltaNetParams {
d_k: shape.d_k,
d_v: shape.d_v,
n_k_heads: shape.n_k_heads,
n_v_heads: shape.n_v_heads,
n_tokens: seq_len,
n_seqs: FORWARD_DISPATCH_N_SEQS,
};
{
let s = arena_ref
.gdn_params_buf
.as_mut_slice::<u32>()
.expect("gdn_params");
s[0] = gdn_params.d_k;
s[1] = gdn_params.d_v;
s[2] = gdn_params.n_k_heads;
s[3] = gdn_params.n_v_heads;
s[4] = gdn_params.n_tokens;
s[5] = gdn_params.n_seqs;
s[6] = 0;
s[7] = 0;
s[8] = 0; }
let ssm_conv_params = SsmConvParams {
channels: qkv_channels,
n_tokens: seq_len,
n_seqs: FORWARD_DISPATCH_N_SEQS,
k_width: shape.conv_kernel,
};
{
let mut enc = device.command_encoder().expect("ref enc ops1-3");
apply_pre_norm_into(
&mut enc,
&mut registry,
&device,
&x_ref,
&gpu_weights.attn_norm,
&arena_ref.x_norm_buf,
&mut arena_ref.pre_norm_params_buf,
seq_len,
shape.hidden_size,
shape.rms_norm_eps,
)
.expect("ref pre_norm");
enc.memory_barrier();
apply_proj_into(
&mut enc,
&mut registry,
&device,
&arena_ref.x_norm_buf,
&gpu_weights.attn_qkv,
&mut arena_ref.qkv_raw_buf,
seq_len,
shape.hidden_size,
qkv_channels,
)
.expect("ref proj_qkv");
apply_proj_into(
&mut enc,
&mut registry,
&device,
&arena_ref.x_norm_buf,
&gpu_weights.attn_gate,
&mut arena_ref.z_buf,
seq_len,
shape.hidden_size,
z_channels,
)
.expect("ref proj_z");
enc.memory_barrier();
dispatch_ssm_conv(
&mut enc,
&mut registry,
device.metal_device(),
&arena_ref.qkv_raw_buf,
&gpu_weights.ssm_conv1d,
&conv_in_ref,
&conv_out_ref,
&arena_ref.qkv_conv_buf,
&arena_ref.ssm_params_buf,
ssm_conv_params,
)
.expect("ref ssm_conv");
enc.commit_labeled("ref.layer.gdn.ops1-3");
}
{
let params = QkvSplitParams {
seq: seq_len,
q_sp: q_sp as u32,
k_sp: k_sp as u32,
v_sp: v_sp as u32,
};
let mut enc = device.command_encoder().expect("ref enc qkv_split");
dispatch_qkv_split_f32(
&mut enc,
&mut registry,
device.metal_device(),
&arena_ref.qkv_conv_buf,
&arena_ref.q_split_buf,
&arena_ref.k_split_buf,
&arena_ref.v_split_buf,
¶ms,
)
.expect("ref qkv_split");
enc.commit_labeled("ref.layer.gdn.qkv_split");
}
let ref_out_buf = {
let mut enc = device.command_encoder().expect("ref enc ops5-9");
apply_l2_norm_per_head_into(
&mut enc,
&mut registry,
&device,
&arena_ref.q_split_buf,
&arena_ref.q_l2_buf,
&mut arena_ref.l2_params_q_buf,
seq_len,
shape.n_k_heads,
shape.d_k,
shape.rms_norm_eps,
)
.expect("ref l2_q");
apply_l2_norm_per_head_into(
&mut enc,
&mut registry,
&device,
&arena_ref.k_split_buf,
&arena_ref.k_normed_buf,
&mut arena_ref.l2_params_k_buf,
seq_len,
shape.n_k_heads,
shape.d_k,
shape.rms_norm_eps,
)
.expect("ref l2_k");
apply_proj_into(
&mut enc,
&mut registry,
&device,
&arena_ref.x_norm_buf,
&gpu_weights.ssm_alpha,
&mut arena_ref.alpha_logit_buf,
seq_len,
shape.hidden_size,
shape.n_v_heads,
)
.expect("ref alpha");
apply_proj_into(
&mut enc,
&mut registry,
&device,
&arena_ref.x_norm_buf,
&gpu_weights.ssm_beta,
&mut arena_ref.beta_logit_buf,
seq_len,
shape.hidden_size,
shape.n_v_heads,
)
.expect("ref beta");
enc.memory_barrier();
scalar_mul_f32(
&mut enc,
&mut registry,
device.metal_device(),
&arena_ref.q_l2_buf,
&arena_ref.q_scaled_buf,
n_q_elems,
q_scale_val,
)
.expect("ref scalar_mul");
dispatch_compute_g_beta(
&mut enc,
&mut registry,
device.metal_device(),
&arena_ref.alpha_logit_buf,
&arena_ref.beta_logit_buf,
&gpu_weights.ssm_dt_bias,
&gpu_weights.ssm_a,
&arena_ref.g_buf,
&arena_ref.beta_buf,
&arena_ref.g_params_buf,
seq_len,
shape.n_v_heads,
)
.expect("ref compute_g_beta");
enc.memory_barrier();
dispatch_gated_delta_net_decode(
&mut enc,
&mut registry,
device.metal_device(),
&arena_ref.q_scaled_buf,
&arena_ref.k_normed_buf,
&arena_ref.v_split_buf,
&arena_ref.g_buf,
&arena_ref.beta_buf,
&state_in_ref,
&arena_ref.attn_out_buf,
&state_out_ref,
&arena_ref.gdn_params_buf,
gdn_params,
)
.expect("ref gdn_decode");
enc.memory_barrier();
dispatch_ssm_norm_gate(
&mut enc,
&mut registry,
device.metal_device(),
&arena_ref.attn_out_buf,
&gpu_weights.ssm_norm,
&arena_ref.z_buf,
&arena_ref.gated_buf,
&arena_ref.op8_params_buf,
rows_op8,
shape.d_v,
)
.expect("ref ssm_norm_gate");
enc.memory_barrier();
let out = apply_proj(
&mut enc,
&mut registry,
&device,
&arena_ref.gated_buf,
&gpu_weights.ssm_out,
seq_len,
z_channels,
shape.hidden_size,
)
.expect("ref out_proj");
enc.commit_labeled("ref.layer.gdn.ops5-9");
out
};
{
let mut flush_enc = device.command_encoder().expect("flush ref");
flush_enc
.commit_and_wait()
.expect("flush commit_and_wait ref");
}
let ref_out = download_f32(&ref_out_buf).expect("download ref");
let _ = (k_sp, q_sp, v_sp);
assert!(
prod_out.iter().any(|&v| v != 0.0),
"production path returned ALL-ZERO output — GPU dispatch \
chain likely failed silently. ref_first_8={:?}",
&ref_out[..8.min(ref_out.len())]
);
assert!(
ref_out.iter().any(|&v| v != 0.0),
"reference path returned ALL-ZERO output — GPU dispatch \
chain likely failed silently. prod_first_8={:?}",
&prod_out[..8.min(prod_out.len())]
);
assert_eq!(
prod_out.len(),
ref_out.len(),
"byte-exact parity: output lengths differ — prod={} ref={}",
prod_out.len(),
ref_out.len(),
);
let mut n_diff = 0usize;
for (i, (&p, &r)) in prod_out.iter().zip(ref_out.iter()).enumerate() {
if p.to_bits() != r.to_bits() {
if n_diff < 5 {
eprintln!(
" byte-exact diff[{i}]: prod={p:.10} ({:#010x}) \
ref={r:.10} ({:#010x})",
p.to_bits(),
r.to_bits()
);
}
n_diff += 1;
}
}
assert_eq!(
n_diff,
0,
"dn_stage_a_byte_exact_parity_with_pre_phase3a FAIL: \
{n_diff}/{} F32 elements differ — Stage-A consolidation is NOT \
byte-identical to the legacy 2-CB structure. The new intra-CB \
memory_barrier between op3 (writes qkv_conv_buf) and qkv_split \
(reads qkv_conv_buf) is wrong or missing. iter89e2-I HOLD.",
prod_out.len(),
);
eprintln!(
"dn_stage_a_byte_exact_parity_with_pre_phase3a: \
0/{} elements differ (byte-exact) seq_len={seq_len}, \
hidden={}, n_k={}, n_v={}, d_k={}, d_v={}",
prod_out.len(),
shape.hidden_size,
shape.n_k_heads,
shape.n_v_heads,
shape.d_k,
shape.d_v,
);
}
}