#![allow(dead_code)]
use anyhow::{anyhow, Result};
pub fn residual_add(a: &mut [f32], b: &[f32]) -> Result<()> {
if a.len() != b.len() {
return Err(anyhow!(
"residual_add: a len {} != b len {}",
a.len(),
b.len()
));
}
for (ai, bi) in a.iter_mut().zip(b.iter()) {
*ai += *bi;
}
Ok(())
}
pub fn position_embed_add(patch_embeds: &mut [f32], pos_embeds: &[f32]) -> Result<()> {
if patch_embeds.len() != pos_embeds.len() {
return Err(anyhow!(
"position_embed_add: patch len {} != pos len {}",
patch_embeds.len(),
pos_embeds.len()
));
}
for (p, pe) in patch_embeds.iter_mut().zip(pos_embeds.iter()) {
*p += *pe;
}
Ok(())
}
pub fn layer_norm_forward(
input: &mut [f32],
gamma: &[f32],
beta: &[f32],
hidden: usize,
eps: f32,
) -> Result<()> {
if hidden == 0 {
return Err(anyhow!("layer_norm_forward: hidden must be > 0"));
}
if input.len() % hidden != 0 {
return Err(anyhow!(
"layer_norm_forward: input len {} not divisible by hidden {}",
input.len(),
hidden
));
}
if gamma.len() != hidden {
return Err(anyhow!(
"layer_norm_forward: gamma len {} != hidden {}",
gamma.len(),
hidden
));
}
if beta.len() != hidden {
return Err(anyhow!(
"layer_norm_forward: beta len {} != hidden {}",
beta.len(),
hidden
));
}
let n_rows = input.len() / hidden;
let inv_hidden = 1.0f32 / (hidden as f32);
for row in 0..n_rows {
let off = row * hidden;
let slice = &mut input[off..off + hidden];
let mut sum = 0.0f32;
for &v in slice.iter() {
sum += v;
}
let mean = sum * inv_hidden;
let mut var_sum = 0.0f32;
for &v in slice.iter() {
let d = v - mean;
var_sum += d * d;
}
let variance = var_sum * inv_hidden;
let inv_std = 1.0f32 / (variance + eps).sqrt();
for (i, v) in slice.iter_mut().enumerate() {
*v = (*v - mean) * inv_std * gamma[i] + beta[i];
}
}
Ok(())
}
pub fn rms_norm_forward(input: &mut [f32], gamma: &[f32], hidden: usize, eps: f32) -> Result<()> {
if hidden == 0 {
return Err(anyhow!("rms_norm_forward: hidden must be > 0"));
}
if input.len() % hidden != 0 {
return Err(anyhow!(
"rms_norm_forward: input len {} not divisible by hidden {}",
input.len(),
hidden
));
}
if gamma.len() != hidden {
return Err(anyhow!(
"rms_norm_forward: gamma len {} != hidden {}",
gamma.len(),
hidden
));
}
let n_rows = input.len() / hidden;
let inv_hidden = 1.0f32 / (hidden as f32);
for row in 0..n_rows {
let off = row * hidden;
let slice = &mut input[off..off + hidden];
let mut sq_sum = 0.0f32;
for &v in slice.iter() {
sq_sum += v * v;
}
let inv_rms = 1.0f32 / (sq_sum * inv_hidden + eps).sqrt();
for (i, v) in slice.iter_mut().enumerate() {
*v = *v * inv_rms * gamma[i];
}
}
Ok(())
}
pub fn per_head_rms_norm_forward(
input: &mut [f32],
gamma: &[f32],
batch: usize,
num_heads: usize,
head_dim: usize,
eps: f32,
) -> Result<()> {
if batch == 0 || num_heads == 0 || head_dim == 0 {
return Err(anyhow!(
"per_head_rms_norm_forward: batch ({}), num_heads ({}), head_dim ({}) must all be > 0",
batch,
num_heads,
head_dim
));
}
let expected_input_len = batch * num_heads * head_dim;
if input.len() != expected_input_len {
return Err(anyhow!(
"per_head_rms_norm_forward: input len {} != batch*num_heads*head_dim = {}*{}*{} = {}",
input.len(),
batch,
num_heads,
head_dim,
expected_input_len
));
}
if gamma.len() != head_dim {
return Err(anyhow!(
"per_head_rms_norm_forward: gamma len {} != head_dim {}",
gamma.len(),
head_dim
));
}
rms_norm_forward(input, gamma, head_dim, eps)
}
pub fn softmax_last_dim(input: &mut [f32], hidden: usize) -> Result<()> {
if hidden == 0 {
return Err(anyhow!("softmax_last_dim: hidden must be > 0"));
}
if input.len() % hidden != 0 {
return Err(anyhow!(
"softmax_last_dim: input len {} not divisible by hidden {}",
input.len(),
hidden
));
}
let n_rows = input.len() / hidden;
for row in 0..n_rows {
let off = row * hidden;
let slice = &mut input[off..off + hidden];
let mut m = f32::NEG_INFINITY;
for &v in slice.iter() {
if v > m {
m = v;
}
}
let mut sum = 0.0f32;
for v in slice.iter_mut() {
let e = (*v - m).exp();
*v = e;
sum += e;
}
let inv = 1.0f32 / sum;
for v in slice.iter_mut() {
*v *= inv;
}
}
Ok(())
}
pub fn silu_in_place(input: &mut [f32]) {
for v in input.iter_mut() {
let x = *v;
*v = x / (1.0 + (-x).exp());
}
}
pub fn elementwise_mul_in_place(a: &mut [f32], b: &[f32]) -> Result<()> {
if a.len() != b.len() {
return Err(anyhow!(
"elementwise_mul_in_place: a len {} != b len {}",
a.len(),
b.len()
));
}
for (ai, bi) in a.iter_mut().zip(b.iter()) {
*ai *= *bi;
}
Ok(())
}
pub fn gelu_tanh_approx(input: &mut [f32]) {
const C: f32 = 0.7978845608028654_f32;
const K: f32 = 0.044715_f32;
for v in input.iter_mut() {
let x = *v;
let inner = C * (x + K * x * x * x);
*v = 0.5 * x * (1.0 + inner.tanh());
}
}
pub fn patch_embed_forward(
pixel_values: &[f32],
patch_embd_weight: &[f32],
patch_embd_bias: Option<&[f32]>,
image_size: u32,
patch_size: u32,
hidden: u32,
) -> Result<Vec<f32>> {
if patch_size == 0 {
return Err(anyhow!("patch_size must be > 0"));
}
if image_size % patch_size != 0 {
return Err(anyhow!(
"image_size ({}) must be divisible by patch_size ({})",
image_size,
patch_size
));
}
let num_patches_side = image_size / patch_size;
let p = patch_size as usize;
let h = hidden as usize;
let hw = (image_size as usize) * (image_size as usize);
let expected_pixels = 3 * hw;
if pixel_values.len() != expected_pixels {
return Err(anyhow!(
"pixel_values len {} != expected 3*{}*{} = {}",
pixel_values.len(),
image_size,
image_size,
expected_pixels
));
}
let expected_w = h * 3 * p * p;
if patch_embd_weight.len() != expected_w {
return Err(anyhow!(
"patch_embd_weight len {} != expected {}*3*{}*{} = {}",
patch_embd_weight.len(),
hidden,
patch_size,
patch_size,
expected_w
));
}
if let Some(b) = patch_embd_bias {
if b.len() != h {
return Err(anyhow!(
"patch_embd_bias len {} != hidden {}",
b.len(),
hidden
));
}
}
let w = image_size as usize;
let nps = num_patches_side as usize;
let num_patches = nps * nps;
let mut out = vec![0f32; num_patches * h];
let ws_oc = 3 * p * p; let ws_ic = p * p; let ws_y = p;
let ps_c = hw; let ps_y = w;
for py in 0..nps {
let y0 = py * p;
for px in 0..nps {
let x0 = px * p;
let patch_idx = py * nps + px;
let out_base = patch_idx * h;
for oc in 0..h {
let mut acc: f32 = match patch_embd_bias {
Some(b) => b[oc],
None => 0.0,
};
let w_base = oc * ws_oc;
for ic in 0..3usize {
let w_ic_base = w_base + ic * ws_ic;
let p_ic_base = ic * ps_c;
for dy in 0..p {
let w_row_base = w_ic_base + dy * ws_y;
let p_row_base = p_ic_base + (y0 + dy) * ps_y + x0;
for dx in 0..p {
acc +=
pixel_values[p_row_base + dx] * patch_embd_weight[w_row_base + dx];
}
}
}
out[out_base + oc] = acc;
}
}
}
Ok(out)
}
pub fn patch_embed_forward_hw(
pixel_values: &[f32],
patch_embd_weight: &[f32],
patch_embd_bias: Option<&[f32]>,
pixel_h: u32,
pixel_w: u32,
patch_size: u32,
hidden: u32,
) -> Result<Vec<f32>> {
if patch_size == 0 {
return Err(anyhow!("patch_size must be > 0"));
}
if pixel_h % patch_size != 0 {
return Err(anyhow!(
"pixel_h ({pixel_h}) must be divisible by patch_size ({patch_size})"
));
}
if pixel_w % patch_size != 0 {
return Err(anyhow!(
"pixel_w ({pixel_w}) must be divisible by patch_size ({patch_size})"
));
}
let p = patch_size as usize;
let h_dim = hidden as usize;
let hw = (pixel_h as usize) * (pixel_w as usize);
let expected_pixels = 3 * hw;
if pixel_values.len() != expected_pixels {
return Err(anyhow!(
"pixel_values len {} != expected 3*{pixel_h}*{pixel_w} = {expected_pixels}",
pixel_values.len()
));
}
let expected_w = h_dim * 3 * p * p;
if patch_embd_weight.len() != expected_w {
return Err(anyhow!(
"patch_embd_weight len {} != expected {hidden}*3*{patch_size}*{patch_size} = {expected_w}",
patch_embd_weight.len()
));
}
if let Some(b) = patch_embd_bias {
if b.len() != h_dim {
return Err(anyhow!(
"patch_embd_bias len {} != hidden {hidden}",
b.len()
));
}
}
let nps_y = (pixel_h / patch_size) as usize;
let nps_x = (pixel_w / patch_size) as usize;
let num_patches = nps_y * nps_x;
let mut out = vec![0f32; num_patches * h_dim];
let ws_oc = 3 * p * p;
let ws_ic = p * p;
let ws_y = p;
let ps_c = hw; let ps_y = pixel_w as usize;
for py in 0..nps_y {
let y0 = py * p;
for px in 0..nps_x {
let x0 = px * p;
let patch_idx = py * nps_x + px;
let out_base = patch_idx * h_dim;
for oc in 0..h_dim {
let mut acc: f32 = match patch_embd_bias {
Some(b) => b[oc],
None => 0.0,
};
let w_base = oc * ws_oc;
for ic in 0..3usize {
let w_ic_base = w_base + ic * ws_ic;
let p_ic_base = ic * ps_c;
for dy in 0..p {
let w_row_base = w_ic_base + dy * ws_y;
let p_row_base = p_ic_base + (y0 + dy) * ps_y + x0;
for dx in 0..p {
acc +=
pixel_values[p_row_base + dx] * patch_embd_weight[w_row_base + dx];
}
}
}
out[out_base + oc] = acc;
}
}
}
Ok(out)
}
pub fn linear_forward(
input: &[f32],
weight: &[f32],
bias: Option<&[f32]>,
batch: usize,
in_features: usize,
out_features: usize,
) -> Result<Vec<f32>> {
if batch == 0 || in_features == 0 || out_features == 0 {
return Err(anyhow!(
"linear_forward: batch ({}), in_features ({}), out_features ({}) must all be > 0",
batch,
in_features,
out_features
));
}
if input.len() != batch * in_features {
return Err(anyhow!(
"linear_forward: input len {} != batch*in_features = {}*{} = {}",
input.len(),
batch,
in_features,
batch * in_features
));
}
if weight.len() != out_features * in_features {
return Err(anyhow!(
"linear_forward: weight len {} != out_features*in_features = {}*{} = {}",
weight.len(),
out_features,
in_features,
out_features * in_features
));
}
if let Some(b) = bias {
if b.len() != out_features {
return Err(anyhow!(
"linear_forward: bias len {} != out_features {}",
b.len(),
out_features
));
}
}
let mut out = vec![0f32; batch * out_features];
for n in 0..batch {
let x_off = n * in_features;
let y_off = n * out_features;
for o in 0..out_features {
let w_off = o * in_features;
let mut acc: f32 = match bias {
Some(b) => b[o],
None => 0.0,
};
for i in 0..in_features {
acc += input[x_off + i] * weight[w_off + i];
}
out[y_off + o] = acc;
}
}
Ok(out)
}
pub fn scaled_dot_product_attention(
q: &[f32],
k: &[f32],
v: &[f32],
batch: usize,
num_heads: usize,
head_dim: usize,
) -> Result<Vec<f32>> {
if batch == 0 || num_heads == 0 || head_dim == 0 {
return Err(anyhow!(
"scaled_dot_product_attention: batch ({}), num_heads ({}), head_dim ({}) must all be > 0",
batch, num_heads, head_dim
));
}
let expected = batch * num_heads * head_dim;
for (name, t) in [("q", q), ("k", k), ("v", v)] {
if t.len() != expected {
return Err(anyhow!(
"scaled_dot_product_attention: {} len {} != batch*num_heads*head_dim = {}*{}*{} = {}",
name,
t.len(),
batch,
num_heads,
head_dim,
expected
));
}
}
let scale = 1.0f32 / (head_dim as f32).sqrt();
let stride_batch = num_heads * head_dim;
let stride_head = head_dim;
let mut out = vec![0f32; expected];
let mut scores = vec![0f32; batch * batch];
for h in 0..num_heads {
for i in 0..batch {
let q_off = i * stride_batch + h * stride_head;
for j in 0..batch {
let k_off = j * stride_batch + h * stride_head;
let mut acc = 0.0f32;
for d in 0..head_dim {
acc += q[q_off + d] * k[k_off + d];
}
scores[i * batch + j] = acc * scale;
}
}
softmax_last_dim(&mut scores, batch)
.map_err(|e| anyhow!("scaled_dot_product_attention softmax: {e}"))?;
for i in 0..batch {
let out_off = i * stride_batch + h * stride_head;
let sc_off = i * batch;
for d in 0..head_dim {
let mut acc = 0.0f32;
for j in 0..batch {
let v_elem = v[j * stride_batch + h * stride_head + d];
acc += scores[sc_off + j] * v_elem;
}
out[out_off + d] = acc;
}
}
}
Ok(out)
}
pub fn qkv_projection_forward(
input: &[f32],
q_weight: &[f32],
k_weight: &[f32],
v_weight: &[f32],
batch: usize,
hidden: usize,
) -> Result<(Vec<f32>, Vec<f32>, Vec<f32>)> {
let q = linear_forward(input, q_weight, None, batch, hidden, hidden)
.map_err(|e| anyhow!("qkv_projection_forward Q: {e}"))?;
let k = linear_forward(input, k_weight, None, batch, hidden, hidden)
.map_err(|e| anyhow!("qkv_projection_forward K: {e}"))?;
let v = linear_forward(input, v_weight, None, batch, hidden, hidden)
.map_err(|e| anyhow!("qkv_projection_forward V: {e}"))?;
Ok((q, k, v))
}
pub fn scale_in_place(x: &mut [f32], c: f32) {
for v in x.iter_mut() {
*v *= c;
}
}
pub fn avg_pool_2x2_spatial(input: &[f32], n_side: usize, hidden: usize) -> Result<Vec<f32>> {
let n_patches = n_side * n_side;
if input.len() != n_patches * hidden {
return Err(anyhow!(
"avg_pool_2x2_spatial: input len {} != n_patches*hidden = {}*{} = {}",
input.len(),
n_patches,
hidden,
n_patches * hidden
));
}
if n_side == 0 || n_side % 2 != 0 {
return Err(anyhow!(
"avg_pool_2x2_spatial: n_side {} must be positive and even",
n_side
));
}
let out_side = n_side / 2;
let out_patches = out_side * out_side;
let mut out = vec![0f32; out_patches * hidden];
let inv4 = 0.25f32;
for oy in 0..out_side {
for ox in 0..out_side {
let iy0 = oy * 2;
let ix0 = ox * 2;
let out_off = (oy * out_side + ox) * hidden;
for d in 0..hidden {
let a = input[(iy0 * n_side + ix0) * hidden + d];
let b = input[(iy0 * n_side + (ix0 + 1)) * hidden + d];
let c = input[((iy0 + 1) * n_side + ix0) * hidden + d];
let d4 = input[((iy0 + 1) * n_side + (ix0 + 1)) * hidden + d];
out[out_off + d] = (a + b + c + d4) * inv4;
}
}
}
Ok(out)
}
pub fn std_bias_scale_in_place(
x: &mut [f32],
bias: &[f32],
scale: &[f32],
hidden: usize,
) -> Result<()> {
if hidden == 0 {
return Err(anyhow!("std_bias_scale_in_place: hidden must be > 0"));
}
if x.len() % hidden != 0 {
return Err(anyhow!(
"std_bias_scale_in_place: x len {} not divisible by hidden {}",
x.len(),
hidden
));
}
if bias.len() != hidden {
return Err(anyhow!(
"std_bias_scale_in_place: bias len {} != hidden {}",
bias.len(),
hidden
));
}
if scale.len() != hidden {
return Err(anyhow!(
"std_bias_scale_in_place: scale len {} != hidden {}",
scale.len(),
hidden
));
}
let n_rows = x.len() / hidden;
for row in 0..n_rows {
let off = row * hidden;
for i in 0..hidden {
x[off + i] = (x[off + i] - bias[i]) * scale[i];
}
}
Ok(())
}
pub fn apply_vit_block_forward(
hidden_states: Vec<f32>,
weights: &super::mmproj_weights::LoadedMmprojWeights,
cfg: &super::mmproj::MmprojConfig,
block_idx: usize,
) -> Result<Vec<f32>> {
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;
if hidden_states.len() % hidden != 0 {
return Err(anyhow!(
"apply_vit_block_forward: input len {} not divisible by hidden {}",
hidden_states.len(),
hidden
));
}
let batch = hidden_states.len() / hidden;
let widen = |suffix: &str| -> Result<Vec<f32>> {
let buf = weights.block_tensor(block_idx, suffix)?;
weights
.tensor_as_f32_owned(buf)
.map_err(|e| anyhow!("block {}: {} tensor_as_f32_owned: {e}", block_idx, suffix))
};
let ln1_w = widen("ln1.weight")?;
let attn_q_w = widen("attn_q.weight")?;
let attn_k_w = widen("attn_k.weight")?;
let attn_v_w = widen("attn_v.weight")?;
let attn_q_norm_w = widen("attn_q_norm.weight")?;
let attn_k_norm_w = widen("attn_k_norm.weight")?;
let attn_output_w = widen("attn_output.weight")?;
let ln2_w = widen("ln2.weight")?;
let ffn_gate_w = widen("ffn_gate.weight")?;
let ffn_up_w = widen("ffn_up.weight")?;
let ffn_down_w = widen("ffn_down.weight")?;
let post_ffw_norm_w = widen("post_ffw_norm.weight")?;
let mut residual = hidden_states;
let mut cur = residual.clone();
rms_norm_forward(&mut cur, &ln1_w, hidden, eps)?;
let (mut q, mut k, v) =
qkv_projection_forward(&cur, &attn_q_w, &attn_k_w, &attn_v_w, batch, hidden)?;
per_head_rms_norm_forward(&mut q, &attn_q_norm_w, batch, num_heads, head_dim, eps)?;
per_head_rms_norm_forward(&mut k, &attn_k_norm_w, batch, num_heads, head_dim, eps)?;
let attn = scaled_dot_product_attention(&q, &k, &v, batch, num_heads, head_dim)?;
let attn_projected = linear_forward(&attn, &attn_output_w, None, batch, hidden, hidden)?;
residual_add(&mut residual, &attn_projected)?;
let mut cur = residual.clone();
rms_norm_forward(&mut cur, &ln2_w, hidden, eps)?;
let mut gate = linear_forward(&cur, &ffn_gate_w, None, batch, hidden, intermediate)?;
let up = linear_forward(&cur, &ffn_up_w, None, batch, hidden, intermediate)?;
silu_in_place(&mut gate);
elementwise_mul_in_place(&mut gate, &up)?;
let mut down = linear_forward(&gate, &ffn_down_w, None, batch, intermediate, hidden)?;
rms_norm_forward(&mut down, &post_ffw_norm_w, hidden, eps)?;
residual_add(&mut residual, &down)?;
Ok(residual)
}
pub fn apply_vit_full_forward(
pixel_values: &[f32],
weights: &super::mmproj_weights::LoadedMmprojWeights,
cfg: &super::mmproj::MmprojConfig,
) -> Result<Vec<f32>> {
let mut hidden_states = patch_embed_from_mmproj_weights(pixel_values, weights, cfg)?;
let n_layers = cfg.num_hidden_layers as usize;
for block_idx in 0..n_layers {
hidden_states = apply_vit_block_forward(hidden_states, weights, cfg, block_idx)?;
}
let n_side = cfg.num_patches_side as usize;
let hidden = cfg.hidden_size as usize;
let mut hidden_states = avg_pool_2x2_spatial(&hidden_states, n_side, hidden)?;
scale_in_place(&mut hidden_states, (hidden as f32).sqrt());
let std_bias_buf = weights
.get("v.std_bias")
.ok_or_else(|| anyhow!("apply_vit_full_forward: missing v.std_bias"))?;
let std_scale_buf = weights
.get("v.std_scale")
.ok_or_else(|| anyhow!("apply_vit_full_forward: missing v.std_scale"))?;
let std_bias: &[f32] = std_bias_buf
.as_slice::<f32>()
.map_err(|e| anyhow!("std_bias as_slice: {e}"))?;
let std_scale: &[f32] = std_scale_buf
.as_slice::<f32>()
.map_err(|e| anyhow!("std_scale as_slice: {e}"))?;
std_bias_scale_in_place(&mut hidden_states, std_bias, std_scale, hidden)?;
let mm0_buf = weights
.get("mm.0.weight")
.ok_or_else(|| anyhow!("apply_vit_full_forward: missing mm.0.weight"))?;
let mm0_owned = weights
.tensor_as_f32_owned(mm0_buf)
.map_err(|e| anyhow!("mm.0 tensor_as_f32_owned: {e}"))?;
let text_hidden = mm0_owned.len() / hidden;
if mm0_owned.len() != text_hidden * hidden {
return Err(anyhow!(
"mm.0.weight len {} not divisible by hidden {}",
mm0_owned.len(),
hidden
));
}
let n_patches_out = hidden_states.len() / hidden;
let mut projected = linear_forward(
&hidden_states,
&mm0_owned,
None,
n_patches_out,
hidden,
text_hidden,
)?;
let ones = vec![1.0f32; text_hidden];
rms_norm_forward(&mut projected, &ones, text_hidden, cfg.layer_norm_eps)?;
Ok(projected)
}
pub fn patch_embed_from_mmproj_weights(
pixel_values: &[f32],
weights: &super::mmproj_weights::LoadedMmprojWeights,
cfg: &super::mmproj::MmprojConfig,
) -> Result<Vec<f32>> {
let patch_buf = weights
.patch_embd_weight()
.map_err(|e| anyhow!("patch_embd_forward: {e}"))?;
let weight_owned = weights
.tensor_as_f32_owned(patch_buf)
.map_err(|e| anyhow!("patch_embd_forward tensor_as_f32_owned: {e}"))?;
patch_embed_forward(
pixel_values,
&weight_owned,
None, cfg.image_size,
cfg.patch_size,
cfg.hidden_size,
)
}
pub fn gemma4v_patch_embed_forward(
patches: &[f32],
weight: &[f32],
n_patches: u32,
inner: u32,
hidden: u32,
) -> Result<Vec<f32>> {
if n_patches == 0 || inner == 0 || hidden == 0 {
return Err(anyhow!(
"gemma4v_patch_embed_forward: n_patches ({n_patches}), inner ({inner}), \
hidden ({hidden}) must all be > 0"
));
}
let n_us = n_patches as usize;
let in_us = inner as usize;
let h_us = hidden as usize;
if patches.len() != n_us * in_us {
return Err(anyhow!(
"gemma4v_patch_embed_forward: patches.len() ({}) != n_patches*inner ({})",
patches.len(),
n_us * in_us
));
}
if weight.len() != h_us * in_us {
return Err(anyhow!(
"gemma4v_patch_embed_forward: weight.len() ({}) != hidden*inner ({})",
weight.len(),
h_us * in_us
));
}
let mut out = vec![0f32; n_us * h_us];
for n in 0..n_us {
let p_base = n * in_us;
let o_base = n * h_us;
for o in 0..h_us {
let w_base = o * in_us;
let mut acc: f32 = 0.0;
for i in 0..in_us {
acc += patches[p_base + i] * weight[w_base + i];
}
out[o_base + o] = acc;
}
}
Ok(out)
}
pub fn gemma4v_position_embed_lookup(
pos_x: &[u32],
pos_y: &[u32],
pe_table: &[f32],
pos_size: u32,
hidden: u32,
) -> Result<Vec<f32>> {
if pos_size == 0 || hidden == 0 {
return Err(anyhow!(
"gemma4v_position_embed_lookup: pos_size ({pos_size}) and hidden ({hidden}) must be > 0"
));
}
if pos_x.len() != pos_y.len() {
return Err(anyhow!(
"gemma4v_position_embed_lookup: pos_x.len() ({}) != pos_y.len() ({})",
pos_x.len(),
pos_y.len()
));
}
let n = pos_x.len();
let h_us = hidden as usize;
let ps_us = pos_size as usize;
let expected = 2 * ps_us * h_us;
if pe_table.len() != expected {
return Err(anyhow!(
"gemma4v_position_embed_lookup: pe_table.len() ({}) != 2*pos_size*hidden ({})",
pe_table.len(),
expected
));
}
let table_x_base = 0;
let table_y_base = ps_us * h_us;
let max_idx = (pos_size - 1) as u32;
let mut out = vec![0f32; n * h_us];
for k in 0..n {
let x_idx = pos_x[k].min(max_idx) as usize;
let y_idx = pos_y[k].min(max_idx) as usize;
let row_x = table_x_base + x_idx * h_us;
let row_y = table_y_base + y_idx * h_us;
let out_base = k * h_us;
for j in 0..h_us {
out[out_base + j] = pe_table[row_x + j] + pe_table[row_y + j];
}
}
Ok(out)
}
pub fn gemma4v_position_embed_add(
patch_embeds: &mut [f32],
pos_x: &[u32],
pos_y: &[u32],
pe_table: &[f32],
pos_size: u32,
hidden: u32,
) -> Result<()> {
let pos_emb = gemma4v_position_embed_lookup(pos_x, pos_y, pe_table, pos_size, hidden)?;
if patch_embeds.len() != pos_emb.len() {
return Err(anyhow!(
"gemma4v_position_embed_add: patch_embeds.len() ({}) != pos_emb.len() ({})",
patch_embeds.len(),
pos_emb.len()
));
}
for (dst, src) in patch_embeds.iter_mut().zip(pos_emb.iter()) {
*dst += *src;
}
Ok(())
}
pub fn gemma_rms_norm_forward(
input: &mut [f32],
weight: &[f32],
hidden: usize,
eps: f32,
) -> Result<()> {
if hidden == 0 {
return Err(anyhow!("gemma_rms_norm_forward: hidden must be > 0"));
}
if input.len() % hidden != 0 {
return Err(anyhow!(
"gemma_rms_norm_forward: input len {} not divisible by hidden {}",
input.len(),
hidden
));
}
if weight.len() != hidden {
return Err(anyhow!(
"gemma_rms_norm_forward: weight len {} != hidden {}",
weight.len(),
hidden
));
}
rms_norm_forward(input, weight, hidden, eps)
}
pub fn gemma_per_head_rms_norm_forward(
input: &mut [f32],
weight: &[f32],
batch: usize,
num_heads: usize,
head_dim: usize,
eps: f32,
) -> Result<()> {
if batch == 0 || num_heads == 0 || head_dim == 0 {
return Err(anyhow!(
"gemma_per_head_rms_norm_forward: batch ({batch}), num_heads ({num_heads}), \
head_dim ({head_dim}) must all be > 0"
));
}
let expected = batch * num_heads * head_dim;
if input.len() != expected {
return Err(anyhow!(
"gemma_per_head_rms_norm_forward: input len {} != batch*num_heads*head_dim = {}",
input.len(),
expected
));
}
if weight.len() != head_dim {
return Err(anyhow!(
"gemma_per_head_rms_norm_forward: weight len {} != head_dim {}",
weight.len(),
head_dim
));
}
gemma_rms_norm_forward(input, weight, head_dim, eps)
}
pub fn v_norm_no_scale_forward(input: &mut [f32], hidden: usize, eps: f32) -> Result<()> {
if hidden == 0 {
return Err(anyhow!("v_norm_no_scale_forward: hidden must be > 0"));
}
if input.len() % hidden != 0 {
return Err(anyhow!(
"v_norm_no_scale_forward: input len {} not divisible by hidden {}",
input.len(),
hidden
));
}
let inv_h = 1.0_f32 / hidden as f32;
let n_rows = input.len() / hidden;
for row in 0..n_rows {
let off = row * hidden;
let slice = &mut input[off..off + hidden];
let mut sq = 0.0_f32;
for &v in slice.iter() {
sq += v * v;
}
let inv_rms = 1.0_f32 / (sq * inv_h + eps).sqrt();
for v in slice.iter_mut() {
*v *= inv_rms;
}
}
Ok(())
}
#[derive(Debug, Clone, Copy, Default)]
pub struct Gemma4ClippableLinearBounds {
pub input_min: Option<f32>,
pub input_max: Option<f32>,
pub output_min: Option<f32>,
pub output_max: Option<f32>,
}
impl Gemma4ClippableLinearBounds {
pub fn any(&self) -> bool {
self.input_min.is_some()
|| self.input_max.is_some()
|| self.output_min.is_some()
|| self.output_max.is_some()
}
pub fn resolved_input(&self) -> (f32, f32) {
(
self.input_min.unwrap_or(f32::NEG_INFINITY),
self.input_max.unwrap_or(f32::INFINITY),
)
}
pub fn resolved_output(&self) -> (f32, f32) {
(
self.output_min.unwrap_or(f32::NEG_INFINITY),
self.output_max.unwrap_or(f32::INFINITY),
)
}
}
pub fn gemma4v_clippable_linear_forward(
input: &[f32],
weight: &[f32],
bounds: &Gemma4ClippableLinearBounds,
batch: usize,
in_features: usize,
out_features: usize,
) -> Result<Vec<f32>> {
if batch == 0 || in_features == 0 || out_features == 0 {
return Err(anyhow!(
"gemma4v_clippable_linear_forward: batch ({batch}), in_features \
({in_features}), out_features ({out_features}) must all be > 0"
));
}
let clamped_input_owned: Option<Vec<f32>> =
if bounds.input_min.is_some() || bounds.input_max.is_some() {
let (mn, mx) = bounds.resolved_input();
if mn > mx {
return Err(anyhow!(
"gemma4v_clippable_linear_forward: input_min ({mn}) > input_max ({mx})"
));
}
let mut v = input.to_vec();
for x in v.iter_mut() {
*x = x.clamp(mn, mx);
}
Some(v)
} else {
None
};
let input_view: &[f32] = clamped_input_owned.as_deref().unwrap_or(input);
let mut y = linear_forward(input_view, weight, None, batch, in_features, out_features)?;
if bounds.output_min.is_some() || bounds.output_max.is_some() {
let (mn, mx) = bounds.resolved_output();
if mn > mx {
return Err(anyhow!(
"gemma4v_clippable_linear_forward: output_min ({mn}) > output_max ({mx})"
));
}
for v in y.iter_mut() {
*v = v.clamp(mn, mx);
}
}
Ok(y)
}
pub fn vision_2d_rope_forward_cpu(
input: &[f32],
seq_len: usize,
n_heads: usize,
head_dim: usize,
pos_x: &[u32],
pos_y: &[u32],
theta: f32,
) -> Result<Vec<f32>> {
if head_dim == 0 || seq_len == 0 || n_heads == 0 {
return Err(anyhow!(
"vision_2d_rope_forward_cpu: head_dim ({head_dim}), seq_len ({seq_len}), \
n_heads ({n_heads}) must all be > 0"
));
}
if head_dim % 4 != 0 {
return Err(anyhow!(
"vision_2d_rope_forward_cpu: head_dim ({head_dim}) must be divisible by 4"
));
}
if pos_x.len() != seq_len {
return Err(anyhow!(
"vision_2d_rope_forward_cpu: pos_x len {} != seq_len {seq_len}",
pos_x.len()
));
}
if pos_y.len() != seq_len {
return Err(anyhow!(
"vision_2d_rope_forward_cpu: pos_y len {} != seq_len {seq_len}",
pos_y.len()
));
}
let n_rows = seq_len * n_heads;
let n_elem = n_rows * head_dim;
if input.len() != n_elem {
return Err(anyhow!(
"vision_2d_rope_forward_cpu: input len {} != seq_len*n_heads*head_dim = {n_elem}",
input.len()
));
}
let d_half = head_dim / 2;
let d_quarter = d_half / 2;
let mut out = vec![0_f32; n_elem];
out.copy_from_slice(input);
for row in 0..n_rows {
let seq_idx = row / n_heads;
let p_x = pos_x[seq_idx] as f32;
let p_y = pos_y[seq_idx] as f32;
let base = row * head_dim;
for i in 0..d_quarter {
let dim_ratio = (2 * i) as f32 / d_half as f32;
let freq = 1.0_f32 / theta.powf(dim_ratio);
let ax = p_x * freq;
let ay = p_y * freq;
let cx = ax.cos();
let sx = ax.sin();
let cy = ay.cos();
let sy = ay.sin();
let x0 = input[base + i];
let x1 = input[base + i + d_quarter];
out[base + i] = x0 * cx - x1 * sx;
out[base + i + d_quarter] = x0 * sx + x1 * cx;
let y0 = input[base + d_half + i];
let y1 = input[base + d_half + i + d_quarter];
out[base + d_half + i] = y0 * cy - y1 * sy;
out[base + d_half + i + d_quarter] = y0 * sy + y1 * cy;
}
}
Ok(out)
}
pub fn gelu_pytorch_tanh_in_place(input: &mut [f32]) {
const SQRT_2_OVER_PI: f32 = 0.797_884_56_f32;
const COEFF: f32 = 0.044_715_f32;
for v in input.iter_mut() {
let x = *v;
let inner = SQRT_2_OVER_PI * (x + COEFF * x * x * x);
*v = 0.5 * x * (1.0 + inner.tanh());
}
}
pub fn repeat_kv_cpu(
input: &[f32],
batch: usize,
num_kv_heads: usize,
num_kv_groups: usize,
head_dim: usize,
) -> Result<Vec<f32>> {
if batch == 0 || num_kv_heads == 0 || num_kv_groups == 0 || head_dim == 0 {
return Err(anyhow!(
"repeat_kv_cpu: batch, num_kv_heads, num_kv_groups, head_dim must all be > 0"
));
}
let expected_in = batch * num_kv_heads * head_dim;
if input.len() != expected_in {
return Err(anyhow!(
"repeat_kv_cpu: input len {} != batch*num_kv_heads*head_dim = {expected_in}",
input.len()
));
}
let num_heads = num_kv_heads * num_kv_groups;
let mut out = vec![0_f32; 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;
for g in 0..num_kv_groups {
let h_idx = k * num_kv_groups + g;
let out_base = (b * num_heads + h_idx) * head_dim;
out[out_base..out_base + head_dim]
.copy_from_slice(&input[in_base..in_base + head_dim]);
}
}
}
Ok(out)
}
#[allow(clippy::too_many_arguments)]
pub fn gemma4v_block_forward(
hidden_states: Vec<f32>,
block_weights: &Gemma4VisionBlockWeights<'_>,
shape: &Gemma4VisionBlockShape,
pos_x: &[u32],
pos_y: &[u32],
) -> Result<Vec<f32>> {
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 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 {
return Err(anyhow!(
"gemma4v_block_forward: zero dim in shape: {shape:?}"
));
}
if num_heads % num_kv_heads != 0 {
return Err(anyhow!(
"gemma4v_block_forward: num_heads ({num_heads}) must be a multiple of num_kv_heads ({num_kv_heads})"
));
}
if hidden_states.len() % hidden != 0 {
return Err(anyhow!(
"gemma4v_block_forward: hidden_states len {} not divisible by hidden {hidden}",
hidden_states.len()
));
}
let batch = hidden_states.len() / hidden;
if pos_x.len() != batch || pos_y.len() != batch {
return Err(anyhow!(
"gemma4v_block_forward: pos_x ({}) / pos_y ({}) lengths must equal batch ({batch})",
pos_x.len(),
pos_y.len()
));
}
let num_kv_groups = num_heads / num_kv_heads;
let q_dim = num_heads * head_dim;
let kv_dim = num_kv_heads * head_dim;
if q_dim != hidden {
return Err(anyhow!(
"gemma4v_block_forward: num_heads*head_dim ({q_dim}) must equal hidden ({hidden})"
));
}
let mut residual = hidden_states;
let mut cur = residual.clone();
gemma_rms_norm_forward(&mut cur, block_weights.input_layernorm, hidden, eps)?;
let mut q = linear_forward(&cur, block_weights.q_proj, None, batch, hidden, q_dim)?;
let mut k = linear_forward(&cur, block_weights.k_proj, None, batch, hidden, kv_dim)?;
let mut v = linear_forward(&cur, block_weights.v_proj, None, batch, hidden, kv_dim)?;
gemma_per_head_rms_norm_forward(
&mut q,
block_weights.q_norm,
batch,
num_heads,
head_dim,
eps,
)?;
gemma_per_head_rms_norm_forward(
&mut k,
block_weights.k_norm,
batch,
num_kv_heads,
head_dim,
eps,
)?;
v_norm_no_scale_forward(&mut v, head_dim, eps)?;
let q = vision_2d_rope_forward_cpu(&q, batch, num_heads, head_dim, pos_x, pos_y, theta)?;
let k = vision_2d_rope_forward_cpu(&k, batch, num_kv_heads, head_dim, pos_x, pos_y, theta)?;
let k_full = repeat_kv_cpu(&k, batch, num_kv_heads, num_kv_groups, head_dim)?;
let v_full = repeat_kv_cpu(&v, batch, num_kv_heads, num_kv_groups, head_dim)?;
let attn = gemma4v_attention_unit_scale(&q, &k_full, &v_full, batch, num_heads, head_dim)?;
let attn_proj = linear_forward(&attn, block_weights.o_proj, None, batch, hidden, hidden)?;
let mut attn_out = attn_proj;
gemma_rms_norm_forward(
&mut attn_out,
block_weights.post_attention_layernorm,
hidden,
eps,
)?;
residual_add(&mut residual, &attn_out)?;
let mut cur = residual.clone();
gemma_rms_norm_forward(
&mut cur,
block_weights.pre_feedforward_layernorm,
hidden,
eps,
)?;
let mut gate = linear_forward(
&cur,
block_weights.gate_proj,
None,
batch,
hidden,
intermediate,
)?;
let up = linear_forward(
&cur,
block_weights.up_proj,
None,
batch,
hidden,
intermediate,
)?;
gelu_pytorch_tanh_in_place(&mut gate);
elementwise_mul_in_place(&mut gate, &up)?;
let mut down = linear_forward(
&gate,
block_weights.down_proj,
None,
batch,
intermediate,
hidden,
)?;
gemma_rms_norm_forward(
&mut down,
block_weights.post_feedforward_layernorm,
hidden,
eps,
)?;
residual_add(&mut residual, &down)?;
Ok(residual)
}
fn gemma4v_attention_unit_scale(
q: &[f32],
k: &[f32],
v: &[f32],
batch: usize,
num_heads: usize,
head_dim: usize,
) -> Result<Vec<f32>> {
if batch == 0 || num_heads == 0 || head_dim == 0 {
return Err(anyhow!(
"gemma4v_attention_unit_scale: zero dim batch={batch} heads={num_heads} d={head_dim}"
));
}
let expected = batch * num_heads * head_dim;
if q.len() != expected || k.len() != expected || v.len() != expected {
return Err(anyhow!(
"gemma4v_attention_unit_scale: shape mismatch q={} k={} v={} expected {expected}",
q.len(),
k.len(),
v.len()
));
}
let mut scores = vec![0_f32; num_heads * batch * batch];
for h in 0..num_heads {
for i in 0..batch {
for j in 0..batch {
let mut acc = 0_f32;
let q_base = i * num_heads * head_dim + h * head_dim;
let k_base = j * num_heads * head_dim + h * head_dim;
for d in 0..head_dim {
acc += q[q_base + d] * k[k_base + d];
}
scores[h * batch * batch + i * batch + j] = acc;
}
}
}
softmax_last_dim(&mut scores, batch)?;
let mut attn = vec![0_f32; batch * num_heads * head_dim];
for h in 0..num_heads {
for i in 0..batch {
for d in 0..head_dim {
let mut acc = 0_f32;
for j in 0..batch {
let v_idx = j * num_heads * head_dim + h * head_dim + d;
let s_idx = h * batch * batch + i * batch + j;
acc += scores[s_idx] * v[v_idx];
}
attn[i * num_heads * head_dim + h * head_dim + d] = acc;
}
}
}
Ok(attn)
}
#[derive(Debug, Clone, Copy)]
pub struct Gemma4VisionBlockShape {
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,
}
#[derive(Debug)]
pub struct Gemma4VisionBlockWeights<'a> {
pub input_layernorm: &'a [f32], pub post_attention_layernorm: &'a [f32], pub pre_feedforward_layernorm: &'a [f32], pub post_feedforward_layernorm: &'a [f32], pub q_proj: &'a [f32], pub k_proj: &'a [f32], pub v_proj: &'a [f32], pub o_proj: &'a [f32], pub q_norm: &'a [f32], pub k_norm: &'a [f32], pub gate_proj: &'a [f32], pub up_proj: &'a [f32], pub down_proj: &'a [f32], }
#[cfg(test)]
mod tests {
use super::*;
fn make_weight<F: FnMut(usize, usize, usize, usize) -> f32>(
hidden: usize,
patch_size: usize,
mut f: F,
) -> Vec<f32> {
let mut w = vec![0f32; hidden * 3 * patch_size * patch_size];
for oc in 0..hidden {
for ic in 0..3 {
for dy in 0..patch_size {
for dx in 0..patch_size {
let idx = oc * 3 * patch_size * patch_size
+ ic * patch_size * patch_size
+ dy * patch_size
+ dx;
w[idx] = f(oc, ic, dy, dx);
}
}
}
}
w
}
fn make_pixels<F: FnMut(usize, usize, usize) -> f32>(image_size: usize, mut f: F) -> Vec<f32> {
let mut p = vec![0f32; 3 * image_size * image_size];
for c in 0..3 {
for y in 0..image_size {
for x in 0..image_size {
let idx = c * image_size * image_size + y * image_size + x;
p[idx] = f(c, y, x);
}
}
}
p
}
#[test]
fn delta_kernel_copies_top_left_pixel_per_patch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let img = 8;
let patch = 4;
let hidden = 3;
let weight = make_weight(hidden, patch, |oc, ic, dy, dx| {
if oc == ic && dy == 0 && dx == 0 {
1.0
} else {
0.0
}
});
let pixels = make_pixels(img, |c, y, x| (c * 100 + y * 10 + x) as f32);
let out = patch_embed_forward(
&pixels,
&weight,
None,
img as u32,
patch as u32,
hidden as u32,
)
.unwrap();
assert_eq!(out.len(), 2 * 2 * hidden);
let p = |patch: usize, ch: usize| patch * hidden + ch;
assert_eq!(out[p(0, 0)], 0.0); assert_eq!(out[p(0, 1)], 100.0); assert_eq!(out[p(0, 2)], 200.0); assert_eq!(out[p(1, 0)], 4.0); assert_eq!(out[p(1, 1)], 104.0);
assert_eq!(out[p(1, 2)], 204.0);
assert_eq!(out[p(2, 0)], 40.0);
assert_eq!(out[p(2, 1)], 140.0);
assert_eq!(out[p(2, 2)], 240.0);
assert_eq!(out[p(3, 0)], 44.0);
assert_eq!(out[p(3, 1)], 144.0);
assert_eq!(out[p(3, 2)], 244.0);
}
#[test]
fn all_ones_kernel_produces_patch_sum() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let img: usize = 4;
let patch: usize = 2;
let hidden: usize = 1;
let weight = vec![1f32; hidden * 3 * patch * patch];
let pixels = make_pixels(img, |c, _y, _x| (c as f32) + 1.0);
let out = patch_embed_forward(
&pixels,
&weight,
None,
img as u32,
patch as u32,
hidden as u32,
)
.unwrap();
assert_eq!(out.len(), 4 * 1); for (i, v) in out.iter().enumerate() {
assert!(
(*v - 24.0).abs() < 1e-5,
"patch {}: expected 24, got {}",
i,
v
);
}
}
#[test]
fn bias_is_added_once_per_output_element() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let img: usize = 4;
let patch: usize = 2;
let hidden: usize = 2;
let weight = vec![0f32; hidden * 3 * patch * patch];
let bias = vec![10f32, 20.0];
let pixels = vec![999f32; 3 * img * img]; let out = patch_embed_forward(
&pixels,
&weight,
Some(&bias),
img as u32,
patch as u32,
hidden as u32,
)
.unwrap();
assert_eq!(out.len(), 4 * 2);
for i in 0..4 {
assert!((out[i * 2 + 0] - 10.0).abs() < 1e-5);
assert!((out[i * 2 + 1] - 20.0).abs() < 1e-5);
}
}
#[test]
fn single_patch_covers_whole_image() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let img: usize = 4;
let patch: usize = 4;
let hidden: usize = 1;
let weight = make_weight(hidden, patch, |_oc, ic, dy, dx| {
if ic == 0 && dy == 0 && dx == 0 {
1.0
} else {
0.0
}
});
let pixels = make_pixels(img, |c, _, _| (c + 1) as f32);
let out = patch_embed_forward(
&pixels,
&weight,
None,
img as u32,
patch as u32,
hidden as u32,
)
.unwrap();
assert_eq!(out.len(), 1);
assert_eq!(out[0], 1.0); }
#[test]
fn rejects_mismatched_pixel_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let weight = vec![0f32; 1 * 3 * 2 * 2];
let pixels = vec![0f32; 10]; let err = patch_embed_forward(&pixels, &weight, None, 4, 2, 1).unwrap_err();
assert!(format!("{err}").contains("pixel_values len"));
}
#[test]
fn rejects_mismatched_weight_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let pixels = vec![0f32; 3 * 4 * 4];
let weight = vec![0f32; 5]; let err = patch_embed_forward(&pixels, &weight, None, 4, 2, 1).unwrap_err();
assert!(format!("{err}").contains("patch_embd_weight len"));
}
#[test]
fn rejects_bias_length_mismatch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let weight = vec![0f32; 2 * 3 * 2 * 2];
let pixels = vec![0f32; 3 * 4 * 4];
let bias = vec![0f32; 99];
let err = patch_embed_forward(&pixels, &weight, Some(&bias), 4, 2, 2).unwrap_err();
assert!(format!("{err}").contains("patch_embd_bias len"));
}
#[test]
fn rejects_non_divisible_image_size() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let weight = vec![0f32; 1 * 3 * 3 * 3];
let pixels = vec![0f32; 3 * 5 * 5];
let err = patch_embed_forward(&pixels, &weight, None, 5, 3, 1).unwrap_err();
assert!(format!("{err}").contains("divisible"));
}
#[test]
fn rejects_zero_patch_size() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let weight = vec![0f32; 0];
let pixels = vec![0f32; 3 * 4 * 4];
let err = patch_embed_forward(&pixels, &weight, None, 4, 0, 1).unwrap_err();
assert!(format!("{err}").contains("patch_size"));
}
#[test]
fn position_embed_add_is_elementwise() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut patch = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0];
let pos = vec![0.1f32, 0.2, 0.3, 0.4, 0.5, 0.6];
position_embed_add(&mut patch, &pos).unwrap();
let expected = [1.1f32, 2.2, 3.3, 4.4, 5.5, 6.6];
for (got, want) in patch.iter().zip(expected.iter()) {
assert!((*got - *want).abs() < 1e-6, "got {got}, want {want}");
}
}
#[test]
fn position_embed_add_rejects_shape_mismatch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut patch = vec![0.0f32; 10];
let pos = vec![0.0f32; 11];
let err = position_embed_add(&mut patch, &pos).unwrap_err();
assert!(format!("{err}").contains("!= pos len"));
}
#[test]
fn position_embed_add_zero_pos_is_identity() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut patch = vec![0.5f32, -0.3, 1.7, -2.1];
let snapshot = patch.clone();
let pos = vec![0.0f32; 4];
position_embed_add(&mut patch, &pos).unwrap();
assert_eq!(patch, snapshot);
}
#[test]
fn layer_norm_constant_row_goes_to_zero_then_beta() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![7.0f32; 4];
let gamma = vec![1.0f32; 4];
let beta = vec![0.1f32, 0.2, 0.3, 0.4];
layer_norm_forward(&mut x, &gamma, &beta, 4, 1e-6).unwrap();
for (got, want) in x.iter().zip(beta.iter()) {
assert!((*got - *want).abs() < 1e-5, "got {got}, want {want}");
}
}
#[test]
fn layer_norm_pytorch_reference_values() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32, 2.0, 3.0, 4.0];
let gamma = vec![1.0f32; 4];
let beta = vec![0.0f32; 4];
layer_norm_forward(&mut x, &gamma, &beta, 4, 1e-5).unwrap();
let expected = [-1.3416408f32, -0.4472136, 0.4472136, 1.3416408];
for (got, want) in x.iter().zip(expected.iter()) {
assert!((*got - *want).abs() < 1e-3, "got {got}, want {want}");
}
}
#[test]
fn layer_norm_applies_affine_scale_and_shift() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32, 2.0, 3.0, 4.0];
let gamma = vec![2.0f32; 4];
let beta = vec![10.0f32, 20.0, 30.0, 40.0];
layer_norm_forward(&mut x, &gamma, &beta, 4, 1e-5).unwrap();
let expected = [7.3167f32, 19.1056, 30.8944, 42.6833];
for (got, want) in x.iter().zip(expected.iter()) {
assert!((*got - *want).abs() < 1e-2, "got {got}, want {want}");
}
}
#[test]
fn layer_norm_normalizes_multiple_rows_independently() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32, 2.0, 3.0, 4.0, 100.0, 200.0, 300.0, 400.0];
let gamma = vec![1.0f32; 4];
let beta = vec![0.0f32; 4];
layer_norm_forward(&mut x, &gamma, &beta, 4, 1e-5).unwrap();
for i in 0..4 {
assert!(
(x[i] - x[4 + i]).abs() < 1e-3,
"row 0 element {} = {} != row 1 = {}",
i,
x[i],
x[4 + i]
);
}
}
#[test]
fn layer_norm_mean_after_is_approximately_zero() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![0.5f32, 1.7, -2.3, 4.1, -0.8, 3.2];
let gamma = vec![1.0f32; 6];
let beta = vec![0.0f32; 6];
layer_norm_forward(&mut x, &gamma, &beta, 6, 1e-5).unwrap();
let mean: f32 = x.iter().sum::<f32>() / 6.0;
assert!(mean.abs() < 1e-5, "post-LN mean = {mean}");
}
#[test]
fn layer_norm_rejects_hidden_zero() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32; 4];
let err = layer_norm_forward(&mut x, &[], &[], 0, 1e-5).unwrap_err();
assert!(format!("{err}").contains("hidden must be > 0"));
}
#[test]
fn layer_norm_rejects_non_divisible_input_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32; 7];
let gamma = vec![1.0f32; 3];
let beta = vec![0.0f32; 3];
let err = layer_norm_forward(&mut x, &gamma, &beta, 3, 1e-5).unwrap_err();
assert!(format!("{err}").contains("not divisible"));
}
#[test]
fn layer_norm_rejects_wrong_gamma_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32; 4];
let gamma = vec![1.0f32; 3];
let beta = vec![0.0f32; 4];
let err = layer_norm_forward(&mut x, &gamma, &beta, 4, 1e-5).unwrap_err();
assert!(format!("{err}").contains("gamma len"));
}
#[test]
fn layer_norm_rejects_wrong_beta_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32; 4];
let gamma = vec![1.0f32; 4];
let beta = vec![0.0f32; 3];
let err = layer_norm_forward(&mut x, &gamma, &beta, 4, 1e-5).unwrap_err();
assert!(format!("{err}").contains("beta len"));
}
#[test]
fn layer_norm_does_not_divide_by_zero_when_variance_is_zero() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![5.0f32; 8];
let gamma = vec![1.0f32; 8];
let beta = vec![0.0f32; 8];
layer_norm_forward(&mut x, &gamma, &beta, 8, 1e-5).unwrap();
for v in &x {
assert!(v.is_finite(), "got non-finite {}", v);
assert!(
v.abs() < 1e-5,
"constant row should normalize to 0, got {v}"
);
}
}
#[test]
fn std_bias_scale_subtracts_and_scales_per_channel() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![5.0f32, 10.0, 15.0, 10.0, 20.0, 30.0];
let bias = vec![1.0f32, 2.0, 3.0];
let scale = vec![10.0f32, 20.0, 30.0];
std_bias_scale_in_place(&mut x, &bias, &scale, 3).unwrap();
assert_eq!(x, vec![40.0, 160.0, 360.0, 90.0, 360.0, 810.0]);
}
#[test]
fn std_bias_scale_zero_bias_unit_scale_is_identity() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![0.5f32, 1.7, -2.3, 4.1];
let snap = x.clone();
let bias = vec![0.0f32; 2];
let scale = vec![1.0f32; 2];
std_bias_scale_in_place(&mut x, &bias, &scale, 2).unwrap();
assert_eq!(x, snap);
}
#[test]
fn std_bias_scale_rejects_hidden_zero() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![0f32; 4];
let err = std_bias_scale_in_place(&mut x, &[], &[], 0).unwrap_err();
assert!(format!("{err}").contains("hidden must be > 0"));
}
#[test]
fn std_bias_scale_rejects_non_divisible_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![0f32; 7];
let bias = vec![0f32; 3];
let scale = vec![1f32; 3];
let err = std_bias_scale_in_place(&mut x, &bias, &scale, 3).unwrap_err();
assert!(format!("{err}").contains("not divisible"));
}
#[test]
fn std_bias_scale_rejects_wrong_bias_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![0f32; 6];
let bias = vec![0f32; 2]; let scale = vec![1f32; 3];
let err = std_bias_scale_in_place(&mut x, &bias, &scale, 3).unwrap_err();
assert!(format!("{err}").contains("bias len"));
}
#[test]
fn std_bias_scale_rejects_wrong_scale_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![0f32; 6];
let bias = vec![0f32; 3];
let scale = vec![1f32; 2]; let err = std_bias_scale_in_place(&mut x, &bias, &scale, 3).unwrap_err();
assert!(format!("{err}").contains("scale len"));
}
#[allow(dead_code)]
fn _retired_cpu_full_forward_stub() {
use super::super::mmproj::MmprojConfig;
use super::super::mmproj_weights::LoadedMmprojWeights;
use mlx_native::gguf::GgufFile;
let path = std::path::Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!(
"skipping: mmproj fixture not found at {}",
GEMMA4_MMPROJ_PATH
);
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let device = mlx_native::MlxDevice::new().expect("device");
let weights = LoadedMmprojWeights::load(&gguf, &cfg, device).expect("load");
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 t0 = std::time::Instant::now();
let out = apply_vit_full_forward(&pixels, &weights, &cfg).expect("full forward");
let elapsed = t0.elapsed();
eprintln!("apply_vit_full_forward: {:?}", elapsed);
let n_patches_out = 49;
let mm0_buf = weights.get("mm.0.weight").expect("mm.0");
let mm0_slice: &[f32] = mm0_buf.as_slice::<f32>().expect("slice");
let text_hidden = mm0_slice.len() / (cfg.hidden_size as usize);
assert_eq!(text_hidden, 2816, "Gemma 4 projector output width");
assert_eq!(out.len(), n_patches_out * text_hidden);
for v in &out {
assert!(v.is_finite(), "non-finite in ViT output: {v}");
}
for p in 0..n_patches_out {
let row = &out[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() < 1e-2,
"patch {p} mean(x²) = {ms}, expected ≈ 1.0 after no-gain RMSNorm"
);
}
let p0 = &out[0..text_hidden];
let p_last = &out[(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();
assert!(
l2 > 1e-3,
"token 0 and token 48 collapsed to identical output"
);
}
#[test]
fn scale_in_place_multiplies_every_element() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32, 2.0, 3.0, 4.0];
scale_in_place(&mut x, 2.5);
assert_eq!(x, vec![2.5, 5.0, 7.5, 10.0]);
}
#[test]
fn scale_in_place_by_one_is_identity() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![0.1f32, -2.3, 7.7];
let snap = x.clone();
scale_in_place(&mut x, 1.0);
assert_eq!(x, snap);
}
#[test]
fn scale_in_place_by_zero_zeros_everything() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32, 2.0, 3.0];
scale_in_place(&mut x, 0.0);
assert_eq!(x, vec![0.0, 0.0, 0.0]);
}
#[test]
fn avg_pool_2x2_averages_each_2x2_block() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let n_side = 4;
let hidden = 2;
let mut input = 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[patch * hidden + 0] = patch as f32;
input[patch * hidden + 1] = (patch as f32) * 10.0;
}
}
let out = avg_pool_2x2_spatial(&input, n_side, hidden).unwrap();
assert_eq!(out.len(), 2 * 2 * hidden);
let p = |patch: usize, ch: usize| patch * hidden + ch;
assert!((out[p(0, 0)] - 2.5).abs() < 1e-6);
assert!((out[p(0, 1)] - 25.0).abs() < 1e-6);
assert!((out[p(1, 0)] - 4.5).abs() < 1e-6);
assert!((out[p(1, 1)] - 45.0).abs() < 1e-6);
assert!((out[p(2, 0)] - 10.5).abs() < 1e-6);
assert!((out[p(2, 1)] - 105.0).abs() < 1e-6);
assert!((out[p(3, 0)] - 12.5).abs() < 1e-6);
assert!((out[p(3, 1)] - 125.0).abs() < 1e-6);
}
#[test]
fn avg_pool_gemma4_shape_14x14_to_7x7() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let n_side = 14;
let hidden = 1152;
let input = vec![1.0f32; n_side * n_side * hidden];
let out = avg_pool_2x2_spatial(&input, n_side, hidden).unwrap();
assert_eq!(out.len(), 7 * 7 * hidden);
for v in &out {
assert!((*v - 1.0).abs() < 1e-6);
}
}
#[test]
fn avg_pool_rejects_non_even_n_side() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let err = avg_pool_2x2_spatial(&[0f32; 27], 3, 3).unwrap_err();
assert!(format!("{err}").contains("positive and even"));
}
#[test]
fn avg_pool_rejects_mismatched_input_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let err = avg_pool_2x2_spatial(&[0f32; 15], 4, 2).unwrap_err();
assert!(format!("{err}").contains("input len"));
}
#[test]
fn apply_vit_block_forward_real_gemma4_block0_matches_inline_chain() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::mmproj::MmprojConfig;
use super::super::mmproj_weights::LoadedMmprojWeights;
use mlx_native::gguf::GgufFile;
let path = std::path::Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!(
"skipping: mmproj fixture not found at {}",
GEMMA4_MMPROJ_PATH
);
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let device = mlx_native::MlxDevice::new().expect("device");
let weights = LoadedMmprojWeights::load(&gguf, &cfg, device).expect("load");
let hidden = cfg.hidden_size as usize;
let num_patches = 196usize;
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 patch_embed =
patch_embed_from_mmproj_weights(&pixels, &weights, &cfg).expect("patch_embed");
let block_out =
apply_vit_block_forward(patch_embed, &weights, &cfg, 0).expect("block_forward");
assert_eq!(block_out.len(), num_patches * hidden);
for v in &block_out {
assert!(v.is_finite(), "non-finite: {v}");
}
let p0 = &block_out[0..hidden];
let p_last = &block_out[(num_patches - 1) * hidden..num_patches * hidden];
let l2: f32 = p0
.iter()
.zip(p_last.iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<f32>()
.sqrt();
assert!(l2 > 1e-3, "wrapper produced stride-bugged output");
}
#[test]
fn apply_vit_block_forward_rejects_mismatched_hidden() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::mmproj::MmprojConfig;
use super::super::mmproj_weights::LoadedMmprojWeights;
use mlx_native::gguf::GgufFile;
let path = std::path::Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!(
"skipping: mmproj fixture not found at {}",
GEMMA4_MMPROJ_PATH
);
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let device = mlx_native::MlxDevice::new().expect("device");
let weights = LoadedMmprojWeights::load(&gguf, &cfg, device).expect("load");
let bad = vec![0f32; 196 * 1152 + 1];
let err = apply_vit_block_forward(bad, &weights, &cfg, 0).unwrap_err();
assert!(format!("{err}").contains("not divisible"));
}
#[test]
fn block_0_full_forward_on_real_gemma4() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::mmproj::MmprojConfig;
use super::super::mmproj_weights::LoadedMmprojWeights;
use mlx_native::gguf::GgufFile;
let path = std::path::Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!(
"skipping: mmproj fixture not found at {}",
GEMMA4_MMPROJ_PATH
);
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let device = mlx_native::MlxDevice::new().expect("device");
let weights = LoadedMmprojWeights::load(&gguf, &cfg, device).expect("load");
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 num_patches = 196usize;
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 residual_stream =
patch_embed_from_mmproj_weights(&pixels, &weights, &cfg).expect("patch_embed");
let mut hidden_states = residual_stream.clone();
let widen = |name: &str| -> Vec<f32> {
let buf = weights.block_tensor(0, name).expect(name);
weights.tensor_as_f32_owned(buf).expect(name)
};
let ln1_w = widen("ln1.weight");
let attn_q_w = widen("attn_q.weight");
let attn_k_w = widen("attn_k.weight");
let attn_v_w = widen("attn_v.weight");
let attn_q_norm_w = widen("attn_q_norm.weight");
let attn_k_norm_w = widen("attn_k_norm.weight");
let attn_output_w = widen("attn_output.weight");
let ln2_w = widen("ln2.weight");
let ffn_gate_w = widen("ffn_gate.weight");
let ffn_up_w = widen("ffn_up.weight");
let ffn_down_w = widen("ffn_down.weight");
let post_ffw_norm_w = widen("post_ffw_norm.weight");
rms_norm_forward(&mut hidden_states, &ln1_w, hidden, 1e-6).unwrap();
let (mut q, mut k, v) = qkv_projection_forward(
&hidden_states,
&attn_q_w,
&attn_k_w,
&attn_v_w,
num_patches,
hidden,
)
.unwrap();
per_head_rms_norm_forward(
&mut q,
&attn_q_norm_w,
num_patches,
num_heads,
head_dim,
1e-6,
)
.unwrap();
per_head_rms_norm_forward(
&mut k,
&attn_k_norm_w,
num_patches,
num_heads,
head_dim,
1e-6,
)
.unwrap();
let attn =
scaled_dot_product_attention(&q, &k, &v, num_patches, num_heads, head_dim).unwrap();
let attn_projected =
linear_forward(&attn, &attn_output_w, None, num_patches, hidden, hidden).unwrap();
let mut post_attn = residual_stream.clone();
residual_add(&mut post_attn, &attn_projected).unwrap();
let mut pre_ffn = post_attn.clone();
rms_norm_forward(&mut pre_ffn, &ln2_w, hidden, 1e-6).unwrap();
let mut gate = linear_forward(
&pre_ffn,
&ffn_gate_w,
None,
num_patches,
hidden,
intermediate,
)
.unwrap();
let up =
linear_forward(&pre_ffn, &ffn_up_w, None, num_patches, hidden, intermediate).unwrap();
silu_in_place(&mut gate);
elementwise_mul_in_place(&mut gate, &up).unwrap();
let activated = gate;
let mut down = linear_forward(
&activated,
&ffn_down_w,
None,
num_patches,
intermediate,
hidden,
)
.expect("ffn_down");
assert_eq!(down.len(), num_patches * hidden);
rms_norm_forward(&mut down, &post_ffw_norm_w, hidden, 1e-6).expect("post_ffw_norm");
let mut block_out = post_attn.clone();
residual_add(&mut block_out, &down).expect("ffn residual");
assert_eq!(block_out.len(), num_patches * hidden);
for val in &block_out {
assert!(val.is_finite(), "block_out non-finite: {val}");
}
let delta_l2: f32 = post_attn
.iter()
.zip(block_out.iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<f32>()
.sqrt();
assert!(delta_l2 > 1e-3, "FFN half was a silent no-op");
let p0 = &block_out[0..hidden];
let p_last = &block_out[(num_patches - 1) * hidden..num_patches * hidden];
let cross_l2: f32 = p0
.iter()
.zip(p_last.iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<f32>()
.sqrt();
assert!(
cross_l2 > 1e-3,
"token 0 and token 195 identical post-block — probable stride bug"
);
}
#[test]
fn silu_zero_is_zero_exactly() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![0.0f32];
silu_in_place(&mut x);
assert_eq!(x[0], 0.0);
}
#[test]
fn silu_reference_values_match_pytorch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let inputs = [-3.0f32, -1.0, -0.5, 0.5, 1.0, 3.0];
let expected = [
-0.14227762_f32,
-0.26894143,
-0.18877034,
0.31122967,
0.73105858,
2.85772238,
];
let mut x = inputs.to_vec();
silu_in_place(&mut x);
for (got, want) in x.iter().zip(expected.iter()) {
assert!((*got - *want).abs() < 1e-5, "got {got} want {want}");
}
}
#[test]
fn silu_large_positive_approaches_x() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![20.0f32];
silu_in_place(&mut x);
assert!((x[0] - 20.0).abs() < 1e-4);
}
#[test]
fn silu_large_negative_approaches_zero() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![-20.0f32];
silu_in_place(&mut x);
assert!(x[0].abs() < 1e-4);
}
#[test]
fn silu_has_local_minimum_near_negative_point_two_eight() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![-1.0f32, -1.28, -1.5];
silu_in_place(&mut x);
assert!(x[1] < x[0] && x[1] < x[2], "got {:?}", x);
}
#[test]
fn elementwise_mul_pairs_each_index() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut a = vec![1.0f32, 2.0, 3.0, 4.0];
let b = vec![5.0f32, 6.0, 7.0, 8.0];
elementwise_mul_in_place(&mut a, &b).unwrap();
assert_eq!(a, vec![5.0, 12.0, 21.0, 32.0]);
}
#[test]
fn elementwise_mul_with_zero_zeros_everything() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut a = vec![1.0f32, 2.0, 3.0];
let b = vec![0f32; 3];
elementwise_mul_in_place(&mut a, &b).unwrap();
assert_eq!(a, vec![0f32; 3]);
}
#[test]
fn elementwise_mul_rejects_shape_mismatch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut a = vec![0f32; 4];
let b = vec![0f32; 3];
let err = elementwise_mul_in_place(&mut a, &b).unwrap_err();
assert!(format!("{err}").contains("len"));
}
#[test]
fn swiglu_gated_activation_on_real_gemma4_ffn() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::mmproj::MmprojConfig;
use super::super::mmproj_weights::LoadedMmprojWeights;
use mlx_native::gguf::GgufFile;
let path = std::path::Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!(
"skipping: mmproj fixture not found at {}",
GEMMA4_MMPROJ_PATH
);
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let device = mlx_native::MlxDevice::new().expect("device");
let weights = LoadedMmprojWeights::load(&gguf, &cfg, device).expect("load");
let hidden = cfg.hidden_size as usize;
let num_heads = cfg.num_attention_heads as usize;
let head_dim = hidden / num_heads;
let num_patches = 196usize;
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 residual_stream =
patch_embed_from_mmproj_weights(&pixels, &weights, &cfg).expect("patch_embed");
let widen = |suffix: &str| -> Vec<f32> {
let buf = weights.block_tensor(0, suffix).expect("block_tensor");
weights
.tensor_as_f32_owned(buf)
.expect("tensor_as_f32_owned")
};
let ln1_w = widen("ln1.weight");
let attn_q_w = widen("attn_q.weight");
let attn_k_w = widen("attn_k.weight");
let attn_v_w = widen("attn_v.weight");
let attn_q_norm_w = widen("attn_q_norm.weight");
let attn_k_norm_w = widen("attn_k_norm.weight");
let attn_output_w = widen("attn_output.weight");
let ln2_w = widen("ln2.weight");
let mut hidden_states = residual_stream.clone();
rms_norm_forward(&mut hidden_states, &ln1_w, hidden, 1e-6).unwrap();
let (mut q, mut k, v) = qkv_projection_forward(
&hidden_states,
&attn_q_w,
&attn_k_w,
&attn_v_w,
num_patches,
hidden,
)
.unwrap();
per_head_rms_norm_forward(
&mut q,
&attn_q_norm_w,
num_patches,
num_heads,
head_dim,
1e-6,
)
.unwrap();
per_head_rms_norm_forward(
&mut k,
&attn_k_norm_w,
num_patches,
num_heads,
head_dim,
1e-6,
)
.unwrap();
let attn =
scaled_dot_product_attention(&q, &k, &v, num_patches, num_heads, head_dim).unwrap();
let attn_projected =
linear_forward(&attn, &attn_output_w, None, num_patches, hidden, hidden).unwrap();
let mut post_attn = residual_stream.clone();
residual_add(&mut post_attn, &attn_projected).unwrap();
let mut pre_ffn = post_attn.clone();
rms_norm_forward(&mut pre_ffn, &ln2_w, hidden, 1e-6).unwrap();
let intermediate = cfg.intermediate_size as usize;
assert_eq!(intermediate, 4304);
let gate_w = widen("ffn_gate.weight");
let up_w = widen("ffn_up.weight");
assert_eq!(gate_w.len(), intermediate * hidden);
assert_eq!(up_w.len(), intermediate * hidden);
let mut gate = linear_forward(&pre_ffn, &gate_w, None, num_patches, hidden, intermediate)
.expect("gate proj");
let up = linear_forward(&pre_ffn, &up_w, None, num_patches, hidden, intermediate)
.expect("up proj");
silu_in_place(&mut gate);
elementwise_mul_in_place(&mut gate, &up).expect("mul");
let activated = gate; assert_eq!(activated.len(), num_patches * intermediate);
for val in &activated {
assert!(val.is_finite(), "non-finite: {val}");
}
let mean: f32 = activated.iter().sum::<f32>() / (activated.len() as f32);
let var: f32 =
activated.iter().map(|v| (v - mean).powi(2)).sum::<f32>() / (activated.len() as f32);
assert!(var > 1e-6, "activated variance too low: {var}");
let p0 = &activated[0..intermediate];
let p_last = &activated[(num_patches - 1) * intermediate..num_patches * intermediate];
let l2_diff: f32 = p0
.iter()
.zip(p_last.iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<f32>()
.sqrt();
assert!(l2_diff > 1e-3, "all patches produce identical gated output");
}
#[test]
fn residual_add_is_elementwise() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut a = vec![1.0f32, 2.0, 3.0, 4.0];
let b = vec![10.0f32, 20.0, 30.0, 40.0];
residual_add(&mut a, &b).unwrap();
assert_eq!(a, vec![11.0, 22.0, 33.0, 44.0]);
}
#[test]
fn residual_add_with_zero_is_identity() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut a = vec![0.5f32, -1.2, 3.7];
let snap = a.clone();
let b = vec![0.0f32; 3];
residual_add(&mut a, &b).unwrap();
assert_eq!(a, snap);
}
#[test]
fn residual_add_rejects_shape_mismatch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut a = vec![0f32; 4];
let b = vec![0f32; 3];
let err = residual_add(&mut a, &b).unwrap_err();
assert!(format!("{err}").contains("len"));
}
#[test]
fn residual_add_is_commutative_against_add_in_place_loop() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let a: Vec<f32> = (0..10).map(|i| i as f32 * 0.3).collect();
let b: Vec<f32> = (0..10).map(|i| i as f32 * 0.7).collect();
let mut via_fn = a.clone();
residual_add(&mut via_fn, &b).unwrap();
let mut via_loop = a.clone();
for (x, y) in via_loop.iter_mut().zip(b.iter()) {
*x += *y;
}
assert_eq!(via_fn, via_loop);
}
#[test]
fn attention_half_block_end_to_end_real_gemma4() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::mmproj::MmprojConfig;
use super::super::mmproj_weights::LoadedMmprojWeights;
use mlx_native::gguf::GgufFile;
let path = std::path::Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!(
"skipping: mmproj fixture not found at {}",
GEMMA4_MMPROJ_PATH
);
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let device = mlx_native::MlxDevice::new().expect("device");
let weights = LoadedMmprojWeights::load(&gguf, &cfg, device).expect("load");
let hidden = cfg.hidden_size as usize;
let num_heads = cfg.num_attention_heads as usize;
let head_dim = hidden / num_heads;
let num_patches = 196usize;
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 residual_stream =
patch_embed_from_mmproj_weights(&pixels, &weights, &cfg).expect("patch_embed");
let widen = |suffix: &str| -> Vec<f32> {
let buf = weights.block_tensor(0, suffix).expect("block_tensor");
weights
.tensor_as_f32_owned(buf)
.expect("tensor_as_f32_owned")
};
let ln1_w = widen("ln1.weight");
let attn_q_w = widen("attn_q.weight");
let attn_k_w = widen("attn_k.weight");
let attn_v_w = widen("attn_v.weight");
let attn_q_norm_w = widen("attn_q_norm.weight");
let attn_k_norm_w = widen("attn_k_norm.weight");
let attn_output_w = widen("attn_output.weight");
let ln2_w = widen("ln2.weight");
let mut hidden_states = residual_stream.clone();
rms_norm_forward(&mut hidden_states, &ln1_w, hidden, 1e-6).expect("rms ln1");
let (mut q, mut k, v) = qkv_projection_forward(
&hidden_states,
&attn_q_w,
&attn_k_w,
&attn_v_w,
num_patches,
hidden,
)
.expect("qkv");
per_head_rms_norm_forward(
&mut q,
&attn_q_norm_w,
num_patches,
num_heads,
head_dim,
1e-6,
)
.expect("per-head Q");
per_head_rms_norm_forward(
&mut k,
&attn_k_norm_w,
num_patches,
num_heads,
head_dim,
1e-6,
)
.expect("per-head K");
let attn = scaled_dot_product_attention(&q, &k, &v, num_patches, num_heads, head_dim)
.expect("attention");
assert_eq!(attn_output_w.len(), hidden * hidden);
let attn_projected =
linear_forward(&attn, &attn_output_w, None, num_patches, hidden, hidden)
.expect("attn_output proj");
let mut post_attn = residual_stream.clone();
residual_add(&mut post_attn, &attn_projected).expect("residual");
let mut pre_ffn = post_attn.clone();
rms_norm_forward(&mut pre_ffn, &ln2_w, hidden, 1e-6).expect("rms ln2");
for t in [&attn_projected, &post_attn, &pre_ffn] {
assert_eq!(t.len(), num_patches * hidden);
for v in t.iter() {
assert!(v.is_finite(), "non-finite: {v}");
}
}
let delta_l2: f32 = residual_stream
.iter()
.zip(post_attn.iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<f32>()
.sqrt();
assert!(
delta_l2 > 1e-3,
"post-attn residual is identical to pre-attn — attn_output all zero?"
);
let ln_delta: f32 = post_attn
.iter()
.zip(pre_ffn.iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<f32>()
.sqrt();
assert!(ln_delta > 1e-3, "ln2 was a no-op");
}
#[test]
fn attention_uniform_qk_averages_v_across_tokens() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 3;
let num_heads = 2;
let head_dim = 4;
let n = batch * num_heads * head_dim;
let q = vec![1f32; n];
let k = vec![1f32; n];
let idx = |b: usize, h: usize, d: usize| b * num_heads * head_dim + h * head_dim + d;
let mut v = vec![0f32; n];
for b in 0..batch {
for h in 0..num_heads {
for d in 0..head_dim {
v[idx(b, h, d)] = if h == 0 {
(b * 4 + d + 1) as f32
} else {
((b * 4 + d + 1) * 10) as f32
};
}
}
}
let out = scaled_dot_product_attention(&q, &k, &v, batch, num_heads, head_dim).unwrap();
let expected_h0 = [5.0f32, 6.0, 7.0, 8.0];
let expected_h1 = [50.0f32, 60.0, 70.0, 80.0];
for b in 0..batch {
for d in 0..head_dim {
let got_h0 = out[idx(b, 0, d)];
let got_h1 = out[idx(b, 1, d)];
assert!(
(got_h0 - expected_h0[d]).abs() < 1e-4,
"h0 b{b} d{d}: {got_h0} vs {}",
expected_h0[d]
);
assert!(
(got_h1 - expected_h1[d]).abs() < 1e-4,
"h1 b{b} d{d}: {got_h1} vs {}",
expected_h1[d]
);
}
}
}
#[test]
fn attention_single_key_dominant_selects_its_value() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 3;
let num_heads = 1;
let head_dim = 4;
let n = batch * num_heads * head_dim;
let mut q = vec![0f32; n];
for b in 0..batch {
q[b * head_dim] = 1.0;
}
let mut k = vec![0f32; n];
let kidx = |tok: usize, dim: usize| tok * head_dim + dim;
k[kidx(0, 0)] = 100.0; k[kidx(1, 1)] = 100.0;
k[kidx(2, 2)] = 100.0;
let v = vec![
7.0, 7.0, 7.0, 7.0, 1.0, 2.0, 3.0, 4.0, 99.0, 99.0, 99.0, 99.0, ];
let out = scaled_dot_product_attention(&q, &k, &v, batch, num_heads, head_dim).unwrap();
for b in 0..batch {
for d in 0..head_dim {
let got = out[b * head_dim + d];
assert!((got - 7.0).abs() < 0.1, "b{b} d{d}: {got} (expected ≈7.0)");
}
}
}
#[test]
fn attention_scale_factor_applied_to_logits() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 4;
let num_heads = 1;
let head_dim = 1024; let n = batch * num_heads * head_dim;
let q = vec![3.0f32; n];
let k = vec![3.0f32; n];
let v = vec![1.0f32; n];
let out = scaled_dot_product_attention(&q, &k, &v, batch, num_heads, head_dim).unwrap();
for val in &out {
assert!(val.is_finite(), "non-finite: {val}");
}
for val in &out {
assert!((*val - 1.0).abs() < 1e-3, "got {val}");
}
}
#[test]
fn attention_rejects_zero_dims() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let err = scaled_dot_product_attention(&[], &[], &[], 0, 1, 1).unwrap_err();
assert!(format!("{err}").contains("must all be > 0"));
}
#[test]
fn attention_rejects_mismatched_q_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 2;
let num_heads = 2;
let head_dim = 2;
let n = batch * num_heads * head_dim;
let q = vec![0f32; n - 1];
let k = vec![0f32; n];
let v = vec![0f32; n];
let err = scaled_dot_product_attention(&q, &k, &v, batch, num_heads, head_dim).unwrap_err();
assert!(format!("{err}").contains("q len"));
}
#[test]
fn attention_rejects_mismatched_k_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 2;
let num_heads = 2;
let head_dim = 2;
let n = batch * num_heads * head_dim;
let q = vec![0f32; n];
let k = vec![0f32; n + 1];
let v = vec![0f32; n];
let err = scaled_dot_product_attention(&q, &k, &v, batch, num_heads, head_dim).unwrap_err();
assert!(format!("{err}").contains("k len"));
}
#[test]
fn attention_end_to_end_real_gemma4_full_self_attention() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::mmproj::MmprojConfig;
use super::super::mmproj_weights::LoadedMmprojWeights;
use mlx_native::gguf::GgufFile;
let path = std::path::Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!(
"skipping: mmproj fixture not found at {}",
GEMMA4_MMPROJ_PATH
);
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let device = mlx_native::MlxDevice::new().expect("device");
let weights = LoadedMmprojWeights::load(&gguf, &cfg, device).expect("load");
let hidden = cfg.hidden_size as usize;
let num_heads = cfg.num_attention_heads as usize;
let head_dim = hidden / num_heads;
let num_patches = 196usize;
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 mut hidden_states =
patch_embed_from_mmproj_weights(&pixels, &weights, &cfg).expect("patch_embed");
let widen = |suffix: &str| -> Vec<f32> {
let buf = weights.block_tensor(0, suffix).expect("block_tensor");
weights
.tensor_as_f32_owned(buf)
.expect("tensor_as_f32_owned")
};
let ln1_w = widen("ln1.weight");
let attn_q_w = widen("attn_q.weight");
let attn_k_w = widen("attn_k.weight");
let attn_v_w = widen("attn_v.weight");
let attn_q_norm_w = widen("attn_q_norm.weight");
let attn_k_norm_w = widen("attn_k_norm.weight");
rms_norm_forward(&mut hidden_states, &ln1_w, hidden, 1e-6).expect("rms ln1");
let (mut q, mut k, v) = qkv_projection_forward(
&hidden_states,
&attn_q_w,
&attn_k_w,
&attn_v_w,
num_patches,
hidden,
)
.expect("qkv");
per_head_rms_norm_forward(
&mut q,
&attn_q_norm_w,
num_patches,
num_heads,
head_dim,
1e-6,
)
.expect("per-head Q");
per_head_rms_norm_forward(
&mut k,
&attn_k_norm_w,
num_patches,
num_heads,
head_dim,
1e-6,
)
.expect("per-head K");
let attn_out = scaled_dot_product_attention(&q, &k, &v, num_patches, num_heads, head_dim)
.expect("attention");
assert_eq!(attn_out.len(), num_patches * hidden);
for val in &attn_out {
assert!(val.is_finite(), "non-finite attention output: {val}");
}
let token_0 = &attn_out[0..hidden];
let token_last = &attn_out[(num_patches - 1) * hidden..num_patches * hidden];
let l2_diff: f32 = token_0
.iter()
.zip(token_last.iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<f32>()
.sqrt();
assert!(
l2_diff > 1e-3,
"token 0 and token 195 attention outputs identical — stride bug?"
);
}
#[test]
fn per_head_rms_norm_with_unit_gain_normalizes_each_head_slice_independently() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 2;
let num_heads = 3;
let head_dim = 4;
let mut input = vec![0f32; batch * num_heads * head_dim];
for b in 0..batch {
for h in 0..num_heads {
for d in 0..head_dim {
input[b * (num_heads * head_dim) + h * head_dim + d] =
((h + 1) * 100) as f32 + (d as f32);
}
}
}
let gamma = vec![1.0f32; head_dim];
per_head_rms_norm_forward(&mut input, &gamma, batch, num_heads, head_dim, 1e-6).unwrap();
for b in 0..batch {
for h in 0..num_heads {
let off = b * (num_heads * head_dim) + h * head_dim;
let slice = &input[off..off + head_dim];
let ms: f32 = slice.iter().map(|v| v * v).sum::<f32>() / (head_dim as f32);
assert!((ms - 1.0).abs() < 1e-3, "head ({},{}): ms = {}", b, h, ms);
}
}
}
#[test]
fn per_head_rms_norm_broadcasts_same_gamma_across_heads() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 1;
let num_heads = 3;
let head_dim = 4;
let mut input = vec![5.0f32; batch * num_heads * head_dim];
let gamma = vec![0.5f32, 1.0, 2.0, 4.0];
per_head_rms_norm_forward(&mut input, &gamma, batch, num_heads, head_dim, 1e-6).unwrap();
for h in 0..num_heads {
let off = h * head_dim;
for (i, g) in gamma.iter().enumerate() {
assert!(
(input[off + i] - g).abs() < 1e-4,
"head {}: got {}, want {}",
h,
input[off + i],
g
);
}
}
}
#[test]
fn per_head_rms_norm_rejects_zero_dims() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let err = per_head_rms_norm_forward(&mut [], &[], 0, 1, 1, 1e-5).unwrap_err();
assert!(format!("{err}").contains("must all be > 0"));
}
#[test]
fn per_head_rms_norm_rejects_mismatched_input_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut input = vec![0f32; 5];
let gamma = vec![1f32; 2];
let err = per_head_rms_norm_forward(&mut input, &gamma, 2, 2, 2, 1e-5).unwrap_err();
assert!(format!("{err}").contains("input len"));
}
#[test]
fn per_head_rms_norm_rejects_wrong_gamma_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut input = vec![0f32; 8];
let gamma = vec![1f32; 3]; let err = per_head_rms_norm_forward(&mut input, &gamma, 2, 2, 2, 1e-5).unwrap_err();
assert!(format!("{err}").contains("gamma len"));
}
#[test]
fn per_head_rms_norm_end_to_end_real_gemma4_chain() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::mmproj::MmprojConfig;
use super::super::mmproj_weights::LoadedMmprojWeights;
use mlx_native::gguf::GgufFile;
let path = std::path::Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!(
"skipping: mmproj fixture not found at {}",
GEMMA4_MMPROJ_PATH
);
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let device = mlx_native::MlxDevice::new().expect("device");
let weights = LoadedMmprojWeights::load(&gguf, &cfg, device).expect("load");
let hidden = cfg.hidden_size as usize;
let num_heads = cfg.num_attention_heads as usize;
let head_dim = hidden / num_heads;
assert_eq!(num_heads, 16);
assert_eq!(head_dim, 72);
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 mut hidden_states =
patch_embed_from_mmproj_weights(&pixels, &weights, &cfg).expect("patch_embed");
let num_patches = 196usize;
assert_eq!(hidden_states.len(), num_patches * hidden);
let widen = |suffix: &str| -> Vec<f32> {
let buf = weights.block_tensor(0, suffix).expect("block_tensor");
weights
.tensor_as_f32_owned(buf)
.expect("tensor_as_f32_owned")
};
let ln1_gamma = widen("ln1.weight");
let q_w = widen("attn_q.weight");
let k_w = widen("attn_k.weight");
let v_w = widen("attn_v.weight");
let q_norm_gamma = widen("attn_q_norm.weight");
let k_norm_gamma = widen("attn_k_norm.weight");
rms_norm_forward(&mut hidden_states, &ln1_gamma, hidden, 1e-6).expect("rms ln1");
let (mut q, mut k, v) =
qkv_projection_forward(&hidden_states, &q_w, &k_w, &v_w, num_patches, hidden)
.expect("qkv");
assert_eq!(q_norm_gamma.len(), head_dim);
assert_eq!(k_norm_gamma.len(), head_dim);
per_head_rms_norm_forward(
&mut q,
&q_norm_gamma,
num_patches,
num_heads,
head_dim,
1e-6,
)
.expect("per-head Q");
per_head_rms_norm_forward(
&mut k,
&k_norm_gamma,
num_patches,
num_heads,
head_dim,
1e-6,
)
.expect("per-head K");
for t in [&q, &k, &v] {
for val in t.iter() {
assert!(val.is_finite(), "non-finite: {val}");
}
}
let q_slice0 = &q[0..head_dim];
let q_ms: f32 = q_slice0.iter().map(|v| v * v).sum::<f32>() / (head_dim as f32);
assert!(
q_ms > 0.01 && q_ms < 100.0,
"Q head-0 mean(x²) outside expected range: {q_ms}"
);
}
#[test]
fn linear_identity_weight_preserves_input() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let d = 4;
let batch = 2;
let mut weight = vec![0f32; d * d];
for i in 0..d {
weight[i * d + i] = 1.0;
}
let input: Vec<f32> = (0..(batch * d)).map(|i| i as f32).collect();
let out = linear_forward(&input, &weight, None, batch, d, d).unwrap();
assert_eq!(out, input);
}
#[test]
fn linear_all_ones_weight_produces_row_sum_per_output() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 3;
let d_in = 4;
let d_out = 2;
let weight = vec![1f32; d_out * d_in];
let input: Vec<f32> = vec![1., 2., 3., 4., 5., 6., 7., 8., 9., 10., 11., 12.];
let out = linear_forward(&input, &weight, None, batch, d_in, d_out).unwrap();
let expected = [10.0, 10.0, 26.0, 26.0, 42.0, 42.0];
for (got, want) in out.iter().zip(expected.iter()) {
assert!((*got - *want).abs() < 1e-5, "got {got} want {want}");
}
}
#[test]
fn linear_applies_bias_per_output_once() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 3;
let d_in = 2;
let d_out = 3;
let weight = vec![0f32; d_out * d_in];
let bias = vec![0.5f32, 1.5, 2.5];
let input = vec![99f32; batch * d_in]; let out = linear_forward(&input, &weight, Some(&bias), batch, d_in, d_out).unwrap();
assert_eq!(out.len(), batch * d_out);
for n in 0..batch {
for o in 0..d_out {
assert!((out[n * d_out + o] - bias[o]).abs() < 1e-6);
}
}
}
#[test]
fn linear_reference_dot_product() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let input = vec![1.0f32, 2.0, 3.0];
let weight = vec![0.1f32, 0.2, 0.3, 0.4, 0.5, 0.6];
let out = linear_forward(&input, &weight, None, 1, 3, 2).unwrap();
assert!((out[0] - 1.4).abs() < 1e-5);
assert!((out[1] - 3.2).abs() < 1e-5);
}
#[test]
fn linear_rejects_mismatched_input_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let weight = vec![0f32; 6];
let input = vec![0f32; 5]; let err = linear_forward(&input, &weight, None, 2, 3, 2).unwrap_err();
assert!(format!("{err}").contains("input len"));
}
#[test]
fn linear_rejects_mismatched_weight_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let weight = vec![0f32; 5]; let input = vec![0f32; 6];
let err = linear_forward(&input, &weight, None, 2, 3, 2).unwrap_err();
assert!(format!("{err}").contains("weight len"));
}
#[test]
fn linear_rejects_wrong_bias_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let weight = vec![0f32; 6];
let input = vec![0f32; 6];
let bias = vec![0f32; 3]; let err = linear_forward(&input, &weight, Some(&bias), 2, 3, 2).unwrap_err();
assert!(format!("{err}").contains("bias len"));
}
#[test]
fn linear_rejects_zero_dims() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let err = linear_forward(&[], &[], None, 0, 1, 1).unwrap_err();
assert!(format!("{err}").contains("must all be > 0"));
}
#[test]
fn qkv_projection_returns_three_tensors_of_expected_shape() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 2;
let hidden = 4;
let input = vec![1f32; batch * hidden];
let q_w = vec![1f32; hidden * hidden];
let k_w = vec![2f32; hidden * hidden];
let v_w = vec![3f32; hidden * hidden];
let (q, k, v) = qkv_projection_forward(&input, &q_w, &k_w, &v_w, batch, hidden).unwrap();
assert_eq!(q.len(), batch * hidden);
assert_eq!(k.len(), batch * hidden);
assert_eq!(v.len(), batch * hidden);
for val in &q {
assert!((*val - 4.0).abs() < 1e-5);
}
for val in &k {
assert!((*val - 8.0).abs() < 1e-5);
}
for val in &v {
assert!((*val - 12.0).abs() < 1e-5);
}
}
#[test]
fn qkv_projection_propagates_shape_errors() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 2;
let hidden = 4;
let input = vec![1f32; batch * hidden];
let q_w = vec![1f32; hidden * hidden];
let k_w = vec![2f32; hidden * hidden - 1]; let v_w = vec![3f32; hidden * hidden];
let err = qkv_projection_forward(&input, &q_w, &k_w, &v_w, batch, hidden).unwrap_err();
let msg = format!("{err}");
assert!(msg.contains("qkv_projection_forward K"), "got: {msg}");
}
#[test]
fn qkv_projection_against_real_gemma4_block0_weights() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::mmproj::MmprojConfig;
use super::super::mmproj_weights::LoadedMmprojWeights;
use mlx_native::gguf::GgufFile;
let path = std::path::Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!(
"skipping: mmproj fixture not found at {}",
GEMMA4_MMPROJ_PATH
);
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let device = mlx_native::MlxDevice::new().expect("device");
let weights = LoadedMmprojWeights::load(&gguf, &cfg, device).expect("load");
let q_w = weights
.tensor_as_f32_owned(weights.block_tensor(0, "attn_q.weight").expect("attn_q"))
.expect("widen attn_q");
let k_w = weights
.tensor_as_f32_owned(weights.block_tensor(0, "attn_k.weight").expect("attn_k"))
.expect("widen attn_k");
let v_w = weights
.tensor_as_f32_owned(weights.block_tensor(0, "attn_v.weight").expect("attn_v"))
.expect("widen attn_v");
let hidden = cfg.hidden_size as usize;
assert_eq!(q_w.len(), hidden * hidden);
assert_eq!(k_w.len(), hidden * hidden);
assert_eq!(v_w.len(), hidden * hidden);
let batch = 196;
let mut input = vec![0f32; batch * hidden];
for p in 0..batch {
for h in 0..hidden {
input[p * hidden + h] = 0.01 + (p as f32) * 0.0001 + (h as f32) * 1e-5;
}
}
let (q, k, v) =
qkv_projection_forward(&input, &q_w, &k_w, &v_w, batch, hidden).expect("qkv");
assert_eq!(q.len(), batch * hidden);
let var = |t: &[f32]| {
let m = t.iter().sum::<f32>() / (t.len() as f32);
t.iter().map(|x| (x - m).powi(2)).sum::<f32>() / (t.len() as f32)
};
let q_var = var(&q);
let k_var = var(&k);
let v_var = var(&v);
assert!(q_var > 1e-6, "Q variance too low: {q_var}");
assert!(k_var > 1e-6, "K variance too low: {k_var}");
assert!(v_var > 1e-6, "V variance too low: {v_var}");
assert!(
(q_var - k_var).abs() > 1e-8 || (q_var - v_var).abs() > 1e-8,
"Q/K/V variances all equal — weight aliasing bug?"
);
for t in [&q, &k, &v] {
for val in t.iter() {
assert!(val.is_finite(), "non-finite in projection output: {val}");
}
}
}
#[test]
fn rms_norm_unit_gamma_normalizes_to_unit_rms() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32, 2.0, 3.0, 4.0];
let gamma = vec![1.0f32; 4];
rms_norm_forward(&mut x, &gamma, 4, 1e-6).unwrap();
let ms: f32 = x.iter().map(|v| v * v).sum::<f32>() / 4.0;
assert!(
(ms - 1.0).abs() < 1e-4,
"mean-squared should be ~1, got {ms}"
);
}
#[test]
fn rms_norm_pytorch_reference_values() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32, 2.0, 3.0, 4.0];
let gamma = vec![1.0f32; 4];
rms_norm_forward(&mut x, &gamma, 4, 1e-5).unwrap();
let expected = [0.36514837f32, 0.73029674, 1.09544511, 1.46059349];
for (got, want) in x.iter().zip(expected.iter()) {
assert!((*got - *want).abs() < 1e-5, "got {got}, want {want}");
}
}
#[test]
fn rms_norm_applies_gain_elementwise() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![5.0f32; 4];
let gamma = vec![0.5f32, 1.0, 2.0, 3.0];
rms_norm_forward(&mut x, &gamma, 4, 1e-6).unwrap();
for (got, want) in x.iter().zip(gamma.iter()) {
assert!((*got - *want).abs() < 1e-4, "got {got}, want {want}");
}
}
#[test]
fn rms_norm_normalizes_rows_independently() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32, 2.0, 3.0, 4.0, 100.0, 200.0, 300.0, 400.0];
let gamma = vec![1.0f32; 4];
rms_norm_forward(&mut x, &gamma, 4, 1e-6).unwrap();
for i in 0..4 {
assert!(
(x[i] - x[4 + i]).abs() < 1e-3,
"row 0 elem {} = {} != row 1 = {}",
i,
x[i],
x[4 + i]
);
}
}
#[test]
fn rms_norm_does_not_divide_by_zero_when_input_is_zero() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![0.0f32; 8];
let gamma = vec![1.0f32; 8];
rms_norm_forward(&mut x, &gamma, 8, 1e-6).unwrap();
for v in &x {
assert!(v.is_finite(), "got non-finite {}", v);
assert!(v.abs() < 1e-5, "zero input should stay zero, got {v}");
}
}
#[test]
fn rms_norm_large_inputs_do_not_overflow() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1e18f32, 2e18, 3e18, 4e18];
let gamma = vec![1.0f32; 4];
rms_norm_forward(&mut x, &gamma, 4, 1e-6).unwrap();
for v in &x {
assert!(v.is_finite(), "got {v} — f32 overflow");
}
let ms: f32 = x.iter().map(|v| v * v).sum::<f32>() / 4.0;
assert!((ms - 1.0).abs() < 1e-3, "ms = {ms}");
}
#[test]
fn rms_norm_rejects_hidden_zero() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32; 4];
let err = rms_norm_forward(&mut x, &[], 0, 1e-5).unwrap_err();
assert!(format!("{err}").contains("hidden must be > 0"));
}
#[test]
fn rms_norm_rejects_non_divisible_input_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32; 7];
let gamma = vec![1.0f32; 3];
let err = rms_norm_forward(&mut x, &gamma, 3, 1e-5).unwrap_err();
assert!(format!("{err}").contains("not divisible"));
}
#[test]
fn rms_norm_rejects_wrong_gamma_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32; 4];
let gamma = vec![1.0f32; 3];
let err = rms_norm_forward(&mut x, &gamma, 4, 1e-5).unwrap_err();
assert!(format!("{err}").contains("gamma len"));
}
#[test]
fn rms_norm_runs_against_real_gemma4_ln1_weights() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::mmproj::MmprojConfig;
use super::super::mmproj_weights::LoadedMmprojWeights;
use mlx_native::gguf::GgufFile;
let path = std::path::Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!(
"skipping: mmproj fixture not found at {}",
GEMMA4_MMPROJ_PATH
);
return;
}
let gguf = GgufFile::open(path).expect("open");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
let device = mlx_native::MlxDevice::new().expect("device");
let weights = LoadedMmprojWeights::load(&gguf, &cfg, device).expect("load");
let ln1 = weights.block_tensor(0, "ln1.weight").expect("ln1");
let gamma: &[f32] = ln1.as_slice::<f32>().expect("as_slice f32");
assert_eq!(gamma.len(), cfg.hidden_size as usize);
let num_patches = 196usize;
let hidden = cfg.hidden_size as usize;
let mut input = vec![0f32; num_patches * hidden];
for p in 0..num_patches {
for h in 0..hidden {
input[p * hidden + h] = 0.01 + (p as f32) * 0.001 + (h as f32) * 0.0001;
}
}
rms_norm_forward(&mut input, gamma, hidden, 1e-6).expect("rms_norm");
for v in &input {
assert!(v.is_finite(), "non-finite post-RMS: {v}");
}
let mean_sq: f32 = input.iter().map(|v| v * v).sum::<f32>() / (input.len() as f32);
assert!(
mean_sq > 1e-4 && mean_sq < 1e6,
"mean_sq {} outside sane range",
mean_sq
);
}
#[test]
fn rms_norm_differs_from_layer_norm_on_nonzero_mean_input() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x_rms = vec![2.0f32, 2.0, 2.0, 2.0];
let mut x_ln = x_rms.clone();
let gamma = vec![1.0f32; 4];
let beta = vec![0.0f32; 4];
rms_norm_forward(&mut x_rms, &gamma, 4, 1e-6).unwrap();
layer_norm_forward(&mut x_ln, &gamma, &beta, 4, 1e-6).unwrap();
for v in &x_rms {
assert!((*v - 1.0).abs() < 1e-5, "RMS output not 1: {v}");
}
for v in &x_ln {
assert!(v.abs() < 1e-5, "LN output not 0: {v}");
}
}
#[test]
fn softmax_sums_to_one_per_row() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32, 2.0, 3.0, 4.0, -1.0, 0.5, 2.0, -3.0];
softmax_last_dim(&mut x, 4).unwrap();
for row in 0..2 {
let s: f32 = x[row * 4..row * 4 + 4].iter().sum();
assert!((s - 1.0).abs() < 1e-5, "row {} sum = {}", row, s);
}
}
#[test]
fn softmax_uniform_input_yields_uniform_output() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![5.0f32; 8];
softmax_last_dim(&mut x, 8).unwrap();
for v in &x {
assert!((*v - 0.125).abs() < 1e-6, "expected 1/8 = 0.125, got {}", v);
}
}
#[test]
fn softmax_reference_values_match_pytorch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32, 2.0, 3.0];
softmax_last_dim(&mut x, 3).unwrap();
assert!((x[0] - 0.09003057).abs() < 1e-6, "x[0] = {}", x[0]);
assert!((x[1] - 0.24472848).abs() < 1e-6, "x[1] = {}", x[1]);
assert!((x[2] - 0.66524094).abs() < 1e-6, "x[2] = {}", x[2]);
}
#[test]
fn softmax_is_numerically_stable_for_large_inputs() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1000.0f32, 999.0, 998.0];
softmax_last_dim(&mut x, 3).unwrap();
assert!((x[0] - 0.6652).abs() < 1e-3, "x[0] = {}", x[0]);
assert!((x[1] - 0.2447).abs() < 1e-3, "x[1] = {}", x[1]);
assert!((x[2] - 0.0900).abs() < 1e-3, "x[2] = {}", x[2]);
}
#[test]
fn softmax_concentrates_on_the_dominant_element() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![0.0f32, 20.0, 0.0, 0.0];
softmax_last_dim(&mut x, 4).unwrap();
assert!(x[1] > 0.999, "dominant element prob = {}", x[1]);
for i in [0, 2, 3] {
assert!(x[i] < 1e-8, "x[{}] = {}", i, x[i]);
}
}
#[test]
fn softmax_rejects_hidden_zero() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32; 4];
let err = softmax_last_dim(&mut x, 0).unwrap_err();
assert!(format!("{err}").contains("hidden must be > 0"));
}
#[test]
fn softmax_rejects_non_divisible_input_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![1.0f32; 7];
let err = softmax_last_dim(&mut x, 3).unwrap_err();
assert!(format!("{err}").contains("not divisible"));
}
#[test]
fn gelu_zero_is_zero() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![0.0f32];
gelu_tanh_approx(&mut x);
assert!(x[0].abs() < 1e-7, "gelu(0) should be 0, got {}", x[0]);
}
#[test]
fn gelu_reference_values_match_pytorch_tanh_approximate() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let inputs = [-3.0f32, -1.0, -0.5, 0.5, 1.0, 3.0];
let expected = [
-0.00363752_f32,
-0.15880802,
-0.15428595,
0.34571406,
0.84119198,
2.99636247,
];
let mut x = inputs.to_vec();
gelu_tanh_approx(&mut x);
for (i, (got, want)) in x.iter().zip(expected.iter()).enumerate() {
assert!(
(*got - *want).abs() < 1e-4,
"i={}: gelu({}) got {} want {}",
i,
inputs[i],
got,
want
);
}
}
#[test]
fn gelu_is_monotonic_on_nonneg_inputs() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x: Vec<f32> = (0..80).map(|i| i as f32 * 0.1).collect();
gelu_tanh_approx(&mut x);
for i in 1..x.len() {
assert!(
x[i] > x[i - 1] - 1e-6,
"gelu not monotone on x>=0 at i={}: {} -> {}",
i,
x[i - 1],
x[i]
);
}
}
#[test]
fn gelu_has_local_minimum_near_negative_point_seven_five() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut probe = vec![-1.0f32, -0.75, -0.5];
gelu_tanh_approx(&mut probe);
assert!(
probe[1] < probe[0] && probe[1] < probe[2],
"expected local min at x=-0.75, got: \
gelu(-1.0)={}, gelu(-0.75)={}, gelu(-0.5)={}",
probe[0],
probe[1],
probe[2]
);
}
#[test]
fn gelu_large_positive_approaches_x() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![10.0f32];
gelu_tanh_approx(&mut x);
assert!((x[0] - 10.0).abs() < 1e-4, "gelu(10) ≈ 10, got {}", x[0]);
}
#[test]
fn gelu_large_negative_approaches_zero() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut x = vec![-10.0f32];
gelu_tanh_approx(&mut x);
assert!(x[0].abs() < 1e-5, "gelu(-10) ≈ 0, got {}", x[0]);
}
#[test]
fn gemma4_shape_does_not_panic() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let img: usize = 56;
let patch: usize = 14;
let hidden: usize = 32;
let weight = vec![0.01f32; hidden * 3 * patch * patch];
let bias = vec![0.1f32; hidden];
let pixels = vec![0.5f32; 3 * img * img];
let out = patch_embed_forward(
&pixels,
&weight,
Some(&bias),
img as u32,
patch as u32,
hidden as u32,
)
.unwrap();
let num_patches = (img / patch) * (img / patch);
assert_eq!(out.len(), num_patches * hidden);
for v in &out {
assert!((*v - 3.04).abs() < 1e-3, "expected 3.04, got {v}");
}
}
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";
#[test]
fn patch_embed_from_real_gemma4_weights_produces_sensible_embeddings() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::mmproj::MmprojConfig;
use super::super::mmproj_weights::LoadedMmprojWeights;
use mlx_native::gguf::GgufFile;
let path = std::path::Path::new(GEMMA4_MMPROJ_PATH);
if !path.exists() {
eprintln!(
"skipping: mmproj fixture not found at {}",
GEMMA4_MMPROJ_PATH
);
return;
}
let gguf = GgufFile::open(path).expect("open gemma4 mmproj");
let cfg = MmprojConfig::from_gguf(&gguf).expect("cfg");
assert_eq!(cfg.image_size, 224);
assert_eq!(cfg.patch_size, 16);
assert_eq!(cfg.hidden_size, 1152);
let num_patches_side = (cfg.image_size / cfg.patch_size) as usize;
assert_eq!(num_patches_side, 14);
let num_patches = num_patches_side * num_patches_side;
assert_eq!(num_patches, 196);
let device = mlx_native::MlxDevice::new().expect("device");
let weights = LoadedMmprojWeights::load(&gguf, &cfg, device).expect("load");
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 {
let idx = c * img * img + y * img + x;
pixels[idx] = ((c + 1) as f32) * 0.1 + (y as f32) * 0.001 + (x as f32) * 0.001;
}
}
}
let out = patch_embed_from_mmproj_weights(&pixels, &weights, &cfg)
.expect("patch_embed_from_mmproj_weights");
assert_eq!(out.len(), num_patches * (cfg.hidden_size as usize));
let mean: f32 = out.iter().sum::<f32>() / (out.len() as f32);
let var: f32 = out.iter().map(|v| (v - mean).powi(2)).sum::<f32>() / (out.len() as f32);
assert!(var > 1e-4, "output variance unexpectedly low: {var}");
let max_abs = out.iter().map(|v| v.abs()).fold(0f32, f32::max);
assert!(max_abs > 1e-3, "output max_abs unexpectedly low: {max_abs}");
let p0 = &out[0..cfg.hidden_size as usize];
let p_last = &out[195 * (cfg.hidden_size as usize)..196 * (cfg.hidden_size as usize)];
let l2_diff: f32 = p0
.iter()
.zip(p_last.iter())
.map(|(a, b)| (a - b).powi(2))
.sum::<f32>()
.sqrt();
assert!(
l2_diff > 1e-3,
"patch 0 and patch 195 are identical — stride bug"
);
}
#[test]
fn gemma4v_patch_embed_shapes_and_math() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let n_patches = 4u32;
let inner = 12u32;
let hidden = 8u32;
let patches: Vec<f32> = (0..(n_patches * inner)).map(|i| i as f32 * 0.01).collect();
let weight: Vec<f32> = (0..(hidden * inner))
.map(|i| ((i % 5) as f32) * 0.1 - 0.2)
.collect();
let out = gemma4v_patch_embed_forward(&patches, &weight, n_patches, inner, hidden).unwrap();
assert_eq!(out.len(), (n_patches * hidden) as usize);
let n = 2usize;
let o = 3usize;
let in_us = inner as usize;
let mut expect: f32 = 0.0;
for i in 0..in_us {
expect += patches[n * in_us + i] * weight[o * in_us + i];
}
let got = out[n * (hidden as usize) + o];
assert!((got - expect).abs() < 1e-5, "got {got} expect {expect}");
}
#[test]
fn gemma4v_patch_embed_rejects_zero_dims() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let err = gemma4v_patch_embed_forward(&[1.0], &[1.0], 0, 1, 1).unwrap_err();
assert!(format!("{err}").contains("must all be > 0"));
let err2 = gemma4v_patch_embed_forward(&[1.0], &[1.0], 1, 0, 1).unwrap_err();
assert!(format!("{err2}").contains("must all be > 0"));
}
#[test]
fn gemma4v_patch_embed_rejects_shape_mismatch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let err = gemma4v_patch_embed_forward(&[1.0; 5], &[1.0; 24], 4, 12, 2).unwrap_err();
assert!(format!("{err}").contains("patches.len()"));
let err2 = gemma4v_patch_embed_forward(&[1.0; 48], &[1.0; 5], 4, 12, 2).unwrap_err();
assert!(format!("{err2}").contains("weight.len()"));
}
#[test]
fn gemma4v_position_embed_lookup_basic() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let pe: Vec<f32> = vec![
10.0, 11.0, 12.0, 13.0, 20.0, 21.0, 22.0, 23.0, 30.0, 31.0, 32.0, 33.0, 40.0, 41.0,
42.0, 43.0, 50.0, 51.0, 52.0, 53.0, 60.0, 61.0, 62.0, 63.0,
];
let pos_x = vec![0u32, 1, 0, 2];
let pos_y = vec![0u32, 0, 1, 2];
let out = gemma4v_position_embed_lookup(&pos_x, &pos_y, &pe, 3, 4).unwrap();
assert_eq!(out.len(), 16);
assert_eq!(&out[0..4], &[50.0, 52.0, 54.0, 56.0]);
assert_eq!(&out[4..8], &[60.0, 62.0, 64.0, 66.0]);
assert_eq!(&out[8..12], &[60.0, 62.0, 64.0, 66.0]);
assert_eq!(&out[12..16], &[90.0, 92.0, 94.0, 96.0]);
}
#[test]
fn gemma4v_position_embed_lookup_clamps_out_of_range() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let pe: Vec<f32> = vec![1.0, 2.0, 3.0, 4.0, 10.0, 20.0, 30.0, 40.0];
let pos_x = vec![99u32];
let pos_y = vec![99u32];
let out = gemma4v_position_embed_lookup(&pos_x, &pos_y, &pe, 2, 2).unwrap();
assert_eq!(out, vec![33.0, 44.0]);
}
#[test]
fn gemma4v_position_embed_lookup_shape_errors() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let pe = vec![0f32; 12];
let err = gemma4v_position_embed_lookup(&[0u32, 1], &[0u32], &pe, 3, 2).unwrap_err();
assert!(format!("{err}").contains("pos_x.len()"));
let err2 = gemma4v_position_embed_lookup(&[0u32], &[0u32], &pe, 7, 2).unwrap_err();
assert!(format!("{err2}").contains("pe_table.len()"));
}
#[test]
fn gemma4v_position_embed_add_in_place() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut patch_embeds = vec![1.0_f32; 8]; let pe: Vec<f32> = vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0, ];
let pos_x = vec![0u32, 1];
let pos_y = vec![1u32, 0];
gemma4v_position_embed_add(&mut patch_embeds, &pos_x, &pos_y, &pe, 2, 4).unwrap();
assert_eq!(&patch_embeds[0..4], &[15.0, 17.0, 19.0, 21.0]);
assert_eq!(&patch_embeds[4..8], &[15.0, 17.0, 19.0, 21.0]);
}
#[test]
fn gemma4v_clippable_linear_forward_no_bounds_matches_plain_linear() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 3usize;
let in_features = 4usize;
let out_features = 2usize;
let input: Vec<f32> = vec![
-1.0, 2.0, 3.0, 4.0, 5.0, -6.0, 7.0, 8.0, -9.0, 10.0, -11.0, 12.0,
];
let weight: Vec<f32> = vec![0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8];
let bounds = Gemma4ClippableLinearBounds::default();
assert!(!bounds.any());
let plain =
linear_forward(&input, &weight, None, batch, in_features, out_features).unwrap();
let clipped = gemma4v_clippable_linear_forward(
&input,
&weight,
&bounds,
batch,
in_features,
out_features,
)
.unwrap();
assert_eq!(
plain, clipped,
"no-bounds clippable_linear must match plain linear"
);
}
#[test]
fn gemma4v_clippable_linear_forward_input_clamp_only() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 1usize;
let in_features = 4usize;
let out_features = 1usize;
let input: Vec<f32> = vec![-100.0, 1.0, 1.0, 1.0];
let weight: Vec<f32> = vec![1.0, 1.0, 1.0, 1.0];
let bounds = Gemma4ClippableLinearBounds {
input_min: Some(-1.0),
input_max: Some(1.0),
output_min: None,
output_max: None,
};
let out = gemma4v_clippable_linear_forward(
&input,
&weight,
&bounds,
batch,
in_features,
out_features,
)
.unwrap();
assert_eq!(out, vec![2.0]);
}
#[test]
fn gemma4v_clippable_linear_forward_both_clamps() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let batch = 1usize;
let in_features = 4usize;
let out_features = 1usize;
let input: Vec<f32> = vec![100.0, 100.0, 100.0, 100.0];
let weight: Vec<f32> = vec![1.0, 1.0, 1.0, 1.0];
let bounds = Gemma4ClippableLinearBounds {
input_min: Some(-2.0),
input_max: Some(2.0),
output_min: Some(-3.0),
output_max: Some(3.0),
};
let out = gemma4v_clippable_linear_forward(
&input,
&weight,
&bounds,
batch,
in_features,
out_features,
)
.unwrap();
assert_eq!(out, vec![3.0]);
}
#[test]
fn gemma4v_clippable_linear_bounds_resolve_default_to_neg_pos_inf() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let bounds = Gemma4ClippableLinearBounds::default();
let (mn_in, mx_in) = bounds.resolved_input();
let (mn_out, mx_out) = bounds.resolved_output();
assert_eq!(mn_in, f32::NEG_INFINITY);
assert_eq!(mx_in, f32::INFINITY);
assert_eq!(mn_out, f32::NEG_INFINITY);
assert_eq!(mx_out, f32::INFINITY);
assert!(!bounds.any());
}
#[test]
fn gemma4v_clippable_linear_forward_rejects_min_gt_max() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let input: Vec<f32> = vec![0.0; 4];
let weight: Vec<f32> = vec![0.0; 4];
let bounds = Gemma4ClippableLinearBounds {
input_min: Some(5.0),
input_max: Some(1.0),
..Default::default()
};
let err = gemma4v_clippable_linear_forward(&input, &weight, &bounds, 1, 4, 1).unwrap_err();
assert!(format!("{err}").contains("input_min"));
}
}