//! Qwen GDN graph composition: fused live-state decode, fused-conv verify and //! masked/unmasked delta update and Conv1d prefill with ragged cache lengths. //! GatedDeltaNet branch selection and ArraysCache advancement use this graph; //! model-level capture commit/repair and product integration remain separate. use super::super::gpu::Buffer; use super::array::{Array, Dtype, Layout}; use super::mlp::QuantizedLinear; use super::stream::{Device, Stream}; use super::{dispatch_geometry, encoder, ops, tensor, views}; pub(super) enum InputProjections<'a> { Fused(&'a QuantizedLinear, [i32; 3]), Separate([&'a QuantizedLinear; 4]), } impl InputProjections<'_> { /// GatedDeltaNet's projection branch, including z's head reshape before /// constructing b/a in the separate case. Fused projections use Split siblings. pub(super) fn apply(&self, input: &Array, stream: Stream) -> Result<[Array; 4], String> { let shape = input.layout().shape().to_vec(); if shape.len() != 3 { return Err("GDN projection requires batch/sequence/hidden dimensions".into()); } let z_shape = [shape[0], shape[1], 48, 128]; match self { Self::Fused(projection, splits) => { let projected = projection.apply(input, stream)?; let [qkv, z, b, a]: [Array; 4] = ops::split(&projected, splits, -1, stream)? .try_into() .map_err(|_| "GDN projection must produce four outputs")?; Ok([qkv, ops::reshape(&z, &z_shape, stream)?, b, a]) } Self::Separate(projections) => { let [qkv, z, b, a] = projections; let qkv = qkv.apply(input, stream)?; let z = ops::reshape(&z.apply(input, stream)?, &z_shape, stream)?; Ok([qkv, z, b.apply(input, stream)?, a.apply(input, stream)?]) } } } } pub(super) struct Weights<'a> { pub(super) input: InputProjections<'a>, pub(super) output: &'a QuantizedLinear, pub(super) conv: &'a Array, pub(super) a_log: &'a Array, pub(super) dt_bias: &'a Array, pub(super) norm: &'a Array, } #[derive(Default)] pub(super) struct Cache { pub(super) conv: Option, pub(super) delta: Option, pub(super) lengths: Option, pub(super) left_padding: Option, pub(super) capture: Option<[Array; 6]>, } impl Cache { /// First pass of Qwen4ExpTextModel.commit_verified_window. The outer model /// must validate every layer before any layer replays or trims its cache. pub(super) fn validate_verified_window( &self, snapshot: Option<&[Option]>, verified: i32, ) -> Result<(), String> { let snapshot = snapshot.ok_or("snapshot_missing")?; let rows = self.capture.as_ref().ok_or("gdn_rows_missing")?; let width = rows[0].layout().shape()[1]; if width != verified { return Err(format!("gdn_rows_width_{width}_vs_{verified}")); } if snapshot.len() < 2 || snapshot[1].is_none() { return Err("gdn_snapshot_short".into()); } Ok(()) } /// GDN branch of the second pass, after model-wide validation. Only the /// kept recurrence is replayed; no projection, convolution or eval here. pub(super) fn replay_verified_prefix( &mut self, snapshot: &[Option], keep: i32, a_log: &Array, dt_bias: &Array, stream: Stream, ) -> Result<(), String> { let captured = self.capture.as_ref().ok_or("gdn_rows_missing")?; if keep < 1 || keep > captured[0].layout().shape()[1] { return Err("invalid verified prefix length".into()); } let state = snapshot .get(1) .and_then(Option::as_ref) .ok_or("gdn_snapshot_short")?; let qkv_shape = captured[0].layout().shape().to_vec(); let conv = match &snapshot[0] { Some(conv) => conv.clone(), None => { let zero = Array::new( &[], captured[0].layout().dtype(), super::scalar_buffer(&[0, 0])?, )?; ops::broadcast_to(&zero, &[qkv_shape[0], 3, qkv_shape[2]], stream)? } }; let prefix = |array: &Array| { let mut stop = array.layout().shape().to_vec(); stop[1] = keep; ops::slice( array, &vec![0; stop.len()], &stop, &vec![1; stop.len()], stream, ) }; let [qkv, q, k, v, a, b] = captured; let updated = delta_update( [&prefix(q)?, &prefix(k)?, &prefix(v)?], [&prefix(a)?, &prefix(b)?, a_log, dt_bias], Some(state), None, stream, )?; let window = ops::concatenate(&[conv, prefix(qkv)?], 1, stream)?; let tail = ops::slice( &window, &[0, -3, 0], &[qkv_shape[0], keep + 3, qkv_shape[2]], &[1, 1, 1], stream, )?; self.conv = Some(ops::contiguous(&tail, false, stream)?); self.delta = Some(updated.delta); self.capture = None; Ok(()) } fn advance(&mut self, rows: i32, stream: Stream) -> Result<(), String> { for array in [&mut self.lengths, &mut self.left_padding] .into_iter() .flatten() { if array.layout().dtype() != Dtype::I32 { return Err("GDN cache metadata requires INT32 arrays".into()); } let n = Array::new( &[], array.layout().dtype(), super::scalar_buffer(&rows.to_le_bytes())?, )?; *array = super::binary::binary(array, &n, super::binary::Binary::Subtract, stream)?; } Ok(()) } } #[derive(Clone, Copy, Default)] pub(super) struct Options { pub(super) step: bool, pub(super) conv_single: bool, pub(super) conv_verify: bool, pub(super) output: bool, pub(super) capture: bool, } /// Installed inference geometry and sigmoid norm. GatedDeltaNet.__call__'s /// branch order, mutations and graph construction, including absent caches. pub(super) fn forward( weights: &Weights<'_>, input: &Array, mask: Option<&Array>, mut cache: Option<&mut Cache>, options: Options, streams: &super::stream::Streams, stream: Stream, ) -> Result { let shape = input.layout().shape().to_vec(); if shape.len() != 3 || shape[0] < 1 || shape[1] < 1 || shape[2] != 2560 { return Err("GDN requires [B,S,2560] input".into()); } let (batch, rows) = (shape[0], shape[1]); let [qkv, z, b, a] = weights.input.apply(input, stream)?; let conv = if let Some(conv) = cache.as_ref().and_then(|c| c.conv.as_ref()) { conv.clone() } else { let scalar = Array::new(&[], Dtype::BF16, super::scalar_buffer(&[0, 0])?)?; ops::broadcast_to(&scalar, &[batch, 3, 10240], stream)? }; let qkv = if let Some(mask) = mask { if mask.layout().shape() != [batch, rows] { return Err("GDN mask shape mismatch".into()); } let zero = Array::new(&[], Dtype::BF16, super::scalar_buffer(&[0, 0])?)?; super::indexing::select(&ops::expand_dims(mask, &[-1], stream)?, &qkv, &zero, stream)? } else { qkv }; let dense_cached = batch == 1 && mask.is_none() && cache.as_ref().is_some_and(|c| c.lengths.is_none()); if rows == 1 && dense_cached && options.step && !options.capture && let Some(delta) = cache.as_ref().and_then(|c| c.delta.as_ref()) && delta.layout().dtype() == Dtype::F32 { let [gated, next_conv, next_delta] = fused_step( [ &ops::reshape(&qkv, &[-1], stream)?, &ops::reshape(&z, &[-1], stream)?, &ops::reshape(&a, &[-1], stream)?, &ops::reshape(&b, &[-1], stream)?, &ops::reshape(&conv, &[3, 10240], stream)?, weights.conv, weights.a_log, weights.dt_bias, delta, weights.norm, ], stream, )?; let cache = cache.as_deref_mut().unwrap(); cache.conv = Some(ops::reshape(&next_conv, &[batch, 3, 10240], stream)?); cache.delta = Some(ops::reshape(&next_delta, &[batch, 48, 128, 128], stream)?); cache.advance(rows, stream)?; return weights .output .apply(&ops::reshape(&gated, &[batch, rows, -1], stream)?, stream); } let fused_conv = dense_cached && ((rows == 1 && options.conv_single) || ((2..=6).contains(&rows) && options.conv_verify)) && supports_conv(streams, stream, rows > 1); let [q, k, v] = if fused_conv { let qkv_shape = if rows == 1 { vec![-1] } else { vec![rows, -1] }; let [q, k, v, next_conv] = fused_conv_norm( &ops::reshape(&qkv, &qkv_shape, stream)?, &ops::reshape(&conv, &[3, 10240], stream)?, weights.conv, stream, )?; cache.as_deref_mut().unwrap().conv = Some(ops::reshape(&next_conv, &[batch, 3, 10240], stream)?); [ ops::reshape(&q, &[batch, rows, 16, 128], stream)?, ops::reshape(&k, &[batch, rows, 16, 128], stream)?, ops::reshape(&v, &[batch, rows, 48, 128], stream)?, ] } else { conv_fallback(weights, &qkv, &conv, cache.as_deref_mut(), stream)? }; if options.capture && let Some(cache) = cache.as_deref_mut() { cache.capture = Some([qkv, q.clone(), k.clone(), v.clone(), a.clone(), b.clone()]); } let recurrent = delta_update( [&q, &k, &v], [&a, &b, weights.a_log, weights.dt_bias], cache.as_ref().and_then(|c| c.delta.as_ref()), mask, stream, )?; if let Some(cache) = cache { cache.delta = Some(recurrent.delta); cache.advance(rows, stream)?; } if batch == 1 && rows == 1 && options.output && weights.output.bits == 4 && matches!(weights.output.group, 32 | 64) { let output = fused_output( [ &ops::reshape(&recurrent.hidden, &[-1], stream)?, &ops::reshape(&z, &[-1], stream)?, weights.norm, &weights.output.weight, &weights.output.scales, &weights.output.biases, ], weights.output.group, stream, )?; return ops::reshape(&output, &[batch, rows, -1], stream); } let normed = super::normalization::sigmoid_rms_norm( &recurrent.hidden, weights.norm, Some(&z), 1e-6, stream, )?; weights .output .apply(&ops::reshape(&normed, &[batch, rows, -1], stream)?, stream) } fn supports_conv(streams: &super::stream::Streams, stream: Stream, multi: bool) -> bool { static SINGLE: std::sync::OnceLock = std::sync::OnceLock::new(); static MULTI: std::sync::OnceLock = std::sync::OnceLock::new(); *(if multi { &MULTI } else { &SINGLE }).get_or_init(|| { let probe = || -> Result<(), String> { let zero = |shape: &[i32]| { let scalar = Array::new(&[], Dtype::BF16, super::scalar_buffer(&[0, 0])?)?; ops::broadcast_to(&scalar, shape, stream) }; let x = zero(if multi { &[2, 10240] } else { &[10240] })?; let outputs = fused_conv_norm(&x, &zero(&[3, 10240])?, &zero(&[10240, 4])?, stream)?; ops::evaluate(streams, &outputs, stream, false) }; if let Err(error) = probe() { eprintln!("MTPLX fused GDN conv disabled on this GPU (multi={multi}): {error}"); false } else { true } }) } pub(super) fn fused_output( inputs: [&Array; 6], group: u32, stream: Stream, ) -> Result { let dtype = inputs[0].layout().dtype(); if !matches!(dtype, Dtype::BF16 | Dtype::F32) || !matches!(group, 32 | 64) || inputs[1].layout().dtype() != Dtype::BF16 || inputs[2].layout().dtype() != Dtype::BF16 || inputs[3].layout().dtype() != Dtype::U32 || inputs[4].layout().dtype() != Dtype::BF16 || inputs[5].layout().dtype() != Dtype::BF16 || inputs .iter() .zip([ 6144, 6144, 128, 2560 * 768, 2560 * 6144 / group as usize, 2560 * 6144 / group as usize, ]) .any(|(a, n)| a.layout().size() != n) { return Err( "fused GDN output requires installed BF16 quantization and BF16/FP32 hidden".into(), ); } Array::make_operation_with_inputs( stream, &inputs.map(Clone::clone), Layout::new(&[2560], dtype)?, ops::Operation::GdnOutput(group), ) } pub(super) fn output_evaluate( inputs: &[Array], outputs: &[Array], group: u32, ) -> Result<(), String> { let name = format!( "kernel_qwen_mtplx_gdn_out_fused_gs{group}_{}", if inputs[0].layout().dtype() == Dtype::F32 { "f32" } else { "bf16" } ); custom_evaluate(inputs, outputs, |bindings| { dispatch_geometry(&name, bindings, &[], [80 * 1024, 1, 1], [1024, 1, 1], true) }) } /// Original gdn_step_fused.py wrapper and its three-output custom primitive. pub(super) fn fused_step(inputs: [&Array; 10], stream: Stream) -> Result<[Array; 3], String> { let sizes = [10240, 6144, 48, 48, 30720, 40960, 48, 48, 786432, 128]; if stream.device() != Device::Gpu || inputs.iter().enumerate().any(|(i, a)| { let layout = a.layout(); layout.size() != sizes[i] || layout.dtype() != if i == 8 { Dtype::F32 } else { Dtype::BF16 } }) { return Err("fused GDN step requires installed BF16 geometry and FP32 delta state".into()); } let [qkv, z, a, b, conv, cw, a_log, dt_bias, delta, norm] = inputs; let cw = ops::reshape(cw, &[10240, 4], stream)?; let inputs = [ qkv.clone(), z.clone(), a.clone(), b.clone(), ops::reshape(conv, &[-1], stream)?, ops::reshape(&cw, &[-1], stream)?, a_log.clone(), dt_bias.clone(), ops::reshape(delta, &[-1], stream)?, norm.clone(), ]; Array::make_operations( stream, &inputs, &[ Layout::new(&[6144], Dtype::BF16)?, Layout::new(&[3, 10240], Dtype::BF16)?, Layout::new(&[48, 128, 128], Dtype::F32)?, ], ops::Operation::GdnStepBf16, )? .try_into() .map_err(|_| "GDN step must produce three outputs".into()) } pub(super) fn step_evaluate(inputs: &[Array], outputs: &[Array]) -> Result<(), String> { custom_evaluate(inputs, outputs, dispatch_step) } pub(super) fn custom_evaluate( inputs: &[Array], outputs: &[Array], dispatch: impl FnOnce(&[super::Binding<'_>]) -> Result<(), String>, ) -> Result<(), String> { // CustomKernel allocates every output first; no donation, then General // copies of non-row-contiguous inputs in their original argument order. for output in outputs { let nbytes = output.layout().nbytes(); output.set_data(Buffer::mtplx_bytes(nbytes as u64)?)?; } let mut copies = Vec::new(); let prepared = inputs .iter() .map(|input| { let layout = input.layout(); if layout.flags().row_contiguous { return Ok(input.clone()); } let copy = Array::new( layout.shape(), layout.dtype(), Buffer::mtplx_bytes(layout.nbytes() as u64)?, )?; drop(layout); views::general_copy_inplace(input, ©)?; copies.push(copy.clone()); Ok(copy) }) .collect::, String>>()?; let input_buffers = prepared.iter().map(Array::buffer).collect::>(); let output_buffers = outputs.iter().map(Array::buffer).collect::>(); let bindings = prepared .iter() .zip(&input_buffers) .enumerate() .map(|(i, (a, b))| { tensor(i as u32, b) .at_byte_offset(a.offset()) .input(a.data_size()) }) .chain( outputs .iter() .zip(&output_buffers) .enumerate() .map(|(i, (a, b))| { tensor((i + inputs.len()) as u32, b) .at_byte_offset(a.offset()) .output(a.data_size()) }), ) .collect::>(); let result = dispatch(&bindings); for copy in &copies { encoder::add_temporaries(std::slice::from_ref(&*copy.buffer())); } result } pub(super) fn dispatch_step(bindings: &[super::Binding<'_>]) -> Result<(), String> { // The Python custom-kernel grid counts THREADS, not threadgroups. dispatch_geometry( "kernel_qwen_mtplx_gdn_step_fused_f32_state_bf16", bindings, &[], [256, 1, 48], [256, 1, 1], true, ) } /// Fused single/verify-row conv wrappers, not the general Conv1d fallback. pub(super) fn fused_conv_norm( qkv: &Array, conv: &Array, weight: &Array, stream: Stream, ) -> Result<[Array; 4], String> { let shape = qkv.layout().shape().to_vec(); let rows = match shape.as_slice() { [10240] => 1, [s, 10240] if (2..=6).contains(s) => *s as u32, _ => return Err("fused conv norm requires one decode row or 2..6 verify rows".into()), }; if stream.device() != Device::Gpu || [qkv, conv, weight] .iter() .any(|a| a.layout().dtype() != Dtype::BF16) || conv.layout().shape() != [3, 10240] || weight.layout().size() != 40960 { return Err("invalid fused conv norm dtype/state/weight".into()); } let cw = ops::reshape(weight, &[10240, 4], stream)?; let inputs = [ if rows == 1 { qkv.clone() } else { ops::reshape(qkv, &[-1], stream)? }, ops::reshape(conv, &[-1], stream)?, ops::reshape(&cw, &[-1], stream)?, ]; let layout = |width| { Layout::new( &if rows == 1 { vec![width] } else { vec![rows as i32, width] }, Dtype::BF16, ) }; Array::make_operations( stream, &inputs, &[ layout(2048)?, layout(2048)?, layout(6144)?, Layout::new(&[3, 10240], Dtype::BF16)?, ], ops::Operation::GdnConvNorm(rows), )? .try_into() .map_err(|_| "conv norm requires four outputs".into()) } pub(super) fn conv_norm_evaluate( inputs: &[Array], outputs: &[Array], rows: u32, ) -> Result<(), String> { let name = if rows == 1 { "kernel_qwen_mtplx_gdn_conv_norm_bf16".into() } else { format!("kernel_qwen_mtplx_gdn_conv_norm_rows_s{rows}_bf16") }; custom_evaluate(inputs, outputs, |bindings| { dispatch_geometry(&name, bindings, &[], [10240, 1, 1], [1024, 1, 1], true) }) } pub(super) fn compute_g( a_log: &Array, a: &Array, dt_bias: &Array, stream: Stream, ) -> Result { let shape = a.layout().shape().to_vec(); if stream.device() != Device::Gpu || shape.last() != Some(&48) || a_log.layout().shape() != [48] || dt_bias.layout().shape() != [48] || [a_log, a, dt_bias] .iter() .any(|a| a.layout().dtype() != Dtype::BF16) { return Err("compute_g requires BF16 per-head inputs".into()); } if [a_log, a, dt_bias].iter().any(|a| a.is_tracer()) { use super::{ binary::{Binary, binary}, unary::{Unary, unary}, }; let decay = unary( &unary( &ops::astype(a_log, Dtype::F32, false, stream)?, Unary::Exp, stream, )?, Unary::Negative, stream, )?; let bias = binary(a, dt_bias, Binary::Add, stream)?; let zero = Array::new(&[], Dtype::BF16, super::scalar_buffer(&[0, 0])?)?; let softplus = binary(&bias, &zero, Binary::LogAddExp, stream)?; return unary( &binary(&decay, &softplus, Binary::Multiply, stream)?, Unary::Exp, stream, ); } Array::make_operation_with_inputs( stream, &[a_log.clone(), a.clone(), dt_bias.clone()], Layout::new(&shape, Dtype::F32)?, ops::Operation::ComputeGBf16, ) } pub(super) struct DeltaOutput { pub(super) hidden: Array, pub(super) delta: Array, #[cfg_attr(not(test), expect(dead_code, reason = "Reference diagnostic API"))] pub(super) g: Array, #[cfg_attr(not(test), expect(dead_code, reason = "Reference diagnostic API"))] pub(super) beta: Array, } pub(super) fn delta_update( qkv: [&Array; 3], gates: [&Array; 4], state: Option<&Array>, mask: Option<&Array>, stream: Stream, ) -> Result { let [q, _, v] = qkv; let [a, b, a_log, dt_bias] = gates; let beta = super::unary::unary(b, super::unary::Unary::Sigmoid, stream)?; let g = compute_g(a_log, a, dt_bias, stream)?; let initial_state = if state.is_none() { let qshape = q.layout().shape().to_vec(); let vshape = v.layout().shape().to_vec(); if qshape.len() != 4 || vshape.len() != 4 { return Err("delta update requires rank-4 Q/V".into()); } let zero = Array::new(&[], Dtype::F32, super::scalar_buffer(&0_f32.to_le_bytes())?)?; Some(ops::broadcast_to( &zero, &[qshape[0], vshape[2], vshape[3], qshape[3]], stream, )?) } else { None }; let [hidden, delta] = delta_kernel( qkv, [&g, &beta], state.or(initial_state.as_ref()).unwrap(), mask, stream, )?; Ok(DeltaOutput { hidden, delta, g, beta, }) } /// Inference custom kernel from the pinned gated_delta.py; masks preserve state. pub(super) fn delta_kernel( qkv: [&Array; 3], gates: [&Array; 2], state: &Array, mask: Option<&Array>, stream: Stream, ) -> Result<[Array; 2], String> { let [q, k, v] = qkv; let [g, beta] = gates; let shape = q.layout().shape().to_vec(); if shape.len() != 4 || shape[0] <= 0 || shape[1] <= 0 || shape[2..] != [16, 128] { return Err("delta kernel requires [B,T,16,128] Q/K".into()); } let (batch, rows) = (shape[0], shape[1]); let hidden_shape = [batch, rows, 48, 128]; if stream.device() != Device::Gpu || k.layout().shape() != shape || v.layout().shape() != hidden_shape || [q, k, v, beta] .iter() .any(|a| a.layout().dtype() != Dtype::BF16) || g.layout().dtype() != Dtype::F32 || state.layout().dtype() != Dtype::F32 || g.layout().shape() != [batch, rows, 48] || beta.layout().shape() != [batch, rows, 48] || state.layout().shape() != [batch, 48, 128, 128] || mask.is_some_and(|m| { m.layout().shape() != [batch, rows] || m.layout().dtype() != Dtype::Bool }) { return Err("invalid delta kernel dtype/shape".into()); } let length = Array::new(&[], Dtype::I32, super::scalar_buffer(&rows.to_le_bytes())?)?; let mut inputs = vec![ q.clone(), k.clone(), v.clone(), g.clone(), beta.clone(), state.clone(), length, ]; if let Some(mask) = mask { inputs.push(mask.clone()); } Array::make_operations( stream, &inputs, &[ Layout::new(&hidden_shape, Dtype::BF16)?, Layout::new(state.layout().shape(), Dtype::F32)?, ], ops::Operation::GatedDeltaBf16(mask.is_some()), )? .try_into() .map_err(|_| "delta kernel requires two outputs".into()) } pub(super) fn delta_evaluate( inputs: &[Array], outputs: &[Array], masked: bool, ) -> Result<(), String> { let batch = inputs[0].layout().shape()[0] as u32; custom_evaluate(inputs, outputs, |bindings| { dispatch_geometry( if masked { "kernel_qwen_mtplx_gated_delta_masked_bf16" } else { "kernel_qwen_mtplx_gated_delta_bf16" }, bindings, &[], [32, 128, batch.checked_mul(48).ok_or("delta grid overflow")?], [32, 4, 1], true, ) }) } #[cfg_attr(not(test), expect(dead_code, reason = "Reference diagnostic API"))] pub(super) struct VerifyOutput { pub(super) hidden: Array, pub(super) conv: Array, pub(super) delta: Array, /// Exact capture-commit inputs, retained only when capture is active. pub(super) capture: Option<[Array; 6]>, // qkv, q, k, v, a, b } /// Installed GDN's conv1d branch: groups=C=O, stride/dilations=1, padding=0. /// Other Convolution branches (including dilated PLE) are not substituted here. fn depthwise_conv1d(input: &Array, weight: &Array, stream: Stream) -> Result { let src = input.layout(); let wt = weight.layout(); if src.shape().len() != 3 || wt.shape().len() != 3 || !matches!(src.dtype(), Dtype::BF16 | Dtype::F16 | Dtype::F32) || src.shape()[2] == 0 || wt.shape()[0] != src.shape()[2] || wt.shape()[2] != 1 || wt.shape()[1] <= 0 || wt.shape()[1] > src.shape()[1] { return Err("depthwise Conv1d requires [B,T,C] and [C,K,1], 0 Result<(), String> { let prepared = prepare_conv_inputs(inputs, output)?; conv1d_dispatch(&prepared, output) } /// Both installed GDN and PLE convolutions use this original allocation/copy order. pub(super) fn prepare_conv_inputs(inputs: &[Array], output: &Array) -> Result, String> { // conv.cpp allocates out before copying input, then weight, in that order. let nbytes = output.layout().nbytes(); output.set_data(Buffer::mtplx_bytes(nbytes as u64)?)?; inputs .iter() .map(|input| { let layout = input.layout(); if layout.flags().row_contiguous { return Ok(input.clone()); } let copy = Array::new( layout.shape(), layout.dtype(), Buffer::mtplx_bytes(layout.nbytes() as u64)?, )?; drop(layout); views::general_copy_inplace(input, ©)?; encoder::add_temporaries(std::slice::from_ref(&*copy.buffer())); Ok(copy) }) .collect::, String>>() } fn conv1d_dispatch(prepared: &[Array], output: &Array) -> Result<(), String> { let [input, weight] = prepared else { return Err("invalid conv arity".into()); }; let layout = input.layout(); let large = layout.size() > i32::MAX as usize || input.data_size() > i32::MAX as usize; let strides = layout.strides().to_vec(); let small_strides = strides.iter().map(|&s| s as i32).collect::>(); let grid = [ layout.shape()[2] as u32, output.layout().shape()[1] as u32, layout.shape()[0] as u32, ]; let width = weight.layout().shape()[1]; let name = format!( "depthwise_conv_1d_{}{}", layout.dtype().kernel_name(), if large { "_large" } else { "" } ); drop(layout); let x = input.buffer(); let w = weight.buffer(); let out = output.buffer(); let bindings = [ tensor(0, &x) .at_byte_offset(input.offset()) .input(input.data_size()), tensor(1, &w) .at_byte_offset(weight.offset()) .input(weight.data_size()), tensor(2, &out) .at_byte_offset(output.offset()) .output(output.data_size()), if large { super::bytes(3, strides.as_slice()) } else { super::bytes(3, small_strides.as_slice()) }, super::bytes(4, &width), ]; dispatch_geometry(&name, &bindings, &[], grid, super::block_dims(grid), true) } /// Reduce::Sum's contiguous last-axis branches used by GDN l2norm and PLE. /// The remaining reduction plans are explicitly outside this primitive. fn l2norm(input: &Array, stream: Stream) -> Result { use super::binary::{Binary, binary}; use super::unary::{Unary, unary}; let xf = ops::astype(input, Dtype::F32, false, stream)?; let square = binary(&xf, &xf, Binary::Multiply, stream)?; let shape = square.layout().shape().to_vec(); if shape.last() != Some(&128) { return Err("installed GDN l2norm requires head width 128".into()); } let sum = super::reduce::sum(&square, &[-1], true, stream)?; let eps = Array::new( &[], Dtype::F32, super::scalar_buffer(&1e-6_f32.to_le_bytes())?, )?; let scale = unary( &binary(&sum, &eps, Binary::Add, stream)?, Unary::Rsqrt, stream, )?; ops::astype( &binary(&xf, &scale, Binary::Multiply, stream)?, input.layout().dtype(), false, stream, ) } /// GatedDeltaNet's fallback after projection, zero-state creation and masking. fn conv_fallback( weights: &Weights<'_>, qkv: &Array, conv: &Array, cache: Option<&mut Cache>, stream: Stream, ) -> Result<[Array; 3], String> { let shape = qkv.layout().shape().to_vec(); let (batch, rows) = (shape[0], shape[1]); let window = ops::concatenate(&[conv.clone(), qkv.clone()], 1, stream)?; if let Some(cache) = cache { let conv = if let Some(lengths) = &cache.lengths { use super::binary::{Binary, binary}; if lengths.layout().shape() != [batch] || lengths.layout().dtype() != Dtype::I32 { return Err("GDN cache lengths require one INT32 value per batch".into()); } let zero = Array::new(&[], Dtype::I32, super::scalar_buffer(&0_i32.to_le_bytes())?)?; let upper = Array::new(&[], Dtype::I32, super::scalar_buffer(&rows.to_le_bytes())?)?; let ends = binary( &binary(lengths, &zero, Binary::Maximum, stream)?, &upper, Binary::Minimum, stream, )?; let ends = ops::expand_dims(&ends, &[-1], stream)?; let range = super::indexing::arange_i32(3, stream)?; let positions = binary(&ends, &range, Binary::Add, stream)?; let positions = ops::expand_dims(&positions, &[-1], stream)?; super::indexing::take_along_axis(&window, &positions, 1, stream)? } else { ops::contiguous( &ops::slice( &window, &[0, -3, 0], &[batch, rows + 3, 10240], &[1, 1, 1], stream, )?, false, stream, )? }; cache.conv = Some(conv); } let activated = super::mlp::silu(&depthwise_conv1d(&window, weights.conv, stream)?, stream)?; let [q, k, v]: [Array; 3] = ops::split(&activated, &[2048, 4096], -1, stream)? .try_into() .map_err(|_| "invalid GDN split")?; let q = ops::reshape(&q, &[batch, rows, 16, 128], stream)?; let k = ops::reshape(&k, &[batch, rows, 16, 128], stream)?; let v = ops::reshape(&v, &[batch, rows, 48, 128], stream)?; // Python array.__rmul__ keeps the array FIRST even for inv_scale * array // (pinned python/src/array.cpp:631). Preserve this order for compile_fuse. let scale = 128_f32.powf(-0.5).to_bits(); let scale = ((scale + 0x7fff + ((scale >> 16) & 1)) >> 16) as u16; let scale = Array::new( &[], Dtype::BF16, super::scalar_buffer(&scale.to_le_bytes())?, )?; let q = super::binary::binary( &l2norm(&q, stream)?, &scale, super::binary::Binary::Multiply, stream, )?; let k = l2norm(&k, stream)?; Ok([q, k, v]) } #[test] #[ignore = "requires Apple Metal; complete GDN Conv1d prefill and fused verify graphs"] fn mtplx_graph_gdn_prefill_and_verify_match_model_reference() { use super::super::{ gpu::Context, qwen_mtplx_tests::{capture_dispatches, pattern, pattern_f32, verify_typed}, }; use super::{allocator, configure_sources, mlp::fixture_linear, stream::Streams}; configure_sources().unwrap(); let _context = Context::open_qwen(0).unwrap(); let streams = Streams::new(Some(allocator::Allocator::new().unwrap())); let gpu = streams.default_stream(Device::Gpu).unwrap(); let in_proj = fixture_linear(16480, 2560, 4, 64, 26); let out_proj = fixture_linear(2560, 6144, 4, 64, 28); let bf16 = |shape: &[i32], salt| { Array::new( shape, Dtype::BF16, pattern(shape.iter().map(|&d| d as u32).product(), salt), ) .unwrap() }; let cw = bf16(&[10240, 4, 1], 8); let a_log = bf16(&[48], 12); let dt = bf16(&[48], 13); let norm = bf16(&[128], 14); let initial_conv = bf16(&[1, 3, 10240], 7); let initial_delta = Array::new(&[1, 48, 128, 128], Dtype::F32, pattern_f32(786432, 15)).unwrap(); let weights = Weights { input: InputProjections::Fused(&in_proj, [10240, 16384, 16432]), output: &out_proj, conv: &cw, a_log: &a_log, dt_bias: &dt, norm: &norm, }; let fixtures = include_str!("../../../../tests/fixtures/mtplx-custom-kernels.jsonl") .lines() .map(|l| serde_json::from_str::(l).unwrap()) .filter(|r| r["kernel"].as_str().unwrap().starts_with("gdn_staged_")) .collect::>(); assert_eq!(fixtures.len(), 130); let mut without_capture = Vec::new(); for capture in [false, true] { let mut cache = Cache::default(); for (case, row) in fixtures.iter().enumerate() { super::eval::execution_tests::observed(|| { let rows = row["rows"].as_u64().unwrap() as i32; let step = row["step"].as_u64().unwrap() as u32; let initial = row["initial"].as_bool().unwrap(); if step == 0 { cache = Cache { conv: initial.then(|| initial_conv.clone()), delta: initial.then(|| initial_delta.clone()), ..Cache::default() }; } let input = bf16(&[1, rows, 2560], 30 + step); let mask = (row["masked"] == true).then(|| { Array::new( &[1, rows], Dtype::Bool, super::scalar_buffer( &(0..rows).map(|n| u8::from(n % 3 != 0)).collect::>(), ) .unwrap(), ) .unwrap() }); let lengths = row["length"].as_i64().map(|n| { Array::new( &[1], Dtype::I32, super::scalar_buffer(&(n as i32).to_le_bytes()).unwrap(), ) .unwrap() }); cache.lengths = lengths.clone(); let enabled = row["fused"] == true; let hidden = forward( &weights, &input, mask.as_ref(), Some(&mut cache), Options { step: false, conv_single: enabled, conv_verify: enabled, output: enabled, capture, }, &streams, gpu, ) .unwrap(); let next = VerifyOutput { hidden, conv: cache.conv.as_ref().unwrap().clone(), delta: cache.delta.as_ref().unwrap().clone(), capture: cache.capture.clone(), }; drop(input); let (result, dispatches) = capture_dispatches(|| { ops::evaluate( &streams, &[next.hidden.clone(), next.conv.clone(), next.delta.clone()], gpu, false, ) }); result.unwrap(); let selects = dispatches .iter() .filter(|r| r.0.contains("_Select")) .collect::>(); assert_eq!(selects.len(), usize::from(mask.is_some())); if let Some(select) = selects.first() { let name = if rows == 1 { "g1_Selectbfloat16" } else { "g2_Selectbfloat16" }; assert_eq!( **select, ( name.into(), [10240, rows as u32, 1], super::block_dims([10240, rows as u32, 1]), true ) ); } let gathers = dispatches .iter() .filter(|r| r.0.contains("gather_axis")) .collect::>(); assert_eq!(gathers.len(), usize::from(lengths.is_some())); if let Some(gather) = gathers.first() { assert_eq!( **gather, ( "kernel_qwen_mtplx_gather_axis_bf16_idxi32_int_10".into(), [10240, 3, 1], super::block_dims([10240, 3, 1]), true ) ); } if row["used_fused"] == false { let conv = dispatches .iter() .filter(|r| r.0.starts_with("depthwise_conv_1d_")) .collect::>(); assert_eq!(conv.len(), 1); assert_eq!( *conv[0], ( "depthwise_conv_1d_bfloat16".into(), [10240, rows as u32, 1], super::block_dims([10240, rows as u32, 1]), true ) ); let sums = dispatches .iter() .filter(|r| r.0.starts_with("row_reduce_")) .collect::>(); assert_eq!(sums.len(), 2); for sum in sums { let (name, groups) = if rows == 1 { ("row_reduce_looped_1_reduce_sumfloat32", 16) } else { ("row_reduce_simple_sumfloat32", rows as u32 * 4) }; assert_eq!(*sum, (name.into(), [32, groups, 1], [32, 1, 1], true)); } } let values = [&next.hidden, &next.conv, &next.delta].map(|a| { let mut bytes = vec![0; a.layout().nbytes()]; a.buffer().read(a.offset(), &mut bytes).unwrap(); bytes }); if let Some([_, q, k, v, _, _]) = &next.capture { assert_eq!(values, without_capture[case]); // The fallback V is a strided view. Pack only for the // fixture reader, AFTER observing the actual forward. let v = ops::contiguous(v, false, gpu).unwrap(); ops::evaluate(&streams, std::slice::from_ref(&v), gpu, false).unwrap(); verify_typed( row["kernel"].as_str().unwrap(), &[ (&next.hidden.buffer(), rows as u32 * 2560, 2), (&next.conv.buffer(), 30720, 2), (&next.delta.buffer(), 786432, 4), (&q.buffer(), rows as u32 * 2048, 2), (&k.buffer(), rows as u32 * 2048, 2), (&v.buffer(), rows as u32 * 6144, 2), ], ); } else { // Compare all three outputs byte-for-byte with the // captured run, whose six outputs have actual receipts. without_capture.push(values); } if let Some(lengths) = &cache.lengths { ops::evaluate(&streams, std::slice::from_ref(lengths), gpu, false).unwrap(); let mut bytes = [0; 4]; lengths.buffer().read(lengths.offset(), &mut bytes).unwrap(); assert_eq!( i32::from_le_bytes(bytes), row["length"].as_i64().unwrap() as i32 - rows ); } }); } } streams.clear_streams().unwrap(); } #[test] #[ignore = "requires Apple Metal; actual MTPLX GDN forward and cache receipts"] fn mtplx_graph_gdn_forward_matches_model_reference() { use super::super::{ gpu::Context, qwen_mtplx_tests::{capture_dispatches, pattern, pattern_f32, verify_typed}, }; use super::{allocator, configure_sources, mlp::fixture_linear, stream::Streams}; configure_sources().unwrap(); let _context = Context::open_qwen(0).unwrap(); let streams = Streams::new(Some(allocator::Allocator::new().unwrap())); let gpu = streams.default_stream(Device::Gpu).unwrap(); let fused = fixture_linear(16480, 2560, 4, 64, 26); let separate: [_; 4] = std::array::from_fn(|i| { fixture_linear([10240, 6144, 48, 48][i], 2560, 4, 64, 64 + i as u32 * 2) }); let mixed: [_; 4] = std::array::from_fn(|i| { let [bits, group] = [[4, 64], [8, 32], [8, 64], [4, 32]][i]; fixture_linear( [10240, 6144, 48, 48][i], 2560, bits, group, 64 + i as u32 * 2, ) }); let out = fixture_linear(2560, 6144, 4, 64, 28); let bf16 = |shape: &[i32], salt| { Array::new( shape, Dtype::BF16, pattern(shape.iter().map(|&d| d as u32).product(), salt), ) .unwrap() }; let cw = bf16(&[10240, 4, 1], 8); let a_log = bf16(&[48], 12); let dt_bias = bf16(&[48], 13); let norm = bf16(&[128], 14); let initial_conv = bf16(&[1, 3, 10240], 7); let initial_delta = Array::new( &[1, 48, 128, 128], Dtype::F32, pattern_f32(48 * 128 * 128, 15), ) .unwrap(); let receipts = include_str!("../../../../tests/fixtures/mtplx-custom-kernels.jsonl") .lines() .map(|line| serde_json::from_str::(line).unwrap()) .filter(|row| row["kernel"].as_str().unwrap().starts_with("gdn_forward_")) .map(|row| (row["kernel"].as_str().unwrap().to_owned(), row)) .collect::>(); assert_eq!(receipts.len(), 264); for layout in ["fused", "separate", "mixed"] { let input = match layout { "fused" => InputProjections::Fused(&fused, [10240, 16384, 16432]), "separate" => InputProjections::Separate(separate.each_ref()), _ => InputProjections::Separate(mixed.each_ref()), }; let weights = Weights { input, output: &out, conv: &cw, a_log: &a_log, dt_bias: &dt_bias, norm: &norm, }; let cases = (0..4) .flat_map(|initial| (0..6).map(move |mode| (1, initial, mode))) .chain([2, 6, 7, 33, 2048].into_iter().flat_map(|rows| { (0..2).flat_map(move |initial| (0..2).map(move |mode| (rows, initial, mode))) })) .collect::>(); if layout == "fused" { let fixtures = include_str!("../../../../tests/fixtures/mtplx-gdn-lifecycle.jsonl") .lines() .map(|line| serde_json::from_str::(line).unwrap()) .collect::>(); assert_eq!(fixtures.len(), 24); let mut cache = Cache::default(); for row in fixtures { super::eval::execution_tests::observed(|| { let [batch, rows, step] = ["batch", "rows", "step"].map(|key| row[key].as_i64().unwrap() as i32); let mode = row["mode"].as_str().unwrap(); if step == 0 { cache = Cache::default(); if mode == "metadata" { let metadata = |values: [i32; 2]| { Array::new( &[batch], Dtype::I32, super::scalar_buffer( &values[..batch as usize] .iter() .flat_map(|value| value.to_le_bytes()) .collect::>(), ) .unwrap(), ) .unwrap() }; cache.lengths = Some(metadata([rows + 1, rows - 1])); cache.left_padding = Some(metadata([1, 0])); } } let old_capture = cache .capture .as_ref() .map(|arrays| arrays.each_ref().map(Array::id)); let mask = (mode == "metadata").then(|| { Array::new( &[batch, rows], Dtype::Bool, super::scalar_buffer( &(0..batch * rows) .map(|n| u8::from(n % 3 != 0)) .collect::>(), ) .unwrap(), ) .unwrap() }); let hidden = forward( &weights, &bf16(&[batch, rows, 2560], 30 + step as u32), mask.as_ref(), (mode != "none").then_some(&mut cache), Options { step: true, conv_single: true, conv_verify: true, output: true, capture: step == 0, }, &streams, gpu, ) .unwrap(); if step == 0 { assert_eq!(cache.capture.is_some(), mode != "none"); } else { assert_eq!(row["capture_preserved"].as_bool(), Some(true)); assert_eq!( cache .capture .as_ref() .map(|arrays| arrays.each_ref().map(Array::id)), old_capture ); } let mut roots = vec![hidden]; if mode != "none" { roots.extend([ cache.conv.as_ref().unwrap().clone(), cache.delta.as_ref().unwrap().clone(), ]); } if mode == "metadata" { for value in [&cache.lengths, &cache.left_padding] { roots.push( ops::astype(value.as_ref().unwrap(), Dtype::F32, false, gpu) .unwrap(), ); } } ops::evaluate(&streams, &roots, gpu, false).unwrap(); let buffers = roots.iter().map(Array::buffer).collect::>(); let outputs = roots .iter() .zip(&buffers) .map(|(array, buffer)| { ( &**buffer, array.layout().size() as u32, if array.layout().dtype() == Dtype::BF16 { 2 } else { 4 }, ) }) .collect::>(); verify_typed(row["kernel"].as_str().unwrap(), &outputs); }); } } if layout == "fused" { let fixtures = include_str!("../../../../tests/fixtures/mtplx-gdn-commit.jsonl") .lines() .map(|line| serde_json::from_str::(line).unwrap()) .collect::>(); assert_eq!(fixtures.len(), 26); for row in fixtures { for deferred in [false, true] { super::eval::execution_tests::observed(|| { let rows = row["rows"].as_i64().unwrap() as i32; let keep = row["keep"].as_i64().unwrap() as i32; let snapshot = [ row["conv_present"] .as_bool() .unwrap() .then(|| initial_conv.clone()), Some(initial_delta.clone()), ]; let mut cache = Cache { conv: snapshot[0].clone(), delta: snapshot[1].clone(), ..Cache::default() }; let options = Options { step: true, conv_single: true, conv_verify: true, output: true, capture: true, }; let _ = forward( &weights, &bf16(&[1, rows, 2560], 30), None, Some(&mut cache), options, &streams, gpu, ) .unwrap(); let mut verify_roots = vec![ cache.conv.as_ref().unwrap().clone(), cache.delta.as_ref().unwrap().clone(), ]; verify_roots.extend(cache.capture.as_ref().unwrap().iter().cloned()); ops::evaluate(&streams, &verify_roots, gpu, false).unwrap(); drop(verify_roots); let metadata = |value: i32| { Array::new( &[1], Dtype::I32, super::scalar_buffer(&value.to_le_bytes()).unwrap(), ) .unwrap() }; cache.lengths = Some(metadata(17)); cache.left_padding = Some(metadata(-3)); let identities = || { [ cache.conv.as_ref().unwrap().id(), cache.delta.as_ref().unwrap().id(), cache.capture.as_ref().unwrap()[0].id(), ] }; let before = identities(); assert_eq!( cache.validate_verified_window(None, rows).unwrap_err(), "snapshot_missing" ); assert_eq!( cache .validate_verified_window(Some(&snapshot[..1]), rows) .unwrap_err(), "gdn_snapshot_short" ); assert_eq!( cache .validate_verified_window(Some(&[None, None]), rows) .unwrap_err(), "gdn_snapshot_short" ); assert_eq!( cache .validate_verified_window(Some(&snapshot), rows + 1) .unwrap_err(), format!("gdn_rows_width_{rows}_vs_{}", rows + 1) ); assert_eq!(before, identities()); for invalid_keep in [0, rows + 1] { assert!( cache .replay_verified_prefix( &snapshot, invalid_keep, &a_log, &dt_bias, gpu ) .is_err() ); assert_eq!( before, [ cache.conv.as_ref().unwrap().id(), cache.delta.as_ref().unwrap().id(), cache.capture.as_ref().unwrap()[0].id() ] ); } cache .validate_verified_window(Some(&snapshot), rows) .unwrap(); let metadata_ids = [ cache.lengths.as_ref().unwrap().id(), cache.left_padding.as_ref().unwrap().id(), ]; let (_, submissions) = capture_dispatches(|| { cache .replay_verified_prefix(&snapshot, keep, &a_log, &dt_bias, gpu) .unwrap() }); assert!( submissions.is_empty(), "commit must only build a lazy graph" ); assert!(cache.capture.is_none()); assert_eq!( metadata_ids, [ cache.lengths.as_ref().unwrap().id(), cache.left_padding.as_ref().unwrap().id() ], "commit does not advance or rewind metadata" ); let mut roots = vec![ cache.conv.as_ref().unwrap().clone(), cache.delta.as_ref().unwrap().clone(), ]; if !deferred { let (_, dispatches) = capture_dispatches(|| { ops::evaluate(&streams, &roots, gpu, false).unwrap() }); assert_eq!( dispatches .iter() .filter(|r| r.0.contains("gated_delta_")) .count(), 1 ); assert!(!dispatches.iter().any(|r| { ["qmv", "qmm", "gemv", "gemm", "conv", "gdn_step"] .iter() .any(|name| r.0.contains(name)) })); } let hidden = forward( &weights, &bf16(&[1, 1, 2560], 31), None, Some(&mut cache), Options { capture: false, ..options }, &streams, gpu, ) .unwrap(); // Python's saved metadata objects observe __isub__'s // descriptor overwrite in the subsequent forward. // Read the current cache, not a pre-forward Rust clone. for value in [&cache.lengths, &cache.left_padding] { roots.push( ops::astype(value.as_ref().unwrap(), Dtype::F32, false, gpu) .unwrap(), ); } roots.extend([ hidden, cache.conv.as_ref().unwrap().clone(), cache.delta.as_ref().unwrap().clone(), ]); ops::evaluate(&streams, &roots, gpu, false).unwrap(); let buffers = roots.iter().map(Array::buffer).collect::>(); let outputs = roots .iter() .zip(&buffers) .map(|(array, buffer)| { ( &**buffer, array.layout().size() as u32, if array.layout().dtype() == Dtype::BF16 { 2 } else { 4 }, ) }) .collect::>(); verify_typed(row["kernel"].as_str().unwrap(), &outputs); let before = [ cache.conv.as_ref().unwrap().id(), cache.delta.as_ref().unwrap().id(), ]; assert_eq!( cache .validate_verified_window(Some(&snapshot), 1) .unwrap_err(), "gdn_rows_missing" ); assert_eq!( before, [ cache.conv.as_ref().unwrap().id(), cache.delta.as_ref().unwrap().id() ] ); }); } } } assert_eq!(cases.len(), 44); for (rows, initial, mode) in cases { for deferred in [false, true] { super::eval::execution_tests::observed(|| { let mut states: Vec = Vec::new(); let mut dispatches = Vec::new(); let mut cache = Cache { conv: matches!(initial, 1 | 3).then(|| initial_conv.clone()), delta: matches!(initial, 1 | 2).then(|| initial_delta.clone()), lengths: (mode == 4).then(|| { Array::new( &[1], Dtype::I32, super::scalar_buffer(&1_i32.to_le_bytes()).unwrap(), ) .unwrap() }), ..Cache::default() }; let enabled = mode != 0; let options = Options { step: enabled, conv_single: enabled && mode != 5, conv_verify: enabled && mode != 5, output: enabled && mode != 5, capture: mode == 2, }; let mask = (mode == 3).then(|| { Array::new( &[1, rows], Dtype::Bool, super::scalar_buffer( &(0..rows).map(|n| u8::from(n % 3 != 0)).collect::>(), ) .unwrap(), ) .unwrap() }); for step in 0..2 { let input = bf16(&[1, rows, 2560], 30 + step); let hidden = forward( &weights, &input, mask.as_ref(), Some(&mut cache), options, &streams, gpu, ) .unwrap(); let next = VerifyOutput { hidden, conv: cache.conv.as_ref().unwrap().clone(), delta: cache.delta.as_ref().unwrap().clone(), capture: cache.capture.clone(), }; drop(input); if !deferred { let (_, calls) = capture_dispatches(|| { ops::evaluate( &streams, &[next.hidden.clone(), next.conv.clone(), next.delta.clone()], gpu, false, ) .unwrap() }); dispatches.extend(calls); } states.push(next); } if deferred { let last = states.last().unwrap(); let (_, calls) = capture_dispatches(|| { ops::evaluate( &streams, &[ states[0].hidden.clone(), last.hidden.clone(), last.conv.clone(), last.delta.clone(), ], gpu, false, ) .unwrap() }); dispatches.extend(calls); } let mut expected = [0; 2]; for (step, state) in states.iter().enumerate() { let name = format!( "gdn_forward_{layout}_r{rows}_initial{initial}_mode{mode}_step{step}" ); let receipt = &receipts[&name]; for (count, key) in expected.iter_mut().zip(["used_step", "used_conv"]) { *count += usize::from(receipt[key].as_bool().unwrap()); } assert_eq!(state.conv.layout().shape(), [1, 3, 10240]); assert_eq!(state.delta.layout().shape(), [1, 48, 128, 128]); verify_typed( &name, &[ (&state.hidden.buffer(), rows as u32 * 2560, 2), (&state.conv.buffer(), 3 * 10240, 2), (&state.delta.buffer(), 48 * 128 * 128, 4), ], ); } for (expected, kernel) in expected .into_iter() .zip(["gdn_step_fused", "gdn_conv_norm"]) { assert_eq!( dispatches .iter() .filter(|call| call.0.contains(kernel)) .count(), expected, "{layout} rows={rows} initial={initial} mode={mode} deferred={deferred} {kernel}" ); } }); } } } streams.clear_streams().unwrap(); } #[test] #[ignore = "requires Apple Metal; 80 actual delta-update receipts and six conv-normalization kernels"] fn mtplx_graph_gdn_staged_primitives_match_reference() { use super::super::{ gpu::Context, qwen_mtplx_tests::{ capture_dispatches, pattern, pattern_f32, pattern_scaled, verify, verify_typed, }, }; use super::{allocator, configure_sources, stream::Streams}; configure_sources().unwrap(); let _context = Context::open_qwen(0).unwrap(); let streams = Streams::new(Some(allocator::Allocator::new().unwrap())); let gpu = streams.default_stream(Device::Gpu).unwrap(); let bf16 = |shape: &[i32], salt, scale| { Array::new( shape, Dtype::BF16, pattern_scaled(shape.iter().map(|&d| d as u32).product(), salt, scale), ) .unwrap() }; for rows in 1..=6 { super::eval::execution_tests::observed(|| { let shape = if rows == 1 { vec![10240] } else { vec![rows, 10240] }; let x = bf16(&shape, 6, 1.); let conv = bf16(&[3, 10240], 7, 1.); let weight = bf16(&[10240, 4, 1], 8, 1.); let output = fused_conv_norm(&x, &conv, &weight, gpu).unwrap(); let (_, receipts) = capture_dispatches(|| { ops::evaluate(&streams, std::slice::from_ref(&output[3]), gpu, false).unwrap(); }); let dispatches = receipts .iter() .filter(|r| r.0.contains("gdn_conv_norm")) .collect::>(); assert_eq!(dispatches.len(), 1); assert_eq!( (dispatches[0].1, dispatches[0].2, dispatches[0].3), ([10240, 1, 1], [1024, 1, 1], true) ); verify( &format!("gdn_s{rows}"), &[ (&output[0].buffer(), rows as u32 * 2048), (&output[1].buffer(), rows as u32 * 2048), (&output[2].buffer(), rows as u32 * 6144), (&output[3].buffer(), 30720), ], ); }); } for group in [32, 64] { for dtype in [Dtype::BF16, Dtype::F32] { super::eval::execution_tests::observed(|| { let x = Array::new( &[6144], dtype, if dtype == Dtype::BF16 { pattern(6144, 18) } else { pattern_f32(6144, 18) }, ) .unwrap(); let z = bf16(&[6144], 9, 1.); let norm = bf16(&[128], 14, 1.); let words = (0_u32..2560 * 768) .flat_map(|i| i.wrapping_mul(2654435761).wrapping_add(12345).to_le_bytes()) .collect::>(); let weight = Array::new( &[2560, 768], Dtype::U32, super::scalar_buffer(&words).unwrap(), ) .unwrap(); let scales = bf16(&[2560, 6144 / group as i32], 16, 1.); let biases = bf16(&[2560, 6144 / group as i32], 17, 1.); let output = fused_output([&x, &z, &norm, &weight, &scales, &biases], group, gpu).unwrap(); let (_, dispatches) = capture_dispatches(|| { ops::evaluate(&streams, std::slice::from_ref(&output), gpu, false).unwrap() }); let tag = if dtype == Dtype::BF16 { "bf16" } else { "f32" }; assert_eq!( dispatches, [( format!("kernel_qwen_mtplx_gdn_out_fused_gs{group}_{tag}"), [81920, 1, 1], [1024, 1, 1], true )] ); verify_typed( &format!("gdn_out_gs{group}_{tag}"), &[( &output.buffer(), 2560, if dtype == Dtype::BF16 { 2 } else { 4 }, )], ); }); } } let fixtures = include_str!("../../../../tests/fixtures/mtplx-custom-kernels.jsonl") .lines() .map(|l| serde_json::from_str::(l).unwrap()) .filter(|r| r["kernel"].as_str().unwrap().starts_with("delta_update_")) .collect::>(); assert_eq!(fixtures.len(), 80); let mut state = None; for row in fixtures { super::eval::execution_tests::observed(|| { let [rows, stride, step] = ["rows", "stride", "step"].map(|k| row[k].as_u64().unwrap() as i32); if step == 0 { state = row["initial"].as_bool().unwrap().then(|| { Array::new(&[1, 48, 128, 128], Dtype::F32, pattern_f32(786432, 15)).unwrap() }); } let q = bf16(&[1, rows, 16, 128], 58 + step as u32, 1. / 16.); let k = bf16(&[1, rows, 16, 128], 59 + step as u32, 1. / 16.); let v = bf16(&[1, rows, 48, 128], 60 + step as u32, 1.); let source = bf16(&[1, rows, stride], 60 + step as u32, 1.); let (a, b) = if stride == 48 { (source.clone(), bf16(&[1, rows, 48], 61 + step as u32, 1.)) } else { ( ops::slice(&source, &[0, 0, 16432], &[1, rows, 16480], &[1, 1, 1], gpu) .unwrap(), ops::slice(&source, &[0, 0, 16384], &[1, rows, 16432], &[1, 1, 1], gpu) .unwrap(), ) }; let a_log = Array::new(&[48], Dtype::BF16, pattern(48, 62)).unwrap(); let dt = Array::new(&[48], Dtype::BF16, pattern(48, 63)).unwrap(); let masked = row["masked"].as_bool().unwrap(); let mask = Buffer::mtplx_bytes(rows as u64).unwrap(); mask.write( 0, &(0..rows).map(|i| u8::from(i % 3 != 0)).collect::>(), ) .unwrap(); let mask = Array::new(&[1, rows], Dtype::Bool, mask).unwrap(); let output = delta_update( [&q, &k, &v], [&a, &b, &a_log, &dt], state.as_ref(), masked.then_some(&mask), gpu, ) .unwrap(); let (_, receipts) = capture_dispatches(|| { ops::evaluate( &streams, &[ output.hidden.clone(), output.delta.clone(), output.g.clone(), output.beta.clone(), ], gpu, false, ) .unwrap(); }); let dispatches = receipts .iter() .filter(|r| r.0.contains("gated_delta_")) .collect::>(); assert_eq!(dispatches.len(), 1); assert_eq!( (dispatches[0].1, dispatches[0].2, dispatches[0].3), ([32, 128, 48], [32, 4, 1], true) ); // A_log/dt_bias [48] and a [1,T,48] enter Compiled directly in // the actual reference graph: even T=1 is strided (collapsed rank 1). let gates = receipts .iter() .filter(|r| r.0.contains("compute_g_")) .collect::>(); assert_eq!(gates.len(), 1); assert_eq!( gates[0].0, format!( "kernel_qwen_mtplx_compute_g_bf16_strided_{}", if rows == 1 { 1 } else { 2 } ) ); assert_eq!((gates[0].1, gates[0].3), ([48, rows as u32, 1], true)); verify_typed( row["kernel"].as_str().unwrap(), &[ (&output.hidden.buffer(), rows as u32 * 6144, 2), (&output.delta.buffer(), 786432, 4), (&output.g.buffer(), rows as u32 * 48, 4), (&output.beta.buffer(), rows as u32 * 48, 2), ], ); state = Some(output.delta); }); } streams.clear_streams().unwrap(); } #[test] #[ignore = "requires Apple Metal; GDN custom primitive strided copies and sibling roots"] fn mtplx_graph_gdn_step_preserves_views_and_siblings() { super::eval::execution_tests::observed(|| { use super::super::{ gpu::Context, qwen_mtplx_tests::{pattern, pattern_f32, verify_typed}, }; use super::{allocator, configure_sources, stream::Streams}; configure_sources().unwrap(); let _context = Context::open_qwen(0).unwrap(); let streams = Streams::new(Some(allocator::Allocator::new().unwrap())); let gpu = streams.default_stream(Device::Gpu).unwrap(); // Every input has a nonzero offset and stride 2. Preserve that view // until the custom primitive's own reference General-copy boundary. let strided = |n: u32, salt, dtype: Dtype| { let storage = if dtype == Dtype::F32 { pattern_f32(n, salt) } else { pattern(n, salt) }; let mut raw = vec![0; n as usize * dtype.itemsize()]; storage.read(0, &mut raw).unwrap(); let mut packed = vec![0; dtype.itemsize()]; for value in raw.chunks_exact(dtype.itemsize()) { packed.extend_from_slice(value); packed.extend(std::iter::repeat_n(0, dtype.itemsize())); } let storage = Buffer::mtplx_bytes(packed.len() as u64).unwrap(); storage.write(0, &packed).unwrap(); let base = Array::new(&[(2 * n + 1) as i32], dtype, storage).unwrap(); ops::slice(&base, &[1], &[(2 * n + 1) as i32], &[2], gpu).unwrap() }; let mut conv = strided(30720, 7, Dtype::BF16); let mut delta = strided(786432, 15, Dtype::F32); let z = strided(6144, 9, Dtype::BF16); let a = strided(48, 10, Dtype::BF16); let b = strided(48, 11, Dtype::BF16); let cw = strided(40960, 8, Dtype::BF16); let a_log = strided(48, 12, Dtype::BF16); let dt_bias = strided(48, 13, Dtype::BF16); let norm = strided(128, 14, Dtype::BF16); for step in 0..2 { let qkv = strided(10240, 6 + step, Dtype::BF16); let [y, ns, nd] = fused_step( [ &qkv, &z, &a, &b, &conv, &cw, &a_log, &dt_bias, &delta, &norm, ], gpu, ) .unwrap(); // Evaluate from the last sibling, not always the hidden output. let (_, dispatches) = super::super::qwen_mtplx_tests::capture_dispatches(|| { ops::evaluate(&streams, std::slice::from_ref(&nd), gpu, false).unwrap(); }); super::super::qwen_mtplx_tests::assert_gdn_step_dispatch(&dispatches); verify_typed( &format!("gdn_step_{step}"), &[ (&y.buffer(), 6144, 2), (&ns.buffer(), 30720, 2), (&nd.buffer(), 786432, 4), ], ); conv = ns; delta = nd; } streams.clear_streams().unwrap(); }); }