use super::VisionError;
use super::checkpoint::VisualMergerWeights;
use super::vit::layer_norm;
use crate::model::qwen35_config::VisionModelConfig;
fn erf(x: f64) -> f64 {
let sign = if x < 0.0 { -1.0 } else { 1.0 };
let x = x.abs();
const A1: f64 = 0.254829592;
const A2: f64 = -0.284496736;
const A3: f64 = 1.421413741;
const A4: f64 = -1.453152027;
const A5: f64 = 1.061405429;
const P: f64 = 0.3275911;
let t = 1.0 / (1.0 + P * x);
let y = 1.0 - (((((A5 * t + A4) * t) + A3) * t + A2) * t + A1) * t * (-x * x).exp();
sign * y
}
#[inline(always)]
fn gelu_exact(x: f32) -> f32 {
let xd = x as f64;
(0.5 * xd * (1.0 + erf(xd / std::f64::consts::SQRT_2))) as f32
}
pub fn qwen35_merger_forward(
weights: &VisualMergerWeights,
cfg: &VisionModelConfig,
pre_merger_hidden: &[f32],
) -> Result<Vec<f32>, VisionError> {
let mut never_cancel = || false;
match qwen35_merger_forward_with_cancel(weights, cfg, pre_merger_hidden, &mut never_cancel)? {
Some(output) => Ok(output),
None => Err(VisionError::InvalidConfig(
"non-cancellable vision merger was cancelled".to_string(),
)),
}
}
fn batch_matvec_with_cancel(
weights: &[f32],
input: &[f32],
batch: usize,
rows: usize,
cols: usize,
should_cancel: &mut dyn FnMut() -> bool,
) -> Result<Option<Vec<f32>>, VisionError> {
let output_len = batch.checked_mul(rows).ok_or_else(|| {
VisionError::InvalidConfig(
"qwen35_merger_forward: matrix output shape overflow".to_string(),
)
})?;
let mut output = vec![0.0f32; output_len];
crate::forward::cpu::validate_gemm_bt(
input.len(),
weights.len(),
output.len(),
batch,
cols,
rows,
"qwen35_merger_forward",
);
for batch_index in 0..batch {
if should_cancel() {
return Ok(None);
}
let input_row = &input[batch_index * cols..(batch_index + 1) * cols];
let output_row = &mut output[batch_index * rows..(batch_index + 1) * rows];
for row in 0..rows {
if row.is_multiple_of(16) && should_cancel() {
return Ok(None);
}
let weight_row = &weights[row * cols..(row + 1) * cols];
let mut accumulator = 0.0f32;
for (weight, value) in weight_row.iter().zip(input_row) {
accumulator += weight * value;
}
output_row[row] = accumulator;
}
}
Ok(Some(output))
}
pub(crate) fn qwen35_merger_forward_with_cancel(
weights: &VisualMergerWeights,
cfg: &VisionModelConfig,
pre_merger_hidden: &[f32],
should_cancel: &mut dyn FnMut() -> bool,
) -> Result<Option<Vec<f32>>, VisionError> {
if should_cancel() {
return Ok(None);
}
let hidden = cfg.hidden_size;
if hidden == 0 || !pre_merger_hidden.len().is_multiple_of(hidden) {
return Err(VisionError::ShapeMismatch {
expected: 0,
actual: pre_merger_hidden.len(),
context: "qwen35_merger_forward: pre_merger_hidden length must be a multiple of \
hidden_size"
.into(),
});
}
let n = pre_merger_hidden.len() / hidden;
let merge_sq = cfg.spatial_merge_size * cfg.spatial_merge_size;
if merge_sq == 0 || !n.is_multiple_of(merge_sq) {
return Err(VisionError::ShapeMismatch {
expected: 0,
actual: n,
context: "qwen35_merger_forward: patch count must be a multiple of \
spatial_merge_size^2"
.into(),
});
}
let mut normed = pre_merger_hidden.to_vec();
for i in 0..n {
if should_cancel() {
return Ok(None);
}
layer_norm(
&mut normed[i * hidden..(i + 1) * hidden],
&weights.norm_weight,
&weights.norm_bias,
1e-6,
);
}
let merge_in = merge_sq * hidden;
let num_visual_tokens = n / merge_sq;
let Some(mut fc1_out) = batch_matvec_with_cancel(
&weights.fc1_weight,
&normed,
num_visual_tokens,
merge_in,
merge_in,
should_cancel,
)?
else {
return Ok(None);
};
for i in 0..num_visual_tokens {
if should_cancel() {
return Ok(None);
}
for j in 0..merge_in {
let idx = i * merge_in + j;
fc1_out[idx] = gelu_exact(fc1_out[idx] + weights.fc1_bias[j]);
}
}
let out_hidden = cfg.out_hidden_size;
let Some(mut out) = batch_matvec_with_cancel(
&weights.fc2_weight,
&fc1_out,
num_visual_tokens,
out_hidden,
merge_in,
should_cancel,
)?
else {
return Ok(None);
};
for i in 0..num_visual_tokens {
if should_cancel() {
return Ok(None);
}
for j in 0..out_hidden {
out[i * out_hidden + j] += weights.fc2_bias[j];
}
}
if should_cancel() {
return Ok(None);
}
Ok(Some(out))
}
#[cfg(test)]
mod tests {
use super::*;
fn tiny_cfg() -> VisionModelConfig {
VisionModelConfig {
depth: 1,
hidden_size: 8,
num_heads: 2,
patch_size: 2,
spatial_merge_size: 2,
out_hidden_size: 6,
temporal_patch_size: 1,
num_position_embeddings: 16,
in_channels: 3,
deepstack_visual_indexes: vec![],
intermediate_size: None,
}
}
fn make_test_merger_weights(cfg: &VisionModelConfig) -> VisualMergerWeights {
let hidden = cfg.hidden_size;
let merge_in = cfg.spatial_merge_size * cfg.spatial_merge_size * hidden;
let out_hidden = cfg.out_hidden_size;
let mut state = 0x9e37_79b9_u32;
let mut next = move || {
state ^= state << 13;
state ^= state >> 17;
state ^= state << 5;
(state as f32 / u32::MAX as f32) * 0.2 - 0.1
};
let mut v = |n: usize| (0..n).map(|_| next()).collect::<Vec<f32>>();
VisualMergerWeights {
fc1_weight: v(merge_in * merge_in),
fc1_bias: v(merge_in),
fc2_weight: v(out_hidden * merge_in),
fc2_bias: v(out_hidden),
norm_weight: vec![1.0; hidden],
norm_bias: vec![0.0; hidden],
}
}
#[test]
fn erf_matches_known_values() {
assert!((erf(0.0) - 0.0).abs() < 1e-6);
assert!((erf(1.0) - 0.842_700_79).abs() < 1e-6);
assert!((erf(-1.0) + 0.842_700_79).abs() < 1e-6);
assert!((erf(2.0) - 0.995_322_3).abs() < 1e-6);
}
#[test]
fn gelu_exact_matches_known_values() {
assert!((gelu_exact(0.0)).abs() < 1e-6);
assert!((gelu_exact(1.0) - 0.841_344_7).abs() < 1e-4);
assert!((gelu_exact(-1.0) + 0.158_655_3).abs() < 1e-4);
}
#[test]
fn merger_forward_output_shape_and_finite() {
let cfg = tiny_cfg();
let weights = make_test_merger_weights(&cfg);
let n = 8; let pre_merger = vec![0.1f32; n * cfg.hidden_size];
let out = qwen35_merger_forward(&weights, &cfg, &pre_merger).expect("merger forward");
let merge_sq = cfg.spatial_merge_size * cfg.spatial_merge_size;
assert_eq!(out.len(), (n / merge_sq) * cfg.out_hidden_size);
assert!(out.iter().all(|v| v.is_finite()));
}
#[test]
fn merger_forward_rejects_hidden_size_mismatch() {
let cfg = tiny_cfg();
let weights = make_test_merger_weights(&cfg);
let bad = vec![0.0f32; 3]; let err = qwen35_merger_forward(&weights, &cfg, &bad).unwrap_err();
assert!(matches!(err, VisionError::ShapeMismatch { .. }));
}
#[test]
fn merger_forward_rejects_patch_count_not_multiple_of_merge_sq() {
let cfg = tiny_cfg();
let weights = make_test_merger_weights(&cfg);
let bad = vec![0.0f32; 3 * cfg.hidden_size];
let err = qwen35_merger_forward(&weights, &cfg, &bad).unwrap_err();
assert!(matches!(err, VisionError::ShapeMismatch { .. }));
}
#[test]
fn merger_forward_cancels_before_shape_validation() {
let cfg = tiny_cfg();
let weights = make_test_merger_weights(&cfg);
let result = qwen35_merger_forward_with_cancel(&weights, &cfg, &[0.0; 3], &mut || true)
.expect("cancellation is not a merger failure");
assert!(result.is_none());
}
#[test]
fn merger_forward_cancels_after_work_started() {
let cfg = tiny_cfg();
let weights = make_test_merger_weights(&cfg);
let pre_merger = vec![0.1f32; 8 * cfg.hidden_size];
let mut polls = 0;
let result = qwen35_merger_forward_with_cancel(&weights, &cfg, &pre_merger, &mut || {
polls += 1;
polls == 12
})
.expect("cancellation is not a merger failure");
assert!(result.is_none());
assert_eq!(polls, 12);
}
#[test]
fn merger_forward_is_deterministic() {
let cfg = tiny_cfg();
let weights = make_test_merger_weights(&cfg);
let n = 8;
let mut state = 0x1234_5678_u32;
let mut next = move || {
state ^= state << 13;
state ^= state >> 17;
state ^= state << 5;
(state as f32 / u32::MAX as f32) * 0.2 - 0.1
};
let pre_merger: Vec<f32> = (0..n * cfg.hidden_size).map(|_| next()).collect();
let out1 = qwen35_merger_forward(&weights, &cfg, &pre_merger).expect("forward 1");
let out2 = qwen35_merger_forward(&weights, &cfg, &pre_merger).expect("forward 2");
assert_eq!(out1, out2);
}
#[test]
fn merger_forward_is_sensitive_to_weight_mutation() {
let cfg = tiny_cfg();
let mut weights = make_test_merger_weights(&cfg);
let n = 8;
let pre_merger = vec![0.05f32; n * cfg.hidden_size];
let baseline = qwen35_merger_forward(&weights, &cfg, &pre_merger).expect("forward");
weights.fc1_weight[0] += 5.0;
let mutated = qwen35_merger_forward(&weights, &cfg, &pre_merger).expect("forward");
assert_ne!(
baseline, mutated,
"weight mutation had no effect on merger output"
);
}
}