use super::VisionError;
use super::checkpoint::VisualMergerWeights;
use super::vit::{batch_matvec, 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 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 {
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 mut fc1_out = batch_matvec(
&weights.fc1_weight,
&normed,
num_visual_tokens,
merge_in,
merge_in,
);
for i in 0..num_visual_tokens {
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 mut out = batch_matvec(
&weights.fc2_weight,
&fc1_out,
num_visual_tokens,
out_hidden,
merge_in,
);
for i in 0..num_visual_tokens {
for j in 0..out_hidden {
out[i * out_hidden + j] += weights.fc2_bias[j];
}
}
Ok(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_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"
);
}
}