#![allow(dead_code)]
use anyhow::{anyhow, Context, Result};
use mlx_native::metal::MTLSize;
use mlx_native::ops::dense_mm_bf16::{dense_matmul_bf16_f32_tensor, DenseMmBf16F32Params};
use mlx_native::ops::dense_mm_f16::{dense_matmul_f16_f32_tensor, DenseMmF16F32Params};
use mlx_native::ops::dense_mm_f32_f32::{dense_matmul_f32_f32_tensor, DenseMmF32F32Params};
use mlx_native::ops::elementwise::{cast, CastDirection};
use mlx_native::ops::elementwise::{elementwise_add, elementwise_mul, scalar_mul_f32};
use mlx_native::ops::encode_helpers::KernelArg;
use mlx_native::ops::gather::dispatch_gather_f32;
use mlx_native::ops::rms_norm::dispatch_rms_norm;
use mlx_native::ops::sigmoid_mul::dispatch_sigmoid_mul;
use mlx_native::ops::softmax::dispatch_softmax;
use mlx_native::ops::transpose::{permute_021_f32, transpose_last2_f16};
use mlx_native::{CommandEncoder, DType, KernelRegistry, MlxBuffer, MlxDevice};
use std::sync::LazyLock;
pub(crate) static VIT_F32_ATTENTION_ACTIVE: LazyLock<bool> = LazyLock::new(|| {
std::env::var("HF2Q_VIT_F32_ATTENTION")
.map(|v| v == "1")
.unwrap_or(false)
});
pub fn vit_linear_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight: &MlxBuffer,
seq_len: u32,
in_features: u32,
out_features: u32,
) -> Result<MlxBuffer> {
if in_features < 32 {
return Err(anyhow!(
"vit_linear_gpu: in_features ({}) must be >= 32",
in_features
));
}
if seq_len == 0 || out_features == 0 {
return Err(anyhow!(
"vit_linear_gpu: seq_len ({}) and out_features ({}) must be > 0",
seq_len,
out_features
));
}
let metal_dev = device.metal_device();
let out_bytes = (seq_len as usize) * (out_features as usize) * 4;
let mut dst = device
.alloc_buffer(
out_bytes,
DType::F32,
vec![seq_len as usize, out_features as usize],
)
.map_err(|e| anyhow!("alloc output: {e}"))?;
match weight.dtype() {
DType::F16 => {
let params = DenseMmF16F32Params {
m: seq_len,
n: out_features,
k: in_features,
src0_batch: 1,
src1_batch: 1,
};
dense_matmul_f16_f32_tensor(
encoder, registry, device, weight, input, &mut dst, ¶ms,
)
.context("vit_linear_gpu: dense_matmul_f16_f32_tensor")?;
}
DType::BF16 => {
let params = DenseMmBf16F32Params {
m: seq_len,
n: out_features,
k: in_features,
src0_batch: 1,
src1_batch: 1,
};
dense_matmul_bf16_f32_tensor(
encoder, registry, device, weight, input, &mut dst, ¶ms,
)
.context("vit_linear_gpu: dense_matmul_bf16_f32_tensor (BF16-native)")?;
}
DType::F32 => {
let n_w = (out_features as usize) * (in_features as usize);
let weight_bf16 = device
.alloc_buffer(
n_w * 2,
DType::BF16,
vec![out_features as usize, in_features as usize],
)
.map_err(|e| anyhow!("alloc weight_bf16: {e}"))?;
cast(
encoder,
registry,
metal_dev,
weight,
&weight_bf16,
n_w,
CastDirection::F32ToBF16,
)
.context("vit_linear_gpu: F32→BF16 cast (legacy F32 path)")?;
encoder.memory_barrier();
let params = DenseMmBf16F32Params {
m: seq_len,
n: out_features,
k: in_features,
src0_batch: 1,
src1_batch: 1,
};
dense_matmul_bf16_f32_tensor(
encoder,
registry,
device,
&weight_bf16,
input,
&mut dst,
¶ms,
)
.context("vit_linear_gpu: dense_matmul_bf16_f32_tensor (F32→BF16 legacy path)")?;
}
other => {
return Err(anyhow!(
"vit_linear_gpu: unsupported weight dtype {other:?} (expected F16/BF16/F32)"
));
}
}
Ok(dst)
}
pub fn vit_rms_norm_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
gain_f32: &MlxBuffer,
rows: u32,
dim: u32,
eps: f32,
) -> Result<MlxBuffer> {
if rows == 0 || dim == 0 {
return Err(anyhow!(
"vit_rms_norm_gpu: rows ({}) and dim ({}) must be > 0",
rows,
dim
));
}
let out_bytes = (rows as usize) * (dim as usize) * 4;
let output = device
.alloc_buffer(out_bytes, DType::F32, vec![rows as usize, dim as usize])
.map_err(|e| anyhow!("vit_rms_norm_gpu: alloc output: {e}"))?;
let params_buf = device
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow!("vit_rms_norm_gpu: alloc params: {e}"))?;
{
let s: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(params_buf.contents_ptr() as *mut f32, 2) };
s[0] = eps;
s[1] = dim as f32;
}
dispatch_rms_norm(
encoder,
registry,
device.metal_device(),
input,
gain_f32,
&output,
¶ms_buf,
rows,
dim,
)
.context("vit_rms_norm_gpu: dispatch_rms_norm")?;
Ok(output)
}
pub fn vit_per_head_rms_norm_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
gain_f32: &MlxBuffer,
batch: u32,
num_heads: u32,
head_dim: u32,
eps: f32,
) -> Result<MlxBuffer> {
if batch == 0 || num_heads == 0 || head_dim == 0 {
return Err(anyhow!(
"vit_per_head_rms_norm_gpu: batch ({}), num_heads ({}), head_dim ({}) must all be > 0",
batch,
num_heads,
head_dim
));
}
let rows = batch
.checked_mul(num_heads)
.ok_or_else(|| anyhow!("vit_per_head_rms_norm_gpu: batch*num_heads overflow"))?;
vit_rms_norm_gpu(
encoder, registry, device, input, gain_f32, rows, head_dim, eps,
)
}
pub fn vit_softmax_last_dim_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
rows: u32,
cols: u32,
) -> Result<MlxBuffer> {
if rows == 0 || cols == 0 {
return Err(anyhow!(
"vit_softmax_last_dim_gpu: rows ({}) and cols ({}) must be > 0",
rows,
cols
));
}
let out_bytes = (rows as usize) * (cols as usize) * 4;
let output = device
.alloc_buffer(out_bytes, DType::F32, vec![rows as usize, cols as usize])
.map_err(|e| anyhow!("vit_softmax_last_dim_gpu: alloc output: {e}"))?;
let params_buf = device
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow!("vit_softmax_last_dim_gpu: alloc params: {e}"))?;
{
let s: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(params_buf.contents_ptr() as *mut f32, 2) };
s[0] = cols as f32;
s[1] = 0.0;
}
dispatch_softmax(
encoder,
registry,
device.metal_device(),
input,
&output,
¶ms_buf,
rows,
cols,
)
.context("vit_softmax_last_dim_gpu: dispatch_softmax")?;
Ok(output)
}
pub fn vit_attention_scores_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
q_seq_major: &MlxBuffer,
k_seq_major: &MlxBuffer,
batch: u32,
num_heads: u32,
head_dim: u32,
scale: f32,
) -> Result<MlxBuffer> {
if head_dim < 32 {
return Err(anyhow!(
"vit_attention_scores_gpu: head_dim ({}) must be >= 32",
head_dim
));
}
if batch == 0 || num_heads == 0 {
return Err(anyhow!(
"vit_attention_scores_gpu: batch ({}) and num_heads ({}) must be > 0",
batch,
num_heads
));
}
let metal_dev = device.metal_device();
let n_qk_elems = (batch as usize) * (num_heads as usize) * (head_dim as usize);
let q_perm = device
.alloc_buffer(
n_qk_elems * 4,
DType::F32,
vec![num_heads as usize, batch as usize, head_dim as usize],
)
.map_err(|e| anyhow!("alloc q_perm: {e}"))?;
let k_perm = device
.alloc_buffer(
n_qk_elems * 4,
DType::F32,
vec![num_heads as usize, batch as usize, head_dim as usize],
)
.map_err(|e| anyhow!("alloc k_perm: {e}"))?;
permute_021_f32(
encoder,
registry,
metal_dev,
q_seq_major,
&q_perm,
batch as usize,
num_heads as usize,
head_dim as usize,
)
.context("permute Q seq→head major")?;
permute_021_f32(
encoder,
registry,
metal_dev,
k_seq_major,
&k_perm,
batch as usize,
num_heads as usize,
head_dim as usize,
)
.context("permute K seq→head major")?;
encoder.memory_barrier();
let n_scores = (num_heads as usize) * (batch as usize) * (batch as usize);
let mut scores = device
.alloc_buffer(
n_scores * 4,
DType::F32,
vec![num_heads as usize, batch as usize, batch as usize],
)
.map_err(|e| anyhow!("alloc scores: {e}"))?;
if *VIT_F32_ATTENTION_ACTIVE {
let params = DenseMmF32F32Params {
m: batch,
n: batch,
k: head_dim,
src0_batch: num_heads,
src1_batch: num_heads,
};
dense_matmul_f32_f32_tensor(
encoder,
registry,
device,
&k_perm,
&q_perm,
&mut scores,
¶ms,
)
.context("attention scores matmul (F32 debug override, HF2Q_VIT_F32_ATTENTION=1)")?;
} else {
let k_f16 = device
.alloc_buffer(
n_qk_elems * 2,
DType::F16,
vec![num_heads as usize, batch as usize, head_dim as usize],
)
.map_err(|e| anyhow!("alloc k_f16: {e}"))?;
cast(
encoder,
registry,
metal_dev,
&k_perm,
&k_f16,
n_qk_elems,
CastDirection::F32ToF16,
)
.context("cast K F32→F16")?;
encoder.memory_barrier();
let params = DenseMmF16F32Params {
m: batch,
n: batch,
k: head_dim,
src0_batch: num_heads,
src1_batch: num_heads,
};
dense_matmul_f16_f32_tensor(
encoder,
registry,
device,
&k_f16,
&q_perm,
&mut scores,
¶ms,
)
.context("attention scores matmul (F16, peer parity)")?;
}
if scale != 1.0 {
encoder.memory_barrier();
scalar_mul_f32(
encoder, registry, metal_dev, &scores, &scores, n_scores, scale,
)
.context("scale scores")?;
}
Ok(scores)
}
pub fn vit_attention_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
q_seq_major: &MlxBuffer,
k_seq_major: &MlxBuffer,
v_seq_major: &MlxBuffer,
batch: u32,
num_heads: u32,
head_dim: u32,
scale: f32,
) -> Result<MlxBuffer> {
if head_dim < 32 {
return Err(anyhow!(
"vit_attention_gpu: head_dim ({}) must be >= 32",
head_dim
));
}
if batch == 0 || num_heads == 0 {
return Err(anyhow!(
"vit_attention_gpu: batch ({}) and num_heads ({}) must be > 0",
batch,
num_heads
));
}
let metal_dev = device.metal_device();
let scores = vit_attention_scores_gpu(
encoder,
registry,
device,
q_seq_major,
k_seq_major,
batch,
num_heads,
head_dim,
scale,
)?;
encoder.memory_barrier();
let n_rows = (num_heads as u64) * (batch as u64);
let softmaxed =
vit_softmax_last_dim_gpu(encoder, registry, device, &scores, n_rows as u32, batch)?;
encoder.memory_barrier();
let n_v = (batch as usize) * (num_heads as usize) * (head_dim as usize);
let v_perm = device
.alloc_buffer(
n_v * 4,
DType::F32,
vec![num_heads as usize, batch as usize, head_dim as usize],
)
.map_err(|e| anyhow!("alloc v_perm: {e}"))?;
permute_021_f32(
encoder,
registry,
metal_dev,
v_seq_major,
&v_perm,
batch as usize,
num_heads as usize,
head_dim as usize,
)
.context("permute V seq→head major")?;
encoder.memory_barrier();
let v_f16 = device
.alloc_buffer(
n_v * 2,
DType::F16,
vec![num_heads as usize, batch as usize, head_dim as usize],
)
.map_err(|e| anyhow!("alloc v_f16: {e}"))?;
cast(
encoder,
registry,
metal_dev,
&v_perm,
&v_f16,
n_v,
CastDirection::F32ToF16,
)
.context("cast V f32→f16")?;
encoder.memory_barrier();
let v_t_f16 = device
.alloc_buffer(
n_v * 2,
DType::F16,
vec![num_heads as usize, head_dim as usize, batch as usize],
)
.map_err(|e| anyhow!("alloc v_t_f16: {e}"))?;
transpose_last2_f16(
encoder,
registry,
metal_dev,
&v_f16,
&v_t_f16,
num_heads as usize,
batch as usize,
head_dim as usize,
)
.context("transpose V last 2 (f16)")?;
encoder.memory_barrier();
let n_attn = (num_heads as usize) * (batch as usize) * (head_dim as usize);
let mut attn_head_major = device
.alloc_buffer(
n_attn * 4,
DType::F32,
vec![num_heads as usize, batch as usize, head_dim as usize],
)
.map_err(|e| anyhow!("alloc attn_head_major: {e}"))?;
let params = DenseMmF16F32Params {
m: batch, n: head_dim, k: batch, src0_batch: num_heads,
src1_batch: num_heads,
};
dense_matmul_f16_f32_tensor(
encoder,
registry,
device,
&v_t_f16,
&softmaxed,
&mut attn_head_major,
¶ms,
)
.context("attention scores @ V matmul (f16)")?;
encoder.memory_barrier();
let attn_seq_major = device
.alloc_buffer(
n_attn * 4,
DType::F32,
vec![batch as usize, num_heads as usize, head_dim as usize],
)
.map_err(|e| anyhow!("alloc attn_seq_major: {e}"))?;
permute_021_f32(
encoder,
registry,
metal_dev,
&attn_head_major,
&attn_seq_major,
num_heads as usize,
batch as usize,
head_dim as usize,
)
.context("permute attn head→seq major")?;
Ok(attn_seq_major)
}
pub fn vit_residual_add_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
a: &MlxBuffer,
b: &MlxBuffer,
n_elements: u32,
) -> Result<MlxBuffer> {
if n_elements == 0 {
return Err(anyhow!("vit_residual_add_gpu: n_elements must be > 0"));
}
let out_shape: Vec<usize> = if a.element_count() == n_elements as usize {
a.shape().to_vec()
} else {
vec![n_elements as usize]
};
let out = device
.alloc_buffer((n_elements as usize) * 4, DType::F32, out_shape)
.map_err(|e| anyhow!("alloc residual add output: {e}"))?;
elementwise_add(
encoder,
registry,
device.metal_device(),
a,
b,
&out,
n_elements as usize,
DType::F32,
)
.context("vit_residual_add_gpu: elementwise_add")?;
Ok(out)
}
pub fn vit_silu_mul_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_silu_mul_gpu: n_elements must be > 0"));
}
let metal_dev = device.metal_device();
let silu_out = device
.alloc_buffer(
(n_elements as usize) * 4,
DType::F32,
vec![n_elements as usize],
)
.map_err(|e| anyhow!("alloc silu_out: {e}"))?;
let params_buf = device
.alloc_buffer(4, DType::F32, vec![1])
.map_err(|e| anyhow!("alloc sigmoid_mul params: {e}"))?;
{
let s: &mut [u32] =
unsafe { std::slice::from_raw_parts_mut(params_buf.contents_ptr() as *mut u32, 1) };
s[0] = n_elements;
}
dispatch_sigmoid_mul(
encoder,
registry,
metal_dev,
gate,
gate,
&silu_out,
¶ms_buf,
n_elements,
)
.context("vit_silu_mul_gpu: sigmoid_mul (silu step)")?;
encoder.memory_barrier();
let out = device
.alloc_buffer(
(n_elements as usize) * 4,
DType::F32,
vec![n_elements as usize],
)
.map_err(|e| anyhow!("alloc final out: {e}"))?;
elementwise_mul(
encoder,
registry,
metal_dev,
&silu_out,
up,
&out,
n_elements as usize,
DType::F32,
)
.context("vit_silu_mul_gpu: elementwise_mul")?;
Ok(out)
}
pub fn apply_vit_block_forward_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
weights: &super::mmproj_weights::LoadedMmprojWeights,
cfg: &super::mmproj::MmprojConfig,
block_idx: usize,
input: &MlxBuffer,
batch: u32,
scale: f32,
) -> Result<MlxBuffer> {
let hidden = cfg.hidden_size;
let num_heads = cfg.num_attention_heads;
let head_dim = hidden / num_heads;
let intermediate = cfg.intermediate_size;
let eps = cfg.layer_norm_eps;
let n_hidden = (batch as usize) * (hidden as usize);
let block = |suffix: &str| -> Result<&MlxBuffer> {
weights
.block_tensor(block_idx, suffix)
.map_err(|e| anyhow!("block {} {}: {e}", block_idx, suffix))
};
let cur = vit_rms_norm_gpu(
encoder,
registry,
device,
input,
block("ln1.weight")?,
batch,
hidden,
eps,
)?;
encoder.memory_barrier();
let q = vit_linear_gpu(
encoder,
registry,
device,
&cur,
block("attn_q.weight")?,
batch,
hidden,
hidden,
)?;
encoder.memory_barrier();
let k = vit_linear_gpu(
encoder,
registry,
device,
&cur,
block("attn_k.weight")?,
batch,
hidden,
hidden,
)?;
encoder.memory_barrier();
let v = vit_linear_gpu(
encoder,
registry,
device,
&cur,
block("attn_v.weight")?,
batch,
hidden,
hidden,
)?;
encoder.memory_barrier();
let q_norm = vit_per_head_rms_norm_gpu(
encoder,
registry,
device,
&q,
block("attn_q_norm.weight")?,
batch,
num_heads,
head_dim,
eps,
)?;
encoder.memory_barrier();
let k_norm = vit_per_head_rms_norm_gpu(
encoder,
registry,
device,
&k,
block("attn_k_norm.weight")?,
batch,
num_heads,
head_dim,
eps,
)?;
encoder.memory_barrier();
let attn = vit_attention_gpu(
encoder, registry, device, &q_norm, &k_norm, &v, batch, num_heads, head_dim, scale,
)?;
encoder.memory_barrier();
let attn_proj = vit_linear_gpu(
encoder,
registry,
device,
&attn,
block("attn_output.weight")?,
batch,
hidden,
hidden,
)?;
encoder.memory_barrier();
let post_attn = vit_residual_add_gpu(
encoder,
registry,
device,
input,
&attn_proj,
n_hidden as u32,
)?;
encoder.memory_barrier();
let pre_ffn = vit_rms_norm_gpu(
encoder,
registry,
device,
&post_attn,
block("ln2.weight")?,
batch,
hidden,
eps,
)?;
encoder.memory_barrier();
let gate = vit_linear_gpu(
encoder,
registry,
device,
&pre_ffn,
block("ffn_gate.weight")?,
batch,
hidden,
intermediate,
)?;
encoder.memory_barrier();
let up = vit_linear_gpu(
encoder,
registry,
device,
&pre_ffn,
block("ffn_up.weight")?,
batch,
hidden,
intermediate,
)?;
encoder.memory_barrier();
let activated = vit_silu_mul_gpu(
encoder,
registry,
device,
&gate,
&up,
(batch as usize * intermediate as usize) as u32,
)?;
encoder.memory_barrier();
let down = vit_linear_gpu(
encoder,
registry,
device,
&activated,
block("ffn_down.weight")?,
batch,
intermediate,
hidden,
)?;
encoder.memory_barrier();
let down_normed = vit_rms_norm_gpu(
encoder,
registry,
device,
&down,
block("post_ffw_norm.weight")?,
batch,
hidden,
eps,
)?;
encoder.memory_barrier();
let block_out = vit_residual_add_gpu(
encoder,
registry,
device,
&post_attn,
&down_normed,
n_hidden as u32,
)?;
Ok(block_out)
}
const VIT_CUSTOM_SHADERS_SOURCE: &str = r#"
#include <metal_stdlib>
using namespace metal;
struct AvgPool2x2Params {
uint n_side;
uint hidden;
};
kernel void vit_avg_pool_2x2_f32(
device const float* input [[buffer(0)]],
device float* output [[buffer(1)]],
constant AvgPool2x2Params& params [[buffer(2)]],
uint3 gid [[thread_position_in_grid]]
) {
uint h = gid.x;
uint ox = gid.y;
uint oy = gid.z;
uint out_side = params.n_side / 2;
if (h >= params.hidden || ox >= out_side || oy >= out_side) return;
uint iy = oy * 2u;
uint ix = ox * 2u;
uint h_stride = params.hidden;
uint row_stride = params.n_side * h_stride;
float a = input[iy * row_stride + ix * h_stride + h];
float b = input[iy * row_stride + (ix + 1u) * h_stride + h];
float c = input[(iy + 1u) * row_stride + ix * h_stride + h];
float d = input[(iy + 1u) * row_stride + (ix + 1u) * h_stride + h];
output[oy * out_side * h_stride + ox * h_stride + h] = (a + b + c + d) * 0.25;
}
// Parameterized k×k spatial avg-pool on a rectangular [n_y, n_x, hidden]
// grid. Output shape is [n_y/k, n_x/k, hidden]. Generalizes
// vit_avg_pool_2x2_f32 (which is the n_x=n_y=n_side, k=2 case).
//
// gemma4v call: n_x = n_patches_x, n_y = n_patches_y, k = n_merge = 3
// (per /opt/llama.cpp/tools/mtmd/clip.cpp:1337).
//
// Layout assumption: input is row-major with rows iterating Y first
// (so row stride = n_x * hidden), matching the
// [n_y, n_x, hidden] reshape in the gemma4v post-blocks pipeline.
struct AvgPoolKxKParams {
uint n_x; // input width in patches
uint n_y; // input height in patches
uint k; // pool kernel edge (= stride; non-overlapping)
uint hidden; // channel dim
};
kernel void vit_avg_pool_kxk_f32(
device const float* input [[buffer(0)]],
device float* output [[buffer(1)]],
constant AvgPoolKxKParams& params [[buffer(2)]],
uint3 gid [[thread_position_in_grid]]
) {
uint h = gid.x;
uint ox = gid.y;
uint oy = gid.z;
uint out_x = params.n_x / params.k;
uint out_y = params.n_y / params.k;
if (h >= params.hidden || ox >= out_x || oy >= out_y) return;
uint iy0 = oy * params.k;
uint ix0 = ox * params.k;
uint h_stride = params.hidden;
uint row_stride = params.n_x * h_stride;
float acc = 0.0f;
// Sum over the k×k input block. k is small (== 3 for gemma4v) so
// the inner double-loop is a handful of ops; the compiler will
// unroll once k becomes a constant via specialization. We keep it
// dynamic so the same kernel handles k=2 (SigLIP path), k=3
// (gemma4v), and any future arch's pool factor.
for (uint dy = 0u; dy < params.k; ++dy) {
for (uint dx = 0u; dx < params.k; ++dx) {
uint ix = ix0 + dx;
uint iy = iy0 + dy;
acc += input[iy * row_stride + ix * h_stride + h];
}
}
float k2 = float(params.k * params.k);
output[oy * out_x * h_stride + ox * h_stride + h] = acc / k2;
}
struct StdBiasScaleParams {
uint hidden;
uint batch;
};
kernel void vit_std_bias_scale_f32(
device const float* input [[buffer(0)]],
device const float* bias [[buffer(1)]],
device const float* scale [[buffer(2)]],
device float* output [[buffer(3)]],
constant StdBiasScaleParams& params [[buffer(4)]],
uint2 gid [[thread_position_in_grid]]
) {
uint h = gid.x;
uint b = gid.y;
if (h >= params.hidden || b >= params.batch) return;
uint idx = b * params.hidden + h;
output[idx] = (input[idx] - bias[h]) * scale[h];
}
// Elementwise in-place scalar clamp. Both min and max are scalar f32
// values supplied via a 2-element params buffer. The convention:
// out[i] = clamp(in[i], min, max)
// Setting min = -FLT_MAX or max = FLT_MAX makes that side a no-op,
// matching llama.cpp's get_scalar default for the gemma4v
// Gemma4ClippableLinear (`tools/mtmd/clip.cpp:1953-1956`).
//
// Caller dispatches once per clamp (input or output side); we don't
// fuse the matmul-and-clamps because the Gemma4ClippableLinear
// composes from existing primitives (clip + matmul + clip) and the
// clamp itself is bandwidth-bound (one pass through the output).
struct ClipInplaceParams {
float min_val;
float max_val;
uint n_elements;
uint _pad; // align to 16 bytes for Metal
};
kernel void vit_clip_inplace_f32(
device const float* input [[buffer(0)]],
device float* output [[buffer(1)]],
constant ClipInplaceParams& params [[buffer(2)]],
uint gid [[thread_position_in_grid]]
) {
if (gid >= params.n_elements) return;
float x = input[gid];
x = max(x, params.min_val);
x = min(x, params.max_val);
output[gid] = x;
}
"#;
#[repr(C)]
#[derive(Clone, Copy)]
struct AvgPool2x2GpuParams {
n_side: u32,
hidden: u32,
}
#[repr(C)]
#[derive(Clone, Copy)]
struct AvgPoolKxKGpuParams {
n_x: u32,
n_y: u32,
k: u32,
hidden: u32,
}
#[repr(C)]
#[derive(Clone, Copy)]
struct StdBiasScaleGpuParams {
hidden: u32,
batch: u32,
}
#[repr(C)]
#[derive(Clone, Copy)]
struct ClipInplaceGpuParams {
min_val: f32,
max_val: f32,
n_elements: u32,
_pad: u32,
}
fn pod_as_bytes<T: Copy>(p: &T) -> &[u8] {
unsafe { std::slice::from_raw_parts(p as *const T as *const u8, std::mem::size_of::<T>()) }
}
pub fn register_vit_custom_shaders(registry: &mut KernelRegistry) {
registry.register_source("vit_avg_pool_2x2_f32", VIT_CUSTOM_SHADERS_SOURCE);
registry.register_source("vit_avg_pool_kxk_f32", VIT_CUSTOM_SHADERS_SOURCE);
registry.register_source("vit_std_bias_scale_f32", VIT_CUSTOM_SHADERS_SOURCE);
registry.register_source("vit_clip_inplace_f32", VIT_CUSTOM_SHADERS_SOURCE);
}
pub fn vit_avg_pool_2x2_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
n_side: u32,
hidden: u32,
) -> Result<MlxBuffer> {
if n_side == 0 || n_side % 2 != 0 {
return Err(anyhow!(
"vit_avg_pool_2x2_gpu: n_side ({}) must be positive and even",
n_side
));
}
if hidden == 0 {
return Err(anyhow!("vit_avg_pool_2x2_gpu: hidden must be > 0"));
}
let out_side = n_side / 2;
let out_n_patches = (out_side as usize) * (out_side as usize);
let out_bytes = out_n_patches * (hidden as usize) * 4;
let output = device
.alloc_buffer(
out_bytes,
DType::F32,
vec![out_side as usize, out_side as usize, hidden as usize],
)
.map_err(|e| anyhow!("alloc avg_pool output: {e}"))?;
let pipeline = registry
.get_pipeline("vit_avg_pool_2x2_f32", device.metal_device())
.map_err(|e| anyhow!("vit_avg_pool_2x2_gpu: get_pipeline: {e}"))?;
let params = AvgPool2x2GpuParams { n_side, hidden };
let bytes = pod_as_bytes(¶ms);
let grid = MTLSize::new(hidden as u64, out_side as u64, out_side as u64);
let tg_x = std::cmp::min(64, hidden as u64);
let tg = MTLSize::new(tg_x, 1, 1);
encoder.encode_with_args(
pipeline,
&[
(0, KernelArg::Buffer(input)),
(1, KernelArg::Buffer(&output)),
(2, KernelArg::Bytes(bytes)),
],
grid,
tg,
);
Ok(output)
}
#[allow(clippy::too_many_arguments)]
pub fn vit_avg_pool_kxk_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
n_x: u32,
n_y: u32,
k: u32,
hidden: u32,
) -> Result<MlxBuffer> {
if n_x == 0 || n_y == 0 || k == 0 {
return Err(anyhow!(
"vit_avg_pool_kxk_gpu: n_x ({n_x}), n_y ({n_y}), k ({k}) must all be > 0"
));
}
if hidden == 0 {
return Err(anyhow!("vit_avg_pool_kxk_gpu: hidden must be > 0"));
}
if n_x % k != 0 || n_y % k != 0 {
return Err(anyhow!(
"vit_avg_pool_kxk_gpu: n_x ({n_x}) and n_y ({n_y}) must both be multiples of k ({k})"
));
}
let out_x = n_x / k;
let out_y = n_y / k;
let out_n = (out_x as usize) * (out_y as usize);
let out_bytes = out_n * (hidden as usize) * 4;
let output = device
.alloc_buffer(
out_bytes,
DType::F32,
vec![out_y as usize, out_x as usize, hidden as usize],
)
.map_err(|e| anyhow!("alloc avg_pool_kxk output: {e}"))?;
let pipeline = registry
.get_pipeline("vit_avg_pool_kxk_f32", device.metal_device())
.map_err(|e| anyhow!("vit_avg_pool_kxk_gpu: get_pipeline: {e}"))?;
let params = AvgPoolKxKGpuParams {
n_x,
n_y,
k,
hidden,
};
let bytes = pod_as_bytes(¶ms);
let grid = MTLSize::new(hidden as u64, out_x as u64, out_y as u64);
let tg_x = std::cmp::min(64, hidden as u64);
let tg = MTLSize::new(tg_x, 1, 1);
encoder.encode_with_args(
pipeline,
&[
(0, KernelArg::Buffer(input)),
(1, KernelArg::Buffer(&output)),
(2, KernelArg::Bytes(bytes)),
],
grid,
tg,
);
Ok(output)
}
pub fn gemma4v_avg_pool_3x3_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
n_x: u32,
n_y: u32,
hidden: u32,
) -> Result<MlxBuffer> {
vit_avg_pool_kxk_gpu(encoder, registry, device, input, n_x, n_y, 3, hidden)
}
pub fn vit_clip_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
n_elements: u32,
min_val: f32,
max_val: f32,
) -> Result<MlxBuffer> {
if n_elements == 0 {
return Err(anyhow!("vit_clip_gpu: n_elements must be > 0"));
}
if min_val.is_nan() || max_val.is_nan() {
return Err(anyhow!(
"vit_clip_gpu: NaN clamp bounds (min={min_val}, max={max_val}) — \
would silently corrupt all elements"
));
}
if min_val > max_val {
return Err(anyhow!(
"vit_clip_gpu: min_val ({min_val}) > max_val ({max_val}) is undefined"
));
}
let out_bytes = (n_elements as usize) * 4;
let output = device
.alloc_buffer(out_bytes, DType::F32, vec![n_elements as usize])
.map_err(|e| anyhow!("alloc clip output: {e}"))?;
let pipeline = registry
.get_pipeline("vit_clip_inplace_f32", device.metal_device())
.map_err(|e| anyhow!("vit_clip_gpu: get_pipeline: {e}"))?;
let params = ClipInplaceGpuParams {
min_val,
max_val,
n_elements,
_pad: 0,
};
let bytes = pod_as_bytes(¶ms);
let grid = MTLSize::new(n_elements as u64, 1, 1);
let tg_x = std::cmp::min(256, n_elements as u64);
let tg = MTLSize::new(tg_x, 1, 1);
encoder.encode_with_args(
pipeline,
&[
(0, KernelArg::Buffer(input)),
(1, KernelArg::Buffer(&output)),
(2, KernelArg::Bytes(bytes)),
],
grid,
tg,
);
Ok(output)
}
pub fn vit_std_bias_scale_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
bias: &MlxBuffer,
scale: &MlxBuffer,
batch: u32,
hidden: u32,
) -> Result<MlxBuffer> {
if batch == 0 || hidden == 0 {
return Err(anyhow!(
"vit_std_bias_scale_gpu: batch ({}) and hidden ({}) must be > 0",
batch,
hidden
));
}
let total = (batch as usize) * (hidden as usize);
let output = device
.alloc_buffer(total * 4, DType::F32, vec![batch as usize, hidden as usize])
.map_err(|e| anyhow!("alloc std_bias_scale output: {e}"))?;
let pipeline = registry
.get_pipeline("vit_std_bias_scale_f32", device.metal_device())
.map_err(|e| anyhow!("vit_std_bias_scale_gpu: get_pipeline: {e}"))?;
let params = StdBiasScaleGpuParams { hidden, batch };
let bytes = pod_as_bytes(¶ms);
let grid = MTLSize::new(hidden as u64, batch as u64, 1);
let tg = MTLSize::new(std::cmp::min(64, hidden as u64), 1, 1);
encoder.encode_with_args(
pipeline,
&[
(0, KernelArg::Buffer(input)),
(1, KernelArg::Buffer(bias)),
(2, KernelArg::Buffer(scale)),
(3, KernelArg::Buffer(&output)),
(4, KernelArg::Bytes(bytes)),
],
grid,
tg,
);
Ok(output)
}
pub fn vit_scale_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
buf: &MlxBuffer,
n_elements: u32,
scalar: f32,
) -> Result<()> {
if n_elements == 0 {
return Err(anyhow!("vit_scale_gpu: n_elements must be > 0"));
}
scalar_mul_f32(
encoder,
registry,
device.metal_device(),
buf,
buf,
n_elements as usize,
scalar,
)
.context("vit_scale_gpu: scalar_mul_f32")
}
pub fn apply_vit_blocks_loop_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
weights: &super::mmproj_weights::LoadedMmprojWeights,
cfg: &super::mmproj::MmprojConfig,
input: &MlxBuffer,
batch: u32,
scale: f32,
) -> Result<MlxBuffer> {
let mut hidden_states = apply_vit_block_forward_gpu(
encoder, registry, device, weights, cfg, 0, input, batch, scale,
)?;
encoder.memory_barrier();
for block_idx in 1..(cfg.num_hidden_layers as usize) {
hidden_states = apply_vit_block_forward_gpu(
encoder,
registry,
device,
weights,
cfg,
block_idx,
&hidden_states,
batch,
scale,
)?;
encoder.memory_barrier();
}
Ok(hidden_states)
}
pub fn apply_vit_full_forward_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
weights: &super::mmproj_weights::LoadedMmprojWeights,
cfg: &super::mmproj::MmprojConfig,
pixel_values: &[f32],
scale: f32,
) -> Result<MlxBuffer> {
use super::vit::patch_embed_forward as patch_embed_cpu;
let hidden = cfg.hidden_size;
let num_patches_side = cfg.num_patches_side;
let n_patches = (num_patches_side as u32) * (num_patches_side as u32);
let patch_embd_buf = weights
.patch_embd_weight()
.map_err(|e| anyhow!("apply_vit_full_forward_gpu: {e}"))?;
let patch_embd_f32_owned = weights
.tensor_as_f32_owned(patch_embd_buf)
.context("apply_vit_full_forward_gpu: patch_embd → f32 widen")?;
let patch_bias_f32_owned: Option<Vec<f32>> = weights
.get("v.patch_embd.bias")
.and_then(|b| weights.tensor_as_f32_owned(b).ok());
let patch_embeds_cpu = patch_embed_cpu(
pixel_values,
&patch_embd_f32_owned,
patch_bias_f32_owned.as_deref(),
cfg.image_size,
cfg.patch_size,
hidden,
)
.context("apply_vit_full_forward_gpu: cpu patch_embed")?;
let n_hidden = (n_patches as usize) * (hidden as usize);
let input_gpu = device
.alloc_buffer(
n_hidden * 4,
DType::F32,
vec![n_patches as usize, hidden as usize],
)
.map_err(|e| anyhow!("alloc input_gpu: {e}"))?;
{
let dst: &mut [f32] = unsafe {
std::slice::from_raw_parts_mut(input_gpu.contents_ptr() as *mut f32, n_hidden)
};
dst.copy_from_slice(&patch_embeds_cpu);
}
let after_blocks = apply_vit_blocks_loop_gpu(
encoder, registry, device, weights, cfg, &input_gpu, n_patches, scale,
)?;
encoder.memory_barrier();
let pooled = vit_avg_pool_2x2_gpu(
encoder,
registry,
device,
&after_blocks,
num_patches_side as u32,
hidden,
)?;
encoder.memory_barrier();
let pooled_n_patches = ((num_patches_side / 2) * (num_patches_side / 2)) as usize;
let pooled_total = pooled_n_patches * (hidden as usize);
vit_scale_gpu(
encoder,
registry,
device,
&pooled,
pooled_total as u32,
(hidden as f32).sqrt(),
)?;
encoder.memory_barrier();
let std_bias = weights
.get("v.std_bias")
.ok_or_else(|| anyhow!("apply_vit_full_forward_gpu: missing v.std_bias"))?;
let std_scale = weights
.get("v.std_scale")
.ok_or_else(|| anyhow!("apply_vit_full_forward_gpu: missing v.std_scale"))?;
let normed = vit_std_bias_scale_gpu(
encoder,
registry,
device,
&pooled,
std_bias,
std_scale,
pooled_n_patches as u32,
hidden,
)?;
encoder.memory_barrier();
let mm0 = weights
.mm_0_weight()
.map_err(|e| anyhow!("apply_vit_full_forward_gpu: mm.0.weight: {e}"))?;
let text_hidden = (mm0.element_count() / (hidden as usize)) as u32;
let projected = vit_linear_gpu(
encoder,
registry,
device,
&normed,
mm0,
pooled_n_patches as u32,
hidden,
text_hidden,
)?;
encoder.memory_barrier();
let ones = device
.alloc_buffer(
(text_hidden as usize) * 4,
DType::F32,
vec![text_hidden as usize],
)
.map_err(|e| anyhow!("alloc ones: {e}"))?;
{
let s: &mut [f32] = unsafe {
std::slice::from_raw_parts_mut(ones.contents_ptr() as *mut f32, text_hidden as usize)
};
for v in s.iter_mut() {
*v = 1.0;
}
}
let final_out = vit_rms_norm_gpu(
encoder,
registry,
device,
&projected,
&ones,
pooled_n_patches as u32,
text_hidden,
cfg.layer_norm_eps,
)?;
Ok(final_out)
}
pub fn warmup_vit_gpu(
weights: &super::mmproj_weights::LoadedMmprojWeights,
cfg: &super::mmproj::MmprojConfig,
) -> Result<()> {
use mlx_native::{GraphExecutor, MlxDevice};
let executor =
GraphExecutor::new(MlxDevice::new().map_err(|e| anyhow!("warmup_vit_gpu device: {e}"))?);
let mut session = executor
.begin()
.map_err(|e| anyhow!("warmup_vit_gpu begin: {e}"))?;
let mut registry = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
mlx_native::ops::sigmoid_mul::register(&mut registry);
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let img = cfg.image_size as usize;
let pixels = vec![0.01f32; 3 * img * img];
let head_dim_f = (cfg.hidden_size / cfg.num_attention_heads) as f32;
let scale = 1.0f32 / head_dim_f.sqrt();
let _output = apply_vit_full_forward_gpu(
session.encoder_mut(),
&mut registry,
device,
weights,
cfg,
&pixels,
scale,
)
.map_err(|e| anyhow!("warmup_vit_gpu forward: {e}"))?;
session
.finish()
.map_err(|e| anyhow!("warmup_vit_gpu finish: {e}"))?;
Ok(())
}
pub fn compute_vision_embeddings_gpu(
images: &[super::PreprocessedImage],
mmproj_weights: &super::mmproj_weights::LoadedMmprojWeights,
mmproj_cfg: &super::mmproj::MmprojConfig,
scale: f32,
) -> Result<Vec<Vec<f32>>> {
use mlx_native::{GraphExecutor, MlxDevice};
let mut out = Vec::with_capacity(images.len());
for (idx, img) in images.iter().enumerate() {
let executor =
GraphExecutor::new(MlxDevice::new().map_err(|e| {
anyhow!("compute_vision_embeddings_gpu image {}: device: {e}", idx)
})?);
let mut session = executor
.begin()
.map_err(|e| anyhow!("compute_vision_embeddings_gpu image {}: begin: {e}", idx))?;
let mut registry = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
mlx_native::ops::sigmoid_mul::register(&mut registry);
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let buf = apply_vit_full_forward_gpu(
session.encoder_mut(),
&mut registry,
device,
mmproj_weights,
mmproj_cfg,
&img.pixel_values,
scale,
)
.map_err(|e| anyhow!("compute_vision_embeddings_gpu image {}: forward: {e}", idx))?;
session
.finish()
.map_err(|e| anyhow!("compute_vision_embeddings_gpu image {}: finish: {e}", idx))?;
let n_patches_out =
((mmproj_cfg.num_patches_side / 2) * (mmproj_cfg.num_patches_side / 2)) as usize;
let mm0 = mmproj_weights
.mm_0_weight()
.map_err(|e| anyhow!("mm.0: {e}"))?;
let text_hidden = mm0.element_count() / (mmproj_cfg.hidden_size as usize);
let total = n_patches_out * text_hidden;
let slice: &[f32] = buf
.as_slice::<f32>()
.map_err(|e| anyhow!("readback: {e}"))?;
if slice.len() != total {
return Err(anyhow!(
"compute_vision_embeddings_gpu image {}: readback len {} != expected {}",
idx,
slice.len(),
total
));
}
out.push(slice.to_vec());
}
Ok(out)
}
#[derive(Debug, Clone)]
pub struct Gemma4vPreprocessedImage {
pub patches: Vec<f32>,
pub pos_x: Vec<u32>,
pub pos_y: Vec<u32>,
pub n_x: u32,
pub n_y: u32,
pub source_label: String,
}
pub fn compute_vision_embeddings_gpu_gemma4v(
images: &[Gemma4vPreprocessedImage],
mmproj_weights: &super::mmproj_weights::LoadedMmprojWeights,
mmproj_cfg: &super::mmproj::MmprojConfig,
) -> Result<Vec<Vec<f32>>> {
use mlx_native::{GraphExecutor, MlxDevice};
let dump_dir = super::vit_dump::resolve_dump_dir()?;
let mut out = Vec::with_capacity(images.len());
for (idx, img) in images.iter().enumerate() {
let executor = GraphExecutor::new(MlxDevice::new().map_err(|e| {
anyhow!(
"compute_vision_embeddings_gpu_gemma4v image {}: device: {e}",
idx
)
})?);
let mut session = executor.begin().map_err(|e| {
anyhow!(
"compute_vision_embeddings_gpu_gemma4v image {}: begin: {e}",
idx
)
})?;
let mut registry = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
mlx_native::ops::sigmoid_mul::register(&mut registry);
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let (buf, collected) = if dump_dir.is_some() {
super::vit_dump::with_dump_collector(|| {
gemma4v_apply_full_forward_gpu(
session.encoder_mut(),
&mut registry,
device,
mmproj_weights,
mmproj_cfg,
&img.patches,
&img.pos_x,
&img.pos_y,
img.n_x,
img.n_y,
)
.map_err(|e| {
anyhow!(
"compute_vision_embeddings_gpu_gemma4v image {}: forward: {e}",
idx
)
})
})?
} else {
let b = gemma4v_apply_full_forward_gpu(
session.encoder_mut(),
&mut registry,
device,
mmproj_weights,
mmproj_cfg,
&img.patches,
&img.pos_x,
&img.pos_y,
img.n_x,
img.n_y,
)
.map_err(|e| {
anyhow!(
"compute_vision_embeddings_gpu_gemma4v image {}: forward: {e}",
idx
)
})?;
(b, Vec::new())
};
session.finish().map_err(|e| {
anyhow!(
"compute_vision_embeddings_gpu_gemma4v image {}: finish: {e}",
idx
)
})?;
if let Some(ref dir) = dump_dir {
let img_dir = if images.len() > 1 {
dir.join(format!("image_{}", idx))
} else {
dir.clone()
};
if !img_dir.exists() {
std::fs::create_dir_all(&img_dir)
.map_err(|e| anyhow!("create dump subdir {}: {e}", img_dir.display()))?;
}
for mirror in super::vit_dump::drain_cpu_mirrors() {
super::vit_dump::write_dump_cpu(&img_dir, &mirror)
.map_err(|e| anyhow!("write CPU dump {}: {e}", mirror.name))?;
}
for (name, buffer) in &collected {
super::vit_dump::write_dump_gpu(&img_dir, name, buffer)
.map_err(|e| anyhow!("write GPU dump {}: {e}", name))?;
}
let audit_entries = super::vit_dump::drain_audit_entries();
if !audit_entries.is_empty() {
super::vit_dump::write_dtype_audit(&img_dir, &audit_entries)
.map_err(|e| anyhow!("write dtype audit JSON: {e}"))?;
}
}
let pooled_n = ((img.n_x / 3) as usize) * ((img.n_y / 3) as usize);
let mm0 = mmproj_weights
.mm_0_weight()
.map_err(|e| anyhow!("mm.0: {e}"))?;
let text_hidden = mm0.element_count() / (mmproj_cfg.hidden_size as usize);
let total = pooled_n * text_hidden;
let slice: &[f32] = buf
.as_slice::<f32>()
.map_err(|e| anyhow!("readback: {e}"))?;
if slice.len() != total {
return Err(anyhow!(
"compute_vision_embeddings_gpu_gemma4v image {}: readback len {} != expected {} \
(pooled_n={}, text_hidden={})",
idx,
slice.len(),
total,
pooled_n,
text_hidden
));
}
out.push(slice.to_vec());
}
Ok(out)
}
#[derive(Debug, Clone)]
pub enum VisionInput {
Siglip49(super::PreprocessedImage),
Gemma4v(Gemma4vPreprocessedImage),
}
pub fn compute_vision_embeddings_gpu_dispatch(
inputs: &[VisionInput],
arch: super::mmproj::ArchProfile,
mmproj_weights: &super::mmproj_weights::LoadedMmprojWeights,
mmproj_cfg: &super::mmproj::MmprojConfig,
scale: f32,
) -> Result<Vec<Vec<f32>>> {
if !arch.is_supported() {
return Err(anyhow!(
"compute_vision_embeddings_gpu_dispatch: arch profile is Unknown — \
cannot dispatch a vision forward"
));
}
for (idx, input) in inputs.iter().enumerate() {
match (&arch, input) {
(super::mmproj::ArchProfile::Gemma4Siglip, VisionInput::Gemma4v(_))
| (super::mmproj::ArchProfile::Gemma4Siglip, VisionInput::Siglip49(_))
| (super::mmproj::ArchProfile::ClipClassic, VisionInput::Siglip49(_)) => {}
(super::mmproj::ArchProfile::ClipClassic, VisionInput::Gemma4v(_)) => {
return Err(anyhow!(
"compute_vision_embeddings_gpu_dispatch: input {idx} is Gemma4v \
but arch is ClipClassic — preprocessing/arch mismatch"
));
}
(super::mmproj::ArchProfile::Qwen3VlSiglip, VisionInput::Siglip49(_)) => {}
(super::mmproj::ArchProfile::Qwen3VlSiglip, VisionInput::Gemma4v(_)) => {
return Err(anyhow!(
"compute_vision_embeddings_gpu_dispatch: input {idx} is Gemma4v \
but arch is Qwen3VlSiglip — preprocessing/arch mismatch (Qwen3-VL \
uses square-fixed-resolution Siglip49 preprocessing in Phase 1)"
));
}
(super::mmproj::ArchProfile::Unknown, _) => unreachable!("guarded above"),
}
}
let mut siglip_idx: Vec<usize> = Vec::new();
let mut siglip_imgs: Vec<super::PreprocessedImage> = Vec::new();
let mut gemma_idx: Vec<usize> = Vec::new();
let mut gemma_imgs: Vec<Gemma4vPreprocessedImage> = Vec::new();
for (idx, input) in inputs.iter().enumerate() {
match input {
VisionInput::Siglip49(p) => {
siglip_idx.push(idx);
siglip_imgs.push(p.clone());
}
VisionInput::Gemma4v(g) => {
gemma_idx.push(idx);
gemma_imgs.push(g.clone());
}
}
}
let mut out: Vec<Option<Vec<f32>>> = (0..inputs.len()).map(|_| None).collect();
if matches!(arch, super::mmproj::ArchProfile::Qwen3VlSiglip) {
let num_position_embeddings: u32 = mmproj_weights
.position_embd_weight()
.map_err(|e| {
anyhow!(
"compute_vision_embeddings_gpu_dispatch (Qwen3VlSiglip): \
v.position_embd.weight required for shape extraction: {e}"
)
})?
.shape()[0] as u32;
let cfg = super::vit_gpu_qwen3vl::Qwen3VlViTConfig::from_mmproj(
mmproj_cfg,
num_position_embeddings,
)?;
let r = super::vit_gpu_qwen3vl::compute_vision_embeddings_gpu_qwen3vl(
inputs,
mmproj_weights,
&cfg,
mmproj_cfg,
)?;
for (i, e) in r.into_iter().enumerate() {
out[i] = Some(e);
}
return out
.into_iter()
.enumerate()
.map(|(i, slot)| {
slot.ok_or_else(|| {
anyhow!("compute_vision_embeddings_gpu_dispatch: slot {i} unfilled")
})
})
.collect();
}
if !siglip_imgs.is_empty() {
let r = compute_vision_embeddings_gpu(&siglip_imgs, mmproj_weights, mmproj_cfg, scale)?;
for (i, e) in siglip_idx.into_iter().zip(r.into_iter()) {
out[i] = Some(e);
}
}
if !gemma_imgs.is_empty() {
let r = compute_vision_embeddings_gpu_gemma4v(&gemma_imgs, mmproj_weights, mmproj_cfg)?;
for (i, e) in gemma_idx.into_iter().zip(r.into_iter()) {
out[i] = Some(e);
}
}
out.into_iter()
.enumerate()
.map(|(i, slot)| {
slot.ok_or_else(|| anyhow!("compute_vision_embeddings_gpu_dispatch: slot {i} unfilled"))
})
.collect()
}
pub fn gemma4v_patch_embed_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
patches: &MlxBuffer,
weight: &MlxBuffer,
n_patches: u32,
inner: u32,
hidden: u32,
) -> Result<MlxBuffer> {
if n_patches == 0 {
return Err(anyhow!("gemma4v_patch_embed_gpu: n_patches must be > 0"));
}
vit_linear_gpu(
encoder, registry, device, patches, weight, n_patches, inner, hidden,
)
.context("gemma4v_patch_embed_gpu: vit_linear_gpu")
}
#[allow(clippy::too_many_arguments)]
pub fn gemma4v_position_embed_lookup_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
pe_table: &MlxBuffer,
pos_x_idx: &MlxBuffer,
pos_y_idx: &MlxBuffer,
n_patches: u32,
pos_size: u32,
hidden: u32,
) -> Result<MlxBuffer> {
if n_patches == 0 || pos_size == 0 || hidden == 0 {
return Err(anyhow!(
"gemma4v_position_embed_lookup_gpu: n_patches ({n_patches}), pos_size ({pos_size}), \
hidden ({hidden}) must all be > 0"
));
}
let row_elems = (pos_size as usize) * (hidden as usize);
let table_bytes = row_elems * 4;
let table_x = pe_table.slice_view(0, row_elems);
let table_y = pe_table.slice_view(table_bytes as u64, row_elems);
let n_us = n_patches as usize;
let h_us = hidden as usize;
let out_bytes = n_us * h_us * 4;
let emb_x = device
.alloc_buffer(out_bytes, DType::F32, vec![n_us, h_us])
.map_err(|e| anyhow!("alloc emb_x: {e}"))?;
let emb_y = device
.alloc_buffer(out_bytes, DType::F32, vec![n_us, h_us])
.map_err(|e| anyhow!("alloc emb_y: {e}"))?;
dispatch_gather_f32(
encoder,
registry,
device.metal_device(),
&table_x,
pos_x_idx,
&emb_x,
pos_size,
hidden,
n_patches,
)
.context("gemma4v_position_embed_lookup_gpu: gather X")?;
dispatch_gather_f32(
encoder,
registry,
device.metal_device(),
&table_y,
pos_y_idx,
&emb_y,
pos_size,
hidden,
n_patches,
)
.context("gemma4v_position_embed_lookup_gpu: gather Y")?;
encoder.memory_barrier();
let out = device
.alloc_buffer(out_bytes, DType::F32, vec![n_us, h_us])
.map_err(|e| anyhow!("alloc pos_emb sum: {e}"))?;
elementwise_add(
encoder,
registry,
device.metal_device(),
&emb_x,
&emb_y,
&out,
n_us * h_us,
DType::F32,
)
.context("gemma4v_position_embed_lookup_gpu: elementwise_add")?;
Ok(out)
}
#[allow(clippy::too_many_arguments)]
pub fn gemma4v_apply_position_embed_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
patch_embeds: &MlxBuffer,
pe_table: &MlxBuffer,
pos_x_idx: &MlxBuffer,
pos_y_idx: &MlxBuffer,
n_patches: u32,
pos_size: u32,
hidden: u32,
) -> Result<MlxBuffer> {
let pos_emb = gemma4v_position_embed_lookup_gpu(
encoder, registry, device, pe_table, pos_x_idx, pos_y_idx, n_patches, pos_size, hidden,
)?;
encoder.memory_barrier();
let n_elem = (n_patches as u64).saturating_mul(hidden as u64);
if n_elem > u32::MAX as u64 {
return Err(anyhow!(
"gemma4v_apply_position_embed_gpu: n_patches*hidden ({n_elem}) exceeds u32::MAX"
));
}
vit_residual_add_gpu(
encoder,
registry,
device,
patch_embeds,
&pos_emb,
n_elem as u32,
)
.context("gemma4v_apply_position_embed_gpu: residual add")
}
use mlx_native::ops::elementwise::elementwise_mul as mlx_elementwise_mul;
use mlx_native::ops::gelu::dispatch_gelu;
use mlx_native::ops::rms_norm::dispatch_rms_norm_no_scale_f32;
use mlx_native::ops::vision_2d_rope::{build_vision_2d_rope_params, dispatch_vision_2d_rope};
#[allow(clippy::too_many_arguments)]
pub fn vit_gemma_rms_norm_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
gain_f32: &MlxBuffer,
rows: u32,
dim: u32,
eps: f32,
) -> Result<MlxBuffer> {
if rows == 0 || dim == 0 {
return Err(anyhow!(
"vit_gemma_rms_norm_gpu: rows ({rows}) and dim ({dim}) must be > 0"
));
}
vit_rms_norm_gpu(encoder, registry, device, input, gain_f32, rows, dim, eps)
.context("vit_gemma_rms_norm_gpu: rms_norm_gpu")
}
#[allow(clippy::too_many_arguments)]
pub fn vit_gemma_per_head_rms_norm_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
gain_f32: &MlxBuffer,
batch: u32,
num_heads: u32,
head_dim: u32,
eps: f32,
) -> Result<MlxBuffer> {
if batch == 0 || num_heads == 0 || head_dim == 0 {
return Err(anyhow!(
"vit_gemma_per_head_rms_norm_gpu: batch ({batch}), num_heads ({num_heads}), \
head_dim ({head_dim}) must all be > 0"
));
}
let rows = batch
.checked_mul(num_heads)
.ok_or_else(|| anyhow!("vit_gemma_per_head_rms_norm_gpu: batch*num_heads overflow"))?;
vit_gemma_rms_norm_gpu(
encoder, registry, device, input, gain_f32, rows, head_dim, eps,
)
}
pub fn vit_v_norm_no_scale_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
batch: u32,
num_kv_heads: u32,
head_dim: u32,
eps: f32,
) -> Result<MlxBuffer> {
if batch == 0 || num_kv_heads == 0 || head_dim == 0 {
return Err(anyhow!(
"vit_v_norm_no_scale_gpu: batch/num_kv_heads/head_dim must all be > 0"
));
}
let rows = batch
.checked_mul(num_kv_heads)
.ok_or_else(|| anyhow!("vit_v_norm_no_scale_gpu: batch*num_kv_heads overflow"))?;
let n_elem = (rows as usize) * (head_dim as usize);
let output = device
.alloc_buffer(
n_elem * 4,
DType::F32,
vec![rows as usize, head_dim as usize],
)
.map_err(|e| anyhow!("vit_v_norm_no_scale_gpu: alloc output: {e}"))?;
let params_buf = device
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow!("vit_v_norm_no_scale_gpu: alloc params: {e}"))?;
{
let s: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(params_buf.contents_ptr() as *mut f32, 2) };
s[0] = eps;
s[1] = head_dim as f32;
}
dispatch_rms_norm_no_scale_f32(
encoder,
registry,
device.metal_device(),
input,
&output,
¶ms_buf,
rows,
head_dim,
)
.context("vit_v_norm_no_scale_gpu: dispatch_rms_norm_no_scale_f32")?;
Ok(output)
}
#[allow(clippy::too_many_arguments)]
pub fn vit_vision_2d_rope_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
pos_x: &MlxBuffer,
pos_y: &MlxBuffer,
seq_len: u32,
n_heads: u32,
head_dim: u32,
theta: f32,
) -> Result<MlxBuffer> {
if seq_len == 0 || n_heads == 0 || head_dim == 0 {
return Err(anyhow!(
"vit_vision_2d_rope_gpu: seq_len/n_heads/head_dim must all be > 0"
));
}
let n_rows = (seq_len as usize) * (n_heads as usize);
let n_elem = n_rows * (head_dim as usize);
let output = device
.alloc_buffer(n_elem * 4, DType::F32, vec![n_rows, head_dim as usize])
.map_err(|e| anyhow!("vit_vision_2d_rope_gpu: alloc output: {e}"))?;
let params = build_vision_2d_rope_params(device, theta, head_dim, n_heads)
.map_err(|e| anyhow!("vit_vision_2d_rope_gpu: build params: {e}"))?;
dispatch_vision_2d_rope(
encoder,
registry,
device.metal_device(),
input,
&output,
¶ms,
pos_x,
pos_y,
seq_len,
n_heads,
head_dim,
)
.context("vit_vision_2d_rope_gpu: dispatch_vision_2d_rope")?;
Ok(output)
}
pub fn vit_repeat_kv_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
batch: u32,
num_kv_heads: u32,
num_kv_groups: u32,
head_dim: u32,
) -> Result<MlxBuffer> {
if batch == 0 || num_kv_heads == 0 || num_kv_groups == 0 || head_dim == 0 {
return Err(anyhow!(
"vit_repeat_kv_gpu: batch/num_kv_heads/num_kv_groups/head_dim must all be > 0"
));
}
let num_heads = num_kv_heads
.checked_mul(num_kv_groups)
.ok_or_else(|| anyhow!("vit_repeat_kv_gpu: num_kv_heads*num_kv_groups overflow"))?;
let n_out_rows = batch
.checked_mul(num_heads)
.ok_or_else(|| anyhow!("vit_repeat_kv_gpu: batch*num_heads overflow"))?;
let n_in_rows = batch
.checked_mul(num_kv_heads)
.ok_or_else(|| anyhow!("vit_repeat_kv_gpu: batch*num_kv_heads overflow"))?;
let idx_buf = device
.alloc_buffer(
n_out_rows as usize * 4,
DType::U32,
vec![n_out_rows as usize],
)
.map_err(|e| anyhow!("vit_repeat_kv_gpu: alloc idx: {e}"))?;
{
let s: &mut [u32] = unsafe {
std::slice::from_raw_parts_mut(idx_buf.contents_ptr() as *mut u32, n_out_rows as usize)
};
for b in 0..batch {
for h in 0..num_heads {
let kv_h = h / num_kv_groups;
s[(b * num_heads + h) as usize] = b * num_kv_heads + kv_h;
}
}
}
let n_out_elem = (n_out_rows as usize) * (head_dim as usize);
let output = device
.alloc_buffer(
n_out_elem * 4,
DType::F32,
vec![n_out_rows as usize, head_dim as usize],
)
.map_err(|e| anyhow!("vit_repeat_kv_gpu: alloc output: {e}"))?;
dispatch_gather_f32(
encoder,
registry,
device.metal_device(),
input,
&idx_buf,
&output,
n_in_rows,
head_dim,
n_out_rows,
)
.context("vit_repeat_kv_gpu: dispatch_gather_f32")?;
Ok(output)
}
#[derive(Debug, Clone, Copy)]
pub struct Gemma4VisionBlockShapeGpu {
pub hidden: u32,
pub num_heads: u32,
pub num_kv_heads: u32,
pub head_dim: u32,
pub intermediate: u32,
pub rms_norm_eps: f32,
pub rope_theta: f32,
}
pub fn vit_gelu_pytorch_tanh_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
n_elements: u32,
) -> Result<MlxBuffer> {
if n_elements == 0 {
return Err(anyhow!("vit_gelu_pytorch_tanh_gpu: n_elements must be > 0"));
}
let output = device
.alloc_buffer(
(n_elements as usize) * 4,
DType::F32,
vec![n_elements as usize],
)
.map_err(|e| anyhow!("vit_gelu_pytorch_tanh_gpu: alloc output: {e}"))?;
dispatch_gelu(encoder, registry, device.metal_device(), input, &output)
.context("vit_gelu_pytorch_tanh_gpu: dispatch_gelu")?;
Ok(output)
}
#[allow(clippy::too_many_arguments)]
pub fn gemma4v_block_forward_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
weights: &super::mmproj_weights::LoadedMmprojWeights,
shape: &Gemma4VisionBlockShapeGpu,
block_idx: usize,
input: &MlxBuffer,
pos_x_idx: &MlxBuffer,
pos_y_idx: &MlxBuffer,
batch: u32,
) -> Result<MlxBuffer> {
let hidden = shape.hidden;
let num_heads = shape.num_heads;
let num_kv_heads = shape.num_kv_heads;
let head_dim = shape.head_dim;
let intermediate = shape.intermediate;
let eps = shape.rms_norm_eps;
let theta = shape.rope_theta;
if hidden == 0
|| num_heads == 0
|| num_kv_heads == 0
|| head_dim == 0
|| intermediate == 0
|| batch == 0
{
return Err(anyhow!(
"gemma4v_block_forward_gpu: zero dim in shape ({shape:?}) or batch ({batch})"
));
}
if num_heads % num_kv_heads != 0 {
return Err(anyhow!(
"gemma4v_block_forward_gpu: num_heads ({num_heads}) must be a multiple of num_kv_heads ({num_kv_heads})"
));
}
let num_kv_groups = num_heads / num_kv_heads;
let q_dim = num_heads
.checked_mul(head_dim)
.ok_or_else(|| anyhow!("gemma4v_block_forward_gpu: num_heads*head_dim overflow"))?;
let kv_dim = num_kv_heads
.checked_mul(head_dim)
.ok_or_else(|| anyhow!("gemma4v_block_forward_gpu: num_kv_heads*head_dim overflow"))?;
if q_dim != hidden {
return Err(anyhow!(
"gemma4v_block_forward_gpu: q_dim ({q_dim}) != hidden ({hidden})"
));
}
let block = |suffix: &str| -> Result<&MlxBuffer> {
weights
.block_tensor(block_idx, suffix)
.map_err(|e| anyhow!("block {} {}: {e}", block_idx, suffix))
};
let n_hidden = batch
.checked_mul(hidden)
.ok_or_else(|| anyhow!("gemma4v_block_forward_gpu: batch*hidden overflow"))?;
let dump_intra =
super::vit_dump::is_armed() && (block_idx <= 1 || block_idx == 25 || block_idx == 26);
let intra_name = |suffix: &str| format!("03_block_{:02}_{}", block_idx, suffix);
let audit_intra = super::vit_dump::is_dtype_audit_armed() && super::vit_dump::is_armed();
let cur = vit_gemma_rms_norm_gpu(
encoder,
registry,
device,
input,
block("ln1.weight")?,
batch,
hidden,
eps,
)?;
encoder.memory_barrier();
if dump_intra {
super::vit_dump::record(&intra_name("01_pre_attn_norm"), &cur);
}
if audit_intra {
super::vit_dump::record_audit(&intra_name("01_pre_attn_norm"), &cur);
}
let q = vit_linear_gpu(
encoder,
registry,
device,
&cur,
block("attn_q.weight")?,
batch,
hidden,
q_dim,
)?;
encoder.memory_barrier();
if audit_intra {
super::vit_dump::record_audit(&intra_name("attn_q_proj"), &q);
super::vit_dump::record_audit(
&intra_name("attn_q_weight_storage"),
block("attn_q.weight")?,
);
}
let k = vit_linear_gpu(
encoder,
registry,
device,
&cur,
block("attn_k.weight")?,
batch,
hidden,
kv_dim,
)?;
encoder.memory_barrier();
if audit_intra {
super::vit_dump::record_audit(&intra_name("attn_k_proj"), &k);
super::vit_dump::record_audit(
&intra_name("attn_k_weight_storage"),
block("attn_k.weight")?,
);
}
let v = vit_linear_gpu(
encoder,
registry,
device,
&cur,
block("attn_v.weight")?,
batch,
hidden,
kv_dim,
)?;
encoder.memory_barrier();
if audit_intra {
super::vit_dump::record_audit(&intra_name("attn_v_proj"), &v);
super::vit_dump::record_audit(
&intra_name("attn_v_weight_storage"),
block("attn_v.weight")?,
);
}
let q = vit_gemma_per_head_rms_norm_gpu(
encoder,
registry,
device,
&q,
block("attn_q_norm.weight")?,
batch,
num_heads,
head_dim,
eps,
)?;
encoder.memory_barrier();
if audit_intra {
super::vit_dump::record_audit(&intra_name("attn_q_normed"), &q);
super::vit_dump::record_audit(
&intra_name("attn_q_norm_weight_storage"),
block("attn_q_norm.weight")?,
);
}
let k = vit_gemma_per_head_rms_norm_gpu(
encoder,
registry,
device,
&k,
block("attn_k_norm.weight")?,
batch,
num_kv_heads,
head_dim,
eps,
)?;
encoder.memory_barrier();
if audit_intra {
super::vit_dump::record_audit(&intra_name("attn_k_normed"), &k);
super::vit_dump::record_audit(
&intra_name("attn_k_norm_weight_storage"),
block("attn_k_norm.weight")?,
);
}
let v = vit_v_norm_no_scale_gpu(
encoder,
registry,
device,
&v,
batch,
num_kv_heads,
head_dim,
eps,
)?;
encoder.memory_barrier();
if audit_intra {
super::vit_dump::record_audit(&intra_name("attn_v_normed_no_scale"), &v);
}
let q = vit_vision_2d_rope_gpu(
encoder, registry, device, &q, pos_x_idx, pos_y_idx, batch, num_heads, head_dim, theta,
)?;
encoder.memory_barrier();
if audit_intra {
super::vit_dump::record_audit(&intra_name("attn_q_rope"), &q);
}
let k = vit_vision_2d_rope_gpu(
encoder,
registry,
device,
&k,
pos_x_idx,
pos_y_idx,
batch,
num_kv_heads,
head_dim,
theta,
)?;
encoder.memory_barrier();
if audit_intra {
super::vit_dump::record_audit(&intra_name("attn_k_rope"), &k);
}
if dump_intra {
super::vit_dump::record(&intra_name("02_q_pos"), &q);
super::vit_dump::record(&intra_name("03_k_pos"), &k);
super::vit_dump::record(&intra_name("04_v_normed"), &v);
}
let k_full = vit_repeat_kv_gpu(
encoder,
registry,
device,
&k,
batch,
num_kv_heads,
num_kv_groups,
head_dim,
)?;
encoder.memory_barrier();
if audit_intra {
super::vit_dump::record_audit(&intra_name("attn_k_full_gqa"), &k_full);
}
let v_full = vit_repeat_kv_gpu(
encoder,
registry,
device,
&v,
batch,
num_kv_heads,
num_kv_groups,
head_dim,
)?;
encoder.memory_barrier();
if audit_intra {
super::vit_dump::record_audit(&intra_name("attn_v_full_gqa"), &v_full);
}
let attn = vit_attention_gpu(
encoder, registry, device, &q, &k_full, &v_full, batch, num_heads, head_dim, 1.0,
)?;
encoder.memory_barrier();
if audit_intra {
super::vit_dump::record_audit(&intra_name("attn_kqv_out"), &attn);
}
if dump_intra {
super::vit_dump::record(&intra_name("05_kqv_out"), &attn);
}
let attn_proj = vit_linear_gpu(
encoder,
registry,
device,
&attn,
block("attn_out.weight")?,
batch,
hidden,
hidden,
)?;
encoder.memory_barrier();
if dump_intra {
super::vit_dump::record(&intra_name("06_attn_out"), &attn_proj);
}
if audit_intra {
super::vit_dump::record_audit(&intra_name("attn_out_proj"), &attn_proj);
super::vit_dump::record_audit(
&intra_name("attn_out_weight_storage"),
block("attn_out.weight")?,
);
}
let attn_out = vit_gemma_rms_norm_gpu(
encoder,
registry,
device,
&attn_proj,
block("attn_post_norm.weight")?,
batch,
hidden,
eps,
)?;
encoder.memory_barrier();
if dump_intra {
super::vit_dump::record(&intra_name("07_attn_post_normed"), &attn_out);
}
if audit_intra {
super::vit_dump::record_audit(&intra_name("attn_post_normed"), &attn_out);
super::vit_dump::record_audit(
&intra_name("attn_post_norm_weight_storage"),
block("attn_post_norm.weight")?,
);
}
let x_mid = vit_residual_add_gpu(encoder, registry, device, input, &attn_out, n_hidden)?;
encoder.memory_barrier();
if dump_intra {
super::vit_dump::record(&intra_name("08_ffn_inp"), &x_mid);
}
if audit_intra {
super::vit_dump::record_audit(&intra_name("ffn_inp_residual"), &x_mid);
}
let cur = vit_gemma_rms_norm_gpu(
encoder,
registry,
device,
&x_mid,
block("ln2.weight")?,
batch,
hidden,
eps,
)?;
encoder.memory_barrier();
if dump_intra {
super::vit_dump::record(&intra_name("09_ffn_inp_normed"), &cur);
}
if audit_intra {
super::vit_dump::record_audit(&intra_name("ffn_inp_normed"), &cur);
super::vit_dump::record_audit(&intra_name("ln2_weight_storage"), block("ln2.weight")?);
}
let gate = vit_linear_gpu(
encoder,
registry,
device,
&cur,
block("ffn_gate.weight")?,
batch,
hidden,
intermediate,
)?;
encoder.memory_barrier();
if audit_intra {
super::vit_dump::record_audit(&intra_name("ffn_gate_proj"), &gate);
super::vit_dump::record_audit(
&intra_name("ffn_gate_weight_storage"),
block("ffn_gate.weight")?,
);
}
let up = vit_linear_gpu(
encoder,
registry,
device,
&cur,
block("ffn_up.weight")?,
batch,
hidden,
intermediate,
)?;
encoder.memory_barrier();
if audit_intra {
super::vit_dump::record_audit(&intra_name("ffn_up_proj"), &up);
super::vit_dump::record_audit(
&intra_name("ffn_up_weight_storage"),
block("ffn_up.weight")?,
);
}
let n_inter = batch
.checked_mul(intermediate)
.ok_or_else(|| anyhow!("gemma4v_block_forward_gpu: batch*intermediate overflow"))?;
let gated = vit_gelu_pytorch_tanh_gpu(encoder, registry, device, &gate, n_inter)?;
encoder.memory_barrier();
if audit_intra {
super::vit_dump::record_audit(&intra_name("ffn_gated_gelu"), &gated);
}
let activated = device
.alloc_buffer((n_inter as usize) * 4, DType::F32, vec![n_inter as usize])
.map_err(|e| anyhow!("gemma4v_block_forward_gpu: alloc activated: {e}"))?;
mlx_elementwise_mul(
encoder,
registry,
device.metal_device(),
&gated,
&up,
&activated,
n_inter as usize,
DType::F32,
)
.context("gemma4v_block_forward_gpu: gate * up")?;
encoder.memory_barrier();
if audit_intra {
super::vit_dump::record_audit(&intra_name("ffn_activated"), &activated);
}
let down = vit_linear_gpu(
encoder,
registry,
device,
&activated,
block("ffn_down.weight")?,
batch,
intermediate,
hidden,
)?;
encoder.memory_barrier();
if dump_intra {
super::vit_dump::record(&intra_name("10_ffn_out"), &down);
}
if audit_intra {
super::vit_dump::record_audit(&intra_name("ffn_down_proj"), &down);
super::vit_dump::record_audit(
&intra_name("ffn_down_weight_storage"),
block("ffn_down.weight")?,
);
}
let down = vit_gemma_rms_norm_gpu(
encoder,
registry,
device,
&down,
block("ffn_post_norm.weight")?,
batch,
hidden,
eps,
)?;
encoder.memory_barrier();
if dump_intra {
super::vit_dump::record(&intra_name("11_ffn_post_normed"), &down);
}
if audit_intra {
super::vit_dump::record_audit(&intra_name("ffn_post_normed"), &down);
super::vit_dump::record_audit(
&intra_name("ffn_post_norm_weight_storage"),
block("ffn_post_norm.weight")?,
);
}
let x_out = vit_residual_add_gpu(encoder, registry, device, &x_mid, &down, n_hidden)?;
if audit_intra {
super::vit_dump::record_audit(&intra_name("layer_out_residual"), &x_out);
}
Ok(x_out)
}
#[allow(clippy::too_many_arguments)]
pub fn gemma4v_clippable_linear_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight_f32: &MlxBuffer,
bounds: &super::vit::Gemma4ClippableLinearBounds,
seq_len: u32,
in_features: u32,
out_features: u32,
) -> Result<MlxBuffer> {
let input_for_matmul: MlxBuffer;
let input_ref: &MlxBuffer = if bounds.input_min.is_some() || bounds.input_max.is_some() {
let (mn, mx) = bounds.resolved_input();
let n = (seq_len as u64).saturating_mul(in_features as u64);
if n > u32::MAX as u64 {
return Err(anyhow!(
"gemma4v_clippable_linear_gpu: input element count ({n}) exceeds u32::MAX"
));
}
input_for_matmul = vit_clip_gpu(encoder, registry, device, input, n as u32, mn, mx)
.context("gemma4v_clippable_linear_gpu: input clamp")?;
encoder.memory_barrier();
&input_for_matmul
} else {
input
};
let projected = vit_linear_gpu(
encoder,
registry,
device,
input_ref,
weight_f32,
seq_len,
in_features,
out_features,
)
.context("gemma4v_clippable_linear_gpu: vit_linear_gpu")?;
if bounds.output_min.is_some() || bounds.output_max.is_some() {
encoder.memory_barrier();
let (mn, mx) = bounds.resolved_output();
let n = (seq_len as u64).saturating_mul(out_features as u64);
if n > u32::MAX as u64 {
return Err(anyhow!(
"gemma4v_clippable_linear_gpu: output element count ({n}) exceeds u32::MAX"
));
}
return vit_clip_gpu(encoder, registry, device, &projected, n as u32, mn, mx)
.context("gemma4v_clippable_linear_gpu: output clamp");
}
Ok(projected)
}
#[allow(clippy::too_many_arguments)]
pub fn gemma4v_apply_full_forward_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
weights: &super::mmproj_weights::LoadedMmprojWeights,
cfg: &super::mmproj::MmprojConfig,
patches: &[f32],
pos_x: &[u32],
pos_y: &[u32],
n_x: u32,
n_y: u32,
) -> Result<MlxBuffer> {
mlx_native::ops::vision_2d_rope::register(registry);
mlx_native::ops::gelu::register(registry);
mlx_native::ops::gather::register(registry);
if n_x == 0 || n_y == 0 {
return Err(anyhow!(
"gemma4v_apply_full_forward_gpu: n_x ({n_x}) and n_y ({n_y}) must be > 0"
));
}
if n_x % 3 != 0 || n_y % 3 != 0 {
return Err(anyhow!(
"gemma4v_apply_full_forward_gpu: n_x ({n_x}) and n_y ({n_y}) must both \
be multiples of 3 (gemma4v pool kernel size)"
));
}
let n_patches = (n_x as u64).saturating_mul(n_y as u64);
if n_patches == 0 || n_patches > u32::MAX as u64 {
return Err(anyhow!(
"gemma4v_apply_full_forward_gpu: n_patches ({n_patches}) overflow u32"
));
}
let n_patches = n_patches as u32;
let hidden = cfg.hidden_size;
let patch_size = cfg.patch_size;
let inner = patch_size
.checked_mul(patch_size)
.and_then(|s| s.checked_mul(3))
.ok_or_else(|| anyhow!("gemma4v_apply_full_forward_gpu: inner overflow"))?;
let expected_patches = (n_patches as usize).saturating_mul(inner as usize);
if patches.len() != expected_patches {
return Err(anyhow!(
"gemma4v_apply_full_forward_gpu: patches len {} != n_patches*inner = {}*{} = {}",
patches.len(),
n_patches,
inner,
expected_patches
));
}
if pos_x.len() != n_patches as usize || pos_y.len() != n_patches as usize {
return Err(anyhow!(
"gemma4v_apply_full_forward_gpu: pos_x ({}) / pos_y ({}) length must equal n_patches ({})",
pos_x.len(),
pos_y.len(),
n_patches
));
}
let patches_buf = device
.alloc_buffer(
expected_patches * 4,
DType::F32,
vec![n_patches as usize, inner as usize],
)
.map_err(|e| anyhow!("alloc patches: {e}"))?;
{
let dst: &mut [f32] = unsafe {
std::slice::from_raw_parts_mut(patches_buf.contents_ptr() as *mut f32, expected_patches)
};
dst.copy_from_slice(patches);
}
let pos_x_buf = device
.alloc_buffer(
(n_patches as usize) * 4,
DType::U32,
vec![n_patches as usize],
)
.map_err(|e| anyhow!("alloc pos_x: {e}"))?;
{
let dst: &mut [u32] = unsafe {
std::slice::from_raw_parts_mut(pos_x_buf.contents_ptr() as *mut u32, n_patches as usize)
};
dst.copy_from_slice(pos_x);
}
let pos_y_buf = device
.alloc_buffer(
(n_patches as usize) * 4,
DType::U32,
vec![n_patches as usize],
)
.map_err(|e| anyhow!("alloc pos_y: {e}"))?;
{
let dst: &mut [u32] = unsafe {
std::slice::from_raw_parts_mut(pos_y_buf.contents_ptr() as *mut u32, n_patches as usize)
};
dst.copy_from_slice(pos_y);
}
if super::vit_dump::is_armed() {
super::vit_dump::record_f32(
"00_post_patchify",
patches,
vec![n_patches as usize, inner as usize],
);
let p = patch_size as usize;
let p2 = p * p;
let h = (n_y as usize) * p;
let w = (n_x as usize) * p;
let mut planar = vec![0f32; 3 * h * w];
for py in 0..(n_y as usize) {
for px in 0..(n_x as usize) {
let patch_idx = py * (n_x as usize) + px;
let row_base = patch_idx * (inner as usize);
for dy in 0..p {
for dx in 0..p {
let pos_in_plane = dy * p + dx;
let yy = py * p + dy;
let xx = px * p + dx;
for c in 0..3 {
planar[c * h * w + yy * w + xx] =
patches[row_base + c * p2 + pos_in_plane];
}
}
}
}
}
super::vit_dump::record_f32("00_pre_patchify", &planar, vec![3, h, w]);
}
let patch_w = weights
.patch_embd_weight()
.map_err(|e| anyhow!("gemma4v_apply_full_forward_gpu: {e}"))?;
let patch_embeds = gemma4v_patch_embed_gpu(
encoder,
registry,
device,
&patches_buf,
patch_w,
n_patches,
inner,
hidden,
)?;
encoder.memory_barrier();
super::vit_dump::record("01_patch_embd", &patch_embeds);
let (pe_table, pos_size, pe_hidden) = weights
.position_embd_table_3d()
.map_err(|e| anyhow!("gemma4v_apply_full_forward_gpu: {e}"))?;
if pe_hidden != hidden {
return Err(anyhow!(
"gemma4v_apply_full_forward_gpu: position_embd hidden ({pe_hidden}) != cfg.hidden_size ({hidden})"
));
}
let after_pos = gemma4v_apply_position_embed_gpu(
encoder,
registry,
device,
&patch_embeds,
pe_table,
&pos_x_buf,
&pos_y_buf,
n_patches,
pos_size,
hidden,
)?;
encoder.memory_barrier();
super::vit_dump::record("02_pos_embd", &after_pos);
let head_dim = hidden / cfg.num_attention_heads;
let attn_k0 = weights
.block_tensor(0, "attn_k.weight")
.map_err(|e| anyhow!("gemma4v_apply_full_forward_gpu: probe num_kv_heads: {e}"))?;
let attn_k0_rows = attn_k0.shape().first().copied().unwrap_or(0) as u32;
if attn_k0_rows == 0 || attn_k0_rows % head_dim != 0 {
return Err(anyhow!(
"gemma4v_apply_full_forward_gpu: cannot infer num_kv_heads from attn_k.weight \
row count ({attn_k0_rows}); head_dim = {head_dim}"
));
}
let num_kv_heads = attn_k0_rows / head_dim;
let shape = Gemma4VisionBlockShapeGpu {
hidden,
num_heads: cfg.num_attention_heads,
num_kv_heads,
head_dim,
intermediate: cfg.intermediate_size,
rms_norm_eps: cfg.layer_norm_eps,
rope_theta: 100.0f32,
};
let mut hidden_states = gemma4v_block_forward_gpu(
encoder, registry, device, weights, &shape, 0, &after_pos, &pos_x_buf, &pos_y_buf,
n_patches,
)?;
encoder.memory_barrier();
if super::vit_dump::is_armed() {
super::vit_dump::record("03_block_00", &hidden_states);
}
for block_idx in 1..(cfg.num_hidden_layers as usize) {
hidden_states = gemma4v_block_forward_gpu(
encoder,
registry,
device,
weights,
&shape,
block_idx,
&hidden_states,
&pos_x_buf,
&pos_y_buf,
n_patches,
)?;
encoder.memory_barrier();
if super::vit_dump::is_armed() {
super::vit_dump::record(&format!("03_block_{:02}", block_idx), &hidden_states);
}
}
let pooled =
gemma4v_avg_pool_3x3_gpu(encoder, registry, device, &hidden_states, n_x, n_y, hidden)?;
encoder.memory_barrier();
super::vit_dump::record("30_final_pool", &pooled);
let pooled_n = ((n_x / 3) as usize) * ((n_y / 3) as usize);
let pooled_total = pooled_n * (hidden as usize);
vit_scale_gpu(
encoder,
registry,
device,
&pooled,
pooled_total as u32,
(hidden as f32).sqrt(),
)?;
encoder.memory_barrier();
super::vit_dump::record("31_pool_sqrt_scale", &pooled);
let std_bias = weights
.get("v.std_bias")
.ok_or_else(|| anyhow!("gemma4v_apply_full_forward_gpu: missing v.std_bias"))?;
let std_scale = weights
.get("v.std_scale")
.ok_or_else(|| anyhow!("gemma4v_apply_full_forward_gpu: missing v.std_scale"))?;
let normed = vit_std_bias_scale_gpu(
encoder,
registry,
device,
&pooled,
std_bias,
std_scale,
pooled_n as u32,
hidden,
)?;
encoder.memory_barrier();
super::vit_dump::record("32_std_bias_scale", &normed);
let mm0 = weights
.mm_0_weight()
.map_err(|e| anyhow!("gemma4v_apply_full_forward_gpu: mm.0.weight: {e}"))?;
let text_hidden = (mm0.element_count() / (hidden as usize)) as u32;
if text_hidden == 0 {
return Err(anyhow!(
"gemma4v_apply_full_forward_gpu: text_hidden derived from mm.0.weight is 0"
));
}
let bounds = weights.mm_0_bounds();
let projected = gemma4v_clippable_linear_gpu(
encoder,
registry,
device,
&normed,
mm0,
&bounds,
pooled_n as u32,
hidden,
text_hidden,
)?;
encoder.memory_barrier();
super::vit_dump::record("33_projector", &projected);
let ones = device
.alloc_buffer(
(text_hidden as usize) * 4,
DType::F32,
vec![text_hidden as usize],
)
.map_err(|e| anyhow!("alloc ones: {e}"))?;
{
let s: &mut [f32] = unsafe {
std::slice::from_raw_parts_mut(ones.contents_ptr() as *mut f32, text_hidden as usize)
};
for v in s.iter_mut() {
*v = 1.0;
}
}
let post_proj_rms = vit_rms_norm_gpu(
encoder,
registry,
device,
&projected,
&ones,
pooled_n as u32,
text_hidden,
cfg.layer_norm_eps,
)?;
super::vit_dump::record("34_post_proj_rms", &post_proj_rms);
Ok(post_proj_rms)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::inference::vision::mmproj::MmprojConfig;
use crate::inference::vision::mmproj_weights::LoadedMmprojWeights;
use crate::inference::vision::vit::{
apply_vit_block_forward as apply_vit_block_forward_cpu, elementwise_mul_in_place,
gemma4v_block_forward as gemma4v_block_forward_cpu,
gemma4v_patch_embed_forward as gemma4v_patch_embed_cpu,
gemma4v_position_embed_lookup as gemma4v_pos_embed_cpu, linear_forward as linear_cpu,
per_head_rms_norm_forward as per_head_rms_cpu, residual_add as residual_add_cpu,
rms_norm_forward as rms_norm_cpu, scaled_dot_product_attention as attention_cpu,
silu_in_place, softmax_last_dim as softmax_cpu,
};
use mlx_native::gguf::GgufFile;
use mlx_native::{GraphExecutor, MlxDevice};
use std::path::Path;
const GEMMA4_MMPROJ_PATH: &str =
"/opt/hf2q/models/gemma-4-26B-A4B-it-ara-abliterated-dwq/gemma-4-26B-A4B-it-ara-abliterated-dwq-mmproj.gguf";
fn upload_f32(device: &MlxDevice, data: &[f32], shape: Vec<usize>) -> MlxBuffer {
let bytes = data.len() * 4;
let buf = device
.alloc_buffer(bytes, DType::F32, shape)
.expect("alloc upload");
let slice: &mut [f32] = unsafe {
std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, data.len())
};
slice.copy_from_slice(data);
buf
}
fn readback_f32(buf: &MlxBuffer, expected_len: usize) -> Vec<f32> {
let slice: &[f32] = buf.as_slice::<f32>().expect("readback as_slice");
assert_eq!(slice.len(), expected_len, "readback length mismatch");
slice.to_vec()
}
#[test]
fn vit_linear_gpu_matches_cpu_reference_on_small_input() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let seq = 4usize;
let in_features = 64usize;
let out_features = 32usize;
let input_cpu: Vec<f32> = (0..seq * in_features)
.map(|i| ((i as f32) * 0.001).sin())
.collect();
let weight_cpu: Vec<f32> = (0..out_features * in_features)
.map(|i| ((i as f32) * 0.01).cos() * 0.1)
.collect();
let expected_cpu = linear_cpu(
&input_cpu,
&weight_cpu,
None,
seq,
in_features,
out_features,
)
.expect("cpu ref");
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let input_buf = upload_f32(executor.device(), &input_cpu, vec![seq, in_features]);
let weight_buf = upload_f32(
executor.device(),
&weight_cpu,
vec![out_features, in_features],
);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_linear_gpu(
session.encoder_mut(),
&mut registry,
device,
&input_buf,
&weight_buf,
seq as u32,
in_features as u32,
out_features as u32,
)
.expect("gpu dispatch");
session.finish().expect("finish");
let got = readback_f32(&out_buf, seq * out_features);
for (i, (g, e)) in got.iter().zip(expected_cpu.iter()).enumerate() {
let diff = (g - e).abs();
assert!(
diff < 1e-2,
"GPU/CPU mismatch at element {i}: gpu={g} cpu={e} diff={diff}"
);
}
let max_diff = got
.iter()
.zip(expected_cpu.iter())
.map(|(g, e)| (g - e).abs())
.fold(0f32, f32::max);
assert!(max_diff < 1e-2, "overall max_diff = {max_diff}");
}
#[test]
fn vit_linear_gpu_rejects_small_in_features() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let input = executor
.device()
.alloc_buffer(4 * 16 * 4, DType::F32, vec![4, 16])
.expect("alloc");
let weight = executor
.device()
.alloc_buffer(32 * 16 * 4, DType::F32, vec![32, 16])
.expect("alloc");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err = vit_linear_gpu(
session.encoder_mut(),
&mut registry,
device,
&input,
&weight,
4,
16,
32,
)
.unwrap_err();
assert!(format!("{err}").contains("in_features"));
}
#[test]
fn vit_linear_gpu_rejects_zero_dims() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let input = executor
.device()
.alloc_buffer(32 * 4, DType::F32, vec![0, 32])
.expect("alloc");
let weight = executor
.device()
.alloc_buffer(32 * 32 * 4, DType::F32, vec![32, 32])
.expect("alloc");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err = vit_linear_gpu(
session.encoder_mut(),
&mut registry,
device,
&input,
&weight,
0,
32,
32,
)
.unwrap_err();
assert!(format!("{err}").contains("> 0"));
}
#[test]
fn vit_linear_gpu_on_real_gemma4_mm0_matches_cpu_at_small_seq() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let path = Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!("skipping: mmproj fixture not found");
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let device = MlxDevice::new().expect("device");
let weights = LoadedMmprojWeights::load(&gguf, &cfg, device).expect("load");
let hidden = cfg.hidden_size as usize;
let seq = 4usize;
let mm0 = weights.mm_0_weight().expect("mm.0");
let text_hidden = mm0.element_count() / hidden;
assert_eq!(text_hidden, 2816);
let weight_cpu: Vec<f32> = weights.tensor_as_f32_owned(mm0).expect("mm.0 widen");
let input_cpu: Vec<f32> = (0..seq * hidden)
.map(|i| ((i as f32) * 1e-4).sin() * 0.1)
.collect();
let expected =
linear_cpu(&input_cpu, &weight_cpu, None, seq, hidden, text_hidden).expect("cpu ref");
let exec_device = MlxDevice::new().expect("device2");
let executor = GraphExecutor::new(exec_device);
let input_buf = upload_f32(executor.device(), &input_cpu, vec![seq, hidden]);
let weight_buf = upload_f32(executor.device(), &weight_cpu, vec![text_hidden, hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_linear_gpu(
session.encoder_mut(),
&mut registry,
device,
&input_buf,
&weight_buf,
seq as u32,
hidden as u32,
text_hidden as u32,
)
.expect("gpu dispatch");
session.finish().expect("finish");
let got = readback_f32(&out_buf, seq * text_hidden);
let mut max_diff = 0f32;
let mut fail_count = 0usize;
for (g, e) in got.iter().zip(expected.iter()) {
let d = (g - e).abs();
if d > max_diff {
max_diff = d;
}
if d > 5e-2 {
fail_count += 1;
}
}
let total = got.len();
let fail_frac = (fail_count as f32) / (total as f32);
assert!(
fail_frac < 0.01,
"too many GPU/CPU mismatches: {}/{} = {:.3}% failed max_diff = {}",
fail_count,
total,
fail_frac * 100.0,
max_diff
);
}
#[test]
fn vit_rms_norm_gpu_matches_cpu_reference_on_small_input() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let rows = 8usize;
let dim = 16usize;
let eps = 1e-6f32;
let input_cpu: Vec<f32> = (0..rows * dim)
.map(|i| ((i as f32) * 0.05).sin() + 0.5)
.collect();
let gain_cpu: Vec<f32> = (0..dim).map(|i| 0.5 + (i as f32) * 0.05).collect();
let mut expected = input_cpu.clone();
rms_norm_cpu(&mut expected, &gain_cpu, dim, eps).expect("cpu ref");
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let input_buf = upload_f32(executor.device(), &input_cpu, vec![rows, dim]);
let gain_buf = upload_f32(executor.device(), &gain_cpu, vec![dim]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_rms_norm_gpu(
session.encoder_mut(),
&mut registry,
device,
&input_buf,
&gain_buf,
rows as u32,
dim as u32,
eps,
)
.expect("rms_norm");
session.finish().expect("finish");
let got = readback_f32(&out_buf, rows * dim);
let max_diff = got
.iter()
.zip(expected.iter())
.map(|(g, e)| (g - e).abs())
.fold(0f32, f32::max);
assert!(max_diff < 1e-4, "rms_norm GPU vs CPU max_diff = {max_diff}");
}
#[test]
fn vit_rms_norm_gpu_rejects_zero_dims() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let input = executor
.device()
.alloc_buffer(16 * 4, DType::F32, vec![4, 4])
.expect("a");
let gain = executor
.device()
.alloc_buffer(4 * 4, DType::F32, vec![4])
.expect("b");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err = vit_rms_norm_gpu(
session.encoder_mut(),
&mut registry,
device,
&input,
&gain,
0,
4,
1e-6,
)
.unwrap_err();
assert!(format!("{err}").contains("must be > 0"));
}
#[test]
fn vit_rms_norm_gpu_on_real_gemma4_ln1_matches_cpu() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let path = Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!("skipping: mmproj fixture not found");
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let device = MlxDevice::new().expect("device");
let weights = LoadedMmprojWeights::load(&gguf, &cfg, device).expect("load");
let hidden = cfg.hidden_size as usize;
let rows = 8usize;
let ln1_buf = weights.block_tensor(0, "ln1.weight").expect("ln1");
let gain_f32: &[f32] = ln1_buf.as_slice::<f32>().expect("ln1 slice");
assert_eq!(gain_f32.len(), hidden);
let gain_cpu: Vec<f32> = gain_f32.to_vec();
let input_cpu: Vec<f32> = (0..rows * hidden)
.map(|i| ((i as f32) * 1e-3).sin() * 0.5)
.collect();
let mut expected = input_cpu.clone();
rms_norm_cpu(&mut expected, &gain_cpu, hidden, cfg.layer_norm_eps).expect("cpu ref");
let exec_dev = MlxDevice::new().expect("device2");
let executor = GraphExecutor::new(exec_dev);
let input_buf = upload_f32(executor.device(), &input_cpu, vec![rows, hidden]);
let gain_buf = upload_f32(executor.device(), &gain_cpu, vec![hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device_inner: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_rms_norm_gpu(
session.encoder_mut(),
&mut registry,
device_inner,
&input_buf,
&gain_buf,
rows as u32,
hidden as u32,
cfg.layer_norm_eps,
)
.expect("rms_norm");
session.finish().expect("finish");
let got = readback_f32(&out_buf, rows * hidden);
let max_diff = got
.iter()
.zip(expected.iter())
.map(|(g, e)| (g - e).abs())
.fold(0f32, f32::max);
assert!(max_diff < 1e-3, "real-data rms_norm max_diff = {max_diff}");
}
#[test]
fn vit_per_head_rms_norm_gpu_matches_cpu_reference() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 4usize;
let num_heads = 8usize;
let head_dim = 16usize;
let total = batch * num_heads * head_dim;
let eps = 1e-6f32;
let input_cpu: Vec<f32> = (0..total).map(|i| ((i as f32) * 0.03).cos()).collect();
let gain_cpu: Vec<f32> = (0..head_dim).map(|i| 1.0 + (i as f32) * 0.1).collect();
let mut expected = input_cpu.clone();
per_head_rms_cpu(&mut expected, &gain_cpu, batch, num_heads, head_dim, eps)
.expect("cpu ref");
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let input_buf = upload_f32(
executor.device(),
&input_cpu,
vec![batch, num_heads, head_dim],
);
let gain_buf = upload_f32(executor.device(), &gain_cpu, vec![head_dim]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_per_head_rms_norm_gpu(
session.encoder_mut(),
&mut registry,
device,
&input_buf,
&gain_buf,
batch as u32,
num_heads as u32,
head_dim as u32,
eps,
)
.expect("per-head rms");
session.finish().expect("finish");
let got = readback_f32(&out_buf, total);
let max_diff = got
.iter()
.zip(expected.iter())
.map(|(g, e)| (g - e).abs())
.fold(0f32, f32::max);
assert!(
max_diff < 1e-4,
"per_head_rms GPU vs CPU max_diff = {max_diff}"
);
}
#[test]
fn vit_softmax_last_dim_gpu_matches_cpu_reference() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let rows = 4usize;
let cols = 8usize;
let input_cpu: Vec<f32> = (0..rows * cols)
.map(|i| ((i as f32) * 0.3).sin() + 0.5)
.collect();
let mut expected = input_cpu.clone();
softmax_cpu(&mut expected, cols).expect("cpu ref");
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let input_buf = upload_f32(executor.device(), &input_cpu, vec![rows, cols]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_softmax_last_dim_gpu(
session.encoder_mut(),
&mut registry,
device,
&input_buf,
rows as u32,
cols as u32,
)
.expect("softmax");
session.finish().expect("finish");
let got = readback_f32(&out_buf, rows * cols);
let max_diff = got
.iter()
.zip(expected.iter())
.map(|(g, e)| (g - e).abs())
.fold(0f32, f32::max);
assert!(max_diff < 1e-5, "softmax GPU vs CPU max_diff = {max_diff}");
for r in 0..rows {
let row_sum: f32 = got[r * cols..(r + 1) * cols].iter().sum();
assert!((row_sum - 1.0).abs() < 1e-4, "row {r} sum = {row_sum}");
}
}
#[test]
fn vit_softmax_last_dim_gpu_numerically_stable_for_large_inputs() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let rows = 1usize;
let cols = 3usize;
let input_cpu = vec![1000.0f32, 999.0, 998.0];
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let input_buf = upload_f32(executor.device(), &input_cpu, vec![rows, cols]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_softmax_last_dim_gpu(
session.encoder_mut(),
&mut registry,
device,
&input_buf,
rows as u32,
cols as u32,
)
.expect("softmax");
session.finish().expect("finish");
let got = readback_f32(&out_buf, rows * cols);
for v in &got {
assert!(v.is_finite(), "non-finite: {v}");
}
assert!((got[0] - 0.6652).abs() < 1e-3, "got[0] = {}", got[0]);
assert!((got[1] - 0.2447).abs() < 1e-3, "got[1] = {}", got[1]);
assert!((got[2] - 0.0900).abs() < 1e-3, "got[2] = {}", got[2]);
}
fn attention_scores_cpu(
q: &[f32],
k: &[f32],
batch: usize,
num_heads: usize,
head_dim: usize,
scale: f32,
) -> Vec<f32> {
let mut out = vec![0f32; num_heads * batch * batch];
let stride_seq = num_heads * head_dim;
for h in 0..num_heads {
for q_pos in 0..batch {
for k_pos in 0..batch {
let mut acc = 0f32;
let q_off = q_pos * stride_seq + h * head_dim;
let k_off = k_pos * stride_seq + h * head_dim;
for d in 0..head_dim {
acc += q[q_off + d] * k[k_off + d];
}
out[h * batch * batch + q_pos * batch + k_pos] = acc * scale;
}
}
}
out
}
#[test]
fn permute_021_f32_seq_to_head_major_round_trips() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 2usize;
let num_heads = 2usize;
let head_dim = 4usize;
let n = batch * num_heads * head_dim;
let input: Vec<f32> = (0..n)
.map(|i| {
let dim = i % head_dim;
let head = (i / head_dim) % num_heads;
let b = i / (head_dim * num_heads);
(b * 100 + head * 10 + dim) as f32
})
.collect();
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let in_buf = upload_f32(executor.device(), &input, vec![batch, num_heads, head_dim]);
let out_buf = executor
.device()
.alloc_buffer(n * 4, DType::F32, vec![num_heads, batch, head_dim])
.expect("alloc out");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
permute_021_f32(
session.encoder_mut(),
&mut registry,
executor.device().metal_device(),
&in_buf,
&out_buf,
batch,
num_heads,
head_dim,
)
.expect("permute");
session.finish().expect("finish");
let got = readback_f32(&out_buf, n);
let layout = |head: usize, b: usize, d: usize| head * batch * head_dim + b * head_dim + d;
for d in 0..head_dim {
assert_eq!(got[layout(0, 0, d)], d as f32, "h0 b0 d{d}");
assert_eq!(got[layout(0, 1, d)], (100 + d) as f32, "h0 b1 d{d}");
assert_eq!(got[layout(1, 0, d)], (10 + d) as f32, "h1 b0 d{d}");
assert_eq!(got[layout(1, 1, d)], (110 + d) as f32, "h1 b1 d{d}");
}
}
#[test]
fn vit_attention_scores_gpu_matches_cpu_reference_on_small_input() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 4usize;
let num_heads = 2usize;
let head_dim = 64usize;
let scale = 0.125f32;
let n = batch * num_heads * head_dim;
let q_cpu: Vec<f32> = (0..n).map(|i| ((i as f32) * 0.05).sin() * 0.3).collect();
let k_cpu: Vec<f32> = (0..n).map(|i| ((i as f32) * 0.07).cos() * 0.3).collect();
let expected = attention_scores_cpu(&q_cpu, &k_cpu, batch, num_heads, head_dim, scale);
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let q_buf = upload_f32(executor.device(), &q_cpu, vec![batch, num_heads, head_dim]);
let k_buf = upload_f32(executor.device(), &k_cpu, vec![batch, num_heads, head_dim]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let scores = vit_attention_scores_gpu(
session.encoder_mut(),
&mut registry,
device,
&q_buf,
&k_buf,
batch as u32,
num_heads as u32,
head_dim as u32,
scale,
)
.expect("scores");
session.finish().expect("finish");
let got = readback_f32(&scores, num_heads * batch * batch);
let mut max_diff = 0f32;
let mut fail_count = 0usize;
for (g, e) in got.iter().zip(expected.iter()) {
let d = (g - e).abs();
if d > max_diff {
max_diff = d;
}
if d > 5e-3 {
fail_count += 1;
}
}
assert!(
fail_count == 0,
"{}/{} elements exceeded 5e-3, max_diff = {}",
fail_count,
got.len(),
max_diff
);
}
#[test]
fn vit_attention_scores_gpu_unit_scale_does_not_apply() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 3usize;
let num_heads = 2usize;
let head_dim = 32usize;
let n = batch * num_heads * head_dim;
let qk_cpu: Vec<f32> = (0..n).map(|i| 0.1 + (i as f32) * 0.01).collect();
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let q_buf = upload_f32(executor.device(), &qk_cpu, vec![batch, num_heads, head_dim]);
let k_buf = upload_f32(executor.device(), &qk_cpu, vec![batch, num_heads, head_dim]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let scores = vit_attention_scores_gpu(
session.encoder_mut(),
&mut registry,
device,
&q_buf,
&k_buf,
batch as u32,
num_heads as u32,
head_dim as u32,
1.0,
)
.expect("scores");
session.finish().expect("finish");
let got = readback_f32(&scores, num_heads * batch * batch);
for h in 0..num_heads {
for q in 0..batch {
let mut expected_norm_sq = 0f32;
for d in 0..head_dim {
let v = qk_cpu[q * num_heads * head_dim + h * head_dim + d];
expected_norm_sq += v * v;
}
let got_diag = got[h * batch * batch + q * batch + q];
let diff = (got_diag - expected_norm_sq).abs();
let rel = diff / expected_norm_sq.max(1e-3);
assert!(
rel < 5e-3,
"diag h={h} q={q}: got {got_diag}, want {expected_norm_sq}, rel {rel}"
);
}
}
}
#[test]
fn vit_attention_gpu_matches_cpu_scaled_dot_product_attention() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 32usize;
let num_heads = 2usize;
let head_dim = 64usize;
let scale = 1.0f32 / (head_dim as f32).sqrt();
let n = batch * num_heads * head_dim;
let q_cpu: Vec<f32> = (0..n).map(|i| ((i as f32) * 0.05).sin() * 0.3).collect();
let k_cpu: Vec<f32> = (0..n).map(|i| ((i as f32) * 0.07).cos() * 0.3).collect();
let v_cpu: Vec<f32> = (0..n)
.map(|i| ((i as f32) * 0.11).sin() * ((i as f32) * 0.13).cos() * 0.5)
.collect();
let expected =
attention_cpu(&q_cpu, &k_cpu, &v_cpu, batch, num_heads, head_dim).expect("cpu ref");
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let q_buf = upload_f32(executor.device(), &q_cpu, vec![batch, num_heads, head_dim]);
let k_buf = upload_f32(executor.device(), &k_cpu, vec![batch, num_heads, head_dim]);
let v_buf = upload_f32(executor.device(), &v_cpu, vec![batch, num_heads, head_dim]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let attn = vit_attention_gpu(
session.encoder_mut(),
&mut registry,
device,
&q_buf,
&k_buf,
&v_buf,
batch as u32,
num_heads as u32,
head_dim as u32,
scale,
)
.expect("attn");
session.finish().expect("finish");
let got = readback_f32(&attn, n);
let mut max_diff = 0f32;
let mut fail_count = 0usize;
for (g, e) in got.iter().zip(expected.iter()) {
let d = (g - e).abs();
if d > max_diff {
max_diff = d;
}
if d > 1e-2 {
fail_count += 1;
}
}
let total = got.len();
let frac = (fail_count as f32) / (total as f32);
assert!(
frac < 0.01,
"{}/{} elements ({:.3}%) exceeded 1e-2, max_diff = {}",
fail_count,
total,
frac * 100.0,
max_diff
);
}
#[test]
fn vit_residual_add_gpu_matches_cpu_reference() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let n = 32usize;
let a_cpu: Vec<f32> = (0..n).map(|i| 0.5 + (i as f32) * 0.1).collect();
let b_cpu: Vec<f32> = (0..n).map(|i| -0.3 + (i as f32) * 0.05).collect();
let mut expected = a_cpu.clone();
residual_add_cpu(&mut expected, &b_cpu).expect("cpu");
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let a_buf = upload_f32(executor.device(), &a_cpu, vec![n]);
let b_buf = upload_f32(executor.device(), &b_cpu, vec![n]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_residual_add_gpu(
session.encoder_mut(),
&mut registry,
device,
&a_buf,
&b_buf,
n as u32,
)
.expect("residual_add");
session.finish().expect("finish");
let got = readback_f32(&out_buf, n);
let max_diff = got
.iter()
.zip(expected.iter())
.map(|(g, e)| (g - e).abs())
.fold(0f32, f32::max);
assert!(max_diff < 1e-6, "residual_add max_diff = {max_diff}");
}
#[test]
fn vit_residual_add_gpu_rejects_zero_n() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let a = executor
.device()
.alloc_buffer(16, DType::F32, vec![4])
.expect("a");
let b = executor
.device()
.alloc_buffer(16, DType::F32, vec![4])
.expect("b");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err = vit_residual_add_gpu(session.encoder_mut(), &mut registry, device, &a, &b, 0)
.unwrap_err();
assert!(format!("{err}").contains("must be > 0"));
}
#[test]
fn vit_silu_mul_gpu_matches_cpu_swiglu_gate() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let n = 64usize;
let gate_cpu: Vec<f32> = (0..n).map(|i| ((i as f32) * 0.07).sin()).collect();
let up_cpu: Vec<f32> = (0..n).map(|i| ((i as f32) * 0.05).cos()).collect();
let mut silu_gate = gate_cpu.clone();
silu_in_place(&mut silu_gate);
let mut expected = silu_gate;
elementwise_mul_in_place(&mut expected, &up_cpu).expect("cpu");
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let gate_buf = upload_f32(executor.device(), &gate_cpu, vec![n]);
let up_buf = upload_f32(executor.device(), &up_cpu, vec![n]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
mlx_native::ops::sigmoid_mul::register(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_silu_mul_gpu(
session.encoder_mut(),
&mut registry,
device,
&gate_buf,
&up_buf,
n as u32,
)
.expect("silu_mul");
session.finish().expect("finish");
let got = readback_f32(&out_buf, n);
let max_diff = got
.iter()
.zip(expected.iter())
.map(|(g, e)| (g - e).abs())
.fold(0f32, f32::max);
assert!(max_diff < 1e-5, "silu_mul max_diff = {max_diff}");
}
#[test]
fn vit_silu_mul_gpu_rejects_zero_n() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let g = executor
.device()
.alloc_buffer(16, DType::F32, vec![4])
.expect("g");
let u = executor
.device()
.alloc_buffer(16, DType::F32, vec![4])
.expect("u");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err =
vit_silu_mul_gpu(session.encoder_mut(), &mut registry, device, &g, &u, 0).unwrap_err();
assert!(format!("{err}").contains("must be > 0"));
}
#[test]
fn iter50_bisect_block_forward_gpu_vs_cpu_real_gemma4() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::inference::vision::vit::{
linear_forward as linear_cpu_fn, per_head_rms_norm_forward as per_head_rms_cpu_fn,
rms_norm_forward as rms_norm_cpu_fn, scaled_dot_product_attention as attention_cpu_fn,
};
let path = Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!("skipping: mmproj fixture not found");
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let weights =
LoadedMmprojWeights::load(&gguf, &cfg, MlxDevice::new().expect("dev")).expect("w");
let hidden = cfg.hidden_size as usize;
let num_heads = cfg.num_attention_heads as usize;
let head_dim = hidden / num_heads;
let _intermediate = cfg.intermediate_size as usize;
let eps = cfg.layer_norm_eps;
let scale = 1.0f32 / (head_dim as f32).sqrt();
let batch = 32usize;
let input_cpu: Vec<f32> = (0..batch * hidden)
.map(|i| ((i as f32) * 1e-4).sin() * 0.05)
.collect();
let block = |suffix: &str| -> Vec<f32> {
let buf = weights.block_tensor(0, suffix).unwrap();
weights.tensor_as_f32_owned(buf).unwrap()
};
let l2 = |a: &[f32], b: &[f32]| -> (f32, f32) {
let mut max_d = 0f32;
let mut sum_sq = 0f32;
for (x, y) in a.iter().zip(b.iter()) {
let d = (x - y).abs();
if d > max_d {
max_d = d;
}
sum_sq += d * d;
}
(max_d, sum_sq.sqrt())
};
let ln1_cpu = block("ln1.weight");
let mut ref_after_ln1 = input_cpu.clone();
rms_norm_cpu_fn(&mut ref_after_ln1, &ln1_cpu, hidden, eps).unwrap();
let exec = GraphExecutor::new(MlxDevice::new().expect("e_a"));
let in_a = upload_f32(exec.device(), &input_cpu, vec![batch, hidden]);
let ln1_a = upload_f32(exec.device(), &ln1_cpu, vec![hidden]);
let mut sess = exec.begin().expect("s");
let mut reg = KernelRegistry::new();
let dr: *const MlxDevice = exec.device() as *const _;
let dev: &MlxDevice = unsafe { &*dr };
let stage_a = vit_rms_norm_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&in_a,
&ln1_a,
batch as u32,
hidden as u32,
eps,
)
.unwrap();
sess.finish().unwrap();
let stage_a_cpu = readback_f32(&stage_a, batch * hidden);
let (md, l2v) = l2(&stage_a_cpu, &ref_after_ln1);
eprintln!(
"STAGE A (ln1 rms_norm) : max_diff = {:.6}, l2 = {:.6}",
md, l2v
);
assert!(md < 1e-3, "STAGE A diverges: max_diff = {md}");
let q_w = block("attn_q.weight");
let ref_q = linear_cpu_fn(&ref_after_ln1, &q_w, None, batch, hidden, hidden).unwrap();
let exec = GraphExecutor::new(MlxDevice::new().expect("e_b"));
let cur_b = upload_f32(exec.device(), &ref_after_ln1, vec![batch, hidden]);
let qw_b = upload_f32(exec.device(), &q_w, vec![hidden, hidden]);
let mut sess = exec.begin().expect("s");
let mut reg = KernelRegistry::new();
let dr: *const MlxDevice = exec.device() as *const _;
let dev: &MlxDevice = unsafe { &*dr };
let stage_b = vit_linear_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&cur_b,
&qw_b,
batch as u32,
hidden as u32,
hidden as u32,
)
.unwrap();
sess.finish().unwrap();
let stage_b_cpu = readback_f32(&stage_b, batch * hidden);
let (md, l2v) = l2(&stage_b_cpu, &ref_q);
eprintln!(
"STAGE B (Q linear, uploaded ln1 ref input): max_diff = {:.6}, l2 = {:.6}",
md, l2v
);
assert!(
md < 0.5,
"STAGE B diverges WAY beyond BF16: max_diff = {md}"
);
let exec = GraphExecutor::new(MlxDevice::new().expect("e_bp"));
let cur_bp = upload_f32(exec.device(), &ref_after_ln1, vec![batch, hidden]);
let mut sess = exec.begin().expect("s");
let mut reg = KernelRegistry::new();
let dr: *const MlxDevice = exec.device() as *const _;
let dev: &MlxDevice = unsafe { &*dr };
let qw_native = weights.block_tensor(0, "attn_q.weight").unwrap();
let stage_bp = vit_linear_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&cur_bp,
qw_native,
batch as u32,
hidden as u32,
hidden as u32,
)
.unwrap();
sess.finish().unwrap();
let stage_bp_cpu = readback_f32(&stage_bp, batch * hidden);
let (md, l2v) = l2(&stage_bp_cpu, &ref_q);
eprintln!(
"STAGE B' (Q linear, native weight buffer): max_diff = {:.6}, l2 = {:.6}",
md, l2v
);
let exec = GraphExecutor::new(MlxDevice::new().expect("e_c"));
let in_c = upload_f32(exec.device(), &input_cpu, vec![batch, hidden]);
let ln1_c = upload_f32(exec.device(), &ln1_cpu, vec![hidden]);
let qw_c = upload_f32(exec.device(), &q_w, vec![hidden, hidden]);
let mut sess = exec.begin().expect("s");
let mut reg = KernelRegistry::new();
let dr: *const MlxDevice = exec.device() as *const _;
let dev: &MlxDevice = unsafe { &*dr };
let cur = vit_rms_norm_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&in_c,
&ln1_c,
batch as u32,
hidden as u32,
eps,
)
.unwrap();
sess.encoder_mut().memory_barrier();
let stage_c = vit_linear_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&cur,
&qw_c,
batch as u32,
hidden as u32,
hidden as u32,
)
.unwrap();
sess.finish().unwrap();
let stage_c_cpu = readback_f32(&stage_c, batch * hidden);
let (md, l2v) = l2(&stage_c_cpu, &ref_q);
eprintln!(
"STAGE C (ln1 → Q linear, single session): max_diff = {:.6}, l2 = {:.6}",
md, l2v
);
let k_w = block("attn_k.weight");
let v_w = block("attn_v.weight");
let qn_w = block("attn_q_norm.weight");
let kn_w = block("attn_k_norm.weight");
let ref_k = linear_cpu_fn(&ref_after_ln1, &k_w, None, batch, hidden, hidden).unwrap();
let ref_v = linear_cpu_fn(&ref_after_ln1, &v_w, None, batch, hidden, hidden).unwrap();
let mut ref_q_norm = ref_q.clone();
per_head_rms_cpu_fn(&mut ref_q_norm, &qn_w, batch, num_heads, head_dim, eps).unwrap();
let mut ref_k_norm = ref_k.clone();
per_head_rms_cpu_fn(&mut ref_k_norm, &kn_w, batch, num_heads, head_dim, eps).unwrap();
let exec = GraphExecutor::new(MlxDevice::new().expect("e_d"));
let in_d = upload_f32(exec.device(), &input_cpu, vec![batch, hidden]);
let ln1_d = upload_f32(exec.device(), &ln1_cpu, vec![hidden]);
let qw_d = upload_f32(exec.device(), &q_w, vec![hidden, hidden]);
let kw_d = upload_f32(exec.device(), &k_w, vec![hidden, hidden]);
let vw_d = upload_f32(exec.device(), &v_w, vec![hidden, hidden]);
let qn_d = upload_f32(exec.device(), &qn_w, vec![head_dim]);
let kn_d = upload_f32(exec.device(), &kn_w, vec![head_dim]);
let mut sess = exec.begin().expect("s");
let mut reg = KernelRegistry::new();
let dr: *const MlxDevice = exec.device() as *const _;
let dev: &MlxDevice = unsafe { &*dr };
let cur = vit_rms_norm_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&in_d,
&ln1_d,
batch as u32,
hidden as u32,
eps,
)
.unwrap();
sess.encoder_mut().memory_barrier();
let q = vit_linear_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&cur,
&qw_d,
batch as u32,
hidden as u32,
hidden as u32,
)
.unwrap();
sess.encoder_mut().memory_barrier();
let k = vit_linear_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&cur,
&kw_d,
batch as u32,
hidden as u32,
hidden as u32,
)
.unwrap();
sess.encoder_mut().memory_barrier();
let v = vit_linear_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&cur,
&vw_d,
batch as u32,
hidden as u32,
hidden as u32,
)
.unwrap();
sess.encoder_mut().memory_barrier();
let q_norm_gpu = vit_per_head_rms_norm_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&q,
&qn_d,
batch as u32,
num_heads as u32,
head_dim as u32,
eps,
)
.unwrap();
sess.encoder_mut().memory_barrier();
let k_norm_gpu = vit_per_head_rms_norm_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&k,
&kn_d,
batch as u32,
num_heads as u32,
head_dim as u32,
eps,
)
.unwrap();
sess.finish().unwrap();
let q_norm_cpu_back = readback_f32(&q_norm_gpu, batch * hidden);
let k_norm_cpu_back = readback_f32(&k_norm_gpu, batch * hidden);
let v_back = readback_f32(&v, batch * hidden);
let (md_q, _) = l2(&q_norm_cpu_back, &ref_q_norm);
let (md_k, _) = l2(&k_norm_cpu_back, &ref_k_norm);
let (md_v, _) = l2(&v_back, &ref_v);
eprintln!("STAGE D (Q-norm) : max_diff = {:.6}", md_q);
eprintln!("STAGE D (K-norm) : max_diff = {:.6}", md_k);
eprintln!("STAGE D (V) : max_diff = {:.6}", md_v);
assert!(md_q < 0.5, "STAGE D Q-norm diverges: max_diff = {md_q}");
assert!(md_k < 0.5, "STAGE D K-norm diverges: max_diff = {md_k}");
let ref_attn =
attention_cpu_fn(&ref_q_norm, &ref_k_norm, &ref_v, batch, num_heads, head_dim).unwrap();
let exec = GraphExecutor::new(MlxDevice::new().expect("e_e"));
let qn_e = upload_f32(exec.device(), &ref_q_norm, vec![batch, num_heads, head_dim]);
let kn_e = upload_f32(exec.device(), &ref_k_norm, vec![batch, num_heads, head_dim]);
let v_e = upload_f32(exec.device(), &ref_v, vec![batch, num_heads, head_dim]);
let mut sess = exec.begin().expect("s");
let mut reg = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut reg);
let dr: *const MlxDevice = exec.device() as *const _;
let dev: &MlxDevice = unsafe { &*dr };
let attn_e = vit_attention_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&qn_e,
&kn_e,
&v_e,
batch as u32,
num_heads as u32,
head_dim as u32,
scale,
)
.unwrap();
sess.finish().unwrap();
let attn_e_back = readback_f32(&attn_e, batch * hidden);
let (md_attn, l2_attn) = l2(&attn_e_back, &ref_attn);
eprintln!(
"STAGE E (attention, CPU-uploaded inputs): max_diff = {:.6}, l2 = {:.6}",
md_attn, l2_attn
);
let stat = |name: &str, v: &[f32]| {
let max_abs = v.iter().map(|x| x.abs()).fold(0f32, f32::max);
let mean_abs = v.iter().map(|x| x.abs()).sum::<f32>() / (v.len() as f32);
eprintln!(
" {}: max_abs = {:.4}, mean_abs = {:.4}",
name, max_abs, mean_abs
);
};
stat("ref_q_norm", &ref_q_norm);
stat("ref_k_norm", &ref_k_norm);
stat("ref_v", &ref_v);
stat("ref_attn (CPU)", &ref_attn);
stat("attn (GPU)", &attn_e_back);
let exec = GraphExecutor::new(MlxDevice::new().expect("e_f"));
let in_f = upload_f32(exec.device(), &input_cpu, vec![batch, hidden]);
let ln1_f = upload_f32(exec.device(), &ln1_cpu, vec![hidden]);
let qw_f = upload_f32(exec.device(), &q_w, vec![hidden, hidden]);
let kw_f = upload_f32(exec.device(), &k_w, vec![hidden, hidden]);
let vw_f = upload_f32(exec.device(), &v_w, vec![hidden, hidden]);
let qn_f = upload_f32(exec.device(), &qn_w, vec![head_dim]);
let kn_f = upload_f32(exec.device(), &kn_w, vec![head_dim]);
let mut sess = exec.begin().expect("s");
let mut reg = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut reg);
let dr: *const MlxDevice = exec.device() as *const _;
let dev: &MlxDevice = unsafe { &*dr };
let cur = vit_rms_norm_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&in_f,
&ln1_f,
batch as u32,
hidden as u32,
eps,
)
.unwrap();
sess.encoder_mut().memory_barrier();
let q = vit_linear_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&cur,
&qw_f,
batch as u32,
hidden as u32,
hidden as u32,
)
.unwrap();
sess.encoder_mut().memory_barrier();
let k = vit_linear_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&cur,
&kw_f,
batch as u32,
hidden as u32,
hidden as u32,
)
.unwrap();
sess.encoder_mut().memory_barrier();
let v = vit_linear_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&cur,
&vw_f,
batch as u32,
hidden as u32,
hidden as u32,
)
.unwrap();
sess.encoder_mut().memory_barrier();
let qn = vit_per_head_rms_norm_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&q,
&qn_f,
batch as u32,
num_heads as u32,
head_dim as u32,
eps,
)
.unwrap();
sess.encoder_mut().memory_barrier();
let kn = vit_per_head_rms_norm_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&k,
&kn_f,
batch as u32,
num_heads as u32,
head_dim as u32,
eps,
)
.unwrap();
sess.encoder_mut().memory_barrier();
let attn_f = vit_attention_gpu(
sess.encoder_mut(),
&mut reg,
dev,
&qn,
&kn,
&v,
batch as u32,
num_heads as u32,
head_dim as u32,
scale,
)
.unwrap();
sess.finish().unwrap();
let attn_f_back = readback_f32(&attn_f, batch * hidden);
let (md_attn_f, l2_attn_f) = l2(&attn_f_back, &ref_attn);
eprintln!(
"STAGE F (full chain through attention): max_diff = {:.6}, l2 = {:.6}",
md_attn_f, l2_attn_f
);
}
#[test]
fn vit_avg_pool_2x2_gpu_matches_cpu_reference() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::inference::vision::vit::avg_pool_2x2_spatial as avg_pool_cpu;
let n_side = 4usize;
let hidden = 2usize;
let mut input_cpu = vec![0f32; n_side * n_side * hidden];
for y in 0..n_side {
for x in 0..n_side {
let patch = y * n_side + x;
input_cpu[patch * hidden + 0] = patch as f32;
input_cpu[patch * hidden + 1] = (patch as f32) * 10.0;
}
}
let expected = avg_pool_cpu(&input_cpu, n_side, hidden).unwrap();
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let in_buf = upload_f32(executor.device(), &input_cpu, vec![n_side, n_side, hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_avg_pool_2x2_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
n_side as u32,
hidden as u32,
)
.expect("avg_pool");
session.finish().expect("finish");
let got = readback_f32(&out_buf, 4 * hidden);
let max_diff = got
.iter()
.zip(expected.iter())
.map(|(g, e)| (g - e).abs())
.fold(0f32, f32::max);
assert!(max_diff < 1e-6, "avg_pool max_diff = {max_diff}");
}
#[test]
fn vit_avg_pool_2x2_gpu_gemma4_production_shape() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let n_side = 14usize;
let hidden = 1152usize;
let total = n_side * n_side * hidden;
let input_cpu: Vec<f32> = vec![1.5f32; total];
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let in_buf = upload_f32(executor.device(), &input_cpu, vec![n_side, n_side, hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_avg_pool_2x2_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
n_side as u32,
hidden as u32,
)
.expect("avg_pool");
session.finish().expect("finish");
let got = readback_f32(&out_buf, 49 * hidden);
for v in &got {
assert!((*v - 1.5).abs() < 1e-6, "expected 1.5, got {v}");
}
}
#[test]
fn vit_avg_pool_2x2_gpu_rejects_odd_n_side() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let in_buf = executor
.device()
.alloc_buffer(27 * 4, DType::F32, vec![3, 3, 3])
.expect("a");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err = vit_avg_pool_2x2_gpu(session.encoder_mut(), &mut registry, device, &in_buf, 3, 3)
.unwrap_err();
assert!(format!("{err}").contains("positive and even"));
}
#[test]
fn vit_avg_pool_kxk_k2_byte_identical_to_2x2_path() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let n_side = 4usize;
let hidden = 2usize;
let mut input_cpu = vec![0f32; n_side * n_side * hidden];
for y in 0..n_side {
for x in 0..n_side {
let patch = y * n_side + x;
input_cpu[patch * hidden + 0] = patch as f32;
input_cpu[patch * hidden + 1] = (patch as f32) * 10.0;
}
}
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let in_buf = upload_f32(executor.device(), &input_cpu, vec![n_side, n_side, hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_2x2 = vit_avg_pool_2x2_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
n_side as u32,
hidden as u32,
)
.expect("2x2");
let out_kxk = vit_avg_pool_kxk_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
n_side as u32,
n_side as u32,
2,
hidden as u32,
)
.expect("kxk");
session.finish().expect("finish");
let r_2x2 = readback_f32(&out_2x2, 4 * hidden);
let r_kxk = readback_f32(&out_kxk, 4 * hidden);
assert_eq!(r_2x2.len(), r_kxk.len());
let max_diff = r_2x2
.iter()
.zip(r_kxk.iter())
.map(|(a, b)| (a - b).abs())
.fold(0f32, f32::max);
assert!(
max_diff < 1e-6,
"kxk(k=2) drifted from 2x2 path: max_diff = {max_diff}"
);
}
#[test]
fn vit_avg_pool_kxk_k3_correctness() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let n_x = 6usize;
let n_y = 6usize;
let k = 3usize;
let hidden = 2usize;
let mut input_cpu = vec![0f32; n_x * n_y * hidden];
for y in 0..n_y {
for x in 0..n_x {
let patch = y * n_x + x;
input_cpu[patch * hidden + 0] = patch as f32;
input_cpu[patch * hidden + 1] = patch as f32 + 100.0;
}
}
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let in_buf = upload_f32(executor.device(), &input_cpu, vec![n_y, n_x, hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_avg_pool_kxk_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
n_x as u32,
n_y as u32,
k as u32,
hidden as u32,
)
.expect("kxk k=3");
session.finish().expect("finish");
let out_x = n_x / k;
let out_y = n_y / k;
let got = readback_f32(&out_buf, out_x * out_y * hidden);
let mut expected = vec![0f32; out_x * out_y * hidden];
for oy in 0..out_y {
for ox in 0..out_x {
for h in 0..hidden {
let mut acc = 0f32;
for dy in 0..k {
for dx in 0..k {
let iy = oy * k + dy;
let ix = ox * k + dx;
let patch = iy * n_x + ix;
acc += input_cpu[patch * hidden + h];
}
}
let out_idx = (oy * out_x + ox) * hidden + h;
expected[out_idx] = acc / ((k * k) as f32);
}
}
}
let max_diff = got
.iter()
.zip(expected.iter())
.map(|(a, b)| (a - b).abs())
.fold(0f32, f32::max);
assert!(max_diff < 1e-4, "k=3 max_diff = {max_diff}");
}
#[test]
fn vit_avg_pool_kxk_dimensions_rectangular() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let n_x = 12u32;
let n_y = 9u32;
let k = 3u32;
let hidden = 4u32;
let total = (n_x * n_y * hidden) as usize;
let input_cpu = vec![2.0f32; total];
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let in_buf = upload_f32(
executor.device(),
&input_cpu,
vec![n_y as usize, n_x as usize, hidden as usize],
);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_avg_pool_kxk_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
n_x,
n_y,
k,
hidden,
)
.expect("kxk rect");
session.finish().expect("finish");
let out_x = (n_x / k) as usize;
let out_y = (n_y / k) as usize;
let got = readback_f32(&out_buf, out_x * out_y * hidden as usize);
for v in &got {
assert!((v - 2.0).abs() < 1e-6, "expected 2.0, got {v}");
}
assert_eq!(out_x, 4);
assert_eq!(out_y, 3);
}
#[test]
fn vit_avg_pool_kxk_rejects_non_divisible_n() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let in_buf = executor
.device()
.alloc_buffer(7 * 9 * 4 * 4, DType::F32, vec![7, 9, 4])
.expect("a");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err = vit_avg_pool_kxk_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
9,
7,
3,
4,
)
.unwrap_err();
assert!(format!("{err}").contains("multiples of k"));
}
#[test]
fn vit_clip_gpu_clamps_to_min_max() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let input_cpu: Vec<f32> = vec![-5.0, -1.0, 0.0, 1.0, 5.0, 100.0];
let in_buf = upload_f32(executor.device(), &input_cpu, vec![input_cpu.len()]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_clip_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
input_cpu.len() as u32,
-2.0,
3.0,
)
.expect("clip");
session.finish().expect("finish");
let got = readback_f32(&out_buf, input_cpu.len());
let expected: Vec<f32> = input_cpu.iter().map(|v| v.clamp(-2.0, 3.0)).collect();
assert_eq!(got, expected);
}
#[test]
fn vit_clip_gpu_no_op_on_neg_inf_pos_inf() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let input_cpu: Vec<f32> = vec![-5.0, -1.0, 0.0, 1.0, 5.0];
let in_buf = upload_f32(executor.device(), &input_cpu, vec![input_cpu.len()]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_clip_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
input_cpu.len() as u32,
f32::NEG_INFINITY,
f32::INFINITY,
)
.expect("clip");
session.finish().expect("finish");
let got = readback_f32(&out_buf, input_cpu.len());
assert_eq!(got, input_cpu);
}
#[test]
fn vit_clip_gpu_rejects_min_greater_than_max() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let in_buf = upload_f32(executor.device(), &[1.0, 2.0, 3.0], vec![3]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err = vit_clip_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
3,
5.0,
1.0,
)
.unwrap_err();
assert!(format!("{err}").contains("min_val") || format!("{err}").contains("> max_val"));
}
#[test]
fn gemma4v_clippable_linear_with_and_without_clamp() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::inference::vision::vit::{
gemma4v_clippable_linear_forward as cpu_clip_linear, Gemma4ClippableLinearBounds,
};
let batch = 4usize;
let in_features = 64usize;
let out_features = 32usize;
let input: Vec<f32> = (0..batch * in_features)
.map(|i| ((i as f32) * 0.013).sin() * 2.0)
.collect();
let weight: Vec<f32> = (0..out_features * in_features)
.map(|i| ((i as f32) * 0.017).cos() * 0.05)
.collect();
let bounds_none = Gemma4ClippableLinearBounds::default();
let bounds_input_only = Gemma4ClippableLinearBounds {
input_min: Some(-0.5),
input_max: Some(0.5),
..Default::default()
};
let bounds_both = Gemma4ClippableLinearBounds {
input_min: Some(-1.0),
input_max: Some(1.0),
output_min: Some(-0.05),
output_max: Some(0.05),
};
for (label, bounds) in &[
("none", &bounds_none),
("input_only", &bounds_input_only),
("both", &bounds_both),
] {
let cpu_out =
cpu_clip_linear(&input, &weight, bounds, batch, in_features, out_features)
.expect("cpu clip linear");
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let in_buf = upload_f32(executor.device(), &input, vec![batch, in_features]);
let w_buf = upload_f32(executor.device(), &weight, vec![out_features, in_features]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = gemma4v_clippable_linear_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
&w_buf,
bounds,
batch as u32,
in_features as u32,
out_features as u32,
)
.expect("gpu clip linear");
session.finish().expect("finish");
let gpu_out = readback_f32(&out_buf, batch * out_features);
assert_eq!(gpu_out.len(), cpu_out.len(), "[{label}] length mismatch");
let max_diff = gpu_out
.iter()
.zip(cpu_out.iter())
.map(|(g, c)| (g - c).abs())
.fold(0f32, f32::max);
assert!(max_diff < 5e-3, "[{label}] gpu/cpu max_diff = {max_diff}");
if let Some(mn) = bounds.output_min {
for v in &gpu_out {
assert!(*v >= mn - 1e-4, "[{label}] {v} < output_min {mn}");
}
}
if let Some(mx) = bounds.output_max {
for v in &gpu_out {
assert!(*v <= mx + 1e-4, "[{label}] {v} > output_max {mx}");
}
}
for v in &gpu_out {
assert!(v.is_finite(), "[{label}] non-finite: {v}");
}
}
}
#[test]
fn vit_std_bias_scale_gpu_matches_cpu_reference() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 2usize;
let hidden = 3usize;
let input_cpu = vec![5.0f32, 10.0, 15.0, 10.0, 20.0, 30.0];
let bias_cpu = vec![1.0f32, 2.0, 3.0];
let scale_cpu = vec![10.0f32, 20.0, 30.0];
let expected = vec![40.0f32, 160.0, 360.0, 90.0, 360.0, 810.0];
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let in_buf = upload_f32(executor.device(), &input_cpu, vec![batch, hidden]);
let bias_buf = upload_f32(executor.device(), &bias_cpu, vec![hidden]);
let scale_buf = upload_f32(executor.device(), &scale_cpu, vec![hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_std_bias_scale_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
&bias_buf,
&scale_buf,
batch as u32,
hidden as u32,
)
.expect("std_bias_scale");
session.finish().expect("finish");
let got = readback_f32(&out_buf, batch * hidden);
for (g, e) in got.iter().zip(expected.iter()) {
assert!((g - e).abs() < 1e-5, "got {g} want {e}");
}
}
#[test]
fn vit_std_bias_scale_gpu_zero_bias_unit_scale_is_identity() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 4usize;
let hidden = 8usize;
let input_cpu: Vec<f32> = (0..batch * hidden).map(|i| (i as f32) * 0.1).collect();
let bias_cpu = vec![0.0f32; hidden];
let scale_cpu = vec![1.0f32; hidden];
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let in_buf = upload_f32(executor.device(), &input_cpu, vec![batch, hidden]);
let bias_buf = upload_f32(executor.device(), &bias_cpu, vec![hidden]);
let scale_buf = upload_f32(executor.device(), &scale_cpu, vec![hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out_buf = vit_std_bias_scale_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
&bias_buf,
&scale_buf,
batch as u32,
hidden as u32,
)
.expect("std_bias_scale");
session.finish().expect("finish");
let got = readback_f32(&out_buf, batch * hidden);
for (g, c) in got.iter().zip(input_cpu.iter()) {
assert!((g - c).abs() < 1e-6);
}
}
#[test]
fn vit_std_bias_scale_gpu_rejects_zero_dims() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let buf = executor
.device()
.alloc_buffer(16, DType::F32, vec![4])
.expect("a");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err = vit_std_bias_scale_gpu(
session.encoder_mut(),
&mut registry,
device,
&buf,
&buf,
&buf,
0,
4,
)
.unwrap_err();
assert!(format!("{err}").contains("must be > 0"));
}
#[test]
fn vit_scale_gpu_multiplies_every_element_in_place() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let n = 64usize;
let scalar = 0.25f32;
let cpu_in: Vec<f32> = (0..n).map(|i| (i as f32) * 0.1).collect();
let expected: Vec<f32> = cpu_in.iter().map(|x| x * scalar).collect();
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let buf = upload_f32(executor.device(), &cpu_in, vec![n]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
vit_scale_gpu(
session.encoder_mut(),
&mut registry,
device,
&buf,
n as u32,
scalar,
)
.expect("scale");
session.finish().expect("finish");
let got = readback_f32(&buf, n);
let max_diff = got
.iter()
.zip(expected.iter())
.map(|(g, e)| (g - e).abs())
.fold(0f32, f32::max);
assert!(max_diff < 1e-6, "scale GPU max_diff = {max_diff}");
}
#[test]
fn vit_scale_gpu_by_unit_is_identity() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let n = 16usize;
let cpu_in: Vec<f32> = (0..n).map(|i| (i as f32) * 0.5 - 1.0).collect();
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let buf = upload_f32(executor.device(), &cpu_in, vec![n]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
vit_scale_gpu(
session.encoder_mut(),
&mut registry,
device,
&buf,
n as u32,
1.0,
)
.expect("scale");
session.finish().expect("finish");
let got = readback_f32(&buf, n);
for (g, c) in got.iter().zip(cpu_in.iter()) {
assert!((g - c).abs() < 1e-6);
}
}
#[test]
fn vit_scale_gpu_rejects_zero_n() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let buf = executor
.device()
.alloc_buffer(16, DType::F32, vec![4])
.expect("a");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err =
vit_scale_gpu(session.encoder_mut(), &mut registry, device, &buf, 0, 1.0).unwrap_err();
assert!(format!("{err}").contains("must be > 0"));
}
#[test]
fn warmup_vit_gpu_compiles_all_kernels_real_gemma4() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let path = Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!("skipping: mmproj fixture not found");
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let weights =
LoadedMmprojWeights::load(&gguf, &cfg, MlxDevice::new().expect("dev")).expect("w");
let t0 = std::time::Instant::now();
warmup_vit_gpu(&weights, &cfg).expect("warmup");
eprintln!("warmup_vit_gpu (cold): {:?}", t0.elapsed());
let t1 = std::time::Instant::now();
warmup_vit_gpu(&weights, &cfg).expect("warmup #2");
eprintln!("warmup_vit_gpu (warm): {:?}", t1.elapsed());
}
#[test]
fn compute_vision_embeddings_gpu_multi_image_real_gemma4() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::inference::vision::PreprocessedImage;
let path = Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!("skipping: mmproj fixture not found");
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let weights =
LoadedMmprojWeights::load(&gguf, &cfg, MlxDevice::new().expect("dev")).expect("w");
let img = cfg.image_size as usize;
let make_image = |seed: u32| -> PreprocessedImage {
let mut pixels = vec![0f32; 3 * img * img];
for c in 0..3 {
for y in 0..img {
for x in 0..img {
pixels[c * img * img + y * img + x] = ((c + 1) as f32) * 0.05
+ (y as f32) * 0.001
+ (x as f32) * 0.001
+ (seed as f32) * 0.0001;
}
}
}
PreprocessedImage {
pixel_values: pixels,
target_size: cfg.image_size,
pixel_w: None,
pixel_h: None,
source_label: format!("synthetic-{seed}"),
}
};
let images = vec![make_image(0), make_image(42)];
let head_dim = (cfg.hidden_size / cfg.num_attention_heads) as f32;
let scale = 1.0f32 / head_dim.sqrt();
let t0 = std::time::Instant::now();
let embeddings =
compute_vision_embeddings_gpu(&images, &weights, &cfg, scale).expect("compute");
eprintln!(
"compute_vision_embeddings_gpu × {} images: {:?}",
images.len(),
t0.elapsed()
);
assert_eq!(embeddings.len(), 2);
assert_eq!(embeddings[0].len(), 49 * 2816);
assert_eq!(embeddings[1].len(), 49 * 2816);
for emb in &embeddings {
for v in emb {
assert!(v.is_finite(), "non-finite: {v}");
}
}
let l2: f32 = embeddings[0]
.iter()
.zip(embeddings[1].iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<f32>()
.sqrt();
eprintln!("inter-image L2 = {:.4}", l2);
assert!(
l2 > 1e-2,
"two different images produced identical embeddings"
);
}
#[test]
fn apply_vit_full_forward_gpu_on_real_gemma4_full_pipeline() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let path = Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!("skipping: mmproj fixture not found");
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let weights =
LoadedMmprojWeights::load(&gguf, &cfg, MlxDevice::new().expect("dev")).expect("w");
assert_eq!(cfg.num_hidden_layers, 27);
assert_eq!(cfg.num_patches_side, 14);
let img = cfg.image_size as usize;
let mut pixels = vec![0f32; 3 * img * img];
for c in 0..3 {
for y in 0..img {
for x in 0..img {
pixels[c * img * img + y * img + x] =
((c + 1) as f32) * 0.05 + (y as f32) * 0.001 + (x as f32) * 0.001;
}
}
}
let head_dim = (cfg.hidden_size / cfg.num_attention_heads) as f32;
let scale = 1.0f32 / head_dim.sqrt();
let executor = GraphExecutor::new(MlxDevice::new().expect("dev2"));
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
mlx_native::ops::sigmoid_mul::register(&mut registry);
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
eprintln!("running full ViT GPU forward...");
let t0 = std::time::Instant::now();
let out = apply_vit_full_forward_gpu(
session.encoder_mut(),
&mut registry,
device,
&weights,
&cfg,
&pixels,
scale,
)
.expect("full forward");
session.finish().expect("finish");
let elapsed = t0.elapsed();
eprintln!("full ViT GPU forward done in {:?}", elapsed);
let mm0 = weights.mm_0_weight().expect("mm.0");
let text_hidden = mm0.element_count() / (cfg.hidden_size as usize);
assert_eq!(text_hidden, 2816, "Gemma 4 projector output width");
let n_patches_out = 49;
let total = n_patches_out * text_hidden;
let got = readback_f32(&out, total);
for v in &got {
assert!(v.is_finite(), "non-finite: {v}");
}
for p in 0..n_patches_out {
let row = &got[p * text_hidden..(p + 1) * text_hidden];
let ms: f32 = row.iter().map(|v| v * v).sum::<f32>() / (text_hidden as f32);
assert!(
(ms - 1.0).abs() < 0.05,
"patch {p} post-norm mean(x²) = {ms}, expected ≈ 1.0"
);
}
let p0 = &got[0..text_hidden];
let p_last = &got[(n_patches_out - 1) * text_hidden..n_patches_out * text_hidden];
let l2: f32 = p0
.iter()
.zip(p_last.iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<f32>()
.sqrt();
eprintln!("cross-token L2 = {:.4}", l2);
assert!(l2 > 1e-2, "all output tokens collapsed — diversity lost");
let max_abs = got.iter().map(|v| v.abs()).fold(0f32, f32::max);
let mean_abs = got.iter().map(|v| v.abs()).sum::<f32>() / (got.len() as f32);
eprintln!(
"output: max_abs = {:.4}, mean_abs = {:.4}",
max_abs, mean_abs
);
}
#[test]
fn apply_vit_blocks_loop_gpu_27_blocks_real_gemma4() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let path = Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!("skipping: mmproj fixture not found");
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let weights =
LoadedMmprojWeights::load(&gguf, &cfg, MlxDevice::new().expect("dev")).expect("w");
assert_eq!(cfg.num_hidden_layers, 27);
let hidden = cfg.hidden_size as usize;
let head_dim = hidden / cfg.num_attention_heads as usize;
let scale = 1.0f32 / (head_dim as f32).sqrt();
let batch = 32usize;
let input_cpu: Vec<f32> = (0..batch * hidden)
.map(|i| ((i as f32) * 1e-4).sin() * 0.05)
.collect();
let executor = GraphExecutor::new(MlxDevice::new().expect("dev2"));
let input_buf = upload_f32(executor.device(), &input_cpu, vec![batch, hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
mlx_native::ops::sigmoid_mul::register(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
eprintln!("running 27-block GPU forward...");
let t0 = std::time::Instant::now();
let out = apply_vit_blocks_loop_gpu(
session.encoder_mut(),
&mut registry,
device,
&weights,
&cfg,
&input_buf,
batch as u32,
scale,
)
.expect("blocks loop");
session.finish().expect("finish");
let elapsed = t0.elapsed();
eprintln!("27-block GPU forward done in {:?}", elapsed);
let got = readback_f32(&out, batch * hidden);
let max_abs = got.iter().map(|x| x.abs()).fold(0f32, f32::max);
let mean_abs = got.iter().map(|x| x.abs()).sum::<f32>() / (got.len() as f32);
eprintln!(
"27-block output: max_abs = {:.4}, mean_abs = {:.4}",
max_abs, mean_abs
);
for v in &got {
assert!(v.is_finite(), "non-finite: {v}");
}
assert!(max_abs > 1e-3, "max_abs collapsed to ~0: {max_abs}");
assert!(max_abs < 1e4, "max_abs blew up: {max_abs}");
assert!(
mean_abs > 1e-4 && mean_abs < 1e3,
"mean_abs out of range: {mean_abs}"
);
let token_0 = &got[0..hidden];
let token_last = &got[(batch - 1) * hidden..batch * hidden];
let l2_diff: f32 = token_0
.iter()
.zip(token_last.iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<f32>()
.sqrt();
eprintln!(
"cross-token L2 (token 0 vs token {}) = {:.4}",
batch - 1,
l2_diff
);
assert!(
l2_diff > 1e-2,
"all tokens collapsed to identical output — diversity lost"
);
}
#[test]
#[ignore = "BF16-saturated-softmax drift vs F32 CPU ref is expected; see memory note"]
fn apply_vit_block_forward_gpu_matches_cpu_on_real_gemma4_block0() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let path = Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!("skipping: mmproj fixture not found");
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let weights =
LoadedMmprojWeights::load(&gguf, &cfg, MlxDevice::new().expect("dev")).expect("load");
let hidden = cfg.hidden_size as usize;
let head_dim = hidden / cfg.num_attention_heads as usize;
let scale = 1.0f32 / (head_dim as f32).sqrt();
let batch = 32usize;
let n_input = batch * hidden;
let input_cpu: Vec<f32> = (0..n_input)
.map(|i| ((i as f32) * 1e-4).sin() * 0.05)
.collect();
eprintln!("running CPU block 0 reference (~25-30s)...");
let cpu_t0 = std::time::Instant::now();
let expected =
apply_vit_block_forward_cpu(input_cpu.clone(), &weights, &cfg, 0).expect("cpu");
eprintln!("CPU block 0 done in {:?}", cpu_t0.elapsed());
let executor = GraphExecutor::new(MlxDevice::new().expect("dev2"));
let input_buf = upload_f32(executor.device(), &input_cpu, vec![batch, hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
mlx_native::ops::sigmoid_mul::register(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
eprintln!("running GPU block 0 forward...");
let gpu_t0 = std::time::Instant::now();
let block_out = apply_vit_block_forward_gpu(
session.encoder_mut(),
&mut registry,
device,
&weights,
&cfg,
0,
&input_buf,
batch as u32,
scale,
)
.expect("gpu block forward");
session.finish().expect("finish");
eprintln!("GPU block 0 done in {:?}", gpu_t0.elapsed());
let got = readback_f32(&block_out, n_input);
let mut max_diff = 0f32;
let mut fail_count = 0usize;
for (g, e) in got.iter().zip(expected.iter()) {
let d = (g - e).abs();
if d > max_diff {
max_diff = d;
}
if d > 5e-2 {
fail_count += 1;
}
}
let total = got.len();
let frac = (fail_count as f32) / (total as f32);
eprintln!(
"block 0 GPU vs CPU: {}/{} = {:.3}% > 5e-2, max_diff = {}",
fail_count,
total,
frac * 100.0,
max_diff
);
assert!(
frac < 0.05,
"{:.3}% of elements exceeded 5e-2, max_diff = {}",
frac * 100.0,
max_diff
);
}
#[test]
fn vit_attention_gpu_isolated_head_dim_72_diagnostic() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 32usize;
let num_heads = 16usize;
let head_dim = 72usize;
let scale = 1.0f32 / (head_dim as f32).sqrt();
let n = batch * num_heads * head_dim;
let q_cpu: Vec<f32> = (0..n).map(|i| ((i as f32) * 0.05).sin() * 0.3).collect();
let k_cpu: Vec<f32> = (0..n).map(|i| ((i as f32) * 0.07).cos() * 0.3).collect();
let v_cpu: Vec<f32> = (0..n).map(|i| ((i as f32) * 0.11).sin() * 0.4).collect();
let expected =
attention_cpu(&q_cpu, &k_cpu, &v_cpu, batch, num_heads, head_dim).expect("cpu");
let device = MlxDevice::new().expect("dev");
let executor = GraphExecutor::new(device);
let q_buf = upload_f32(executor.device(), &q_cpu, vec![batch, num_heads, head_dim]);
let k_buf = upload_f32(executor.device(), &k_cpu, vec![batch, num_heads, head_dim]);
let v_buf = upload_f32(executor.device(), &v_cpu, vec![batch, num_heads, head_dim]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let attn = vit_attention_gpu(
session.encoder_mut(),
&mut registry,
device,
&q_buf,
&k_buf,
&v_buf,
batch as u32,
num_heads as u32,
head_dim as u32,
scale,
)
.expect("attn");
session.finish().expect("finish");
let got = readback_f32(&attn, n);
let mut max_diff = 0f32;
for (g, e) in got.iter().zip(expected.iter()) {
let d = (g - e).abs();
if d > max_diff {
max_diff = d;
}
}
eprintln!(
"head_dim=72 isolated attention: max_diff = {}, expected_sample = {}",
max_diff, expected[0]
);
if max_diff > 1.0 {
eprintln!("HYPOTHESIS CONFIRMED: dense_mm_bf16 kernel requires K % 32 == 0");
}
}
#[test]
fn vit_attention_gpu_rejects_small_head_dim() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let q = executor
.device()
.alloc_buffer(2 * 2 * 16 * 4, DType::F32, vec![2, 2, 16])
.expect("q");
let k = executor
.device()
.alloc_buffer(2 * 2 * 16 * 4, DType::F32, vec![2, 2, 16])
.expect("k");
let v = executor
.device()
.alloc_buffer(2 * 2 * 16 * 4, DType::F32, vec![2, 2, 16])
.expect("v");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err = vit_attention_gpu(
session.encoder_mut(),
&mut registry,
device,
&q,
&k,
&v,
2,
2,
16,
1.0,
)
.unwrap_err();
assert!(format!("{err}").contains("head_dim"));
}
#[test]
fn vit_attention_scores_gpu_rejects_small_head_dim() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let q = executor
.device()
.alloc_buffer(2 * 2 * 16 * 4, DType::F32, vec![2, 2, 16])
.expect("a");
let k = executor
.device()
.alloc_buffer(2 * 2 * 16 * 4, DType::F32, vec![2, 2, 16])
.expect("b");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err = vit_attention_scores_gpu(
session.encoder_mut(),
&mut registry,
device,
&q,
&k,
2,
2,
16,
1.0,
)
.unwrap_err();
assert!(format!("{err}").contains("head_dim"));
}
#[test]
fn vit_softmax_last_dim_gpu_rejects_zero_dims() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let input = executor
.device()
.alloc_buffer(16 * 4, DType::F32, vec![4, 4])
.expect("a");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err =
vit_softmax_last_dim_gpu(session.encoder_mut(), &mut registry, device, &input, 0, 4)
.unwrap_err();
assert!(format!("{err}").contains("must be > 0"));
}
#[test]
fn vit_per_head_rms_norm_gpu_rejects_zero_dims() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let input = executor
.device()
.alloc_buffer(64 * 4, DType::F32, vec![64])
.expect("a");
let gain = executor
.device()
.alloc_buffer(8 * 4, DType::F32, vec![8])
.expect("b");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err = vit_per_head_rms_norm_gpu(
session.encoder_mut(),
&mut registry,
device,
&input,
&gain,
0,
8,
8,
1e-6,
)
.unwrap_err();
assert!(format!("{err}").contains("must all be > 0"));
}
fn upload_u32(device: &MlxDevice, data: &[u32], shape: Vec<usize>) -> MlxBuffer {
let bytes = data.len() * 4;
let buf = device
.alloc_buffer(bytes, DType::F32, shape)
.expect("alloc upload u32");
let slice: &mut [u32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut u32, data.len()) };
slice.copy_from_slice(data);
buf
}
#[test]
fn gemma4v_patch_embed_cpu_gpu_parity() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let n_patches = 4usize;
let inner = 64usize;
let hidden = 32usize;
let patches_cpu: Vec<f32> = (0..n_patches * inner)
.map(|i| ((i as f32) * 0.003).sin() * 0.5)
.collect();
let weight_cpu: Vec<f32> = (0..hidden * inner)
.map(|i| ((i as f32) * 0.011).cos() * 0.2)
.collect();
let expect_cpu = gemma4v_patch_embed_cpu(
&patches_cpu,
&weight_cpu,
n_patches as u32,
inner as u32,
hidden as u32,
)
.expect("cpu ref");
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let patches_buf = upload_f32(executor.device(), &patches_cpu, vec![n_patches, inner]);
let weight_buf = upload_f32(executor.device(), &weight_cpu, vec![hidden, inner]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out = gemma4v_patch_embed_gpu(
session.encoder_mut(),
&mut registry,
device,
&patches_buf,
&weight_buf,
n_patches as u32,
inner as u32,
hidden as u32,
)
.expect("gpu dispatch");
session.finish().expect("finish");
let got = readback_f32(&out, n_patches * hidden);
for (i, (g, e)) in got.iter().zip(expect_cpu.iter()).enumerate() {
assert!(
(g - e).abs() < 1e-2,
"patch_embed parity mismatch at {i}: gpu={g} cpu={e}"
);
}
}
#[test]
fn gemma4v_patch_embed_gpu_rejects_zero_n() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let p = executor
.device()
.alloc_buffer(64 * 4, DType::F32, vec![64])
.expect("a");
let w = executor
.device()
.alloc_buffer(64 * 4, DType::F32, vec![64])
.expect("b");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err = gemma4v_patch_embed_gpu(
session.encoder_mut(),
&mut registry,
device,
&p,
&w,
0,
64,
32,
)
.unwrap_err();
assert!(format!("{err}").contains("n_patches must be > 0"));
}
#[test]
fn gemma4v_position_embed_cpu_gpu_parity() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let pos_size = 4usize;
let hidden = 8usize;
let pe_cpu: Vec<f32> = (0..2 * pos_size * hidden)
.map(|i| (i as f32) * 0.1 - 5.0)
.collect();
let pos_x: Vec<u32> = vec![0, 1, 2, 3, 0, 2];
let pos_y: Vec<u32> = vec![0, 0, 1, 2, 3, 1];
let n_patches = pos_x.len() as u32;
let expect_cpu =
gemma4v_pos_embed_cpu(&pos_x, &pos_y, &pe_cpu, pos_size as u32, hidden as u32)
.expect("cpu pos lookup");
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let pe_buf = upload_f32(executor.device(), &pe_cpu, vec![2, pos_size, hidden]);
let posx_buf = upload_u32(executor.device(), &pos_x, vec![pos_x.len()]);
let posy_buf = upload_u32(executor.device(), &pos_y, vec![pos_y.len()]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
mlx_native::ops::gather::register(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out = gemma4v_position_embed_lookup_gpu(
session.encoder_mut(),
&mut registry,
device,
&pe_buf,
&posx_buf,
&posy_buf,
n_patches,
pos_size as u32,
hidden as u32,
)
.expect("gpu pos lookup");
session.finish().expect("finish");
let got = readback_f32(&out, (n_patches as usize) * hidden);
for (i, (g, e)) in got.iter().zip(expect_cpu.iter()).enumerate() {
assert!(
(g - e).abs() < 1e-5,
"pos_embed parity at {i}: gpu={g} cpu={e}"
);
}
}
#[test]
fn gemma4v_position_embed_lookup_gpu_rejects_zero_dims() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let pe = executor
.device()
.alloc_buffer(64 * 4, DType::F32, vec![64])
.expect("a");
let px = executor
.device()
.alloc_buffer(4 * 4, DType::F32, vec![4])
.expect("b");
let py = executor
.device()
.alloc_buffer(4 * 4, DType::F32, vec![4])
.expect("c");
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let err = gemma4v_position_embed_lookup_gpu(
session.encoder_mut(),
&mut registry,
device,
&pe,
&px,
&py,
0,
4,
8,
)
.unwrap_err();
assert!(format!("{err}").contains("must all be > 0"));
}
#[test]
fn gemma4v_apply_position_embed_gpu_adds_pe_to_patch_embeds() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let pos_size = 2usize;
let hidden = 4usize;
let pe_cpu: Vec<f32> = vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 10.0, 20.0, 30.0, 40.0, 50.0, 60.0, 70.0, 80.0,
];
let pos_x: Vec<u32> = vec![0, 1, 0];
let pos_y: Vec<u32> = vec![1, 0, 0];
let n_patches = pos_x.len() as u32;
let patch_embeds_cpu: Vec<f32> = vec![1.0; (n_patches as usize) * hidden];
let pos_emb =
gemma4v_pos_embed_cpu(&pos_x, &pos_y, &pe_cpu, pos_size as u32, hidden as u32).unwrap();
let mut expect: Vec<f32> = patch_embeds_cpu.clone();
for (e, s) in expect.iter_mut().zip(pos_emb.iter()) {
*e += *s;
}
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let pe_buf = upload_f32(executor.device(), &pe_cpu, vec![2, pos_size, hidden]);
let pe_table = pe_buf;
let patch_buf = upload_f32(
executor.device(),
&patch_embeds_cpu,
vec![n_patches as usize, hidden],
);
let posx_buf = upload_u32(executor.device(), &pos_x, vec![pos_x.len()]);
let posy_buf = upload_u32(executor.device(), &pos_y, vec![pos_y.len()]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
mlx_native::ops::gather::register(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let out = gemma4v_apply_position_embed_gpu(
session.encoder_mut(),
&mut registry,
device,
&patch_buf,
&pe_table,
&posx_buf,
&posy_buf,
n_patches,
pos_size as u32,
hidden as u32,
)
.expect("gpu apply pos");
session.finish().expect("finish");
let got = readback_f32(&out, (n_patches as usize) * hidden);
for (i, (g, e)) in got.iter().zip(expect.iter()).enumerate() {
assert!(
(g - e).abs() < 1e-5,
"apply_pos parity at {i}: gpu={g} cpu={e}"
);
}
}
use crate::inference::vision::vit::{Gemma4VisionBlockShape, Gemma4VisionBlockWeights};
fn synth_gemma4v_block_fixture() -> (
Gemma4VisionBlockShape,
usize, // batch (== num_patches == seq_len)
Vec<f32>, // input_layernorm.weight [hidden]
Vec<f32>, // post_attention_layernorm [hidden]
Vec<f32>, // pre_feedforward_layernorm [hidden]
Vec<f32>, // post_feedforward_layernorm [hidden]
Vec<f32>, // q_proj [num_heads*head_dim, hidden]
Vec<f32>, // k_proj [num_kv_heads*head_dim, hidden]
Vec<f32>, // v_proj [num_kv_heads*head_dim, hidden]
Vec<f32>, // o_proj [hidden, num_heads*head_dim]
Vec<f32>, // q_norm [head_dim]
Vec<f32>, // k_norm [head_dim]
Vec<f32>, // gate_proj [intermediate, hidden]
Vec<f32>, // up_proj [intermediate, hidden]
Vec<f32>, // down_proj [hidden, intermediate]
Vec<u32>, // pos_x
Vec<u32>, // pos_y
Vec<f32>, // input [batch, hidden]
) {
let hidden = 128usize;
let num_heads = 4usize;
let num_kv_heads = 2usize;
let head_dim = 32usize;
let intermediate = 64usize;
let batch = 32usize;
let shape = Gemma4VisionBlockShape {
hidden: hidden as u32,
num_heads: num_heads as u32,
num_kv_heads: num_kv_heads as u32,
head_dim: head_dim as u32,
intermediate: intermediate as u32,
rms_norm_eps: 1e-6,
rope_theta: 100.0,
};
let mk = |seed: f32, n: usize| -> Vec<f32> {
(0..n)
.map(|i| ((i as f32) * seed + 0.13).sin() * 0.05)
.collect()
};
let input_ln = mk(0.011, hidden);
let post_attn_ln = mk(0.013, hidden);
let pre_ff_ln = mk(0.017, hidden);
let post_ff_ln = mk(0.019, hidden);
let q_w = mk(0.021, num_heads * head_dim * hidden);
let k_w = mk(0.023, num_kv_heads * head_dim * hidden);
let v_w = mk(0.025, num_kv_heads * head_dim * hidden);
let o_w = mk(0.027, hidden * num_heads * head_dim);
let q_n = mk(0.031, head_dim);
let k_n = mk(0.033, head_dim);
let g_w = mk(0.041, intermediate * hidden);
let u_w = mk(0.043, intermediate * hidden);
let d_w = mk(0.045, hidden * intermediate);
let pos_x: Vec<u32> = (0..batch as u32).collect();
let pos_y: Vec<u32> = (0..batch as u32).rev().collect();
let input: Vec<f32> = (0..batch * hidden)
.map(|i| ((i as f32) * 0.07).cos() * 0.04)
.collect();
(
shape,
batch,
input_ln,
post_attn_ln,
pre_ff_ln,
post_ff_ln,
q_w,
k_w,
v_w,
o_w,
q_n,
k_n,
g_w,
u_w,
d_w,
pos_x,
pos_y,
input,
)
}
#[test]
fn gemma4v_block_forward_4_rmsnorm_count_is_exactly_four() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let (
shape,
batch,
input_ln,
post_attn_ln,
pre_ff_ln,
post_ff_ln,
q_w,
k_w,
v_w,
o_w,
q_n,
k_n,
g_w,
u_w,
d_w,
pos_x,
pos_y,
input,
) = synth_gemma4v_block_fixture();
let baseline = gemma4v_block_forward_cpu(
input.clone(),
&Gemma4VisionBlockWeights {
input_layernorm: &input_ln,
post_attention_layernorm: &post_attn_ln,
pre_feedforward_layernorm: &pre_ff_ln,
post_feedforward_layernorm: &post_ff_ln,
q_proj: &q_w,
k_proj: &k_w,
v_proj: &v_w,
o_proj: &o_w,
q_norm: &q_n,
k_norm: &k_n,
gate_proj: &g_w,
up_proj: &u_w,
down_proj: &d_w,
},
&shape,
&pos_x,
&pos_y,
)
.expect("baseline");
assert_eq!(baseline.len(), batch * shape.hidden as usize);
for which in 0..4usize {
let mut iln = input_ln.clone();
let mut pal = post_attn_ln.clone();
let mut pre = pre_ff_ln.clone();
let mut post = post_ff_ln.clone();
match which {
0 => iln[0] += 0.5,
1 => pal[0] += 0.5,
2 => pre[0] += 0.5,
_ => post[0] += 0.5,
}
let perturbed_out = gemma4v_block_forward_cpu(
input.clone(),
&Gemma4VisionBlockWeights {
input_layernorm: &iln,
post_attention_layernorm: &pal,
pre_feedforward_layernorm: &pre,
post_feedforward_layernorm: &post,
q_proj: &q_w,
k_proj: &k_w,
v_proj: &v_w,
o_proj: &o_w,
q_norm: &q_n,
k_norm: &k_n,
gate_proj: &g_w,
up_proj: &u_w,
down_proj: &d_w,
},
&shape,
&pos_x,
&pos_y,
)
.expect("perturbed");
let max_d = baseline
.iter()
.zip(perturbed_out.iter())
.map(|(a, b)| (a - b).abs())
.fold(0f32, f32::max);
assert!(
max_d > 1e-6,
"norm gain index {which} did not affect output (4-RMSNorm not wired correctly): max|delta| = {max_d}"
);
}
}
#[test]
fn gemma4v_block_forward_gqa_dimensions() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::inference::vision::vit::repeat_kv_cpu;
let batch = 3usize;
let num_kv_heads = 2usize;
let num_kv_groups = 3usize;
let head_dim = 4usize;
let n_in = batch * num_kv_heads * head_dim;
let input: Vec<f32> = (0..n_in).map(|i| i as f32).collect();
let out = repeat_kv_cpu(&input, batch, num_kv_heads, num_kv_groups, head_dim).unwrap();
let num_heads = num_kv_heads * num_kv_groups; assert_eq!(out.len(), batch * num_heads * head_dim);
for b in 0..batch {
for k in 0..num_kv_heads {
let in_base = (b * num_kv_heads + k) * head_dim;
let in_slice = &input[in_base..in_base + head_dim];
for g in 0..num_kv_groups {
let h = k * num_kv_groups + g;
let out_base = (b * num_heads + h) * head_dim;
let out_slice = &out[out_base..out_base + head_dim];
assert_eq!(
in_slice, out_slice,
"repeat_kv at batch={b} kv_head={k} group={g} (out head={h})"
);
}
}
}
}
#[test]
fn gemma4v_block_forward_cpu_gpu_parity() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let (
shape,
batch,
input_ln,
post_attn_ln,
pre_ff_ln,
post_ff_ln,
q_w,
k_w,
v_w,
o_w,
q_n,
k_n,
g_w,
u_w,
d_w,
pos_x,
pos_y,
input,
) = synth_gemma4v_block_fixture();
let hidden = shape.hidden as usize;
let num_heads = shape.num_heads as usize;
let num_kv_heads = shape.num_kv_heads as usize;
let head_dim = shape.head_dim as usize;
let intermediate = shape.intermediate as usize;
let bw_cpu = Gemma4VisionBlockWeights {
input_layernorm: &input_ln,
post_attention_layernorm: &post_attn_ln,
pre_feedforward_layernorm: &pre_ff_ln,
post_feedforward_layernorm: &post_ff_ln,
q_proj: &q_w,
k_proj: &k_w,
v_proj: &v_w,
o_proj: &o_w,
q_norm: &q_n,
k_norm: &k_n,
gate_proj: &g_w,
up_proj: &u_w,
down_proj: &d_w,
};
let cpu_out = gemma4v_block_forward_cpu(input.clone(), &bw_cpu, &shape, &pos_x, &pos_y)
.expect("cpu forward");
let device = MlxDevice::new().expect("device");
let mut tensors: std::collections::HashMap<String, MlxBuffer> =
std::collections::HashMap::new();
let put = |tensors: &mut std::collections::HashMap<String, MlxBuffer>,
dev: &MlxDevice,
key: String,
data: &[f32],
shape: Vec<usize>| {
let bytes = data.len() * 4;
let buf = dev
.alloc_buffer(bytes, DType::F32, shape)
.expect("alloc tensor");
let slice: &mut [f32] = unsafe {
std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, data.len())
};
slice.copy_from_slice(data);
tensors.insert(key, buf);
};
let block_key = |suffix: &str| format!("v.blk.0.{}", suffix);
put(
&mut tensors,
&device,
block_key("ln1.weight"),
&input_ln,
vec![hidden],
);
put(
&mut tensors,
&device,
block_key("attn_post_norm.weight"),
&post_attn_ln,
vec![hidden],
);
put(
&mut tensors,
&device,
block_key("ln2.weight"),
&pre_ff_ln,
vec![hidden],
);
put(
&mut tensors,
&device,
block_key("ffn_post_norm.weight"),
&post_ff_ln,
vec![hidden],
);
put(
&mut tensors,
&device,
block_key("attn_q.weight"),
&q_w,
vec![num_heads * head_dim, hidden],
);
put(
&mut tensors,
&device,
block_key("attn_k.weight"),
&k_w,
vec![num_kv_heads * head_dim, hidden],
);
put(
&mut tensors,
&device,
block_key("attn_v.weight"),
&v_w,
vec![num_kv_heads * head_dim, hidden],
);
put(
&mut tensors,
&device,
block_key("attn_out.weight"),
&o_w,
vec![hidden, num_heads * head_dim],
);
put(
&mut tensors,
&device,
block_key("attn_q_norm.weight"),
&q_n,
vec![head_dim],
);
put(
&mut tensors,
&device,
block_key("attn_k_norm.weight"),
&k_n,
vec![head_dim],
);
put(
&mut tensors,
&device,
block_key("ffn_gate.weight"),
&g_w,
vec![intermediate, hidden],
);
put(
&mut tensors,
&device,
block_key("ffn_up.weight"),
&u_w,
vec![intermediate, hidden],
);
put(
&mut tensors,
&device,
block_key("ffn_down.weight"),
&d_w,
vec![hidden, intermediate],
);
let weights = LoadedMmprojWeights::from_tensors_for_test(tensors, device);
let device = MlxDevice::new().expect("device2");
let executor = GraphExecutor::new(device);
let in_buf = upload_f32(executor.device(), &input, vec![batch, hidden]);
let upload_u32_typed = |dev: &MlxDevice, data: &[u32]| -> MlxBuffer {
let buf = dev
.alloc_buffer(data.len() * 4, DType::U32, vec![data.len()])
.expect("alloc u32 typed");
let slice: &mut [u32] = unsafe {
std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut u32, data.len())
};
slice.copy_from_slice(data);
buf
};
let px_buf = upload_u32_typed(executor.device(), &pos_x);
let py_buf = upload_u32_typed(executor.device(), &pos_y);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
mlx_native::ops::sigmoid_mul::register(&mut registry);
mlx_native::ops::softmax::register(&mut registry);
mlx_native::ops::vision_2d_rope::register(&mut registry);
mlx_native::ops::gelu::register(&mut registry);
mlx_native::ops::gather::register(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let shape_gpu = Gemma4VisionBlockShapeGpu {
hidden: shape.hidden,
num_heads: shape.num_heads,
num_kv_heads: shape.num_kv_heads,
head_dim: shape.head_dim,
intermediate: shape.intermediate,
rms_norm_eps: shape.rms_norm_eps,
rope_theta: shape.rope_theta,
};
let out = gemma4v_block_forward_gpu(
session.encoder_mut(),
&mut registry,
device,
&weights,
&shape_gpu,
0,
&in_buf,
&px_buf,
&py_buf,
batch as u32,
)
.expect("gpu block forward");
session.finish().expect("finish");
let gpu_out = readback_f32(&out, batch * hidden);
let dot: f32 = cpu_out.iter().zip(gpu_out.iter()).map(|(a, b)| a * b).sum();
let na: f32 = cpu_out.iter().map(|v| v * v).sum::<f32>().sqrt();
let nb: f32 = gpu_out.iter().map(|v| v * v).sum::<f32>().sqrt();
let cos = dot / (na * nb).max(1e-30);
let max_abs = cpu_out
.iter()
.zip(gpu_out.iter())
.map(|(a, b)| (a - b).abs())
.fold(0f32, f32::max);
eprintln!(
"gemma4v_block parity: cos={cos:.6}, max|abs|={max_abs:.6}, |cpu|={na:.4}, |gpu|={nb:.4}"
);
assert!(
cos > 0.999,
"gemma4v_block_forward CPU↔GPU cosine = {cos} < 0.999 (max|abs| = {max_abs})"
);
}
#[allow(clippy::too_many_arguments)]
fn synth_gemma4v_full_fixture(
n_x: u32,
n_y: u32,
) -> (
LoadedMmprojWeights,
MmprojConfig,
Vec<f32>,
Vec<u32>,
Vec<u32>,
) {
let hidden: u32 = 128;
let num_heads: u32 = 4;
let num_kv_heads: u32 = 2;
let head_dim: u32 = hidden / num_heads; let intermediate: u32 = 64;
let patch_size: u32 = 4; let inner: u32 = patch_size * patch_size * 3;
let num_layers: u32 = 2;
let text_hidden: u32 = 64;
let pos_size: u32 = (n_x.max(n_y)).max(64);
let n_patches = n_x * n_y;
let mk = |seed: f32, n: usize| -> Vec<f32> {
(0..n)
.map(|i| ((i as f32) * seed + 0.13).sin() * 0.05)
.collect()
};
let device = MlxDevice::new().expect("synth dev");
let put = |tensors: &mut std::collections::HashMap<String, MlxBuffer>,
dev: &MlxDevice,
key: String,
data: &[f32],
shape: Vec<usize>| {
let bytes = data.len() * 4;
let buf = dev
.alloc_buffer(bytes, DType::F32, shape)
.expect("alloc tensor");
let slice: &mut [f32] = unsafe {
std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, data.len())
};
slice.copy_from_slice(data);
tensors.insert(key, buf);
};
let mut tensors: std::collections::HashMap<String, MlxBuffer> =
std::collections::HashMap::new();
let patch_w = mk(0.011, (hidden * inner) as usize);
put(
&mut tensors,
&device,
"v.patch_embd.weight".to_string(),
&patch_w,
vec![hidden as usize, inner as usize],
);
let pe_table = mk(0.013, (2 * pos_size * hidden) as usize);
put(
&mut tensors,
&device,
"v.position_embd.weight".to_string(),
&pe_table,
vec![2usize, pos_size as usize, hidden as usize],
);
let std_bias = mk(0.014, hidden as usize);
put(
&mut tensors,
&device,
"v.std_bias".to_string(),
&std_bias,
vec![hidden as usize],
);
let std_scale: Vec<f32> = (0..hidden).map(|_| 1.0).collect();
put(
&mut tensors,
&device,
"v.std_scale".to_string(),
&std_scale,
vec![hidden as usize],
);
let mm0 = mk(0.017, (text_hidden * hidden) as usize);
put(
&mut tensors,
&device,
"mm.0.weight".to_string(),
&mm0,
vec![text_hidden as usize, hidden as usize],
);
for layer_idx in 0..num_layers {
let prefix = format!("v.blk.{}.", layer_idx);
put(
&mut tensors,
&device,
format!("{prefix}ln1.weight"),
&mk(0.021 + layer_idx as f32 * 0.001, hidden as usize),
vec![hidden as usize],
);
put(
&mut tensors,
&device,
format!("{prefix}attn_post_norm.weight"),
&mk(0.022 + layer_idx as f32 * 0.001, hidden as usize),
vec![hidden as usize],
);
put(
&mut tensors,
&device,
format!("{prefix}ln2.weight"),
&mk(0.023 + layer_idx as f32 * 0.001, hidden as usize),
vec![hidden as usize],
);
put(
&mut tensors,
&device,
format!("{prefix}ffn_post_norm.weight"),
&mk(0.024 + layer_idx as f32 * 0.001, hidden as usize),
vec![hidden as usize],
);
put(
&mut tensors,
&device,
format!("{prefix}attn_q.weight"),
&mk(
0.031 + layer_idx as f32 * 0.001,
(num_heads * head_dim * hidden) as usize,
),
vec![(num_heads * head_dim) as usize, hidden as usize],
);
put(
&mut tensors,
&device,
format!("{prefix}attn_k.weight"),
&mk(
0.033 + layer_idx as f32 * 0.001,
(num_kv_heads * head_dim * hidden) as usize,
),
vec![(num_kv_heads * head_dim) as usize, hidden as usize],
);
put(
&mut tensors,
&device,
format!("{prefix}attn_v.weight"),
&mk(
0.035 + layer_idx as f32 * 0.001,
(num_kv_heads * head_dim * hidden) as usize,
),
vec![(num_kv_heads * head_dim) as usize, hidden as usize],
);
put(
&mut tensors,
&device,
format!("{prefix}attn_out.weight"),
&mk(
0.037 + layer_idx as f32 * 0.001,
(hidden * num_heads * head_dim) as usize,
),
vec![hidden as usize, (num_heads * head_dim) as usize],
);
put(
&mut tensors,
&device,
format!("{prefix}attn_q_norm.weight"),
&mk(0.041 + layer_idx as f32 * 0.001, head_dim as usize),
vec![head_dim as usize],
);
put(
&mut tensors,
&device,
format!("{prefix}attn_k_norm.weight"),
&mk(0.043 + layer_idx as f32 * 0.001, head_dim as usize),
vec![head_dim as usize],
);
put(
&mut tensors,
&device,
format!("{prefix}ffn_gate.weight"),
&mk(
0.051 + layer_idx as f32 * 0.001,
(intermediate * hidden) as usize,
),
vec![intermediate as usize, hidden as usize],
);
put(
&mut tensors,
&device,
format!("{prefix}ffn_up.weight"),
&mk(
0.053 + layer_idx as f32 * 0.001,
(intermediate * hidden) as usize,
),
vec![intermediate as usize, hidden as usize],
);
put(
&mut tensors,
&device,
format!("{prefix}ffn_down.weight"),
&mk(
0.055 + layer_idx as f32 * 0.001,
(hidden * intermediate) as usize,
),
vec![hidden as usize, intermediate as usize],
);
}
let weights = LoadedMmprojWeights::from_tensors_for_test(tensors, device);
let patches: Vec<f32> = (0..(n_patches * inner) as usize)
.map(|i| ((i as f32) * 0.07).cos() * 0.04)
.collect();
let mut pos_x: Vec<u32> = Vec::with_capacity(n_patches as usize);
let mut pos_y: Vec<u32> = Vec::with_capacity(n_patches as usize);
for y in 0..n_y {
for x in 0..n_x {
pos_x.push(x);
pos_y.push(y);
}
}
let cfg = MmprojConfig {
image_size: 0, patch_size,
num_patches_side: 0, hidden_size: hidden,
intermediate_size: intermediate,
num_attention_heads: num_heads,
num_hidden_layers: num_layers,
layer_norm_eps: 1e-6,
projector: crate::inference::vision::mmproj::ProjectorType::Mlp,
image_mean: [0.5, 0.5, 0.5],
image_std: [0.5, 0.5, 0.5],
spatial_merge_size: None,
projection_dim: None,
deepstack_indexes: None,
};
(weights, cfg, patches, pos_x, pos_y)
}
#[test]
fn gemma4v_apply_full_forward_gpu_synthetic_n_36() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let n_x: u32 = 6;
let n_y: u32 = 6;
let (weights, cfg, patches, pos_x, pos_y) = synth_gemma4v_full_fixture(n_x, n_y);
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
mlx_native::ops::sigmoid_mul::register(&mut registry);
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device_ref: &MlxDevice = unsafe { &*device_ref };
let buf = gemma4v_apply_full_forward_gpu(
session.encoder_mut(),
&mut registry,
device_ref,
&weights,
&cfg,
&patches,
&pos_x,
&pos_y,
n_x,
n_y,
)
.expect("full forward synthetic");
session.finish().expect("finish");
let pooled_n = ((n_x / 3) * (n_y / 3)) as usize; let text_hidden = 64usize;
let expected_len = pooled_n * text_hidden;
let slice: &[f32] = buf.as_slice::<f32>().expect("readback");
assert_eq!(slice.len(), expected_len, "output len mismatch");
for v in slice {
assert!(v.is_finite(), "non-finite output: {v}");
}
}
#[test]
fn gemma4v_apply_full_forward_gpu_synthetic_rectangular_n_54() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let n_x: u32 = 9;
let n_y: u32 = 6;
let (weights, cfg, patches, pos_x, pos_y) = synth_gemma4v_full_fixture(n_x, n_y);
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
mlx_native::ops::softmax::register(&mut registry);
mlx_native::ops::sigmoid_mul::register(&mut registry);
register_vit_custom_shaders(&mut registry);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device_ref: &MlxDevice = unsafe { &*device_ref };
let buf = gemma4v_apply_full_forward_gpu(
session.encoder_mut(),
&mut registry,
device_ref,
&weights,
&cfg,
&patches,
&pos_x,
&pos_y,
n_x,
n_y,
)
.expect("full forward synthetic rect");
session.finish().expect("finish");
let pooled_n = ((n_x / 3) * (n_y / 3)) as usize; let text_hidden = 64usize;
let expected_len = pooled_n * text_hidden;
let slice: &[f32] = buf.as_slice::<f32>().expect("readback");
assert_eq!(slice.len(), expected_len, "rect output len mismatch");
for v in slice {
assert!(v.is_finite(), "non-finite rect output: {v}");
}
}
#[test]
fn compute_vision_embeddings_gpu_dispatch_gemma4v_routes_correctly() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::inference::vision::mmproj::ArchProfile;
let n_x: u32 = 6;
let n_y: u32 = 6;
let (weights, cfg, patches, pos_x, pos_y) = synth_gemma4v_full_fixture(n_x, n_y);
let img = Gemma4vPreprocessedImage {
patches,
pos_x,
pos_y,
n_x,
n_y,
source_label: "synthetic".to_string(),
};
let inputs = vec![VisionInput::Gemma4v(img)];
let result = compute_vision_embeddings_gpu_dispatch(
&inputs,
ArchProfile::Gemma4Siglip,
&weights,
&cfg,
1.0,
)
.expect("dispatch");
assert_eq!(result.len(), 1);
let pooled_n = ((n_x / 3) * (n_y / 3)) as usize;
let text_hidden = 64usize;
assert_eq!(result[0].len(), pooled_n * text_hidden);
for v in &result[0] {
assert!(v.is_finite());
}
}
#[test]
fn compute_vision_embeddings_gpu_dispatch_rejects_arch_mismatch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::inference::vision::mmproj::ArchProfile;
let n_x: u32 = 6;
let n_y: u32 = 6;
let (weights, cfg, patches, pos_x, pos_y) = synth_gemma4v_full_fixture(n_x, n_y);
let img = Gemma4vPreprocessedImage {
patches,
pos_x,
pos_y,
n_x,
n_y,
source_label: "synthetic".to_string(),
};
let inputs = vec![VisionInput::Gemma4v(img)];
let err = compute_vision_embeddings_gpu_dispatch(
&inputs,
ArchProfile::ClipClassic,
&weights,
&cfg,
1.0,
)
.unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains("Gemma4v") && msg.contains("ClipClassic"),
"expected arch-mismatch error message, got: {msg}"
);
}
#[test]
fn compute_vision_embeddings_gpu_dispatch_empty_inputs_ok() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::inference::vision::mmproj::ArchProfile;
let n_x: u32 = 6;
let n_y: u32 = 6;
let (weights, cfg, _, _, _) = synth_gemma4v_full_fixture(n_x, n_y);
let inputs: Vec<VisionInput> = Vec::new();
let result = compute_vision_embeddings_gpu_dispatch(
&inputs,
ArchProfile::Gemma4Siglip,
&weights,
&cfg,
1.0,
)
.expect("empty dispatch");
assert!(result.is_empty());
}
#[test]
#[ignore]
fn iter124_parity_probe_dump_four_dots_real_gemma4() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if super::super::vit_dump::resolve_dump_dir()
.expect("resolve dump dir")
.is_none()
{
eprintln!("skip: HF2Q_VIT_DUMP unset (parity probe runner is opt-in)");
return;
}
let mmproj_path = Path::new(GEMMA4_MMPROJ_PATH);
if !mmproj_path.exists() {
eprintln!("skip: mmproj fixture not found at {}", GEMMA4_MMPROJ_PATH);
return;
}
let img_path =
Path::new("/opt/hf2q/tests/fixtures/vision/four_dots_in_corners_128x128.png");
let img_bytes = std::fs::read(img_path).expect("read four-dots fixture");
let pre = crate::inference::vision::preprocess::preprocess_gemma4v(
&img_bytes,
&crate::inference::vision::preprocess::GEMMA4V_PREPROCESS_DEFAULT,
)
.expect("preprocess");
let gguf = GgufFile::open(mmproj_path).expect("open mmproj");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let device = MlxDevice::new().expect("device");
let weights = LoadedMmprojWeights::load(&gguf, &cfg, device).expect("load mmproj");
let img = Gemma4vPreprocessedImage {
patches: pre.patches,
pos_x: pre.pos_x,
pos_y: pre.pos_y,
n_x: pre.n_x,
n_y: pre.n_y,
source_label: "four_dots_in_corners_128x128.png".to_string(),
};
let out = compute_vision_embeddings_gpu_gemma4v(std::slice::from_ref(&img), &weights, &cfg)
.expect("forward");
assert_eq!(out.len(), 1);
assert!(!out[0].is_empty());
let dir = super::super::vit_dump::resolve_dump_dir()
.expect("resolve dump dir")
.expect("dump dir set");
eprintln!(
"iter124 parity probe: dumps written to {} (n_x={}, n_y={})",
dir.display(),
img.n_x,
img.n_y
);
}
}