use anyhow::{anyhow, Context, Result};
use mlx_native::ops::add_bias_row_2d::{
dispatch_add_bias_row_2d_f32, register as register_add_bias_row_2d,
};
use mlx_native::ops::bilinear_resize_2d::{
dispatch_bilinear_resize_2d_f32, register as register_bilinear_resize_2d,
};
use mlx_native::ops::block_merge_2x2::{
dispatch_block_merge_2x2_f32, register as register_block_merge_2x2,
};
use mlx_native::ops::dense_mm_f32_f32::{dense_matmul_f32_f32_tensor, DenseMmF32F32Params};
use mlx_native::ops::elementwise::{elementwise_add, elementwise_mul};
use mlx_native::ops::feature_concat::{
dispatch_feature_concat_f32, register as register_feature_concat,
};
use mlx_native::ops::gelu::dispatch_gelu;
use mlx_native::ops::im2col_2d_3ch::{
dispatch_im2col_2d_3ch_f32, register as register_im2col_2d_3ch,
};
use mlx_native::ops::rope_multi::{dispatch_rope_multi_cached, RopeMultiMode, RopeMultiParams};
use mlx_native::{CommandEncoder, DType, KernelRegistry, MlxBuffer, MlxDevice};
use crate::inference::models::bert::bert_gpu::{
bert_bias_add_gpu, bert_layer_norm_gpu, register_bert_custom_shaders,
};
use super::mmproj::MmprojConfig;
use super::mmproj_weights::LoadedMmprojWeights;
use super::vit_gpu::{
register_vit_custom_shaders, vit_attention_gpu, vit_linear_gpu, vit_residual_add_gpu,
VisionInput,
};
#[derive(Debug, Clone, PartialEq)]
pub struct Qwen3VlViTConfig {
pub n_layer: u32,
pub n_embd: u32,
pub n_head: u32,
pub intermediate_size: u32,
pub patch_size: u32,
pub spatial_merge_size: u32,
pub out_hidden_size: u32,
pub num_position_embeddings: u32,
pub deepstack_indexes: Vec<u32>,
pub eps: f32,
}
impl Qwen3VlViTConfig {
pub fn from_mmproj(cfg: &MmprojConfig, num_position_embeddings: u32) -> Result<Self> {
let spatial_merge_size = cfg.spatial_merge_size.ok_or_else(|| {
anyhow!(
"Qwen3VlViTConfig: MmprojConfig.spatial_merge_size is None — \
a Qwen3-VL mmproj GGUF MUST write 'clip.vision.spatial_merge_size'. \
Refusing to silently fall back to a default."
)
})?;
let out_hidden_size = cfg.projection_dim.ok_or_else(|| {
anyhow!(
"Qwen3VlViTConfig: MmprojConfig.projection_dim is None — \
a Qwen3-VL mmproj GGUF MUST write 'clip.vision.projection_dim'. \
Refusing to silently fall back to a default."
)
})?;
let deepstack_indexes = cfg.deepstack_indexes.clone().ok_or_else(|| {
anyhow!(
"Qwen3VlViTConfig: MmprojConfig.deepstack_indexes is None — \
a Qwen3-VL mmproj GGUF MUST write 'clip.vision.is_deepstack_layers' \
as Bool[block_count]. Refusing to silently fall back."
)
})?;
if num_position_embeddings == 0 {
return Err(anyhow!(
"Qwen3VlViTConfig: num_position_embeddings = 0 — must be > 0 \
(sourced from v.position_embd.weight tensor outer axis = \
hf2q `shape()[0]` = ggml `ne[1]`)"
));
}
for &idx in &deepstack_indexes {
if idx >= cfg.num_hidden_layers {
return Err(anyhow!(
"Qwen3VlViTConfig: deepstack_indexes contains {} which is \
>= num_hidden_layers {} — this should have been caught at \
mmproj parse time",
idx,
cfg.num_hidden_layers
));
}
}
Ok(Self {
n_layer: cfg.num_hidden_layers,
n_embd: cfg.hidden_size,
n_head: cfg.num_attention_heads,
intermediate_size: cfg.intermediate_size,
patch_size: cfg.patch_size,
spatial_merge_size,
out_hidden_size,
num_position_embeddings,
deepstack_indexes,
eps: cfg.layer_norm_eps,
})
}
pub fn n_image_tokens(&self, image_size: u32) -> Result<u32> {
let stride = self.patch_size * self.spatial_merge_size;
if stride == 0 {
return Err(anyhow!(
"Qwen3VlViTConfig::n_image_tokens: stride (patch_size * spatial_merge_size) = 0"
));
}
if image_size % stride != 0 {
return Err(anyhow!(
"Qwen3VlViTConfig::n_image_tokens: image_size {} not divisible by \
patch_size * spatial_merge_size = {}",
image_size,
stride
));
}
let side = image_size / stride;
Ok(side * side)
}
pub fn augmented_embed_dim(&self) -> u32 {
self.out_hidden_size * (1 + self.deepstack_indexes.len() as u32)
}
}
#[cfg(test)]
pub(crate) fn qwen3vl_dual_conv_patch_embed_cpu(
pixel_values: &[f32],
weight_0: &[f32],
weight_1: &[f32],
bias: Option<&[f32]>,
image_size: u32,
patch_size: u32,
hidden: u32,
) -> Result<Vec<f32>> {
qwen3vl_dual_conv_patch_embed_cpu_hw(
pixel_values,
weight_0,
weight_1,
bias,
image_size,
image_size,
patch_size,
hidden,
)
}
#[cfg(test)]
pub(crate) fn qwen3vl_dual_conv_patch_embed_cpu_hw(
pixel_values: &[f32],
weight_0: &[f32],
weight_1: &[f32],
bias: Option<&[f32]>,
pixel_h: u32,
pixel_w: u32,
patch_size: u32,
hidden: u32,
) -> Result<Vec<f32>> {
use super::vit::patch_embed_forward_hw;
let out_0 = patch_embed_forward_hw(
pixel_values,
weight_0,
bias,
pixel_h,
pixel_w,
patch_size,
hidden,
)
.context("qwen3vl_dual_conv_patch_embed_cpu_hw: stem 0")?;
let out_1 = patch_embed_forward_hw(
pixel_values,
weight_1,
None, pixel_h,
pixel_w,
patch_size,
hidden,
)
.context("qwen3vl_dual_conv_patch_embed_cpu_hw: stem 1")?;
if out_0.len() != out_1.len() {
return Err(anyhow!(
"qwen3vl_dual_conv_patch_embed_cpu_hw: stem outputs length mismatch \
({} vs {}) — patch_embed_forward_hw shape contract violated",
out_0.len(),
out_1.len()
));
}
let mut summed = out_0;
for (a, b) in summed.iter_mut().zip(out_1.iter()) {
*a += *b;
}
Ok(summed)
}
#[cfg(test)]
pub(crate) fn qwen3vl_2x2_block_merge_reshape(
input: &[f32],
nx: usize,
ny: usize,
n_embd: usize,
) -> Result<Vec<f32>> {
if nx == 0 || ny == 0 || n_embd == 0 {
return Err(anyhow!(
"qwen3vl_2x2_block_merge_reshape: nx ({nx}), ny ({ny}), n_embd ({n_embd}) \
must all be > 0"
));
}
if nx % 2 != 0 || ny % 2 != 0 {
return Err(anyhow!(
"qwen3vl_2x2_block_merge_reshape: nx ({nx}) and ny ({ny}) must both be \
even (2×2 block merge); enforced upstream by image_size % \
(patch_size * spatial_merge_size) check"
));
}
let expected = ny * nx * n_embd;
if input.len() != expected {
return Err(anyhow!(
"qwen3vl_2x2_block_merge_reshape: input.len() ({}) != ny*nx*n_embd ({})",
input.len(),
expected
));
}
let mut out = vec![0f32; expected];
let half_x = nx / 2;
for by in 0..(ny / 2) {
for bx in 0..half_x {
let block_id = by * half_x + bx;
for y_in in 0..2 {
for x_in in 0..2 {
let src_y = by * 2 + y_in;
let src_x = bx * 2 + x_in;
let src_off = (src_y * nx + src_x) * n_embd;
let within = y_in * 2 + x_in;
let dst_p = block_id * 4 + within;
let dst_off = dst_p * n_embd;
out[dst_off..dst_off + n_embd]
.copy_from_slice(&input[src_off..src_off + n_embd]);
}
}
}
}
Ok(out)
}
#[allow(clippy::too_many_arguments)]
fn qwen3vl_stage_a_dispatch(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
pixel_values: &[f32],
weight_0: &[f32],
weight_1: &[f32],
bias: Option<&[f32]>,
pos_embd_table: &[f32],
num_position_embeddings: u32,
pixel_h: u32,
pixel_w: u32,
patch_size: u32,
hidden: u32,
) -> Result<MlxBuffer> {
let ctx = "qwen3vl_stage_a_dispatch";
if patch_size == 0 || pixel_h == 0 || pixel_w == 0 {
return Err(anyhow!(
"{ctx}: patch_size ({patch_size}), pixel_h ({pixel_h}), \
pixel_w ({pixel_w}) must all be > 0"
));
}
if pixel_h % patch_size != 0 || pixel_w % patch_size != 0 {
return Err(anyhow!(
"{ctx}: pixel grid ({pixel_h}x{pixel_w}) must be divisible by \
patch_size ({patch_size})"
));
}
if num_position_embeddings == 0 {
return Err(anyhow!("{ctx}: num_position_embeddings must be > 0"));
}
let trained_n = (num_position_embeddings as f64).sqrt() as u32;
if trained_n.saturating_mul(trained_n) != num_position_embeddings {
return Err(anyhow!(
"{ctx}: num_position_embeddings ({num_position_embeddings}) is not \
a perfect square (trained pos-embd table must be square per \
clip.cpp:277)"
));
}
let nps_x = (pixel_w / patch_size) as usize;
let nps_y = (pixel_h / patch_size) as usize;
let p2 = (patch_size as usize) * (patch_size as usize);
let k_total = 3 * p2;
let num_patches = nps_y * nps_x;
let h = hidden as usize;
let expected_pixels = 3 * (pixel_h as usize) * (pixel_w as usize);
let expected_w = h * k_total;
let expected_pos_table = (num_position_embeddings as usize) * h;
if pixel_values.len() != expected_pixels {
return Err(anyhow!(
"{ctx}: pixel_values.len() ({}) != 3*pixel_h*pixel_w ({expected_pixels})",
pixel_values.len()
));
}
if weight_0.len() != expected_w {
return Err(anyhow!(
"{ctx}: weight_0.len() ({}) != hidden*3*p² ({expected_w})",
weight_0.len()
));
}
if weight_1.len() != expected_w {
return Err(anyhow!(
"{ctx}: weight_1.len() ({}) != hidden*3*p² ({expected_w})",
weight_1.len()
));
}
if let Some(b) = bias {
if b.len() != h {
return Err(anyhow!("{ctx}: bias.len() ({}) != hidden ({h})", b.len()));
}
}
if pos_embd_table.len() != expected_pos_table {
return Err(anyhow!(
"{ctx}: pos_embd_table.len() ({}) != num_position_embeddings*hidden \
({expected_pos_table})",
pos_embd_table.len()
));
}
let mut pixels_buf = device
.alloc_buffer(
pixel_values.len() * 4,
DType::F32,
vec![3, pixel_h as usize, pixel_w as usize],
)
.map_err(|e| anyhow!("alloc pixels: {e}"))?;
pixels_buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("pixels mut: {e}"))?
.copy_from_slice(pixel_values);
let mut w0_buf = device
.alloc_buffer(weight_0.len() * 4, DType::F32, vec![h, k_total])
.map_err(|e| anyhow!("alloc weight_0: {e}"))?;
w0_buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("weight_0 mut: {e}"))?
.copy_from_slice(weight_0);
let mut w1_buf = device
.alloc_buffer(weight_1.len() * 4, DType::F32, vec![h, k_total])
.map_err(|e| anyhow!("alloc weight_1: {e}"))?;
w1_buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("weight_1 mut: {e}"))?
.copy_from_slice(weight_1);
let im2col_buf = device
.alloc_buffer(
num_patches * k_total * 4,
DType::F32,
vec![num_patches, k_total],
)
.map_err(|e| anyhow!("alloc im2col: {e}"))?;
let mut dst0_buf = device
.alloc_buffer(num_patches * h * 4, DType::F32, vec![num_patches, h])
.map_err(|e| anyhow!("alloc dst0: {e}"))?;
let mut dst1_buf = device
.alloc_buffer(num_patches * h * 4, DType::F32, vec![num_patches, h])
.map_err(|e| anyhow!("alloc dst1: {e}"))?;
let summed_buf = device
.alloc_buffer(num_patches * h * 4, DType::F32, vec![num_patches, h])
.map_err(|e| anyhow!("alloc summed: {e}"))?;
dispatch_im2col_2d_3ch_f32(
encoder,
registry,
device.metal_device(),
&pixels_buf,
&im2col_buf,
pixel_h,
pixel_w,
patch_size,
)
.map_err(|e| anyhow!("im2col dispatch: {e}"))?;
encoder.memory_barrier();
let mm_params = DenseMmF32F32Params {
m: num_patches as u32,
n: hidden,
k: k_total as u32,
src0_batch: 1,
src1_batch: 1,
};
dense_matmul_f32_f32_tensor(
encoder,
registry,
device,
&w0_buf,
&im2col_buf,
&mut dst0_buf,
&mm_params,
)
.map_err(|e| anyhow!("matmul stem 0: {e}"))?;
encoder.memory_barrier();
dense_matmul_f32_f32_tensor(
encoder,
registry,
device,
&w1_buf,
&im2col_buf,
&mut dst1_buf,
&mm_params,
)
.map_err(|e| anyhow!("matmul stem 1: {e}"))?;
encoder.memory_barrier();
elementwise_add(
encoder,
registry,
device.metal_device(),
&dst0_buf,
&dst1_buf,
&summed_buf,
num_patches * h,
DType::F32,
)
.map_err(|e| anyhow!("elementwise_add: {e}"))?;
encoder.memory_barrier();
let patches_pre_buf = if let Some(bias_slice) = bias {
let mut bias_buf = device
.alloc_buffer(bias_slice.len() * 4, DType::F32, vec![h])
.map_err(|e| anyhow!("alloc bias: {e}"))?;
bias_buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("bias mut: {e}"))?
.copy_from_slice(bias_slice);
let biased_buf = device
.alloc_buffer(num_patches * h * 4, DType::F32, vec![num_patches, h])
.map_err(|e| anyhow!("alloc biased: {e}"))?;
dispatch_add_bias_row_2d_f32(
encoder,
registry,
device.metal_device(),
&summed_buf,
&bias_buf,
&biased_buf,
num_patches as u32,
hidden,
)
.map_err(|e| anyhow!("bias broadcast: {e}"))?;
encoder.memory_barrier();
biased_buf
} else {
summed_buf
};
let mut pos_src_buf = device
.alloc_buffer(
pos_embd_table.len() * 4,
DType::F32,
vec![trained_n as usize, trained_n as usize, h],
)
.map_err(|e| anyhow!("alloc pos_embd src: {e}"))?;
pos_src_buf
.as_mut_slice::<f32>()
.map_err(|e| anyhow!("pos_embd src mut: {e}"))?
.copy_from_slice(pos_embd_table);
let pos_resized_buf = device
.alloc_buffer(num_patches * h * 4, DType::F32, vec![nps_y, nps_x, h])
.map_err(|e| anyhow!("alloc pos_embd resized: {e}"))?;
dispatch_bilinear_resize_2d_f32(
encoder,
registry,
device.metal_device(),
&pos_src_buf,
&pos_resized_buf,
trained_n,
nps_x as u32,
nps_y as u32,
hidden,
)
.map_err(|e| anyhow!("bilinear_resize_2d_f32 dispatch: {e}"))?;
encoder.memory_barrier();
let summed_with_pos_buf = device
.alloc_buffer(num_patches * h * 4, DType::F32, vec![num_patches, h])
.map_err(|e| anyhow!("alloc summed_with_pos: {e}"))?;
elementwise_add(
encoder,
registry,
device.metal_device(),
&patches_pre_buf,
&pos_resized_buf,
&summed_with_pos_buf,
num_patches * h,
DType::F32,
)
.map_err(|e| anyhow!("elementwise_add(patches+pos_embd): {e}"))?;
encoder.memory_barrier();
let merged_buf = device
.alloc_buffer(num_patches * h * 4, DType::F32, vec![num_patches, h])
.map_err(|e| anyhow!("alloc merged: {e}"))?;
dispatch_block_merge_2x2_f32(
encoder,
registry,
device.metal_device(),
&summed_with_pos_buf,
&merged_buf,
nps_x as u32,
nps_y as u32,
hidden,
)
.map_err(|e| anyhow!("block_merge_2x2_f32 dispatch: {e}"))?;
encoder.memory_barrier();
Ok(merged_buf)
}
#[cfg(test)]
pub(crate) fn qwen3vl_resize_position_embeddings_bilinear(
pos_embd_table: &[f32],
num_position_embeddings: u32,
n_embd: u32,
target_n_x: u32,
target_n_y: u32,
) -> Result<Vec<f32>> {
if num_position_embeddings == 0 || n_embd == 0 || target_n_x == 0 || target_n_y == 0 {
return Err(anyhow!(
"qwen3vl_resize_position_embeddings_bilinear: num_position_embeddings \
({num_position_embeddings}), n_embd ({n_embd}), target_n_x \
({target_n_x}), target_n_y ({target_n_y}) must all be > 0"
));
}
let trained_n = (num_position_embeddings as f64).sqrt() as u32;
if trained_n.saturating_mul(trained_n) != num_position_embeddings {
return Err(anyhow!(
"qwen3vl_resize_position_embeddings_bilinear: \
num_position_embeddings ({}) is not a perfect square — \
trained position-embedding table must be a square 2-D grid \
(per clip.cpp:277)",
num_position_embeddings
));
}
let expected_table = (num_position_embeddings as usize) * (n_embd as usize);
if pos_embd_table.len() != expected_table {
return Err(anyhow!(
"qwen3vl_resize_position_embeddings_bilinear: \
pos_embd_table.len() ({}) != num_position_embeddings*n_embd ({})",
pos_embd_table.len(),
expected_table
));
}
if trained_n == target_n_x && trained_n == target_n_y {
return Ok(pos_embd_table.to_vec());
}
let trained = trained_n as i64;
let target_x = target_n_x as i64;
let target_y = target_n_y as i64;
let h = n_embd as usize;
let mut out = vec![0f32; (target_y as usize) * (target_x as usize) * h];
let sf_x = (target_x as f32) / (trained as f32);
let sf_y = (target_y as f32) / (trained as f32);
let pixel_offset: f32 = 0.5;
let triangle_filter = |x: f32| -> f32 { (1.0 - x.abs()).max(0.0) };
let support_x = (1.0 / sf_x).max(1.0);
let invscale_x = 1.0 / support_x;
let support_y = (1.0 / sf_y).max(1.0);
let invscale_y = 1.0 / support_y;
for y_dst in 0..target_y {
let y = ((y_dst as f32) + pixel_offset) / sf_y;
let y_min = ((y - support_y + pixel_offset).max(0.0)) as i64;
let y_max = ((y + support_y + pixel_offset).min(trained as f32)) as i64;
for x_dst in 0..target_x {
let x = ((x_dst as f32) + pixel_offset) / sf_x;
let x_min = ((x - support_x + pixel_offset).max(0.0)) as i64;
let x_max = ((x + support_x + pixel_offset).min(trained as f32)) as i64;
let dst_off = ((y_dst as usize) * (target_x as usize) + (x_dst as usize)) * h;
let mut total_weight = 0.0f32;
for sy in y_min..y_max {
let weight_y = triangle_filter(((sy as f32) - y + pixel_offset) * invscale_y);
for sx in x_min..x_max {
let weight_x = triangle_filter(((sx as f32) - x + pixel_offset) * invscale_x);
let weight = weight_x * weight_y;
if weight <= 0.0 {
continue;
}
let src_off = ((sy as usize) * (trained as usize) + (sx as usize)) * h;
for k in 0..h {
out[dst_off + k] += pos_embd_table[src_off + k] * weight;
}
total_weight += weight;
}
}
if total_weight > 0.0 {
let inv = 1.0 / total_weight;
for k in 0..h {
out[dst_off + k] *= inv;
}
}
}
}
Ok(out)
}
#[cfg(test)]
fn upload_f32_to_gpu(device: &MlxDevice, data: &[f32], shape: Vec<usize>) -> Result<MlxBuffer> {
let n = data.len();
let buf = device
.alloc_buffer(n * 4, DType::F32, shape)
.map_err(|e| anyhow!("upload_f32_to_gpu: alloc: {e}"))?;
let dst: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, n) };
dst.copy_from_slice(data);
Ok(buf)
}
fn upload_i32_to_gpu(device: &MlxDevice, data: &[i32], shape: Vec<usize>) -> Result<MlxBuffer> {
let n = data.len();
let buf = device
.alloc_buffer(n * 4, DType::I32, shape)
.map_err(|e| anyhow!("upload_i32_to_gpu: alloc: {e}"))?;
let dst: &mut [i32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut i32, n) };
dst.copy_from_slice(data);
Ok(buf)
}
fn build_qwen3vl_2d_rope_positions(
device: &MlxDevice,
n_x: u32,
n_y: u32,
block_merged_order: bool,
) -> Result<MlxBuffer> {
if n_x == 0 || n_y == 0 {
return Err(anyhow!(
"build_qwen3vl_2d_rope_positions: n_x ({n_x}) and n_y ({n_y}) must be > 0"
));
}
let n_pos = (n_x as usize) * (n_y as usize);
let mut data = vec![0i32; 4 * n_pos];
if block_merged_order {
if n_x % 2 != 0 || n_y % 2 != 0 {
return Err(anyhow!(
"build_qwen3vl_2d_rope_positions(block_merged_order=true): \
n_x ({n_x}) and n_y ({n_y}) must both be even"
));
}
let half_x = (n_x / 2) as usize;
for by in 0..((n_y / 2) as usize) {
for bx in 0..half_x {
let block_id = by * half_x + bx;
for y_in in 0..2usize {
for x_in in 0..2usize {
let src_y = 2 * by + y_in;
let src_x = 2 * bx + x_in;
let token_index = block_id * 4 + (y_in * 2 + x_in);
data[token_index] = src_y as i32; data[n_pos + token_index] = src_x as i32; }
}
}
}
} else {
for y in 0..(n_y as usize) {
for x in 0..(n_x as usize) {
let token_index = y * (n_x as usize) + x;
data[token_index] = y as i32; data[n_pos + token_index] = x as i32; }
}
}
upload_i32_to_gpu(device, &data, vec![4 * n_pos])
}
fn vit_qwen3vl_2d_rope_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
qkv_buf: &MlxBuffer,
positions: &MlxBuffer,
n_pos: u32,
n_heads: u32,
head_dim: u32,
freq_base: f32,
) -> Result<MlxBuffer> {
if n_pos == 0 || n_heads == 0 || head_dim == 0 {
return Err(anyhow!(
"vit_qwen3vl_2d_rope_gpu: n_pos ({n_pos}), n_heads ({n_heads}), \
head_dim ({head_dim}) must all be > 0"
));
}
if head_dim % 4 != 0 {
return Err(anyhow!(
"vit_qwen3vl_2d_rope_gpu: head_dim ({head_dim}) must be a multiple \
of 4 (per-section counts = head_dim/4 per qwen3vl.cpp:14)"
));
}
let n_dims_quarter = head_dim / 4;
let n_elements = (n_pos as usize) * (n_heads as usize) * (head_dim as usize);
let out = device
.alloc_buffer(
n_elements * 4,
DType::F32,
vec![n_pos as usize, n_heads as usize, head_dim as usize],
)
.map_err(|e| anyhow!("vit_qwen3vl_2d_rope_gpu: alloc: {e}"))?;
let params = RopeMultiParams {
head_dim,
rope_dim: head_dim,
n_heads,
seq_len: n_pos,
freq_base,
mode: RopeMultiMode::Vision,
sections: [
n_dims_quarter,
n_dims_quarter,
n_dims_quarter,
n_dims_quarter,
],
};
dispatch_rope_multi_cached(encoder, registry, device, qkv_buf, &out, positions, params)
.map_err(|e| anyhow!("vit_qwen3vl_2d_rope_gpu: dispatch_rope_multi_cached: {e}"))?;
Ok(out)
}
fn vit_qwen3vl_gelu_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
n_elements: u32,
) -> Result<MlxBuffer> {
if n_elements == 0 {
return Err(anyhow!("vit_qwen3vl_gelu_gpu: n_elements must be > 0"));
}
let out = device
.alloc_buffer(
(n_elements as usize) * 4,
DType::F32,
vec![n_elements as usize],
)
.map_err(|e| anyhow!("vit_qwen3vl_gelu_gpu: alloc: {e}"))?;
dispatch_gelu(encoder, registry, device.metal_device(), input, &out)
.map_err(|e| anyhow!("vit_qwen3vl_gelu_gpu: dispatch_gelu: {e}"))?;
Ok(out)
}
fn vit_qwen3vl_geglu_split_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
gate: &MlxBuffer,
up: &MlxBuffer,
n_elements: u32,
) -> Result<MlxBuffer> {
if n_elements == 0 {
return Err(anyhow!(
"vit_qwen3vl_geglu_split_gpu: n_elements must be > 0"
));
}
let gelu_out = vit_qwen3vl_gelu_gpu(encoder, registry, device, gate, n_elements)?;
encoder.memory_barrier();
let out = device
.alloc_buffer(
(n_elements as usize) * 4,
DType::F32,
vec![n_elements as usize],
)
.map_err(|e| anyhow!("vit_qwen3vl_geglu_split_gpu: alloc out: {e}"))?;
elementwise_mul(
encoder,
registry,
device.metal_device(),
&gelu_out,
up,
&out,
n_elements as usize,
DType::F32,
)
.map_err(|e| anyhow!("vit_qwen3vl_geglu_split_gpu: elementwise_mul: {e}"))?;
Ok(out)
}
#[allow(clippy::too_many_arguments)]
fn apply_qwen3vl_block_forward_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
weights: &LoadedMmprojWeights,
cfg: &Qwen3VlViTConfig,
block_idx: usize,
input: &MlxBuffer,
positions: &MlxBuffer,
n_pos: u32,
scale: f32,
freq_base: f32,
) -> Result<MlxBuffer> {
let hidden = cfg.n_embd;
let n_heads = cfg.n_head;
if n_heads == 0 || hidden % n_heads != 0 {
return Err(anyhow!(
"apply_qwen3vl_block_forward_gpu: hidden ({hidden}) must be divisible \
by n_heads ({n_heads})"
));
}
let head_dim = hidden / n_heads;
let intermediate = cfg.intermediate_size;
let eps = cfg.eps;
let n_hidden_total = (n_pos as usize) * (hidden as usize);
let block_w = |suffix: &str| -> Result<&MlxBuffer> {
weights
.block_tensor(block_idx, suffix)
.map_err(|e| anyhow!("qwen3vl block {} {}: {e}", block_idx, suffix))
};
let ln1 = bert_layer_norm_gpu(
encoder,
registry,
device,
input,
block_w("ln1.weight")?,
block_w("ln1.bias")?,
eps,
n_pos,
hidden,
)
.context("apply_qwen3vl_block_forward_gpu: ln1")?;
encoder.memory_barrier();
let q = vit_linear_gpu(
encoder,
registry,
device,
&ln1,
block_w("attn_q.weight")?,
n_pos,
hidden,
hidden,
)
.context("apply_qwen3vl_block_forward_gpu: q proj")?;
encoder.memory_barrier();
let q = bert_bias_add_gpu(
encoder,
registry,
device,
&q,
block_w("attn_q.bias")?,
n_pos,
hidden,
)
.context("apply_qwen3vl_block_forward_gpu: q bias")?;
encoder.memory_barrier();
let k = vit_linear_gpu(
encoder,
registry,
device,
&ln1,
block_w("attn_k.weight")?,
n_pos,
hidden,
hidden,
)
.context("apply_qwen3vl_block_forward_gpu: k proj")?;
encoder.memory_barrier();
let k = bert_bias_add_gpu(
encoder,
registry,
device,
&k,
block_w("attn_k.bias")?,
n_pos,
hidden,
)
.context("apply_qwen3vl_block_forward_gpu: k bias")?;
encoder.memory_barrier();
let v = vit_linear_gpu(
encoder,
registry,
device,
&ln1,
block_w("attn_v.weight")?,
n_pos,
hidden,
hidden,
)
.context("apply_qwen3vl_block_forward_gpu: v proj")?;
encoder.memory_barrier();
let v = bert_bias_add_gpu(
encoder,
registry,
device,
&v,
block_w("attn_v.bias")?,
n_pos,
hidden,
)
.context("apply_qwen3vl_block_forward_gpu: v bias")?;
encoder.memory_barrier();
let q_rot = vit_qwen3vl_2d_rope_gpu(
encoder, registry, device, &q, positions, n_pos, n_heads, head_dim, freq_base,
)
.context("apply_qwen3vl_block_forward_gpu: rope Q")?;
encoder.memory_barrier();
let k_rot = vit_qwen3vl_2d_rope_gpu(
encoder, registry, device, &k, positions, n_pos, n_heads, head_dim, freq_base,
)
.context("apply_qwen3vl_block_forward_gpu: rope K")?;
encoder.memory_barrier();
let attn = vit_attention_gpu(
encoder, registry, device, &q_rot, &k_rot, &v, n_pos, n_heads, head_dim, scale,
)
.context("apply_qwen3vl_block_forward_gpu: attention")?;
encoder.memory_barrier();
let attn_proj = vit_linear_gpu(
encoder,
registry,
device,
&attn,
block_w("attn_out.weight")?,
n_pos,
hidden,
hidden,
)
.context("apply_qwen3vl_block_forward_gpu: out proj")?;
encoder.memory_barrier();
let attn_proj = bert_bias_add_gpu(
encoder,
registry,
device,
&attn_proj,
block_w("attn_out.bias")?,
n_pos,
hidden,
)
.context("apply_qwen3vl_block_forward_gpu: out bias")?;
encoder.memory_barrier();
let post_attn = vit_residual_add_gpu(
encoder,
registry,
device,
input,
&attn_proj,
n_hidden_total as u32,
)
.context("apply_qwen3vl_block_forward_gpu: residual 1")?;
encoder.memory_barrier();
let ln2 = bert_layer_norm_gpu(
encoder,
registry,
device,
&post_attn,
block_w("ln2.weight")?,
block_w("ln2.bias")?,
eps,
n_pos,
hidden,
)
.context("apply_qwen3vl_block_forward_gpu: ln2")?;
encoder.memory_barrier();
let ffn_gate_w_opt = weights.block_tensor(block_idx, "ffn_gate.weight").ok();
let up = vit_linear_gpu(
encoder,
registry,
device,
&ln2,
block_w("ffn_up.weight")?,
n_pos,
hidden,
intermediate,
)
.context("apply_qwen3vl_block_forward_gpu: ffn_up proj")?;
encoder.memory_barrier();
let up = bert_bias_add_gpu(
encoder,
registry,
device,
&up,
block_w("ffn_up.bias")?,
n_pos,
intermediate,
)
.context("apply_qwen3vl_block_forward_gpu: ffn_up bias")?;
encoder.memory_barrier();
let activated = if let Some(ffn_gate_w) = ffn_gate_w_opt {
let gate = vit_linear_gpu(
encoder,
registry,
device,
&ln2,
ffn_gate_w,
n_pos,
hidden,
intermediate,
)
.context("apply_qwen3vl_block_forward_gpu: ffn_gate proj")?;
encoder.memory_barrier();
let gate = bert_bias_add_gpu(
encoder,
registry,
device,
&gate,
block_w("ffn_gate.bias")?,
n_pos,
intermediate,
)
.context("apply_qwen3vl_block_forward_gpu: ffn_gate bias")?;
encoder.memory_barrier();
let act = vit_qwen3vl_geglu_split_gpu(
encoder,
registry,
device,
&gate,
&up,
((n_pos as usize) * (intermediate as usize)) as u32,
)
.context("apply_qwen3vl_block_forward_gpu: geglu split")?;
encoder.memory_barrier();
act
} else {
let act = vit_qwen3vl_gelu_gpu(
encoder,
registry,
device,
&up,
((n_pos as usize) * (intermediate as usize)) as u32,
)
.context("apply_qwen3vl_block_forward_gpu: gelu")?;
encoder.memory_barrier();
act
};
let down = vit_linear_gpu(
encoder,
registry,
device,
&activated,
block_w("ffn_down.weight")?,
n_pos,
intermediate,
hidden,
)
.context("apply_qwen3vl_block_forward_gpu: ffn_down proj")?;
encoder.memory_barrier();
let down = bert_bias_add_gpu(
encoder,
registry,
device,
&down,
block_w("ffn_down.bias")?,
n_pos,
hidden,
)
.context("apply_qwen3vl_block_forward_gpu: ffn_down bias")?;
encoder.memory_barrier();
let block_out = vit_residual_add_gpu(
encoder,
registry,
device,
&post_attn,
&down,
n_hidden_total as u32,
)
.context("apply_qwen3vl_block_forward_gpu: residual 2")?;
Ok(block_out)
}
#[allow(clippy::too_many_arguments)]
fn apply_qwen3vl_deepstack_head_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
weights: &LoadedMmprojWeights,
cfg: &Qwen3VlViTConfig,
deepstack_layer_idx: u32,
block_out: &MlxBuffer,
n_pos: u32,
) -> Result<MlxBuffer> {
let merge_factor = cfg.spatial_merge_size * cfg.spatial_merge_size;
if merge_factor == 0 {
return Err(anyhow!(
"apply_qwen3vl_deepstack_head_gpu: spatial_merge_size² = 0"
));
}
if n_pos % merge_factor != 0 {
return Err(anyhow!(
"apply_qwen3vl_deepstack_head_gpu: n_pos ({}) must be divisible by \
spatial_merge_size² ({}) — image_size guard at compute entry should \
have caught this",
n_pos,
merge_factor
));
}
let n_image_tokens = n_pos / merge_factor;
let merged_hidden = cfg.n_embd * merge_factor;
let il = deepstack_layer_idx;
let ds_name = |suffix: &str| format!("v.deepstack.{il}.{suffix}");
let norm_w = weights
.get(&ds_name("norm.weight"))
.ok_or_else(|| anyhow!("deepstack head missing '{}'", ds_name("norm.weight")))?;
let norm_b = weights
.get(&ds_name("norm.bias"))
.ok_or_else(|| anyhow!("deepstack head missing '{}'", ds_name("norm.bias")))?;
let fc1_w = weights
.get(&ds_name("fc1.weight"))
.ok_or_else(|| anyhow!("deepstack head missing '{}'", ds_name("fc1.weight")))?;
let fc1_b = weights
.get(&ds_name("fc1.bias"))
.ok_or_else(|| anyhow!("deepstack head missing '{}'", ds_name("fc1.bias")))?;
let fc2_w = weights
.get(&ds_name("fc2.weight"))
.ok_or_else(|| anyhow!("deepstack head missing '{}'", ds_name("fc2.weight")))?;
let fc2_b = weights
.get(&ds_name("fc2.bias"))
.ok_or_else(|| anyhow!("deepstack head missing '{}'", ds_name("fc2.bias")))?;
let fc1_shape = fc1_w.shape();
if fc1_shape.len() != 2 || fc1_shape[1] as u32 != merged_hidden {
return Err(anyhow!(
"apply_qwen3vl_deepstack_head_gpu: deepstack.{il}.fc1.weight shape {:?} \
must be [fc1_out, n_embd*spatial_merge²={merged_hidden}]",
fc1_shape
));
}
let fc1_out = fc1_shape[0] as u32;
let fc2_shape = fc2_w.shape();
if fc2_shape.len() != 2 || fc2_shape[1] as u32 != fc1_out {
return Err(anyhow!(
"apply_qwen3vl_deepstack_head_gpu: deepstack.{il}.fc2.weight shape {:?} \
must be [lm_hidden, fc1_out={fc1_out}]",
fc2_shape
));
}
let lm_hidden = fc2_shape[0] as u32;
if lm_hidden != cfg.out_hidden_size {
return Err(anyhow!(
"apply_qwen3vl_deepstack_head_gpu: deepstack.{il}.fc2.weight output dim \
({lm_hidden}) != cfg.out_hidden_size ({}) — projector head and \
deepstack head must agree on lm_hidden so the LM-side split contract \
at qwen3vl.cpp:97 stays consistent",
cfg.out_hidden_size
));
}
let normed = bert_layer_norm_gpu(
encoder,
registry,
device,
block_out,
norm_w,
norm_b,
cfg.eps,
n_image_tokens,
merged_hidden,
)
.context("apply_qwen3vl_deepstack_head_gpu: norm")?;
encoder.memory_barrier();
let fc1_proj = vit_linear_gpu(
encoder,
registry,
device,
&normed,
fc1_w,
n_image_tokens,
merged_hidden,
fc1_out,
)
.context("apply_qwen3vl_deepstack_head_gpu: fc1 proj")?;
encoder.memory_barrier();
let fc1_out_buf = bert_bias_add_gpu(
encoder,
registry,
device,
&fc1_proj,
fc1_b,
n_image_tokens,
fc1_out,
)
.context("apply_qwen3vl_deepstack_head_gpu: fc1 bias")?;
encoder.memory_barrier();
let activated = vit_qwen3vl_gelu_gpu(
encoder,
registry,
device,
&fc1_out_buf,
n_image_tokens * fc1_out,
)
.context("apply_qwen3vl_deepstack_head_gpu: gelu")?;
encoder.memory_barrier();
let fc2_proj = vit_linear_gpu(
encoder,
registry,
device,
&activated,
fc2_w,
n_image_tokens,
fc1_out,
lm_hidden,
)
.context("apply_qwen3vl_deepstack_head_gpu: fc2 proj")?;
encoder.memory_barrier();
let head_out = bert_bias_add_gpu(
encoder,
registry,
device,
&fc2_proj,
fc2_b,
n_image_tokens,
lm_hidden,
)
.context("apply_qwen3vl_deepstack_head_gpu: fc2 bias")?;
Ok(head_out)
}
fn apply_qwen3vl_main_projector_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
weights: &LoadedMmprojWeights,
cfg: &Qwen3VlViTConfig,
post_ln_out: &MlxBuffer,
n_pos: u32,
) -> Result<MlxBuffer> {
let merge_factor = cfg.spatial_merge_size * cfg.spatial_merge_size;
if merge_factor == 0 {
return Err(anyhow!(
"apply_qwen3vl_main_projector_gpu: spatial_merge_size² = 0"
));
}
if n_pos % merge_factor != 0 {
return Err(anyhow!(
"apply_qwen3vl_main_projector_gpu: n_pos ({}) must be divisible by \
spatial_merge_size² ({})",
n_pos,
merge_factor
));
}
let n_image_tokens = n_pos / merge_factor;
let merged_hidden = cfg.n_embd * merge_factor;
let mm0_w = weights
.mm_0_weight()
.context("apply_qwen3vl_main_projector_gpu: mm.0.weight")?;
let mm0_b = weights
.get(super::mmproj::TENSOR_MM_0_BIAS)
.ok_or_else(|| {
anyhow!(
"apply_qwen3vl_main_projector_gpu: missing '{}' (Qwen3-VL projector \
requires mm.0.bias per clip.cpp:1844-1850 / qwen3vl.cpp:179-180)",
super::mmproj::TENSOR_MM_0_BIAS
)
})?;
let mm2_w = weights
.mm_2_weight()
.context("apply_qwen3vl_main_projector_gpu: mm.2.weight")?;
let mm2_b = weights
.get(super::mmproj::TENSOR_MM_2_BIAS)
.ok_or_else(|| {
anyhow!(
"apply_qwen3vl_main_projector_gpu: missing '{}' (Qwen3-VL projector \
requires mm.2.bias per clip.cpp:1844-1850 / qwen3vl.cpp:182-183)",
super::mmproj::TENSOR_MM_2_BIAS
)
})?;
let mm0_shape = mm0_w.shape();
if mm0_shape.len() != 2 || mm0_shape[1] as u32 != merged_hidden {
return Err(anyhow!(
"apply_qwen3vl_main_projector_gpu: mm.0.weight shape {:?} must be \
[mm0_out, n_embd*spatial_merge²={merged_hidden}]",
mm0_shape
));
}
let mm0_out = mm0_shape[0] as u32;
let mm2_shape = mm2_w.shape();
if mm2_shape.len() != 2 || mm2_shape[1] as u32 != mm0_out {
return Err(anyhow!(
"apply_qwen3vl_main_projector_gpu: mm.2.weight shape {:?} must be \
[lm_hidden, mm0_out={mm0_out}]",
mm2_shape
));
}
let lm_hidden = mm2_shape[0] as u32;
if lm_hidden != cfg.out_hidden_size {
return Err(anyhow!(
"apply_qwen3vl_main_projector_gpu: mm.2.weight output dim ({lm_hidden}) \
!= cfg.out_hidden_size ({}) — projector head and deepstack heads must \
agree on lm_hidden so the LM-side split contract at qwen3vl.cpp:97 \
stays consistent",
cfg.out_hidden_size
));
}
let mm0_proj = vit_linear_gpu(
encoder,
registry,
device,
post_ln_out,
mm0_w,
n_image_tokens,
merged_hidden,
mm0_out,
)
.context("apply_qwen3vl_main_projector_gpu: mm.0 proj")?;
encoder.memory_barrier();
let mm0_out_buf = bert_bias_add_gpu(
encoder,
registry,
device,
&mm0_proj,
mm0_b,
n_image_tokens,
mm0_out,
)
.context("apply_qwen3vl_main_projector_gpu: mm.0 bias")?;
encoder.memory_barrier();
let activated = vit_qwen3vl_gelu_gpu(
encoder,
registry,
device,
&mm0_out_buf,
n_image_tokens * mm0_out,
)
.context("apply_qwen3vl_main_projector_gpu: gelu")?;
encoder.memory_barrier();
let mm2_proj = vit_linear_gpu(
encoder,
registry,
device,
&activated,
mm2_w,
n_image_tokens,
mm0_out,
lm_hidden,
)
.context("apply_qwen3vl_main_projector_gpu: mm.2 proj")?;
encoder.memory_barrier();
let proj_out = bert_bias_add_gpu(
encoder,
registry,
device,
&mm2_proj,
mm2_b,
n_image_tokens,
lm_hidden,
)
.context("apply_qwen3vl_main_projector_gpu: mm.2 bias")?;
Ok(proj_out)
}
#[cfg(test)]
fn qwen3vl_concat_augmented_embed_cpu(
chunks: &[Vec<f32>],
n_image_tokens: usize,
lm_hidden: usize,
) -> Result<Vec<f32>> {
if chunks.is_empty() {
return Err(anyhow!(
"qwen3vl_concat_augmented_embed_cpu: chunks must contain at least the \
base projector output (slot 0); got empty slice"
));
}
let expected_chunk_len = n_image_tokens * lm_hidden;
for (i, chunk) in chunks.iter().enumerate() {
if chunk.len() != expected_chunk_len {
return Err(anyhow!(
"qwen3vl_concat_augmented_embed_cpu: chunk[{i}] length {} != \
n_image_tokens*lm_hidden = {n_image_tokens}*{lm_hidden} = {expected_chunk_len}",
chunk.len()
));
}
}
let total_chunks = chunks.len();
let row_stride = total_chunks * lm_hidden;
let mut out = vec![0f32; n_image_tokens * row_stride];
for t in 0..n_image_tokens {
for (c, chunk) in chunks.iter().enumerate() {
let src_off = t * lm_hidden;
let dst_off = t * row_stride + c * lm_hidden;
out[dst_off..dst_off + lm_hidden].copy_from_slice(&chunk[src_off..src_off + lm_hidden]);
}
}
Ok(out)
}
pub fn compute_vision_embeddings_gpu_qwen3vl(
inputs: &[VisionInput],
mmproj_weights: &LoadedMmprojWeights,
cfg: &Qwen3VlViTConfig,
mmproj_cfg: &MmprojConfig,
) -> Result<Vec<Vec<f32>>> {
if inputs.len() != 1 {
return Err(anyhow!(
"compute_vision_embeddings_gpu_qwen3vl: per-image entry point \
expects exactly 1 input image; got {} (multi-image batching \
is the engine-seam's responsibility — call once per image)",
inputs.len()
));
}
let input = &inputs[0];
let (pixel_values, pixel_w, pixel_h): (&[f32], u32, u32) = match input {
VisionInput::Siglip49(p) => {
let (w, h) = p.pixel_grid();
(&p.pixel_values, w, h)
}
VisionInput::Gemma4v(_) => {
return Err(anyhow!(
"compute_vision_embeddings_gpu_qwen3vl: Qwen3-VL preprocessing \
must produce a Siglip49 payload, got Gemma4v — caller-side \
family-branch routing is broken (dispatch should have rejected \
this earlier)"
));
}
};
let stride = cfg.patch_size * cfg.spatial_merge_size;
if stride == 0 {
return Err(anyhow!(
"compute_vision_embeddings_gpu_qwen3vl: patch_size ({}) * \
spatial_merge_size ({}) = 0",
cfg.patch_size,
cfg.spatial_merge_size
));
}
if pixel_w == 0 || pixel_h == 0 {
return Err(anyhow!(
"compute_vision_embeddings_gpu_qwen3vl: pixel grid ({}x{}) has \
zero dimension",
pixel_w,
pixel_h
));
}
if pixel_w % stride != 0 || pixel_h % stride != 0 {
return Err(anyhow!(
"compute_vision_embeddings_gpu_qwen3vl: per-image pixel grid \
({}x{}) — both dimensions must be a multiple of patch_size ({}) * \
spatial_merge_size ({}) = {} (per qwen3vl.cpp:19-20 GGML_ASSERT). \
Preprocessing should have stride-aligned the smart-resize output \
before passing to the ViT.",
pixel_w,
pixel_h,
cfg.patch_size,
cfg.spatial_merge_size,
stride
));
}
let expected_pixels = 3 * (pixel_h as usize) * (pixel_w as usize);
if pixel_values.len() != expected_pixels {
return Err(anyhow!(
"compute_vision_embeddings_gpu_qwen3vl: pixel_values.len() ({}) != \
3 * pixel_h * pixel_w = 3 * {} * {} = {} — preprocessing contract \
violated (mmproj_cfg.image_size={} for reference)",
pixel_values.len(),
pixel_h,
pixel_w,
expected_pixels,
mmproj_cfg.image_size
));
}
let n_x_pre = pixel_w / cfg.patch_size;
let n_y_pre = pixel_h / cfg.patch_size;
if n_x_pre == 0 || n_y_pre == 0 {
return Err(anyhow!(
"compute_vision_embeddings_gpu_qwen3vl: post-conv grid \
({}x{}) has zero dimension (pixel grid {}x{}, patch_size {})",
n_x_pre,
n_y_pre,
pixel_w,
pixel_h,
cfg.patch_size,
));
}
let n_pos_merged = (n_x_pre as usize) * (n_y_pre as usize);
let n_x_merged = n_x_pre; let n_y_merged = n_y_pre;
let patch_embd_buf = mmproj_weights
.patch_embd_weight()
.map_err(|e| anyhow!("compute_vision_embeddings_gpu_qwen3vl: {e}"))?;
let patch_embd_f32 = mmproj_weights
.tensor_as_f32_owned(patch_embd_buf)
.context("compute_vision_embeddings_gpu_qwen3vl: patch_embd → f32 widen")?;
let patch_embd_1_buf = mmproj_weights.get("v.patch_embd.weight.1").ok_or_else(|| {
anyhow!(
"compute_vision_embeddings_gpu_qwen3vl: missing '{}' (Qwen3-VL ViT \
dual-stem patch embedding requires both `v.patch_embd.weight` and \
`v.patch_embd.weight.1`; see qwen3vl.cpp:17, 23-25)",
"v.patch_embd.weight.1",
)
})?;
let patch_embd_1_f32 = mmproj_weights
.tensor_as_f32_owned(patch_embd_1_buf)
.context("compute_vision_embeddings_gpu_qwen3vl: patch_embd.1 → f32 widen")?;
let patch_bias_f32: Option<Vec<f32>> = mmproj_weights
.get(super::mmproj::TENSOR_PATCH_EMBD_BIAS)
.and_then(|b| mmproj_weights.tensor_as_f32_owned(b).ok());
let pos_embd_buf = mmproj_weights
.position_embd_weight()
.map_err(|e| anyhow!("compute_vision_embeddings_gpu_qwen3vl: {e}"))?;
let pos_embd_f32 = mmproj_weights
.tensor_as_f32_owned(pos_embd_buf)
.context("compute_vision_embeddings_gpu_qwen3vl: pos_embd → f32 widen")?;
use mlx_native::GraphExecutor;
let executor = GraphExecutor::new(
MlxDevice::new()
.map_err(|e| anyhow!("compute_vision_embeddings_gpu_qwen3vl: device: {e}"))?,
);
let mut session = executor
.begin()
.map_err(|e| anyhow!("compute_vision_embeddings_gpu_qwen3vl: begin: {e}"))?;
let mut registry = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
mlx_native::ops::sigmoid_mul::register(&mut registry);
mlx_native::ops::rope_multi::register(&mut registry);
mlx_native::ops::gelu::register(&mut registry);
register_vit_custom_shaders(&mut registry);
register_bert_custom_shaders(&mut registry);
register_feature_concat(&mut registry);
register_im2col_2d_3ch(&mut registry);
register_add_bias_row_2d(&mut registry);
register_bilinear_resize_2d(&mut registry);
register_block_merge_2x2(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let input_gpu = qwen3vl_stage_a_dispatch(
session.encoder_mut(),
&mut registry,
device,
pixel_values,
&patch_embd_f32,
&patch_embd_1_f32,
patch_bias_f32.as_deref(),
&pos_embd_f32,
cfg.num_position_embeddings,
pixel_h,
pixel_w,
cfg.patch_size,
cfg.n_embd,
)
.context("compute_vision_embeddings_gpu_qwen3vl: Stage A (GPU, single session)")?;
let positions = build_qwen3vl_2d_rope_positions(
device, n_x_merged, n_y_merged,
true, )
.context("compute_vision_embeddings_gpu_qwen3vl: build rope positions")?;
let head_dim = (cfg.n_embd / cfg.n_head) as f32;
let attn_scale = 1.0_f32 / head_dim.sqrt();
let freq_base = 10000.0_f32;
let mut hidden_states = input_gpu;
if let (Some(pre_ln_w), Some(pre_ln_b)) = (
mmproj_weights.get("v.pre_ln.weight"),
mmproj_weights.get("v.pre_ln.bias"),
) {
hidden_states = bert_layer_norm_gpu(
session.encoder_mut(),
&mut registry,
device,
&hidden_states,
pre_ln_w,
pre_ln_b,
cfg.eps,
n_pos_merged as u32,
cfg.n_embd,
)
.context("compute_vision_embeddings_gpu_qwen3vl: pre-LN")?;
session.encoder_mut().memory_barrier();
}
let mut deepstack_outputs: Vec<MlxBuffer> = Vec::with_capacity(cfg.deepstack_indexes.len());
let deepstack_set: std::collections::HashSet<u32> =
cfg.deepstack_indexes.iter().copied().collect();
for block_idx in 0..(cfg.n_layer as usize) {
let block_out = apply_qwen3vl_block_forward_gpu(
session.encoder_mut(),
&mut registry,
device,
mmproj_weights,
cfg,
block_idx,
&hidden_states,
&positions,
n_pos_merged as u32,
attn_scale,
freq_base,
)
.with_context(|| format!("compute_vision_embeddings_gpu_qwen3vl: block {block_idx}"))?;
session.encoder_mut().memory_barrier();
if deepstack_set.contains(&(block_idx as u32)) {
let head_out = apply_qwen3vl_deepstack_head_gpu(
session.encoder_mut(),
&mut registry,
device,
mmproj_weights,
cfg,
block_idx as u32,
&block_out,
n_pos_merged as u32,
)
.with_context(|| {
format!(
"compute_vision_embeddings_gpu_qwen3vl: deepstack head at block {block_idx}"
)
})?;
session.encoder_mut().memory_barrier();
deepstack_outputs.push(head_out);
}
hidden_states = block_out;
}
if deepstack_outputs.len() != cfg.deepstack_indexes.len() {
return Err(anyhow!(
"compute_vision_embeddings_gpu_qwen3vl: deepstack_outputs.len() ({}) != \
cfg.deepstack_indexes.len() ({}) — per-block loop missed a flagged \
layer (should not happen given deepstack_set construction)",
deepstack_outputs.len(),
cfg.deepstack_indexes.len()
));
}
let post_ln_w = mmproj_weights
.post_ln_weight()
.map_err(|e| anyhow!("compute_vision_embeddings_gpu_qwen3vl: post_ln.weight: {e}"))?;
let post_ln_b = mmproj_weights
.get(super::mmproj::TENSOR_POST_LN_BIAS)
.ok_or_else(|| {
anyhow!(
"compute_vision_embeddings_gpu_qwen3vl: missing '{}' (Qwen3-VL ViT \
final LayerNorm requires both `v.post_ln.weight` and \
`v.post_ln.bias`; see qwen3vl.cpp:171-173)",
super::mmproj::TENSOR_POST_LN_BIAS,
)
})?;
let final_out = bert_layer_norm_gpu(
session.encoder_mut(),
&mut registry,
device,
&hidden_states,
post_ln_w,
post_ln_b,
cfg.eps,
n_pos_merged as u32,
cfg.n_embd,
)
.context("compute_vision_embeddings_gpu_qwen3vl: post-LN")?;
session.encoder_mut().memory_barrier();
let main_out = apply_qwen3vl_main_projector_gpu(
session.encoder_mut(),
&mut registry,
device,
mmproj_weights,
cfg,
&final_out,
n_pos_merged as u32,
)
.context("compute_vision_embeddings_gpu_qwen3vl: main projector")?;
session.encoder_mut().memory_barrier();
let merge_factor = (cfg.spatial_merge_size as usize).pow(2);
if merge_factor == 0 || n_pos_merged % merge_factor != 0 {
return Err(anyhow!(
"compute_vision_embeddings_gpu_qwen3vl: n_pos_merged ({}) not divisible \
by spatial_merge² ({}) — image_size guard at function entry should \
have caught this",
n_pos_merged,
merge_factor
));
}
let n_image_tokens = n_pos_merged / merge_factor;
let lm_hidden = cfg.out_hidden_size as usize;
let augmented_dim = cfg.augmented_embed_dim() as usize;
let expected_augmented_len = n_image_tokens * augmented_dim;
let augmented_buf = device
.alloc_buffer(
expected_augmented_len * 4,
DType::F32,
vec![n_image_tokens, augmented_dim],
)
.map_err(|e| anyhow!("compute_vision_embeddings_gpu_qwen3vl: alloc augmented buf: {e}"))?;
dispatch_feature_concat_f32(
session.encoder_mut(),
&mut registry,
device.metal_device(),
&main_out,
&augmented_buf,
n_image_tokens as u32,
lm_hidden as u32,
0,
augmented_dim as u32,
)
.with_context(|| {
format!(
"compute_vision_embeddings_gpu_qwen3vl: K5 feature_concat (main) — \
n_image_tokens={n_image_tokens} lm_hidden={lm_hidden} augmented_dim={augmented_dim}"
)
})?;
session.encoder_mut().memory_barrier();
for (i, head_out) in deepstack_outputs.iter().enumerate() {
let dst_offset = ((i + 1) * lm_hidden) as u32;
dispatch_feature_concat_f32(
session.encoder_mut(),
&mut registry,
device.metal_device(),
head_out,
&augmented_buf,
n_image_tokens as u32,
lm_hidden as u32,
dst_offset,
augmented_dim as u32,
)
.with_context(|| {
format!(
"compute_vision_embeddings_gpu_qwen3vl: K5 feature_concat (deepstack {i}) — \
dst_offset={dst_offset}"
)
})?;
session.encoder_mut().memory_barrier();
}
session
.finish()
.map_err(|e| anyhow!("compute_vision_embeddings_gpu_qwen3vl: finish: {e}"))?;
let augmented_slice: &[f32] = augmented_buf
.as_slice::<f32>()
.map_err(|e| anyhow!("compute_vision_embeddings_gpu_qwen3vl: augmented readback: {e}"))?;
if augmented_slice.len() < expected_augmented_len {
return Err(anyhow!(
"compute_vision_embeddings_gpu_qwen3vl: augmented embed len {} < \
expected {expected_augmented_len} (n_image_tokens={n_image_tokens}, \
augmented_dim={augmented_dim})",
augmented_slice.len()
));
}
Ok(vec![augmented_slice[..expected_augmented_len].to_vec()])
}
#[cfg(test)]
mod tests {
use super::*;
use crate::inference::vision::mmproj::ProjectorType;
fn synth_qwen3vl_mmproj_cfg(
num_layers: u32,
deepstack_indexes: Option<Vec<u32>>,
spatial_merge_size: Option<u32>,
projection_dim: Option<u32>,
) -> MmprojConfig {
MmprojConfig {
image_size: 768,
patch_size: 16,
num_patches_side: 48,
hidden_size: 1024,
intermediate_size: 4304,
num_attention_heads: 16,
num_hidden_layers: num_layers,
layer_norm_eps: 1e-6,
projector: ProjectorType::Qwen3VlMerger,
image_mean: [0.5, 0.5, 0.5],
image_std: [0.5, 0.5, 0.5],
spatial_merge_size,
projection_dim,
deepstack_indexes,
}
}
#[test]
fn qwen3vl_vit_config_from_mmproj_round_trip() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = synth_qwen3vl_mmproj_cfg(24, Some(vec![5, 11, 17]), Some(2), Some(2048));
let vit = Qwen3VlViTConfig::from_mmproj(&cfg, 2304)
.expect("from_mmproj on a fully-populated Qwen3-VL config must succeed");
assert_eq!(vit.n_layer, 24);
assert_eq!(vit.n_embd, 1024);
assert_eq!(vit.n_head, 16);
assert_eq!(vit.intermediate_size, 4304);
assert_eq!(vit.patch_size, 16);
assert_eq!(vit.spatial_merge_size, 2);
assert_eq!(vit.out_hidden_size, 2048);
assert_eq!(vit.num_position_embeddings, 2304);
assert_eq!(vit.deepstack_indexes, vec![5, 11, 17]);
assert!((vit.eps - 1e-6).abs() < 1e-12);
assert_eq!(vit.augmented_embed_dim(), 8192);
assert_eq!(vit.n_image_tokens(768).unwrap(), 576);
}
#[test]
fn qwen3vl_vit_config_fails_on_missing_deepstack_indexes() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = synth_qwen3vl_mmproj_cfg(
24,
None, Some(2),
Some(2048),
);
let err = Qwen3VlViTConfig::from_mmproj(&cfg, 2304)
.expect_err("from_mmproj must reject a config with deepstack_indexes=None");
let msg = format!("{err}");
assert!(
msg.contains("deepstack_indexes"),
"error must name the missing 'deepstack_indexes' field; got: {msg}"
);
assert!(
msg.contains("is_deepstack_layers") || msg.contains("None"),
"error must hint at the missing GGUF metadata key or the None state; got: {msg}"
);
}
#[test]
fn qwen3vl_vit_config_fails_on_missing_spatial_merge_size() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = synth_qwen3vl_mmproj_cfg(
24,
Some(vec![5, 11, 17]),
None, Some(2048),
);
let err = Qwen3VlViTConfig::from_mmproj(&cfg, 2304)
.expect_err("from_mmproj must reject a config with spatial_merge_size=None");
assert!(format!("{err}").contains("spatial_merge_size"));
}
#[test]
fn qwen3vl_vit_config_fails_on_missing_projection_dim() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = synth_qwen3vl_mmproj_cfg(
24,
Some(vec![5, 11, 17]),
Some(2),
None, );
let err = Qwen3VlViTConfig::from_mmproj(&cfg, 2304)
.expect_err("from_mmproj must reject a config with projection_dim=None");
assert!(format!("{err}").contains("projection_dim"));
}
#[test]
fn qwen3vl_vit_config_fails_on_zero_num_position_embeddings() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = synth_qwen3vl_mmproj_cfg(24, Some(vec![5, 11, 17]), Some(2), Some(2048));
let err = Qwen3VlViTConfig::from_mmproj(&cfg, 0)
.expect_err("from_mmproj must reject num_position_embeddings=0");
assert!(format!("{err}").contains("num_position_embeddings"));
}
#[test]
fn qwen3vl_vit_config_fails_on_out_of_range_deepstack_index() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = synth_qwen3vl_mmproj_cfg(24, Some(vec![5, 11, 24]), Some(2), Some(2048));
let err = Qwen3VlViTConfig::from_mmproj(&cfg, 2304)
.expect_err("from_mmproj must reject deepstack_indexes >= num_hidden_layers");
let msg = format!("{err}");
assert!(
msg.contains("deepstack_indexes") && msg.contains("num_hidden_layers"),
"error must name both fields; got: {msg}"
);
}
#[test]
fn qwen3vl_vit_config_accepts_empty_deepstack_indexes() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = synth_qwen3vl_mmproj_cfg(24, Some(vec![]), Some(2), Some(2048));
let vit = Qwen3VlViTConfig::from_mmproj(&cfg, 2304)
.expect("Some(empty Vec) is a valid deepstack_indexes value");
assert_eq!(vit.deepstack_indexes, Vec::<u32>::new());
assert_eq!(vit.augmented_embed_dim(), 2048);
}
#[test]
fn compute_vision_embeddings_gpu_qwen3vl_returns_scaffold_error() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = synth_qwen3vl_mmproj_cfg(24, Some(vec![5, 11, 17]), Some(2), Some(2048));
let vit_cfg = Qwen3VlViTConfig::from_mmproj(&cfg, 2304).unwrap();
let weights = make_empty_loaded_mmproj_weights();
let result = compute_vision_embeddings_gpu_qwen3vl(&[], &weights, &vit_cfg, &cfg);
let err = result
.expect_err("compute_vision_embeddings_gpu_qwen3vl must reject an empty input slice");
let msg = format!("{err}");
assert!(
msg.contains("expects exactly 1 input image"),
"input-validation must self-identify as the per-image entry \
point reject; got: {msg}"
);
assert!(
msg.contains("got 0"),
"error message must name the actual count (0 here); got: {msg}"
);
}
fn make_empty_loaded_mmproj_weights() -> LoadedMmprojWeights {
let device = mlx_native::MlxDevice::new().expect("MlxDevice::new() for unit test");
LoadedMmprojWeights::empty(device)
}
#[test]
fn qwen3vl_dual_conv_patch_embedding_synthetic() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::vit::patch_embed_forward;
let image_size: u32 = 16;
let patch_size: u32 = 16;
let hidden: u32 = 64;
let inner: usize = 3 * (patch_size as usize) * (patch_size as usize);
let n_w: usize = (hidden as usize) * inner;
let pixel_values: Vec<f32> = (0..(3 * (image_size as usize) * (image_size as usize)))
.map(|i| ((i as f32) * 0.0137 + 0.13).fract())
.collect();
let weight_0: Vec<f32> = vec![1.0; n_w];
let weight_1: Vec<f32> = vec![0.0; n_w];
let dual = qwen3vl_dual_conv_patch_embed_cpu(
&pixel_values,
&weight_0,
&weight_1,
None,
image_size,
patch_size,
hidden,
)
.expect("dual-conv prelude must succeed for matching weight shapes");
let solo_stem_0 = patch_embed_forward(
&pixel_values,
&weight_0,
None,
image_size,
patch_size,
hidden,
)
.expect("stem-0 reference patch_embed_forward");
assert_eq!(dual.len(), hidden as usize);
assert_eq!(solo_stem_0.len(), hidden as usize);
for (i, (d, s)) in dual.iter().zip(solo_stem_0.iter()).enumerate() {
assert_eq!(d.to_bits(), s.to_bits(), "dual[{i}]={d} != stem0[{i}]={s}");
}
let pixel_sum: f32 = pixel_values.iter().copied().sum();
for (i, &v) in dual.iter().enumerate() {
assert!(
(v - pixel_sum).abs() < 1e-3,
"dual[{i}]={v} != Σ pixel_values ≈ {pixel_sum}"
);
}
}
#[test]
fn qwen3vl_position_embedding_no_resize_when_trained_size_matches() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let trained_n: u32 = 4; let n_embd: u32 = 8;
let num_pos = trained_n * trained_n; let table: Vec<f32> = (0..(num_pos as usize) * (n_embd as usize))
.map(|i| (i as f32) * 0.0625 - 1.0)
.collect();
let resized = qwen3vl_resize_position_embeddings_bilinear(
&table, num_pos, n_embd, trained_n, trained_n,
)
.expect("equal-size resize must succeed via fast path");
assert_eq!(resized.len(), table.len());
for (i, (a, b)) in resized.iter().zip(table.iter()).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"fast-path must be byte-exact at index {i}: {a} != {b}"
);
}
}
#[test]
fn qwen3vl_position_embedding_bilinear_interpolation_when_resize() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let trained_n: u32 = 2;
let n_embd: u32 = 1;
let num_pos = trained_n * trained_n; let table: Vec<f32> = vec![1.0, 2.0, 3.0, 4.0];
let target_n: u32 = 3;
let resized = qwen3vl_resize_position_embeddings_bilinear(
&table, num_pos, n_embd, target_n, target_n,
)
.expect("2×2 → 3×3 bilinear must succeed");
assert_eq!(resized.len(), 9);
for (i, v) in resized.iter().enumerate() {
assert!(v.is_finite(), "resized[{i}] = {v} is not finite");
}
let idx = |y: usize, x: usize| (y * (target_n as usize) + x) * (n_embd as usize);
let eps = 1e-6_f32;
assert!(
(resized[idx(0, 0)] - 1.0).abs() < eps,
"dst(0,0) expected 1.0, got {}",
resized[idx(0, 0)]
);
assert!(
(resized[idx(1, 1)] - 2.5).abs() < eps,
"dst(1,1) expected 2.5, got {}",
resized[idx(1, 1)]
);
assert!(
(resized[idx(2, 2)] - 4.0).abs() < eps,
"dst(2,2) expected 4.0, got {}",
resized[idx(2, 2)]
);
assert!(
(resized[idx(0, 1)] - 1.5).abs() < eps,
"dst(0,1) expected 1.5, got {}",
resized[idx(0, 1)]
);
assert!(
(resized[idx(1, 2)] - 3.0).abs() < eps,
"dst(1,2) expected 3.0, got {}",
resized[idx(1, 2)]
);
}
#[test]
fn qwen3vl_position_embedding_antialias_downsample_diverges_from_plain_bilinear() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let trained_n: u32 = 4;
let n_embd: u32 = 1;
let num_pos = trained_n * trained_n; let mut table: Vec<f32> = Vec::with_capacity(num_pos as usize);
for y in 0..trained_n {
for x in 0..trained_n {
table.push((y as f32) * 10.0 + (x as f32));
}
}
let target_n: u32 = 2;
let resized = qwen3vl_resize_position_embeddings_bilinear(
&table, num_pos, n_embd, target_n, target_n,
)
.expect("4×4 → 2×2 downsample must succeed");
assert_eq!(resized.len(), 4);
for v in &resized {
assert!(v.is_finite(), "antialias result {v} must be finite");
}
let dst_00 = resized[0];
assert!(
(dst_00 - 5.5).abs() > 0.05,
"Codex Phase-2c sabotage check: plain bilinear dst(0,0) ≈ 5.5; \
antialias bilinear MUST diverge measurably (>= 0.05) on this \
downsample fixture. Got dst(0,0) = {dst_00}; if this is ~5.5, \
the antialias semantic was reverted to plain bilinear. \
Reference: /opt/llama.cpp/ggml/src/ggml-cpu/ops.cpp:7578-7637 \
(BILINEAR | ANTIALIAS triangle-filter accumulation)."
);
let dst_11 = resized[3];
assert!(
dst_11 > dst_00,
"antialias dst(1,1) ({dst_11}) must be > dst(0,0) ({dst_00}) \
on this monotone source pattern (sanity)"
);
}
#[test]
fn qwen3vl_dispatch_real_num_position_embeddings_extracted() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use mlx_native::DType;
use std::collections::HashMap;
let device = mlx_native::MlxDevice::new().expect("MlxDevice::new()");
let num_pos: usize = 2304;
let hidden: usize = 1024;
let buf = device
.alloc_buffer(num_pos * hidden * 4, DType::F32, vec![num_pos, hidden])
.expect("alloc position_embd_weight");
let mut tensors: HashMap<String, mlx_native::MlxBuffer> = HashMap::new();
tensors.insert(super::super::mmproj::TENSOR_POS_EMBD.to_string(), buf);
let weights = LoadedMmprojWeights::from_tensors_for_test(tensors, device);
let pe_buf = weights.position_embd_weight().expect("pos embed present");
let extracted_count = pe_buf.shape()[0] as u32;
assert_eq!(
extracted_count, 2304,
"dispatch site must extract num_position_embeddings = 2304 from the \
outer (count) axis of v.position_embd.weight; got {extracted_count}"
);
let mmproj_cfg = synth_qwen3vl_mmproj_cfg(24, Some(vec![5, 11, 17]), Some(2), Some(2048));
let vit_cfg =
Qwen3VlViTConfig::from_mmproj(&mmproj_cfg, extracted_count).expect("from_mmproj");
assert_eq!(vit_cfg.num_position_embeddings, 2304);
assert_ne!(
vit_cfg.num_position_embeddings, 1,
"regression: 4c.1 sentinel `num_position_embeddings = 1` must NOT \
reach Qwen3VlViTConfig — 4c.2 closes that drift trap"
);
}
#[test]
fn qwen3vl_compute_reaches_patch_embed_after_4c3() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::inference::vision::vit_gpu::VisionInput;
use crate::inference::vision::PreprocessedImage;
let mmproj_cfg = synth_qwen3vl_mmproj_cfg(24, Some(vec![5, 11, 17]), Some(2), Some(2048));
let image_size: u32 = 32;
let mut mmproj_cfg_small = mmproj_cfg.clone();
mmproj_cfg_small.image_size = image_size;
let vit_cfg = Qwen3VlViTConfig::from_mmproj(&mmproj_cfg_small, 4).unwrap();
let pixel_values = vec![0.0f32; 3 * (image_size as usize) * (image_size as usize)];
let img = PreprocessedImage {
pixel_values,
target_size: image_size,
pixel_w: None,
pixel_h: None,
source_label: "synthetic-4c3".to_string(),
};
let inputs = vec![VisionInput::Siglip49(img)];
let weights = make_empty_loaded_mmproj_weights();
let result =
compute_vision_embeddings_gpu_qwen3vl(&inputs, &weights, &vit_cfg, &mmproj_cfg_small);
let err = result.expect_err(
"4c.3: with a valid 1-input batch but empty weights, the function \
must surface a missing-tensor error from the patch-embed step \
rather than the (now-retired) 4c.2 stub message",
);
let msg = format!("{err}");
assert!(
msg.contains("v.patch_embd.weight"),
"error must name the missing patch_embd tensor; got: {msg}"
);
assert!(
!msg.contains("not yet implemented"),
"4c.3 must NOT keep the 4c.2 stub error message; got: {msg}"
);
}
#[test]
fn qwen3vl_2x2_block_merge_reshape_known_pattern() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let nx = 4;
let ny = 4;
let n_embd = 1;
let input: Vec<f32> = (0..(nx * ny * n_embd) as u32).map(|i| i as f32).collect();
let out = qwen3vl_2x2_block_merge_reshape(&input, nx, ny, n_embd)
.expect("4×4 reshape must succeed");
assert_eq!(out[0], 0.0, "block0 patch at within=0 (y_in=0,x_in=0)");
assert_eq!(out[1], 1.0, "block0 patch at within=1 (y_in=0,x_in=1)");
assert_eq!(out[2], 4.0, "block0 patch at within=2 (y_in=1,x_in=0)");
assert_eq!(out[3], 5.0, "block0 patch at within=3 (y_in=1,x_in=1)");
assert_eq!(out[4], 2.0);
assert_eq!(out[5], 3.0);
assert_eq!(out[6], 6.0);
assert_eq!(out[7], 7.0);
assert_eq!(out[12], 10.0);
assert_eq!(out[13], 11.0);
assert_eq!(out[14], 14.0);
assert_eq!(out[15], 15.0);
}
#[allow(clippy::too_many_arguments, dead_code)]
fn build_synth_qwen3vl_weights(
device: mlx_native::MlxDevice,
n_layers: u32,
hidden: u32,
intermediate: u32,
n_x: u32,
ln_gain: f32,
ln_bias: f32,
proj_w: f32,
ffn_w: f32,
patch_size: u32,
) -> LoadedMmprojWeights {
build_synth_qwen3vl_weights_with_deepstack(
device,
n_layers,
hidden,
intermediate,
n_x,
ln_gain,
ln_bias,
proj_w,
ffn_w,
patch_size,
&[],
32, )
}
#[allow(clippy::too_many_arguments)]
fn build_synth_qwen3vl_weights_with_deepstack(
device: mlx_native::MlxDevice,
n_layers: u32,
hidden: u32,
intermediate: u32,
n_x: u32,
ln_gain: f32,
ln_bias: f32,
proj_w: f32,
ffn_w: f32,
patch_size: u32,
deepstack_indexes: &[u32],
out_hidden_size: u32,
) -> LoadedMmprojWeights {
use mlx_native::DType;
use std::collections::HashMap;
let h = hidden as usize;
let inter = intermediate as usize;
let p = patch_size as usize;
let num_pos_emb = (n_x as usize) * (n_x as usize);
let mut tensors: HashMap<String, mlx_native::MlxBuffer> = HashMap::new();
let alloc_f32 = |bytes_count: usize, shape: Vec<usize>| -> mlx_native::MlxBuffer {
device
.alloc_buffer(bytes_count * 4, DType::F32, shape)
.expect("alloc")
};
let fill_f32 = |buf: &mlx_native::MlxBuffer, val: f32, n: usize| {
let s: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, n) };
for v in s.iter_mut() {
*v = val;
}
};
let patch_n = h * 3 * p * p;
let pe0 = alloc_f32(patch_n, vec![h, 3, p, p]);
fill_f32(&pe0, 0.0, patch_n);
tensors.insert(super::super::mmproj::TENSOR_PATCH_EMBD.to_string(), pe0);
let pe1 = alloc_f32(patch_n, vec![h, 3, p, p]);
fill_f32(&pe1, 0.0, patch_n);
tensors.insert("v.patch_embd.weight.1".to_string(), pe1);
let pos_n = num_pos_emb * h;
let pos = alloc_f32(pos_n, vec![num_pos_emb, h]);
fill_f32(&pos, 0.0, pos_n);
tensors.insert(super::super::mmproj::TENSOR_POS_EMBD.to_string(), pos);
let post_w = alloc_f32(h, vec![h]);
fill_f32(&post_w, ln_gain, h);
tensors.insert(
super::super::mmproj::TENSOR_POST_LN_WEIGHT.to_string(),
post_w,
);
let post_b = alloc_f32(h, vec![h]);
fill_f32(&post_b, ln_bias, h);
tensors.insert(
super::super::mmproj::TENSOR_POST_LN_BIAS.to_string(),
post_b,
);
for il in 0..n_layers as usize {
let blk = format!("v.blk.{il}");
for which in ["ln1", "ln2"] {
let w = alloc_f32(h, vec![h]);
fill_f32(&w, ln_gain, h);
tensors.insert(format!("{blk}.{which}.weight"), w);
let b = alloc_f32(h, vec![h]);
fill_f32(&b, ln_bias, h);
tensors.insert(format!("{blk}.{which}.bias"), b);
}
let proj_n = h * h;
for which in ["attn_q", "attn_k", "attn_v", "attn_out"] {
let w = alloc_f32(proj_n, vec![h, h]);
fill_f32(&w, proj_w, proj_n);
tensors.insert(format!("{blk}.{which}.weight"), w);
let b = alloc_f32(h, vec![h]);
fill_f32(&b, 0.0, h);
tensors.insert(format!("{blk}.{which}.bias"), b);
}
let up_n = inter * h;
for which in ["ffn_gate", "ffn_up"] {
let w = alloc_f32(up_n, vec![inter, h]);
fill_f32(&w, ffn_w, up_n);
tensors.insert(format!("{blk}.{which}.weight"), w);
let b = alloc_f32(inter, vec![inter]);
fill_f32(&b, 0.0, inter);
tensors.insert(format!("{blk}.{which}.bias"), b);
}
let down_n = h * inter;
let dw = alloc_f32(down_n, vec![h, inter]);
fill_f32(&dw, ffn_w, down_n);
tensors.insert(format!("{blk}.ffn_down.weight"), dw);
let db = alloc_f32(h, vec![h]);
fill_f32(&db, 0.0, h);
tensors.insert(format!("{blk}.ffn_down.bias"), db);
}
let lm_h = out_hidden_size as usize;
let merge_factor: usize = 4;
let merged = h * merge_factor;
let mm0_n = inter * merged;
let mm0_w = alloc_f32(mm0_n, vec![inter, merged]);
fill_f32(&mm0_w, proj_w, mm0_n);
tensors.insert(super::super::mmproj::TENSOR_MM_0_WEIGHT.to_string(), mm0_w);
let mm0_b = alloc_f32(inter, vec![inter]);
fill_f32(&mm0_b, 0.0, inter);
tensors.insert(super::super::mmproj::TENSOR_MM_0_BIAS.to_string(), mm0_b);
let mm2_n = lm_h * inter;
let mm2_w = alloc_f32(mm2_n, vec![lm_h, inter]);
fill_f32(&mm2_w, proj_w, mm2_n);
tensors.insert(super::super::mmproj::TENSOR_MM_2_WEIGHT.to_string(), mm2_w);
let mm2_b = alloc_f32(lm_h, vec![lm_h]);
fill_f32(&mm2_b, 0.0, lm_h);
tensors.insert(super::super::mmproj::TENSOR_MM_2_BIAS.to_string(), mm2_b);
for &il in deepstack_indexes {
let il_us = il as usize;
let nw = alloc_f32(merged, vec![merged]);
fill_f32(&nw, ln_gain, merged);
tensors.insert(format!("v.deepstack.{il_us}.norm.weight"), nw);
let nb = alloc_f32(merged, vec![merged]);
fill_f32(&nb, ln_bias, merged);
tensors.insert(format!("v.deepstack.{il_us}.norm.bias"), nb);
let fc1_n = inter * merged;
let fc1_w = alloc_f32(fc1_n, vec![inter, merged]);
fill_f32(&fc1_w, proj_w, fc1_n);
tensors.insert(format!("v.deepstack.{il_us}.fc1.weight"), fc1_w);
let fc1_b = alloc_f32(inter, vec![inter]);
fill_f32(&fc1_b, 0.0, inter);
tensors.insert(format!("v.deepstack.{il_us}.fc1.bias"), fc1_b);
let fc2_n = lm_h * inter;
let fc2_w = alloc_f32(fc2_n, vec![lm_h, inter]);
fill_f32(&fc2_w, proj_w, fc2_n);
tensors.insert(format!("v.deepstack.{il_us}.fc2.weight"), fc2_w);
let fc2_b = alloc_f32(lm_h, vec![lm_h]);
fill_f32(&fc2_b, 0.0, lm_h);
tensors.insert(format!("v.deepstack.{il_us}.fc2.bias"), fc2_b);
}
LoadedMmprojWeights::from_tensors_for_test(tensors, device)
}
fn synth_qwen3vl_block_cfg(n_layers: u32) -> (Qwen3VlViTConfig, MmprojConfig) {
let mut cfg = synth_qwen3vl_mmproj_cfg(n_layers, Some(vec![]), Some(2), Some(32));
cfg.image_size = 128;
cfg.patch_size = 16;
cfg.num_patches_side = 8;
cfg.hidden_size = 32;
cfg.intermediate_size = 64;
cfg.num_attention_heads = 1;
let vit_cfg =
Qwen3VlViTConfig::from_mmproj(&cfg, 64).expect("synth_qwen3vl_block_cfg: from_mmproj");
(vit_cfg, cfg)
}
#[allow(clippy::too_many_arguments, dead_code)]
fn build_synth_qwen3vl_weights_split_block_vs_head(
device: mlx_native::MlxDevice,
n_layers: u32,
hidden: u32,
intermediate: u32,
n_x: u32,
ln_gain: f32,
ln_bias: f32,
block_proj_w: f32,
block_ffn_w: f32,
head_proj_w: f32,
patch_size: u32,
deepstack_indexes: &[u32],
out_hidden_size: u32,
pos_embd_pattern: Option<&[f32]>,
) -> LoadedMmprojWeights {
use mlx_native::DType;
use std::collections::HashMap;
let h = hidden as usize;
let inter = intermediate as usize;
let p = patch_size as usize;
let lm_h = out_hidden_size as usize;
let merge_factor: usize = 4;
let merged = h * merge_factor;
let num_pos_emb = (n_x as usize) * (n_x as usize);
let mut tensors: HashMap<String, mlx_native::MlxBuffer> = HashMap::new();
let alloc_f32 = |bytes_count: usize, shape: Vec<usize>| -> mlx_native::MlxBuffer {
device
.alloc_buffer(bytes_count * 4, DType::F32, shape)
.expect("alloc")
};
let fill_f32 = |buf: &mlx_native::MlxBuffer, val: f32, n: usize| {
let s: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, n) };
for v in s.iter_mut() {
*v = val;
}
};
let patch_n = h * 3 * p * p;
let pe0 = alloc_f32(patch_n, vec![h, 3, p, p]);
fill_f32(&pe0, 0.0, patch_n);
tensors.insert(super::super::mmproj::TENSOR_PATCH_EMBD.to_string(), pe0);
let pe1 = alloc_f32(patch_n, vec![h, 3, p, p]);
fill_f32(&pe1, 0.0, patch_n);
tensors.insert("v.patch_embd.weight.1".to_string(), pe1);
let pos_n = num_pos_emb * h;
let pos = alloc_f32(pos_n, vec![num_pos_emb, h]);
if let Some(pattern) = pos_embd_pattern {
assert_eq!(
pattern.len(),
pos_n,
"pos_embd_pattern.len() ({}) must equal num_pos_emb*hidden ({})",
pattern.len(),
pos_n
);
let dst: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(pos.contents_ptr() as *mut f32, pos_n) };
dst.copy_from_slice(pattern);
} else {
fill_f32(&pos, 0.0, pos_n);
}
tensors.insert(super::super::mmproj::TENSOR_POS_EMBD.to_string(), pos);
let post_w = alloc_f32(h, vec![h]);
fill_f32(&post_w, ln_gain, h);
tensors.insert(
super::super::mmproj::TENSOR_POST_LN_WEIGHT.to_string(),
post_w,
);
let post_b = alloc_f32(h, vec![h]);
fill_f32(&post_b, ln_bias, h);
tensors.insert(
super::super::mmproj::TENSOR_POST_LN_BIAS.to_string(),
post_b,
);
for il in 0..n_layers as usize {
let blk = format!("v.blk.{il}");
for which in ["ln1", "ln2"] {
let w = alloc_f32(h, vec![h]);
fill_f32(&w, ln_gain, h);
tensors.insert(format!("{blk}.{which}.weight"), w);
let b = alloc_f32(h, vec![h]);
fill_f32(&b, ln_bias, h);
tensors.insert(format!("{blk}.{which}.bias"), b);
}
let proj_n = h * h;
for which in ["attn_q", "attn_k", "attn_v", "attn_out"] {
let w = alloc_f32(proj_n, vec![h, h]);
fill_f32(&w, block_proj_w, proj_n);
tensors.insert(format!("{blk}.{which}.weight"), w);
let b = alloc_f32(h, vec![h]);
fill_f32(&b, 0.0, h);
tensors.insert(format!("{blk}.{which}.bias"), b);
}
let up_n = inter * h;
for which in ["ffn_gate", "ffn_up"] {
let w = alloc_f32(up_n, vec![inter, h]);
fill_f32(&w, block_ffn_w, up_n);
tensors.insert(format!("{blk}.{which}.weight"), w);
let b = alloc_f32(inter, vec![inter]);
fill_f32(&b, 0.0, inter);
tensors.insert(format!("{blk}.{which}.bias"), b);
}
let down_n = h * inter;
let dw = alloc_f32(down_n, vec![h, inter]);
fill_f32(&dw, block_ffn_w, down_n);
tensors.insert(format!("{blk}.ffn_down.weight"), dw);
let db = alloc_f32(h, vec![h]);
fill_f32(&db, 0.0, h);
tensors.insert(format!("{blk}.ffn_down.bias"), db);
}
let mm0_n = inter * merged;
let mm0_w = alloc_f32(mm0_n, vec![inter, merged]);
fill_f32(&mm0_w, head_proj_w, mm0_n);
tensors.insert(super::super::mmproj::TENSOR_MM_0_WEIGHT.to_string(), mm0_w);
let mm0_b = alloc_f32(inter, vec![inter]);
fill_f32(&mm0_b, 0.0, inter);
tensors.insert(super::super::mmproj::TENSOR_MM_0_BIAS.to_string(), mm0_b);
let mm2_n = lm_h * inter;
let mm2_w = alloc_f32(mm2_n, vec![lm_h, inter]);
fill_f32(&mm2_w, head_proj_w, mm2_n);
tensors.insert(super::super::mmproj::TENSOR_MM_2_WEIGHT.to_string(), mm2_w);
let mm2_b = alloc_f32(lm_h, vec![lm_h]);
fill_f32(&mm2_b, 0.0, lm_h);
tensors.insert(super::super::mmproj::TENSOR_MM_2_BIAS.to_string(), mm2_b);
for &il in deepstack_indexes {
let il_us = il as usize;
let nw = alloc_f32(merged, vec![merged]);
fill_f32(&nw, ln_gain, merged);
tensors.insert(format!("v.deepstack.{il_us}.norm.weight"), nw);
let nb = alloc_f32(merged, vec![merged]);
fill_f32(&nb, ln_bias, merged);
tensors.insert(format!("v.deepstack.{il_us}.norm.bias"), nb);
let fc1_n = inter * merged;
let fc1_w = alloc_f32(fc1_n, vec![inter, merged]);
fill_f32(&fc1_w, head_proj_w, fc1_n);
tensors.insert(format!("v.deepstack.{il_us}.fc1.weight"), fc1_w);
let fc1_b = alloc_f32(inter, vec![inter]);
fill_f32(&fc1_b, 0.0, inter);
tensors.insert(format!("v.deepstack.{il_us}.fc1.bias"), fc1_b);
let fc2_n = lm_h * inter;
let fc2_w = alloc_f32(fc2_n, vec![lm_h, inter]);
fill_f32(&fc2_w, head_proj_w, fc2_n);
tensors.insert(format!("v.deepstack.{il_us}.fc2.weight"), fc2_w);
let fc2_b = alloc_f32(lm_h, vec![lm_h]);
fill_f32(&fc2_b, 0.0, lm_h);
tensors.insert(format!("v.deepstack.{il_us}.fc2.bias"), fc2_b);
}
LoadedMmprojWeights::from_tensors_for_test(tensors, device)
}
fn synth_zero_pixel_inputs(
image_size: u32,
) -> Vec<crate::inference::vision::vit_gpu::VisionInput> {
use crate::inference::vision::vit_gpu::VisionInput;
use crate::inference::vision::PreprocessedImage;
let pixel_values = vec![0.0f32; 3 * (image_size as usize) * (image_size as usize)];
vec![VisionInput::Siglip49(PreprocessedImage {
pixel_values,
target_size: image_size,
pixel_w: None,
pixel_h: None,
source_label: "synthetic-4c3-test".to_string(),
})]
}
#[test]
fn qwen3vl_per_block_forward_synthetic_2_blocks() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use mlx_native::GraphExecutor;
use std::collections::HashMap;
let n_embd: u32 = 32;
let n_head: u32 = 1;
let intermediate: u32 = 64;
let n_pos: u32 = 32;
let eps: f32 = 1e-6;
let device = mlx_native::MlxDevice::new().expect("MlxDevice");
let mmproj_cfg = MmprojConfig {
image_size: 64,
patch_size: 16,
num_patches_side: 4,
hidden_size: n_embd,
intermediate_size: intermediate,
num_attention_heads: n_head,
num_hidden_layers: 1,
layer_norm_eps: eps,
projector: ProjectorType::Qwen3VlMerger,
image_mean: [0.5, 0.5, 0.5],
image_std: [0.5, 0.5, 0.5],
spatial_merge_size: Some(2),
projection_dim: Some(32),
deepstack_indexes: Some(vec![]),
};
let vit_cfg = Qwen3VlViTConfig::from_mmproj(&mmproj_cfg, 16).expect("from_mmproj");
let n_total = (n_pos * n_embd) as usize;
let input_data: Vec<f32> = (0..n_pos)
.flat_map(|r| (0..n_embd).map(move |k| 0.5 + (r as f32) * 0.1 + (k as f32) * 0.01))
.collect();
assert_eq!(input_data.len(), n_total);
let input_buf =
upload_f32_to_gpu(&device, &input_data, vec![n_pos as usize, n_embd as usize])
.expect("upload input");
let alloc_f32 = |bytes_count: usize, shape: Vec<usize>| -> mlx_native::MlxBuffer {
device
.alloc_buffer(bytes_count * 4, mlx_native::DType::F32, shape)
.expect("alloc")
};
let fill_f32 = |buf: &mlx_native::MlxBuffer, val: f32, n: usize| {
let s: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, n) };
for v in s.iter_mut() {
*v = val;
}
};
let h = n_embd as usize;
let inter = intermediate as usize;
let mut tensors: HashMap<String, mlx_native::MlxBuffer> = HashMap::new();
for which in ["ln1", "ln2"] {
let w = alloc_f32(h, vec![h]);
fill_f32(&w, 1.0, h);
tensors.insert(format!("v.blk.0.{which}.weight"), w);
let b = alloc_f32(h, vec![h]);
fill_f32(&b, 0.0, h);
tensors.insert(format!("v.blk.0.{which}.bias"), b);
}
let proj_n = h * h;
for which in ["attn_q", "attn_k", "attn_v", "attn_out"] {
let w = alloc_f32(proj_n, vec![h, h]);
fill_f32(&w, 0.0, proj_n);
tensors.insert(format!("v.blk.0.{which}.weight"), w);
let b = alloc_f32(h, vec![h]);
fill_f32(&b, 0.0, h);
tensors.insert(format!("v.blk.0.{which}.bias"), b);
}
let up_n = inter * h;
for which in ["ffn_gate", "ffn_up"] {
let w = alloc_f32(up_n, vec![inter, h]);
fill_f32(&w, 0.0, up_n);
tensors.insert(format!("v.blk.0.{which}.weight"), w);
let b = alloc_f32(inter, vec![inter]);
fill_f32(&b, 0.0, inter);
tensors.insert(format!("v.blk.0.{which}.bias"), b);
}
let down_n = h * inter;
let dw = alloc_f32(down_n, vec![h, inter]);
fill_f32(&dw, 0.0, down_n);
tensors.insert("v.blk.0.ffn_down.weight".to_string(), dw);
let db = alloc_f32(h, vec![h]);
fill_f32(&db, 0.0, h);
tensors.insert("v.blk.0.ffn_down.bias".to_string(), db);
let weights = LoadedMmprojWeights::from_tensors_for_test(tensors, device);
let executor =
GraphExecutor::new(mlx_native::MlxDevice::new().expect("MlxDevice for executor"));
let mut session = executor.begin().expect("begin");
let mut registry = mlx_native::KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
mlx_native::ops::sigmoid_mul::register(&mut registry);
mlx_native::ops::rope_multi::register(&mut registry);
mlx_native::ops::gelu::register(&mut registry);
register_vit_custom_shaders(&mut registry);
register_bert_custom_shaders(&mut registry);
let device_ref: *const mlx_native::MlxDevice = executor.device() as *const _;
let device_borrow: &mlx_native::MlxDevice = unsafe { &*device_ref };
let positions = build_qwen3vl_2d_rope_positions(
device_borrow,
8, 4, true,
)
.expect("build positions");
let head_dim = (n_embd / n_head) as f32;
let scale = 1.0_f32 / head_dim.sqrt();
let block_out = apply_qwen3vl_block_forward_gpu(
session.encoder_mut(),
&mut registry,
device_borrow,
&weights,
&vit_cfg,
0,
&input_buf,
&positions,
n_pos,
scale,
10000.0,
)
.expect("apply_qwen3vl_block_forward_gpu must succeed");
session.finish().expect("finish");
let got: &[f32] = block_out.as_slice::<f32>().expect("readback");
assert_eq!(got.len(), n_total);
let input_max_abs = input_data.iter().fold(0.0_f32, |m, &v| m.max(v.abs()));
let got_max_abs = got.iter().fold(0.0_f32, |m, &v| m.max(v.abs()));
assert!(
got_max_abs > input_max_abs * 0.95,
"RESIDUAL-PRESENCE PIN: residual chain must propagate the \
non-uniform input — got_max_abs={got_max_abs:.4} should be \
≈ input_max_abs={input_max_abs:.4}. SABOTAGE: comment out \
either `vit_residual_add_gpu` call in \
apply_qwen3vl_block_forward_gpu (Stage 6 OR Stage 9) → \
output collapses to 0 and this assertion fails."
);
let mut max_diff = 0.0_f32;
for (i, (&g, &x)) in got.iter().zip(input_data.iter()).enumerate() {
assert!(g.is_finite(), "block_out[{i}] = {g} not finite");
max_diff = max_diff.max((g - x).abs());
}
assert!(
max_diff < 1e-3,
"RESIDUAL-IDENTITY PIN: with proj_w=0, block_out should equal \
input within FP tolerance; max element-wise diff = {max_diff:.6}. \
Sabotage of either residual breaks this identity (got=0, input=non-zero)."
);
}
#[test]
fn qwen3vl_2d_rope_vision_mode_consumed() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use mlx_native::ops::rope_multi::{
dispatch_rope_multi_cached, RopeMultiMode, RopeMultiParams,
};
use mlx_native::{DType, GraphExecutor};
let device = mlx_native::MlxDevice::new().expect("MlxDevice");
let head_dim: u32 = 4;
let n_heads: u32 = 1;
let n_pos: u32 = 2;
let n_dims_quarter = head_dim / 4;
let n_elements = (n_pos as usize) * (n_heads as usize) * (head_dim as usize);
let q_data: Vec<f32> = vec![
1.0, 1.0, 0.0, 0.0, 1.0, 1.0, 0.0, 0.0,
]; let positions_data: Vec<i32> = vec![
0, 0, 0, 1, 0, 0, 99, 99,
];
let q_buf = upload_f32_to_gpu(
&device,
&q_data,
vec![n_pos as usize, n_heads as usize, head_dim as usize],
)
.expect("upload q");
let positions_buf = upload_i32_to_gpu(&device, &positions_data, vec![positions_data.len()])
.expect("upload positions");
let executor = GraphExecutor::new(device);
let mut session = executor.begin().expect("begin");
let mut registry = mlx_native::KernelRegistry::new();
mlx_native::ops::rope_multi::register(&mut registry);
let device_ref: *const mlx_native::MlxDevice = executor.device() as *const _;
let device: &mlx_native::MlxDevice = unsafe { &*device_ref };
let out = device
.alloc_buffer(
n_elements * 4,
DType::F32,
vec![n_pos as usize, n_heads as usize, head_dim as usize],
)
.expect("alloc out");
let params = RopeMultiParams {
head_dim,
rope_dim: head_dim,
n_heads,
seq_len: n_pos,
freq_base: 10000.0,
mode: RopeMultiMode::Vision,
sections: [
n_dims_quarter,
n_dims_quarter,
n_dims_quarter,
n_dims_quarter,
],
};
dispatch_rope_multi_cached(
session.encoder_mut(),
&mut registry,
device,
&q_buf,
&out,
&positions_buf,
params,
)
.expect("rope dispatch must accept Vision mode (= 24) without error");
session.finish().expect("finish");
let got: Vec<f32> = out.as_slice::<f32>().expect("readback").to_vec();
assert!(
(got[0] - 1.0).abs() < 1e-5 && (got[1] - 1.0).abs() < 1e-5,
"pos 0 real parts must stay 1.0, got [{}, {}]",
got[0],
got[1]
);
assert!(
got[2].abs() < 1e-5 && got[3].abs() < 1e-5,
"pos 0 imag parts must stay 0.0, got [{}, {}]",
got[2],
got[3]
);
assert!(
(got[4] - 1.0).abs() < 1e-5,
"pos 1 pair0 real (y axis, theta=0) must stay 1.0, got {}",
got[4]
);
assert!(
got[6].abs() < 1e-5,
"pos 1 pair0 imag (y axis, theta=0) must stay 0.0, got {}",
got[6]
);
let expected_cos = 1.0_f32.cos();
let expected_sin = 1.0_f32.sin();
assert!(
(got[5] - expected_cos).abs() < 1e-4,
"pos 1 pair1 real (x axis, theta=1) must be cos(1) ≈ {expected_cos}, got {}",
got[5]
);
assert!(
(got[7] - expected_sin).abs() < 1e-4,
"pos 1 pair1 imag (x axis, theta=1) must be sin(1) ≈ {expected_sin}, got {}",
got[7]
);
assert!(
(got[5] - 1.0).abs() > 0.4,
"pos 1 pair1 must have rotated by ~1 rad — non-zero theta means \
section 1 (x) was consumed; got {} (still ≈ 1.0)",
got[5]
);
}
#[test]
fn qwen3vl_attention_is_bidirectional_no_causal_mask() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use mlx_native::{DType, GraphExecutor};
let batch: u32 = 32;
let num_heads: u32 = 1;
let head_dim: u32 = 32;
let scale = 1.0_f32 / (head_dim as f32).sqrt();
let n = (batch as usize) * (num_heads as usize) * (head_dim as usize);
let last = (batch - 1) as usize;
let mut q_data = vec![0.0f32; n];
for d in 0..(head_dim as usize) {
q_data[d] = 10.0; }
let mut k_data = vec![0.0f32; n];
let k_last_off = last * (head_dim as usize);
for d in 0..(head_dim as usize) {
k_data[k_last_off + d] = 10.0;
}
let mut v_data = vec![0.0f32; n];
let v_last_off = last * (head_dim as usize);
for d in 0..(head_dim as usize) {
v_data[v_last_off + d] = 1.0;
}
let device = mlx_native::MlxDevice::new().expect("device");
let q_buf = upload_f32_to_gpu(
&device,
&q_data,
vec![batch as usize, num_heads as usize, head_dim as usize],
)
.expect("q");
let k_buf = upload_f32_to_gpu(
&device,
&k_data,
vec![batch as usize, num_heads as usize, head_dim as usize],
)
.expect("k");
let v_buf = upload_f32_to_gpu(
&device,
&v_data,
vec![batch as usize, num_heads as usize, head_dim as usize],
)
.expect("v");
let executor = GraphExecutor::new(device);
let mut session = executor.begin().expect("begin");
let mut registry = mlx_native::KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
let device_ref: *const mlx_native::MlxDevice = executor.device() as *const _;
let device: &mlx_native::MlxDevice = unsafe { &*device_ref };
let _ = DType::F32;
let attn = vit_attention_gpu(
session.encoder_mut(),
&mut registry,
device,
&q_buf,
&k_buf,
&v_buf,
batch,
num_heads,
head_dim,
scale,
)
.expect("vit_attention_gpu");
session.finish().expect("finish");
let out: &[f32] = attn.as_slice::<f32>().expect("attn readback");
assert_eq!(out.len(), n);
for d in 0..(head_dim as usize) {
let v = out[d];
assert!(
v > 0.5,
"bidirectional attention: token 0 output[{d}] must be \
close to V[3]=1.0 (NOT V[0]=0.0 as causal would give); got {v}",
);
}
}
#[test]
fn qwen3vl_mlp_uses_gelu_not_silu() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use mlx_native::GraphExecutor;
let n = 64usize;
let device = mlx_native::MlxDevice::new().expect("device");
let gate = upload_f32_to_gpu(&device, &vec![1.0f32; n], vec![n]).expect("gate");
let up = upload_f32_to_gpu(&device, &vec![1.0f32; n], vec![n]).expect("up");
let executor = GraphExecutor::new(device);
let mut session = executor.begin().expect("begin");
let mut registry = mlx_native::KernelRegistry::new();
mlx_native::ops::gelu::register(&mut registry);
let device_ref: *const mlx_native::MlxDevice = executor.device() as *const _;
let device: &mlx_native::MlxDevice = unsafe { &*device_ref };
let activated = vit_qwen3vl_geglu_split_gpu(
session.encoder_mut(),
&mut registry,
device,
&gate,
&up,
n as u32,
)
.expect("geglu_split");
session.finish().expect("finish");
let out: &[f32] = activated.as_slice::<f32>().expect("readback");
let expected_gelu_x_1 = 0.8413_f32; let silu_x_1 = 0.7311_f32; for &v in out.iter().take(8) {
assert!(
(v - expected_gelu_x_1).abs() < 1e-2,
"MLP must use GELU activation: GELU(1.0) ≈ {expected_gelu_x_1}, \
got {v}; SILU(1.0) ≈ {silu_x_1} (different by ≈ 0.11)",
);
let dist_to_silu = (v - silu_x_1).abs();
let dist_to_gelu = (v - expected_gelu_x_1).abs();
assert!(
dist_to_gelu < dist_to_silu,
"value {v} is closer to SILU({silu_x_1}) than GELU({expected_gelu_x_1}) — \
wrong activation dispatched"
);
}
}
#[test]
fn qwen3vl_per_block_residual_present() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use mlx_native::GraphExecutor;
use std::collections::HashMap;
let n_embd: u32 = 32;
let n_head: u32 = 1;
let intermediate: u32 = 64;
let n_pos: u32 = 32;
let eps: f32 = 1e-6;
let device = mlx_native::MlxDevice::new().expect("MlxDevice");
let mmproj_cfg = MmprojConfig {
image_size: 64,
patch_size: 16,
num_patches_side: 4,
hidden_size: n_embd,
intermediate_size: intermediate,
num_attention_heads: n_head,
num_hidden_layers: 2, layer_norm_eps: eps,
projector: ProjectorType::Qwen3VlMerger,
image_mean: [0.5, 0.5, 0.5],
image_std: [0.5, 0.5, 0.5],
spatial_merge_size: Some(2),
projection_dim: Some(32),
deepstack_indexes: Some(vec![]),
};
let vit_cfg = Qwen3VlViTConfig::from_mmproj(&mmproj_cfg, 16).expect("from_mmproj");
let n_total = (n_pos * n_embd) as usize;
let input_data: Vec<f32> = (0..n_pos)
.flat_map(|r| (0..n_embd).map(move |k| 1.0 + (r as f32) * 0.05 - (k as f32) * 0.02))
.collect();
let input_buf =
upload_f32_to_gpu(&device, &input_data, vec![n_pos as usize, n_embd as usize])
.expect("upload input");
let alloc_f32 = |bytes_count: usize, shape: Vec<usize>| -> mlx_native::MlxBuffer {
device
.alloc_buffer(bytes_count * 4, mlx_native::DType::F32, shape)
.expect("alloc")
};
let fill_f32 = |buf: &mlx_native::MlxBuffer, val: f32, n: usize| {
let s: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, n) };
for v in s.iter_mut() {
*v = val;
}
};
let h = n_embd as usize;
let inter = intermediate as usize;
let mut tensors: HashMap<String, mlx_native::MlxBuffer> = HashMap::new();
for il in 0..2usize {
let blk = format!("v.blk.{il}");
for which in ["ln1", "ln2"] {
let w = alloc_f32(h, vec![h]);
fill_f32(&w, 1.0, h);
tensors.insert(format!("{blk}.{which}.weight"), w);
let b = alloc_f32(h, vec![h]);
fill_f32(&b, 0.0, h);
tensors.insert(format!("{blk}.{which}.bias"), b);
}
let proj_n = h * h;
for which in ["attn_q", "attn_k", "attn_v", "attn_out"] {
let w = alloc_f32(proj_n, vec![h, h]);
fill_f32(&w, 0.0, proj_n);
tensors.insert(format!("{blk}.{which}.weight"), w);
let b = alloc_f32(h, vec![h]);
fill_f32(&b, 0.0, h);
tensors.insert(format!("{blk}.{which}.bias"), b);
}
let up_n = inter * h;
for which in ["ffn_gate", "ffn_up"] {
let w = alloc_f32(up_n, vec![inter, h]);
fill_f32(&w, 0.0, up_n);
tensors.insert(format!("{blk}.{which}.weight"), w);
let b = alloc_f32(inter, vec![inter]);
fill_f32(&b, 0.0, inter);
tensors.insert(format!("{blk}.{which}.bias"), b);
}
let down_n = h * inter;
let dw = alloc_f32(down_n, vec![h, inter]);
fill_f32(&dw, 0.0, down_n);
tensors.insert(format!("{blk}.ffn_down.weight"), dw);
let db = alloc_f32(h, vec![h]);
fill_f32(&db, 0.0, h);
tensors.insert(format!("{blk}.ffn_down.bias"), db);
}
let weights = LoadedMmprojWeights::from_tensors_for_test(tensors, device);
let executor =
GraphExecutor::new(mlx_native::MlxDevice::new().expect("MlxDevice for executor"));
let mut session = executor.begin().expect("begin");
let mut registry = mlx_native::KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
mlx_native::ops::sigmoid_mul::register(&mut registry);
mlx_native::ops::rope_multi::register(&mut registry);
mlx_native::ops::gelu::register(&mut registry);
register_vit_custom_shaders(&mut registry);
register_bert_custom_shaders(&mut registry);
let device_ref: *const mlx_native::MlxDevice = executor.device() as *const _;
let device_borrow: &mlx_native::MlxDevice = unsafe { &*device_ref };
let positions =
build_qwen3vl_2d_rope_positions(device_borrow, 8, 4, true).expect("build positions");
let head_dim = (n_embd / n_head) as f32;
let scale = 1.0_f32 / head_dim.sqrt();
let mut hidden = input_buf;
for block_idx in 0..2usize {
hidden = apply_qwen3vl_block_forward_gpu(
session.encoder_mut(),
&mut registry,
device_borrow,
&weights,
&vit_cfg,
block_idx,
&hidden,
&positions,
n_pos,
scale,
10000.0,
)
.with_context(|| format!("block {block_idx}"))
.expect("apply_qwen3vl_block_forward_gpu");
session.encoder_mut().memory_barrier();
}
session.finish().expect("finish");
let got: &[f32] = hidden.as_slice::<f32>().expect("readback");
assert_eq!(got.len(), n_total);
let input_max_abs = input_data.iter().fold(0.0_f32, |m, &v| m.max(v.abs()));
let got_max_abs = got.iter().fold(0.0_f32, |m, &v| m.max(v.abs()));
assert!(
got_max_abs > input_max_abs * 0.95,
"RESIDUAL-PRESENCE PIN (2-block chain): output max_abs={got_max_abs:.4} \
must be ≈ input max_abs={input_max_abs:.4}. SABOTAGE: comment out \
either residual in apply_qwen3vl_block_forward_gpu → output \
collapses to 0 from the sabotaged block onward."
);
let mut max_diff = 0.0_f32;
for (i, (&g, &x)) in got.iter().zip(input_data.iter()).enumerate() {
assert!(g.is_finite(), "block_out[{i}] = {g} not finite");
max_diff = max_diff.max((g - x).abs());
}
assert!(
max_diff < 1e-3,
"RESIDUAL-IDENTITY PIN (2-block): residual chain across 2 blocks \
must preserve input within FP tolerance; max diff = {max_diff:.6}"
);
}
#[test]
fn qwen3vl_compute_returns_ok_augmented_for_2_block_synthetic() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = mlx_native::MlxDevice::new().expect("MlxDevice");
let (mut vit_cfg, mmproj_cfg) = synth_qwen3vl_block_cfg(2);
vit_cfg.deepstack_indexes = vec![0];
let weights = build_synth_qwen3vl_weights_with_deepstack(
device,
vit_cfg.n_layer,
vit_cfg.n_embd,
vit_cfg.intermediate_size,
(vit_cfg.num_position_embeddings as f64).sqrt() as u32,
1.0, 0.0, 0.1, 0.1, vit_cfg.patch_size,
&vit_cfg.deepstack_indexes,
vit_cfg.out_hidden_size,
);
let inputs = synth_zero_pixel_inputs(mmproj_cfg.image_size);
let result =
compute_vision_embeddings_gpu_qwen3vl(&inputs, &weights, &vit_cfg, &mmproj_cfg)
.expect("4c.4 dispatch must return Ok for a complete synthetic 2-block fixture");
assert_eq!(result.len(), 1, "single-image input → single-image output");
let out = &result[0];
let merge_factor = (vit_cfg.spatial_merge_size as usize).pow(2);
let n_pos_merged = (mmproj_cfg.image_size as usize / mmproj_cfg.patch_size as usize).pow(2);
let n_image_tokens = n_pos_merged / merge_factor;
let expected_len = n_image_tokens * (vit_cfg.augmented_embed_dim() as usize);
assert_eq!(
out.len(),
expected_len,
"4c.4 returns the AUGMENTED [n_image_tokens, augmented_embed_dim] \
shape: n_image_tokens={n_image_tokens}, augmented_embed_dim={} \
(= lm_hidden {} * (1 + N_deepstack {})); got len={}",
vit_cfg.augmented_embed_dim(),
vit_cfg.out_hidden_size,
vit_cfg.deepstack_indexes.len(),
out.len()
);
let row_stride = (1 + vit_cfg.deepstack_indexes.len()) * vit_cfg.out_hidden_size as usize;
assert_eq!(
row_stride,
vit_cfg.augmented_embed_dim() as usize,
"LM-split contract: augmented_embed_dim must equal \
(1+N_deepstack)*lm_hidden — pinned by qwen3vl.cpp:97 nb[1] stride"
);
for (i, &v) in out.iter().enumerate() {
assert!(v.is_finite(), "augmented[{i}] = {v} not finite");
}
assert!(
ProjectorType::Qwen3VlMerger.is_supported(),
"4c.5 LANDED: ProjectorType::Qwen3VlMerger.is_supported() must \
return true now that the LM-side hook + loader extension + \
validator extension are wired"
);
}
#[test]
fn qwen3vl_deepstack_head_layernorm_then_mlp_synthetic() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use mlx_native::{DType, GraphExecutor};
use std::collections::HashMap;
let n_image_tokens: u32 = 1;
let n_embd: u32 = 8;
let merge_factor: u32 = 4;
let merged_hidden = n_embd * merge_factor; let fc1_out: u32 = 32;
let lm_hidden: u32 = 32;
let n_pos = n_image_tokens * merge_factor; let eps: f32 = 1e-6;
let device = mlx_native::MlxDevice::new().expect("MlxDevice");
let mmproj_cfg = MmprojConfig {
image_size: 32,
patch_size: 16,
num_patches_side: 2,
hidden_size: n_embd,
intermediate_size: fc1_out,
num_attention_heads: 1,
num_hidden_layers: 1,
layer_norm_eps: eps,
projector: ProjectorType::Qwen3VlMerger,
image_mean: [0.5, 0.5, 0.5],
image_std: [0.5, 0.5, 0.5],
spatial_merge_size: Some(merge_factor.isqrt()),
projection_dim: Some(lm_hidden),
deepstack_indexes: Some(vec![0]),
};
let vit_cfg = Qwen3VlViTConfig::from_mmproj(&mmproj_cfg, 4)
.expect("from_mmproj for deepstack head test");
let input_data: Vec<f32> = (0..(n_pos * n_embd) as usize)
.map(|i| (i + 1) as f32)
.collect();
let input_buf =
upload_f32_to_gpu(&device, &input_data, vec![n_pos as usize, n_embd as usize])
.expect("upload input");
let alloc_f32 = |bytes_count: usize, shape: Vec<usize>| -> mlx_native::MlxBuffer {
device
.alloc_buffer(bytes_count * 4, DType::F32, shape)
.expect("alloc")
};
let fill_f32 = |buf: &mlx_native::MlxBuffer, val: f32, n: usize| {
let s: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, n) };
for v in s.iter_mut() {
*v = val;
}
};
let merged_h = merged_hidden as usize;
let fo = fc1_out as usize;
let lh = lm_hidden as usize;
let mut tensors: HashMap<String, mlx_native::MlxBuffer> = HashMap::new();
let nw = alloc_f32(merged_h, vec![merged_h]);
fill_f32(&nw, 1.0, merged_h);
tensors.insert("v.deepstack.0.norm.weight".to_string(), nw);
let nb = alloc_f32(merged_h, vec![merged_h]);
fill_f32(&nb, 0.0, merged_h);
tensors.insert("v.deepstack.0.norm.bias".to_string(), nb);
let fc1_n = fo * merged_h;
let fc1_w = alloc_f32(fc1_n, vec![fo, merged_h]);
fill_f32(&fc1_w, 0.5, fc1_n);
tensors.insert("v.deepstack.0.fc1.weight".to_string(), fc1_w);
let fc1_b = alloc_f32(fo, vec![fo]);
fill_f32(&fc1_b, 1.0, fo);
tensors.insert("v.deepstack.0.fc1.bias".to_string(), fc1_b);
let fc2_n = lh * fo;
let fc2_w = alloc_f32(fc2_n, vec![lh, fo]);
fill_f32(&fc2_w, 0.5, fc2_n);
tensors.insert("v.deepstack.0.fc2.weight".to_string(), fc2_w);
let fc2_b = alloc_f32(lh, vec![lh]);
fill_f32(&fc2_b, 0.0, lh);
tensors.insert("v.deepstack.0.fc2.bias".to_string(), fc2_b);
let weights = LoadedMmprojWeights::from_tensors_for_test(tensors, device);
let executor =
GraphExecutor::new(mlx_native::MlxDevice::new().expect("MlxDevice for executor"));
let mut session = executor.begin().expect("begin");
let mut registry = mlx_native::KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
mlx_native::ops::sigmoid_mul::register(&mut registry);
mlx_native::ops::rope_multi::register(&mut registry);
mlx_native::ops::gelu::register(&mut registry);
register_vit_custom_shaders(&mut registry);
register_bert_custom_shaders(&mut registry);
let device_ref: *const mlx_native::MlxDevice = executor.device() as *const _;
let device_borrow: &mlx_native::MlxDevice = unsafe { &*device_ref };
let head_out = apply_qwen3vl_deepstack_head_gpu(
session.encoder_mut(),
&mut registry,
device_borrow,
&weights,
&vit_cfg,
0, &input_buf,
n_pos,
)
.expect("apply_qwen3vl_deepstack_head_gpu must succeed");
session.finish().expect("finish");
let got: &[f32] = head_out.as_slice::<f32>().expect("readback");
assert_eq!(got.len(), (n_image_tokens * lm_hidden) as usize);
let n_input = (n_pos * n_embd) as f32; let mean = (1.0 + n_input) / 2.0; let var = ((n_input * n_input) - 1.0) / 12.0; let _stddev = (var + eps).sqrt(); let gelu_at_one: f32 = 0.5
* 1.0
* (1.0
+ (((2.0_f32) / std::f32::consts::PI).sqrt() * (1.0 + 0.044715 * 1.0_f32)).tanh());
let expected_per_element = 0.5 * (fc1_out as f32) * gelu_at_one;
for (i, &v) in got.iter().enumerate() {
assert!(
(v - expected_per_element).abs() < 5e-2,
"deepstack head[{i}] expected ≈ {expected_per_element:.4}, got {v:.4} \
(LN-zero-mean → fc1=0 → +bias=1 → GELU(1)≈{gelu_at_one:.4} → fc2*0.5*32)"
);
}
let _ = mean;
let _ = var;
}
#[test]
fn qwen3vl_spatial_merger_2x2_concat_along_channel() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use mlx_native::GraphExecutor;
let n_pos: u32 = 8;
let n_embd: u32 = 8;
let merge_factor: u32 = 4;
let n_image_tokens = n_pos / merge_factor; let merged_hidden = n_embd * merge_factor; let device = mlx_native::MlxDevice::new().expect("MlxDevice");
let mut input_data = vec![0f32; (n_pos * n_embd) as usize];
for r in 0..n_pos {
for c in 0..n_embd {
input_data[(r * n_embd + c) as usize] = r as f32;
}
}
let input_buf =
upload_f32_to_gpu(&device, &input_data, vec![n_pos as usize, n_embd as usize])
.expect("upload input");
let weight_n = (merged_hidden * merged_hidden) as usize;
let mut weight_data = vec![0f32; weight_n];
for i in 0..merged_hidden {
weight_data[(i * merged_hidden + i) as usize] = 1.0;
}
let weight_buf = upload_f32_to_gpu(
&device,
&weight_data,
vec![merged_hidden as usize, merged_hidden as usize],
)
.expect("upload weight");
let executor = GraphExecutor::new(device);
let mut session = executor.begin().expect("begin");
let mut registry = mlx_native::KernelRegistry::new();
register_vit_custom_shaders(&mut registry);
register_bert_custom_shaders(&mut registry);
let device_ref: *const mlx_native::MlxDevice = executor.device() as *const _;
let device_borrow: &mlx_native::MlxDevice = unsafe { &*device_ref };
let out = vit_linear_gpu(
session.encoder_mut(),
&mut registry,
device_borrow,
&input_buf,
&weight_buf,
n_image_tokens,
merged_hidden,
merged_hidden,
)
.expect("vit_linear_gpu identity-merger");
session.finish().expect("finish");
let got: &[f32] = out.as_slice::<f32>().expect("readback");
assert_eq!(
got.len(),
(n_image_tokens * merged_hidden) as usize,
"merger output: [n_image_tokens=2, merged_hidden=32] = 64 elements"
);
for token_idx in 0..n_image_tokens as usize {
for source_row_within_block in 0..merge_factor as usize {
let source_row_global =
token_idx * (merge_factor as usize) + source_row_within_block;
let source_value = source_row_global as f32;
for c in 0..n_embd as usize {
let dst_idx = token_idx * (merged_hidden as usize)
+ source_row_within_block * (n_embd as usize)
+ c;
assert!(
(got[dst_idx] - source_value).abs() < 1e-5,
"merger reinterpret: token {token_idx}, source row \
{source_row_within_block}, channel {c} → dst[{dst_idx}] \
expected {source_value}, got {}",
got[dst_idx]
);
}
}
}
}
#[test]
fn qwen3vl_main_projector_2layer_gelu_synthetic() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use mlx_native::{DType, GraphExecutor};
use std::collections::HashMap;
let n_image_tokens: u32 = 1;
let n_embd: u32 = 8;
let merge_factor: u32 = 4;
let merged_hidden = n_embd * merge_factor; let mm0_out: u32 = 32;
let lm_hidden: u32 = 32;
let n_pos = n_image_tokens * merge_factor; let device = mlx_native::MlxDevice::new().expect("MlxDevice");
let mmproj_cfg = MmprojConfig {
image_size: 32,
patch_size: 16,
num_patches_side: 2,
hidden_size: n_embd,
intermediate_size: mm0_out,
num_attention_heads: 1,
num_hidden_layers: 1,
layer_norm_eps: 1e-6,
projector: ProjectorType::Qwen3VlMerger,
image_mean: [0.5, 0.5, 0.5],
image_std: [0.5, 0.5, 0.5],
spatial_merge_size: Some(merge_factor.isqrt()),
projection_dim: Some(lm_hidden),
deepstack_indexes: Some(vec![]),
};
let vit_cfg = Qwen3VlViTConfig::from_mmproj(&mmproj_cfg, 4)
.expect("from_mmproj for main projector test");
let input_data: Vec<f32> = (0..(n_pos * n_embd) as usize)
.map(|i| (i + 1) as f32)
.collect();
let input_buf =
upload_f32_to_gpu(&device, &input_data, vec![n_pos as usize, n_embd as usize])
.expect("upload input");
let alloc_f32 = |bytes_count: usize, shape: Vec<usize>| -> mlx_native::MlxBuffer {
device
.alloc_buffer(bytes_count * 4, DType::F32, shape)
.expect("alloc")
};
let fill_f32 = |buf: &mlx_native::MlxBuffer, val: f32, n: usize| {
let s: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, n) };
for v in s.iter_mut() {
*v = val;
}
};
let merged_h = merged_hidden as usize;
let mo = mm0_out as usize;
let lh = lm_hidden as usize;
let mut tensors: HashMap<String, mlx_native::MlxBuffer> = HashMap::new();
let mm0_n = mo * merged_h;
let mm0_w = alloc_f32(mm0_n, vec![mo, merged_h]);
fill_f32(&mm0_w, 0.5, mm0_n);
tensors.insert(super::super::mmproj::TENSOR_MM_0_WEIGHT.to_string(), mm0_w);
let mm0_b = alloc_f32(mo, vec![mo]);
fill_f32(&mm0_b, 1.0, mo);
tensors.insert(super::super::mmproj::TENSOR_MM_0_BIAS.to_string(), mm0_b);
let mm2_n = lh * mo;
let mm2_w = alloc_f32(mm2_n, vec![lh, mo]);
fill_f32(&mm2_w, 0.5, mm2_n);
tensors.insert(super::super::mmproj::TENSOR_MM_2_WEIGHT.to_string(), mm2_w);
let mm2_b = alloc_f32(lh, vec![lh]);
fill_f32(&mm2_b, 0.0, lh);
tensors.insert(super::super::mmproj::TENSOR_MM_2_BIAS.to_string(), mm2_b);
let weights = LoadedMmprojWeights::from_tensors_for_test(tensors, device);
let executor =
GraphExecutor::new(mlx_native::MlxDevice::new().expect("MlxDevice for executor"));
let mut session = executor.begin().expect("begin");
let mut registry = mlx_native::KernelRegistry::new();
mlx_native::ops::gelu::register(&mut registry);
register_vit_custom_shaders(&mut registry);
register_bert_custom_shaders(&mut registry);
let device_ref: *const mlx_native::MlxDevice = executor.device() as *const _;
let device_borrow: &mlx_native::MlxDevice = unsafe { &*device_ref };
let proj_out = apply_qwen3vl_main_projector_gpu(
session.encoder_mut(),
&mut registry,
device_borrow,
&weights,
&vit_cfg,
&input_buf,
n_pos,
)
.expect("apply_qwen3vl_main_projector_gpu must succeed");
session.finish().expect("finish");
let got: &[f32] = proj_out.as_slice::<f32>().expect("readback");
assert_eq!(got.len(), (n_image_tokens * lm_hidden) as usize);
let row_sum: f32 = (1..=32).sum::<i32>() as f32; let mm0_pre_gelu = 0.5 * row_sum + 1.0; let gelu_pre = 0.5
* mm0_pre_gelu
* (1.0
+ ((2.0_f32 / std::f32::consts::PI).sqrt()
* (mm0_pre_gelu + 0.044715 * mm0_pre_gelu.powi(3)))
.tanh());
let expected = 0.5 * (mm0_out as f32) * gelu_pre;
for (i, &v) in got.iter().enumerate() {
assert!(
(v - expected).abs() / expected.abs() < 1e-2,
"main projector[{i}] expected ≈ {expected:.2}, got {v:.2} \
(mm0(=0.5*528+1=265) → GELU({mm0_pre_gelu:.0})≈{gelu_pre:.0} \
→ mm2(=0.5*32*GELU)≈{expected:.0})"
);
assert!(
v > 1000.0 && v < 5000.0,
"main projector[{i}] = {v:.2} is out of expected 1000..5000 \
range — both mm.0 AND mm.2 layers must contribute"
);
}
}
#[test]
fn qwen3vl_deepstack_head_taps_correct_block_index() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = mlx_native::MlxDevice::new().expect("MlxDevice");
let mut mmproj_cfg = synth_qwen3vl_mmproj_cfg(
3,
Some(vec![2]), Some(2),
Some(32),
);
mmproj_cfg.image_size = 128;
mmproj_cfg.patch_size = 16;
mmproj_cfg.num_patches_side = 8;
mmproj_cfg.hidden_size = 32;
mmproj_cfg.intermediate_size = 64;
mmproj_cfg.num_attention_heads = 1;
let vit_cfg = Qwen3VlViTConfig::from_mmproj(&mmproj_cfg, 64).expect("from_mmproj");
let weights = build_synth_qwen3vl_weights_with_deepstack(
device,
vit_cfg.n_layer,
vit_cfg.n_embd,
vit_cfg.intermediate_size,
(vit_cfg.num_position_embeddings as f64).sqrt() as u32,
1.0, 0.0, 0.05, 0.05, vit_cfg.patch_size,
&vit_cfg.deepstack_indexes,
vit_cfg.out_hidden_size,
);
let inputs = synth_zero_pixel_inputs(mmproj_cfg.image_size);
let result =
compute_vision_embeddings_gpu_qwen3vl(&inputs, &weights, &vit_cfg, &mmproj_cfg)
.expect("forward must succeed with single-flag fixture");
let out = &result[0];
let merge_factor = (vit_cfg.spatial_merge_size as usize).pow(2);
let n_pos_merged = (mmproj_cfg.image_size as usize / mmproj_cfg.patch_size as usize).pow(2);
let n_image_tokens = n_pos_merged / merge_factor;
let one_plus_flagged = 1 + vit_cfg.deepstack_indexes.len(); let one_plus_layers = 1 + vit_cfg.n_layer as usize; let expected_len = n_image_tokens * (vit_cfg.out_hidden_size as usize) * one_plus_flagged;
let wrong_len = n_image_tokens * (vit_cfg.out_hidden_size as usize) * one_plus_layers;
assert_eq!(
out.len(),
expected_len,
"head-tap-index pin: augmented_embed length must match \
n_image_tokens * lm_hidden * (1 + |deepstack_indexes|) = \
{n_image_tokens} * {} * {one_plus_flagged} = {expected_len}",
vit_cfg.out_hidden_size
);
assert_ne!(
out.len(),
wrong_len,
"head-tap-index pin: if augmented_embed length matched \
n_image_tokens * lm_hidden * (1 + n_layer) = {wrong_len}, \
that would imply a head was applied at EVERY block, \
violating qwen3vl.cpp:150's `if (layer.has_deepstack())` gate"
);
}
#[test]
fn qwen3vl_augmented_embed_shape_matches_lm_split_contract() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = mlx_native::MlxDevice::new().expect("MlxDevice");
let mut mmproj_cfg = synth_qwen3vl_mmproj_cfg(3, Some(vec![0, 1, 2]), Some(2), Some(32));
mmproj_cfg.image_size = 128;
mmproj_cfg.patch_size = 16;
mmproj_cfg.num_patches_side = 8;
mmproj_cfg.hidden_size = 32;
mmproj_cfg.intermediate_size = 64;
mmproj_cfg.num_attention_heads = 1;
let vit_cfg = Qwen3VlViTConfig::from_mmproj(&mmproj_cfg, 64).expect("from_mmproj");
let weights = build_synth_qwen3vl_weights_with_deepstack(
device,
vit_cfg.n_layer,
vit_cfg.n_embd,
vit_cfg.intermediate_size,
(vit_cfg.num_position_embeddings as f64).sqrt() as u32,
1.0,
0.0,
0.05,
0.05,
vit_cfg.patch_size,
&vit_cfg.deepstack_indexes,
vit_cfg.out_hidden_size,
);
let inputs = synth_zero_pixel_inputs(mmproj_cfg.image_size);
let result =
compute_vision_embeddings_gpu_qwen3vl(&inputs, &weights, &vit_cfg, &mmproj_cfg)
.expect("forward must succeed with 3-flag fixture");
let out = &result[0];
let merge_factor = (vit_cfg.spatial_merge_size as usize).pow(2);
let n_pos_merged = (mmproj_cfg.image_size as usize / mmproj_cfg.patch_size as usize).pow(2);
let n_image_tokens = n_pos_merged / merge_factor;
let lm_hidden = vit_cfg.out_hidden_size as usize;
let n_deepstack = vit_cfg.deepstack_indexes.len();
let row_stride_floats = lm_hidden * (1 + n_deepstack);
let expected_total = n_image_tokens * row_stride_floats;
assert_eq!(
out.len(),
expected_total,
"augmented total = n_image_tokens ({n_image_tokens}) * row_stride_floats \
({row_stride_floats}) = {expected_total}"
);
assert_eq!(
row_stride_floats,
vit_cfg.augmented_embed_dim() as usize,
"row_stride_floats == augmented_embed_dim — pinned by qwen3vl.cpp:97 nb[1]"
);
for token in 0..n_image_tokens {
let row_base = token * row_stride_floats;
for n in 0..(1 + n_deepstack) {
let chunk_start = row_base + n * lm_hidden;
let chunk_end = chunk_start + lm_hidden;
assert!(
chunk_end <= out.len(),
"chunk {n} of token {token} would extend past augmented buffer"
);
let chunk = &out[chunk_start..chunk_end];
assert_eq!(chunk.len(), lm_hidden);
for &v in chunk {
assert!(v.is_finite(), "chunk value not finite");
}
}
}
for il in 0..n_deepstack {
let lm_offset_floats = (il + 1) * lm_hidden;
assert!(
lm_offset_floats < row_stride_floats,
"LM-split offset {lm_offset_floats} for il={il} must be < row_stride \
{row_stride_floats}"
);
}
}
#[test]
fn qwen3vl_cpu_concat_augmented_embed_byte_exact() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let n_image_tokens = 2usize;
let lm_hidden = 4usize;
let chunks: Vec<Vec<f32>> = (0..3)
.map(|c| vec![((c + 1) * 10) as f32; n_image_tokens * lm_hidden])
.collect();
let out = qwen3vl_concat_augmented_embed_cpu(&chunks, n_image_tokens, lm_hidden)
.expect("concat must succeed for valid input");
let row_stride = chunks.len() * lm_hidden; assert_eq!(out.len(), n_image_tokens * row_stride);
for t in 0..n_image_tokens {
let row_base = t * row_stride;
for k in 0..lm_hidden {
assert_eq!(out[row_base + k], 10.0);
}
for k in 0..lm_hidden {
assert_eq!(out[row_base + lm_hidden + k], 20.0);
}
for k in 0..lm_hidden {
assert_eq!(out[row_base + 2 * lm_hidden + k], 30.0);
}
}
let err = qwen3vl_concat_augmented_embed_cpu(&[], 1, 4).unwrap_err();
assert!(format!("{err}").contains("at least the base"));
let bad = vec![vec![1.0f32; 8], vec![1.0f32; 7]];
let err = qwen3vl_concat_augmented_embed_cpu(&bad, 2, 4).unwrap_err();
assert!(format!("{err}").contains("length"));
}
#[test]
fn qwen3vl_merger_is_supported_after_wedge_4c5() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
assert!(
ProjectorType::Qwen3VlMerger.is_supported(),
"Wedge-4c.5 LANDED: ProjectorType::Qwen3VlMerger.is_supported() must return true"
);
assert!(
super::super::mmproj::ArchProfile::Qwen3VlSiglip.is_supported(),
"Wedge-4c.5 LANDED: ArchProfile::Qwen3VlSiglip.is_supported() must return true \
so the validator + handler-side mmproj.arch.is_supported() check accepts \
text-only chat against Qwen3-VL GGUFs"
);
}
#[test]
fn qwen3vl_dispatch_routes_to_compute_vision_embeddings_gpu_qwen3vl() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::inference::vision::vit_gpu::{
compute_vision_embeddings_gpu_dispatch, VisionInput,
};
let device = mlx_native::MlxDevice::new().expect("MlxDevice");
let (vit_cfg, mmproj_cfg) = synth_qwen3vl_block_cfg(2);
let weights = build_synth_qwen3vl_weights_with_deepstack(
device,
vit_cfg.n_layer,
vit_cfg.n_embd,
vit_cfg.intermediate_size,
(vit_cfg.num_position_embeddings as f64).sqrt() as u32,
1.0,
0.0,
0.1,
0.1,
vit_cfg.patch_size,
&vit_cfg.deepstack_indexes,
vit_cfg.out_hidden_size,
);
let _ = VisionInput::Siglip49; let inputs = synth_zero_pixel_inputs(mmproj_cfg.image_size);
let result = compute_vision_embeddings_gpu_dispatch(
&inputs,
super::super::mmproj::ArchProfile::Qwen3VlSiglip,
&weights,
&mmproj_cfg,
1.0, );
match result {
Ok(out) => {
assert_eq!(out.len(), 1, "single image → single output");
let n_pos_merged =
(mmproj_cfg.image_size as usize / mmproj_cfg.patch_size as usize).pow(2);
let merge_factor = (vit_cfg.spatial_merge_size as usize).pow(2);
let n_image_tokens = n_pos_merged / merge_factor;
let expected_len = n_image_tokens * (vit_cfg.out_hidden_size as usize);
assert_eq!(
out[0].len(),
expected_len,
"dispatch routing must produce the augmented [n_image_tokens, \
augmented_embed_dim] shape per the Qwen3-VL contract \
(synthetic empty-deepstack fixture, so augmented_embed_dim = \
out_hidden_size)"
);
}
Err(e) => {
let msg = format!("{e:#}");
assert!(
msg.contains("qwen3vl") || msg.contains("Qwen3VlSiglip"),
"dispatch must route Qwen3VlSiglip to the qwen3vl module; \
got error from a non-qwen3vl path: {msg}"
);
}
}
}
#[test]
fn qwen3vl_pos_embed_resize_to_rectangular_target_iter225() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let trained_n: u32 = 8;
let n_embd: u32 = 2;
let num_pos = trained_n * trained_n; let table: Vec<f32> = (0..(num_pos as usize) * (n_embd as usize))
.map(|i| (i as f32) * 0.01 - 0.5)
.collect();
let target_n_x: u32 = 8;
let target_n_y: u32 = 4;
let resized = qwen3vl_resize_position_embeddings_bilinear(
&table, num_pos, n_embd, target_n_x, target_n_y,
)
.expect("rectangular resize 8×8 → 8×4 must succeed");
let expected_len = (target_n_x as usize) * (target_n_y as usize) * (n_embd as usize);
assert_eq!(resized.len(), expected_len);
for (i, v) in resized.iter().enumerate() {
assert!(v.is_finite(), "resized[{i}] = {v} not finite");
}
let any_nonzero = resized.iter().any(|v| v.abs() > 1e-6);
assert!(
any_nonzero,
"all-zero output suggests bilinear blend collapsed"
);
}
#[test]
fn qwen3vl_compute_accepts_rectangular_input_iter225() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = mlx_native::MlxDevice::new().expect("MlxDevice");
let (mut vit_cfg, mut mmproj_cfg) = synth_qwen3vl_block_cfg(2);
vit_cfg.deepstack_indexes = vec![0];
let weights = build_synth_qwen3vl_weights_with_deepstack(
device,
vit_cfg.n_layer,
vit_cfg.n_embd,
vit_cfg.intermediate_size,
(vit_cfg.num_position_embeddings as f64).sqrt() as u32,
1.0,
0.0,
0.1,
0.1,
vit_cfg.patch_size,
&vit_cfg.deepstack_indexes,
vit_cfg.out_hidden_size,
);
let pixel_w: u32 = 128;
let pixel_h: u32 = 64;
mmproj_cfg.image_size = pixel_h; let pixel_values = vec![0.0f32; 3 * (pixel_h as usize) * (pixel_w as usize)];
let img = crate::inference::vision::PreprocessedImage {
pixel_values,
target_size: pixel_h,
pixel_w: Some(pixel_w),
pixel_h: Some(pixel_h),
source_label: "phase2-rect-fixture".to_string(),
};
let inputs = vec![crate::inference::vision::vit_gpu::VisionInput::Siglip49(
img,
)];
let result =
compute_vision_embeddings_gpu_qwen3vl(&inputs, &weights, &vit_cfg, &mmproj_cfg)
.expect("Phase-2 rectangular ViT forward must succeed");
assert_eq!(result.len(), 1);
let n_image_tokens = 8usize;
let expected_len = n_image_tokens * (vit_cfg.augmented_embed_dim() as usize);
assert_eq!(
result[0].len(),
expected_len,
"Phase-2: augmented embed = n_image_tokens * \
augmented_embed_dim = {n_image_tokens} * {} = {expected_len}",
vit_cfg.augmented_embed_dim()
);
for (i, &v) in result[0].iter().enumerate() {
assert!(v.is_finite(), "augmented[{i}] = {v} not finite");
}
}
#[test]
fn qwen3vl_compute_rejects_misaligned_rectangular_input_iter225() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = mlx_native::MlxDevice::new().expect("MlxDevice");
let (vit_cfg, mut mmproj_cfg) = synth_qwen3vl_block_cfg(1);
let weights = build_synth_qwen3vl_weights_with_deepstack(
device,
vit_cfg.n_layer,
vit_cfg.n_embd,
vit_cfg.intermediate_size,
(vit_cfg.num_position_embeddings as f64).sqrt() as u32,
1.0,
0.0,
0.1,
0.1,
vit_cfg.patch_size,
&[],
vit_cfg.out_hidden_size,
);
let pixel_w: u32 = 48;
let pixel_h: u32 = 32;
mmproj_cfg.image_size = pixel_h;
let pixel_values = vec![0.0f32; 3 * (pixel_h as usize) * (pixel_w as usize)];
let img = crate::inference::vision::PreprocessedImage {
pixel_values,
target_size: pixel_h,
pixel_w: Some(pixel_w),
pixel_h: Some(pixel_h),
source_label: "phase2-misaligned-fixture".to_string(),
};
let inputs = vec![crate::inference::vision::vit_gpu::VisionInput::Siglip49(
img,
)];
let err = compute_vision_embeddings_gpu_qwen3vl(&inputs, &weights, &vit_cfg, &mmproj_cfg)
.expect_err("misaligned pixel_w must fail loud");
let msg = format!("{err}");
assert!(
msg.contains("must be a multiple") && msg.contains("48"),
"error must name the offending dim and stride contract; got: {msg}"
);
}
#[test]
fn qwen3vl_compute_backward_compat_square_input_iter225() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = mlx_native::MlxDevice::new().expect("MlxDevice");
let (mut vit_cfg, mmproj_cfg) = synth_qwen3vl_block_cfg(2);
vit_cfg.deepstack_indexes = vec![0];
let weights = build_synth_qwen3vl_weights_with_deepstack(
device,
vit_cfg.n_layer,
vit_cfg.n_embd,
vit_cfg.intermediate_size,
(vit_cfg.num_position_embeddings as f64).sqrt() as u32,
1.0,
0.0,
0.1,
0.1,
vit_cfg.patch_size,
&vit_cfg.deepstack_indexes,
vit_cfg.out_hidden_size,
);
let inputs = synth_zero_pixel_inputs(mmproj_cfg.image_size);
let result =
compute_vision_embeddings_gpu_qwen3vl(&inputs, &weights, &vit_cfg, &mmproj_cfg)
.expect("backward-compat square input must succeed");
let merge_factor = (vit_cfg.spatial_merge_size as usize).pow(2);
let n_pos_merged = (mmproj_cfg.image_size as usize / mmproj_cfg.patch_size as usize).pow(2);
let n_image_tokens = n_pos_merged / merge_factor;
let expected_len = n_image_tokens * (vit_cfg.augmented_embed_dim() as usize);
assert_eq!(result[0].len(), expected_len);
}
#[test]
fn qwen3vl_compute_square_byte_deterministic_iter225() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device_1 = mlx_native::MlxDevice::new().expect("MlxDevice 1");
let device_2 = mlx_native::MlxDevice::new().expect("MlxDevice 2");
let (mut vit_cfg, mmproj_cfg) = synth_qwen3vl_block_cfg(2);
vit_cfg.deepstack_indexes = vec![0];
let weights_1 = build_synth_qwen3vl_weights_with_deepstack(
device_1,
vit_cfg.n_layer,
vit_cfg.n_embd,
vit_cfg.intermediate_size,
(vit_cfg.num_position_embeddings as f64).sqrt() as u32,
1.0,
0.0,
0.1,
0.1,
vit_cfg.patch_size,
&vit_cfg.deepstack_indexes,
vit_cfg.out_hidden_size,
);
let weights_2 = build_synth_qwen3vl_weights_with_deepstack(
device_2,
vit_cfg.n_layer,
vit_cfg.n_embd,
vit_cfg.intermediate_size,
(vit_cfg.num_position_embeddings as f64).sqrt() as u32,
1.0,
0.0,
0.1,
0.1,
vit_cfg.patch_size,
&vit_cfg.deepstack_indexes,
vit_cfg.out_hidden_size,
);
let inputs_1 = synth_zero_pixel_inputs(mmproj_cfg.image_size);
let inputs_2 = synth_zero_pixel_inputs(mmproj_cfg.image_size);
let r1 =
compute_vision_embeddings_gpu_qwen3vl(&inputs_1, &weights_1, &vit_cfg, &mmproj_cfg)
.expect("first call");
let r2 =
compute_vision_embeddings_gpu_qwen3vl(&inputs_2, &weights_2, &vit_cfg, &mmproj_cfg)
.expect("second call");
assert_eq!(r1.len(), 1);
assert_eq!(r2.len(), 1);
assert_eq!(r1[0].len(), r2[0].len());
for (i, (a, b)) in r1[0].iter().zip(r2[0].iter()).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"iter-225 must be byte-deterministic on square inputs at \
element {i}: a={a} b={b}"
);
}
}
#[test]
fn qwen3vl_phase2_image_grid_n_tokens_invariant_iter225() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let landscape_n_x: u32 = 24;
let landscape_n_y: u32 = 13;
let landscape_n_tokens = landscape_n_x * landscape_n_y;
assert_eq!(landscape_n_tokens, 312);
assert!(landscape_n_x <= 24, "canonical ≤ canvas grid");
assert!(landscape_n_y <= 24, "canonical ≤ canvas grid");
let portrait_n_x: u32 = 13;
let portrait_n_y: u32 = 24;
let portrait_n_tokens = portrait_n_x * portrait_n_y;
assert_eq!(portrait_n_tokens, 312);
let square_side: u32 = 24;
assert_eq!(square_side * square_side, 576);
assert_ne!(
(landscape_n_x, landscape_n_y),
(portrait_n_x, portrait_n_y),
"landscape and portrait grids must be distinct even when \
n_image_tokens matches — proves that explicit grid \
threading (not factorization) is required"
);
}
#[test]
fn patch_embed_forward_hw_square_matches_legacy_and_rect_doubles_iter225() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::inference::vision::vit::{patch_embed_forward, patch_embed_forward_hw};
let patch_size: u32 = 16;
let hidden: u32 = 32;
let inner = 3 * (patch_size as usize) * (patch_size as usize);
let n_w = (hidden as usize) * inner;
let weight: Vec<f32> = (0..n_w).map(|i| (i as f32) * 0.001 - 0.5).collect();
let bias: Vec<f32> = (0..hidden as usize).map(|i| (i as f32) * 0.01).collect();
let pixel_values_sq: Vec<f32> = (0..(3 * 16 * 16))
.map(|i| ((i as f32) * 0.01).sin())
.collect();
let legacy = patch_embed_forward(
&pixel_values_sq,
&weight,
Some(&bias),
16,
patch_size,
hidden,
)
.expect("legacy square");
let phase2_sq = patch_embed_forward_hw(
&pixel_values_sq,
&weight,
Some(&bias),
16,
16,
patch_size,
hidden,
)
.expect("phase-2 square");
assert_eq!(legacy.len(), phase2_sq.len());
for (i, (a, b)) in legacy.iter().zip(phase2_sq.iter()).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"square 16×16 must be byte-exact at index {i}: legacy={a} phase2={b}"
);
}
let pixel_values_rect: Vec<f32> = (0..(3 * 16 * 32))
.map(|i| ((i as f32) * 0.01).cos())
.collect();
let phase2_rect = patch_embed_forward_hw(
&pixel_values_rect,
&weight,
Some(&bias),
16,
32,
patch_size,
hidden,
)
.expect("phase-2 rect");
assert_eq!(phase2_rect.len(), 2 * (hidden as usize));
}
fn fnv1a64_of_f32_slice(xs: &[f32]) -> u64 {
const FNV_OFFSET: u64 = 0xcbf2_9ce4_8422_2325;
const FNV_PRIME: u64 = 0x0000_0100_0000_01b3;
let mut h: u64 = FNV_OFFSET;
for x in xs {
for &b in &x.to_bits().to_le_bytes() {
h ^= b as u64;
h = h.wrapping_mul(FNV_PRIME);
}
}
h
}
fn override_tensor_with_seeded_sin(
weights: &LoadedMmprojWeights,
name: &str,
scale: f32,
offset: f32,
) {
let buf = weights
.get(name)
.unwrap_or_else(|| panic!("ADR-021 baseline: missing tensor '{name}'"));
let n = buf.byte_len() / 4;
let s: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, n) };
for (i, v) in s.iter_mut().enumerate() {
*v = (((i as f32) + offset) * 0.017_3_f32).sin() * scale;
}
}
#[test]
fn adr021_iter1a_e2e_byte_pinned_baseline_2026_05_07() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = mlx_native::MlxDevice::new().expect("MlxDevice");
let (mut vit_cfg, mmproj_cfg) = synth_qwen3vl_block_cfg(2);
vit_cfg.deepstack_indexes = vec![0];
let weights = build_synth_qwen3vl_weights_with_deepstack(
device,
vit_cfg.n_layer,
vit_cfg.n_embd,
vit_cfg.intermediate_size,
(vit_cfg.num_position_embeddings as f64).sqrt() as u32,
1.0, 0.0, 0.0, 0.0, vit_cfg.patch_size,
&vit_cfg.deepstack_indexes,
vit_cfg.out_hidden_size,
);
override_tensor_with_seeded_sin(&weights, "v.patch_embd.weight", 0.05, 1.0);
override_tensor_with_seeded_sin(&weights, "v.patch_embd.weight.1", 0.05, 313.0);
override_tensor_with_seeded_sin(&weights, "v.position_embd.weight", 0.03, 727.0);
override_tensor_with_seeded_sin(&weights, "v.post_ln.weight", 1.0, 11.0);
override_tensor_with_seeded_sin(&weights, "v.post_ln.bias", 0.01, 17.0);
for il in 0..(vit_cfg.n_layer as usize) {
let blk = format!("v.blk.{il}");
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.ln1.weight"),
1.0,
(101 + il) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.ln1.bias"),
0.01,
(151 + il) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.ln2.weight"),
1.0,
(201 + il) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.ln2.bias"),
0.01,
(251 + il) as f32,
);
for which in ["attn_q", "attn_k", "attn_v", "attn_out"] {
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.{which}.weight"),
0.04,
(301 + il * 7) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.{which}.bias"),
0.01,
(401 + il * 7) as f32,
);
}
for which in ["ffn_gate", "ffn_up"] {
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.{which}.weight"),
0.03,
(501 + il * 5) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.{which}.bias"),
0.01,
(601 + il * 5) as f32,
);
}
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.ffn_down.weight"),
0.03,
(701 + il * 3) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.ffn_down.bias"),
0.01,
(801 + il * 3) as f32,
);
}
override_tensor_with_seeded_sin(&weights, "mm.0.weight", 0.02, 9001.0);
override_tensor_with_seeded_sin(&weights, "mm.0.bias", 0.01, 9011.0);
override_tensor_with_seeded_sin(&weights, "mm.2.weight", 0.02, 9101.0);
override_tensor_with_seeded_sin(&weights, "mm.2.bias", 0.01, 9111.0);
for &il in &vit_cfg.deepstack_indexes {
let il_us = il as usize;
override_tensor_with_seeded_sin(
&weights,
&format!("v.deepstack.{il_us}.norm.weight"),
1.0,
(10001 + il_us) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("v.deepstack.{il_us}.norm.bias"),
0.01,
(10101 + il_us) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("v.deepstack.{il_us}.fc1.weight"),
0.02,
(10201 + il_us) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("v.deepstack.{il_us}.fc1.bias"),
0.01,
(10301 + il_us) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("v.deepstack.{il_us}.fc2.weight"),
0.02,
(10401 + il_us) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("v.deepstack.{il_us}.fc2.bias"),
0.01,
(10501 + il_us) as f32,
);
}
let pixel_h: u32 = mmproj_cfg.image_size;
let pixel_w: u32 = mmproj_cfg.image_size;
let n_px = 3 * (pixel_h as usize) * (pixel_w as usize);
let pixel_values: Vec<f32> = (0..n_px)
.map(|i| (((i as f32) * 0.011_7_f32).sin() * 0.5).clamp(-1.0, 1.0))
.collect();
let img = crate::inference::vision::PreprocessedImage {
pixel_values,
target_size: pixel_h,
pixel_w: Some(pixel_w),
pixel_h: Some(pixel_h),
source_label: "adr021-iter1a-baseline".to_string(),
};
let inputs = vec![crate::inference::vision::vit_gpu::VisionInput::Siglip49(
img,
)];
let result =
compute_vision_embeddings_gpu_qwen3vl(&inputs, &weights, &vit_cfg, &mmproj_cfg)
.expect("ADR-021 iter-1a baseline must succeed");
assert_eq!(result.len(), 1);
let out = &result[0];
let expected_len = 16 * (vit_cfg.augmented_embed_dim() as usize);
assert_eq!(
out.len(),
expected_len,
"ADR-021 iter-1a baseline: augmented embed must be n_image_tokens(16) \
* augmented_embed_dim({}) = {expected_len}",
vit_cfg.augmented_embed_dim()
);
let h = fnv1a64_of_f32_slice(out);
let mut first8 = [0u32; 8];
let mut last8 = [0u32; 8];
for i in 0..8 {
first8[i] = out[i].to_bits();
last8[i] = out[out.len() - 8 + i].to_bits();
}
eprintln!("=== ADR-021 iter-1a baseline ===");
eprintln!("len = {}", out.len());
eprintln!("fnv1a64 = 0x{:016x}", h);
eprintln!("first8 = {:08x?}", first8);
eprintln!("last8 = {:08x?}", last8);
const EXPECTED_LEN: usize = 1024;
const EXPECTED_FNV1A64: u64 = 0x7da7_f3ad_353c_585b;
const EXPECTED_FIRST8: [u32; 8] = [
0x3b43_3e0f,
0x3b15_c486,
0x3ba1_1b92,
0x3c05_1b52,
0x3c0b_aa1a,
0x3bbf_01b9,
0x3b50_238e,
0x3b6f_5af2,
];
const EXPECTED_LAST8: [u32; 8] = [
0xbb22_572b,
0x384c_57e0,
0x3aa7_c065,
0x3837_f620,
0xbb09_8b6f,
0xbb2a_2262,
0xba36_da0c,
0x3adc_9039,
];
assert_eq!(out.len(), EXPECTED_LEN);
if EXPECTED_FNV1A64 != 0 {
assert_eq!(
h, EXPECTED_FNV1A64,
"ADR-021 iter-1a byte-pinned hash drift — re-run \
from a clean main, capture the printed fnv1a64, \
update EXPECTED_FNV1A64 if the drift is intentional"
);
assert_eq!(first8, EXPECTED_FIRST8, "first8 bit drift");
assert_eq!(last8, EXPECTED_LAST8, "last8 bit drift");
}
}
#[test]
fn adr021_iter5b_ac3_substitute_three_shapes_2026_05_07() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let run_at_shape = |pixel_h: u32, pixel_w: u32| -> (Vec<f32>, u64) {
let device = mlx_native::MlxDevice::new().expect("MlxDevice");
let (mut vit_cfg, mut mmproj_cfg) = synth_qwen3vl_block_cfg(2);
vit_cfg.deepstack_indexes = vec![0];
mmproj_cfg.image_size = pixel_h;
let weights = build_synth_qwen3vl_weights_with_deepstack(
device,
vit_cfg.n_layer,
vit_cfg.n_embd,
vit_cfg.intermediate_size,
(vit_cfg.num_position_embeddings as f64).sqrt() as u32,
1.0,
0.0,
0.0,
0.0,
vit_cfg.patch_size,
&vit_cfg.deepstack_indexes,
vit_cfg.out_hidden_size,
);
override_tensor_with_seeded_sin(&weights, "v.patch_embd.weight", 0.05, 1.0);
override_tensor_with_seeded_sin(&weights, "v.patch_embd.weight.1", 0.05, 313.0);
override_tensor_with_seeded_sin(&weights, "v.position_embd.weight", 0.03, 727.0);
override_tensor_with_seeded_sin(&weights, "v.post_ln.weight", 1.0, 11.0);
override_tensor_with_seeded_sin(&weights, "v.post_ln.bias", 0.01, 17.0);
for il in 0..(vit_cfg.n_layer as usize) {
let blk = format!("v.blk.{il}");
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.ln1.weight"),
1.0,
(101 + il) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.ln1.bias"),
0.01,
(151 + il) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.ln2.weight"),
1.0,
(201 + il) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.ln2.bias"),
0.01,
(251 + il) as f32,
);
for which in ["attn_q", "attn_k", "attn_v", "attn_out"] {
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.{which}.weight"),
0.04,
(301 + il * 7) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.{which}.bias"),
0.01,
(401 + il * 7) as f32,
);
}
for which in ["ffn_gate", "ffn_up"] {
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.{which}.weight"),
0.03,
(501 + il * 5) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.{which}.bias"),
0.01,
(601 + il * 5) as f32,
);
}
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.ffn_down.weight"),
0.03,
(701 + il * 3) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("{blk}.ffn_down.bias"),
0.01,
(801 + il * 3) as f32,
);
}
override_tensor_with_seeded_sin(&weights, "mm.0.weight", 0.02, 9001.0);
override_tensor_with_seeded_sin(&weights, "mm.0.bias", 0.01, 9011.0);
override_tensor_with_seeded_sin(&weights, "mm.2.weight", 0.02, 9101.0);
override_tensor_with_seeded_sin(&weights, "mm.2.bias", 0.01, 9111.0);
for &il in &vit_cfg.deepstack_indexes {
let il_us = il as usize;
override_tensor_with_seeded_sin(
&weights,
&format!("v.deepstack.{il_us}.norm.weight"),
1.0,
(10001 + il_us) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("v.deepstack.{il_us}.norm.bias"),
0.01,
(10101 + il_us) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("v.deepstack.{il_us}.fc1.weight"),
0.02,
(10201 + il_us) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("v.deepstack.{il_us}.fc1.bias"),
0.01,
(10301 + il_us) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("v.deepstack.{il_us}.fc2.weight"),
0.02,
(10401 + il_us) as f32,
);
override_tensor_with_seeded_sin(
&weights,
&format!("v.deepstack.{il_us}.fc2.bias"),
0.01,
(10501 + il_us) as f32,
);
}
let n_px = 3 * (pixel_h as usize) * (pixel_w as usize);
let pixel_values: Vec<f32> = (0..n_px)
.map(|i| (((i as f32) * 0.011_7_f32).sin() * 0.5).clamp(-1.0, 1.0))
.collect();
let img = crate::inference::vision::PreprocessedImage {
pixel_values,
target_size: pixel_h,
pixel_w: Some(pixel_w),
pixel_h: Some(pixel_h),
source_label: format!("adr021-iter5b-{pixel_h}x{pixel_w}"),
};
let inputs = vec![crate::inference::vision::vit_gpu::VisionInput::Siglip49(
img,
)];
let result =
compute_vision_embeddings_gpu_qwen3vl(&inputs, &weights, &vit_cfg, &mmproj_cfg)
.expect("ADR-021 iter-5b run must succeed");
assert_eq!(result.len(), 1);
let out = result.into_iter().next().unwrap();
let h = fnv1a64_of_f32_slice(&out);
(out, h)
};
let (square_out, square_hash) = run_at_shape(128, 128);
let (wide_out, wide_hash) = run_at_shape(64, 128);
let (tall_out, tall_hash) = run_at_shape(128, 64);
eprintln!("=== ADR-021 iter-5b AC-3 substitute ===");
eprintln!(
" square 128x128 : len={} fnv1a64=0x{:016x}",
square_out.len(),
square_hash
);
eprintln!(
" wide 64x128 : len={} fnv1a64=0x{:016x}",
wide_out.len(),
wide_hash
);
eprintln!(
" tall 128x64 : len={} fnv1a64=0x{:016x}",
tall_out.len(),
tall_hash
);
for (label, out) in &[
("square", &square_out),
("wide", &wide_out),
("tall", &tall_out),
] {
for (i, &v) in out.iter().enumerate() {
assert!(v.is_finite(), "{label}[{i}] = {v} is not finite");
}
}
assert_ne!(
square_out.len(),
wide_out.len(),
"square (128x128) and wide (64x128) must produce different \
output lengths because n_image_tokens differs"
);
assert_ne!(
square_out.len(),
tall_out.len(),
"square (128x128) and tall (128x64) must produce different \
output lengths because n_image_tokens differs"
);
assert_eq!(
wide_out.len(),
tall_out.len(),
"wide (64x128) and tall (128x64) share n_image_tokens=8 \
post-merge but with different per-patch ordering"
);
assert_ne!(
wide_hash, tall_hash,
"wide (64x128) and tall (128x64) hashes must differ — \
same n_image_tokens but different patch ordering"
);
let (rerun_wide_out, rerun_wide_hash) = run_at_shape(64, 128);
assert_eq!(wide_out.len(), rerun_wide_out.len());
assert_eq!(
wide_hash, rerun_wide_hash,
"wide rerun must be byte-deterministic"
);
for (i, (a, b)) in wide_out.iter().zip(rerun_wide_out.iter()).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"wide rerun byte drift at element {i}: a={a} b={b}"
);
}
}
#[test]
fn adr021_iter5c_ac6_perf_stage_a_gpu_vs_cpu_2026_05_07() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use std::time::Instant;
const PIXEL_H: u32 = 256;
const PIXEL_W: u32 = 256;
const PATCH_SIZE: u32 = 16;
const HIDDEN: u32 = 1280;
const TRAINED_N: u32 = 16;
let n_px = (3 * PIXEL_H * PIXEL_W) as usize;
let pixels: Vec<f32> = (0..n_px)
.map(|i| (((i as f32) * 0.011_7_f32).sin() * 0.5).clamp(-1.0, 1.0))
.collect();
let n_w = (HIDDEN as usize) * 3 * (PATCH_SIZE as usize) * (PATCH_SIZE as usize);
let weight_0: Vec<f32> = (0..n_w)
.map(|i| ((i as f32 + 1.0) * 0.001_3_f32).sin() * 0.05)
.collect();
let weight_1: Vec<f32> = (0..n_w)
.map(|i| ((i as f32 + 313.0) * 0.001_3_f32).sin() * 0.05)
.collect();
let bias: Vec<f32> = (0..(HIDDEN as usize))
.map(|i| ((i as f32 + 7.0) * 0.013_3_f32).sin() * 0.01)
.collect();
let n_pos = (TRAINED_N as usize) * (TRAINED_N as usize);
let pos_embd: Vec<f32> = (0..(n_pos * HIDDEN as usize))
.map(|i| ((i as f32 + 727.0) * 0.001_3_f32).sin() * 0.03)
.collect();
let nps_x = (PIXEL_W / PATCH_SIZE) as usize;
let nps_y = (PIXEL_H / PATCH_SIZE) as usize;
let n_pos_pre = nps_x * nps_y;
let cpu_oracle = || -> Vec<f32> {
let patches_pre = qwen3vl_dual_conv_patch_embed_cpu_hw(
&pixels,
&weight_0,
&weight_1,
Some(&bias),
PIXEL_H,
PIXEL_W,
PATCH_SIZE,
HIDDEN,
)
.expect("CPU dual conv");
let pos_resized = qwen3vl_resize_position_embeddings_bilinear(
&pos_embd,
(TRAINED_N * TRAINED_N) as u32,
HIDDEN,
nps_x as u32,
nps_y as u32,
)
.expect("CPU bilinear resize");
let mut summed = patches_pre;
for (a, b) in summed.iter_mut().zip(pos_resized.iter()) {
*a += *b;
}
qwen3vl_2x2_block_merge_reshape(&summed, nps_x, nps_y, HIDDEN as usize)
.expect("CPU block merge")
};
let gpu_pipeline = || -> Vec<f32> {
use mlx_native::GraphExecutor;
let executor = GraphExecutor::new(MlxDevice::new().expect("device"));
let mut session = executor.begin().expect("begin");
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let mut registry = KernelRegistry::new();
register_im2col_2d_3ch(&mut registry);
register_add_bias_row_2d(&mut registry);
register_bilinear_resize_2d(&mut registry);
register_block_merge_2x2(&mut registry);
let merged_buf = qwen3vl_stage_a_dispatch(
session.encoder_mut(),
&mut registry,
device,
&pixels,
&weight_0,
&weight_1,
Some(&bias),
&pos_embd,
(TRAINED_N * TRAINED_N) as u32,
PIXEL_H,
PIXEL_W,
PATCH_SIZE,
HIDDEN,
)
.expect("GPU stage A");
session.finish().expect("finish");
merged_buf.as_slice::<f32>().expect("readback").to_vec()
};
let _ = cpu_oracle();
let _ = gpu_pipeline();
let mut cpu_times = Vec::with_capacity(5);
let mut gpu_times = Vec::with_capacity(5);
for _ in 0..5 {
let t0 = Instant::now();
let _cpu_out = cpu_oracle();
cpu_times.push(t0.elapsed());
let t0 = Instant::now();
let _gpu_out = gpu_pipeline();
gpu_times.push(t0.elapsed());
}
cpu_times.sort();
gpu_times.sort();
let cpu_min = cpu_times[0];
let gpu_min = gpu_times[0];
let cpu_out = cpu_oracle();
let gpu_out = gpu_pipeline();
assert_eq!(cpu_out.len(), gpu_out.len());
let max_abs_diff = cpu_out
.iter()
.zip(gpu_out.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
let cpu_us = cpu_min.as_secs_f64() * 1e6;
let gpu_us = gpu_min.as_secs_f64() * 1e6;
let speedup = cpu_us / gpu_us;
eprintln!("=== ADR-021 iter-5c AC-6 perf (Stage A, 256x256 / hidden=1280) ===");
eprintln!(" CPU oracle : {:>10.1} µs (best-of-5)", cpu_us);
eprintln!(" GPU pipeline: {:>10.1} µs (best-of-5)", gpu_us);
eprintln!(" speedup : {:>10.2}× (GPU vs CPU)", speedup);
eprintln!(" max abs diff: {:e}", max_abs_diff);
eprintln!(" output len : {}", cpu_out.len());
assert_eq!(
cpu_out.len(),
gpu_out.len(),
"Stage A output length mismatch"
);
assert_eq!(
cpu_out.len(),
n_pos_pre * (HIDDEN as usize),
"Stage A output length should equal n_pos_pre * hidden"
);
assert!(
max_abs_diff < 5e-2,
"Stage A GPU vs CPU max abs diff {} exceeds 5e-2 tolerance",
max_abs_diff
);
assert!(
gpu_us < cpu_us,
"Stage A GPU pipeline ({:.1}µs) is slower than CPU oracle ({:.1}µs)",
gpu_us,
cpu_us
);
}
}