use super::VisionError;
#[derive(Debug, Clone)]
pub struct MlpMergerWeights {
pub w1: Vec<f32>,
pub b1: Vec<f32>,
pub w2: Vec<f32>,
pub b2: Vec<f32>,
}
pub struct MlpMerger {
weights: MlpMergerWeights,
d_vit: usize,
d_hidden: usize,
d_model: usize,
merge_size: usize,
}
impl MlpMerger {
pub fn new(
weights: MlpMergerWeights,
d_vit: usize,
d_hidden: usize,
d_model: usize,
merge_size: usize,
) -> Result<Self, VisionError> {
if merge_size == 0 {
return Err(VisionError::InvalidConfig("merge_size must be > 0".into()));
}
let d_in = d_vit * merge_size * merge_size;
let expected_w1 = d_hidden * d_in;
let expected_b1 = d_hidden;
let expected_w2 = d_model * d_hidden;
let expected_b2 = d_model;
if weights.w1.len() != expected_w1 {
return Err(VisionError::ShapeMismatch {
expected: expected_w1,
actual: weights.w1.len(),
context: "MlpMerger w1".into(),
});
}
if weights.b1.len() != expected_b1 {
return Err(VisionError::ShapeMismatch {
expected: expected_b1,
actual: weights.b1.len(),
context: "MlpMerger b1".into(),
});
}
if weights.w2.len() != expected_w2 {
return Err(VisionError::ShapeMismatch {
expected: expected_w2,
actual: weights.w2.len(),
context: "MlpMerger w2".into(),
});
}
if weights.b2.len() != expected_b2 {
return Err(VisionError::ShapeMismatch {
expected: expected_b2,
actual: weights.b2.len(),
context: "MlpMerger b2".into(),
});
}
Ok(Self {
weights,
d_vit,
d_hidden,
d_model,
merge_size,
})
}
pub fn merge_and_project(
&self,
vit_out: &[f32],
raw_patches: usize,
) -> Result<Vec<f32>, VisionError> {
let merge_sq = self.merge_size * self.merge_size;
if raw_patches % merge_sq != 0 {
return Err(VisionError::ShapeMismatch {
expected: 0, actual: raw_patches % merge_sq,
context: format!(
"raw_patches {raw_patches} must be divisible by merge_size^2={merge_sq}"
),
});
}
let expected_len = raw_patches * self.d_vit;
if vit_out.len() != expected_len {
return Err(VisionError::ShapeMismatch {
expected: expected_len,
actual: vit_out.len(),
context: "vit_out length".into(),
});
}
let merged_patches = raw_patches / merge_sq;
let patches_per_side = (raw_patches as f64).sqrt() as usize;
if patches_per_side * patches_per_side != raw_patches {
return Err(VisionError::InvalidConfig(format!(
"raw_patches {raw_patches} must be a perfect square for spatial merge"
)));
}
let merged_per_side = patches_per_side / self.merge_size;
let d_in = self.d_vit * merge_sq;
let mut output = vec![0.0f32; merged_patches * self.d_model];
for gy in 0..merged_per_side {
for gx in 0..merged_per_side {
let group_idx = gy * merged_per_side + gx;
let mut concat = vec![0.0f32; d_in];
let mut concat_pos = 0usize;
for dy in 0..self.merge_size {
for dx in 0..self.merge_size {
let py = gy * self.merge_size + dy;
let px = gx * self.merge_size + dx;
let patch_idx = py * patches_per_side + px;
let src = &vit_out[patch_idx * self.d_vit..(patch_idx + 1) * self.d_vit];
concat[concat_pos..concat_pos + self.d_vit].copy_from_slice(src);
concat_pos += self.d_vit;
}
}
let mut h1 = vec![0.0f32; self.d_hidden];
for r in 0..self.d_hidden {
let row = &self.weights.w1[r * d_in..(r + 1) * d_in];
let mut acc = self.weights.b1[r];
for (a, c) in row.iter().zip(concat.iter()) {
acc += a * c;
}
h1[r] = gelu(acc);
}
let out_slice =
&mut output[group_idx * self.d_model..(group_idx + 1) * self.d_model];
for r in 0..self.d_model {
let row = &self.weights.w2[r * self.d_hidden..(r + 1) * self.d_hidden];
let mut acc = self.weights.b2[r];
for (a, h) in row.iter().zip(h1.iter()) {
acc += a * h;
}
out_slice[r] = acc;
}
}
}
Ok(output)
}
}
#[inline(always)]
fn gelu(x: f32) -> f32 {
let c = (2.0_f32 / std::f32::consts::PI).sqrt();
0.5 * x * (1.0 + (c * (x + 0.044715 * x * x * x)).tanh())
}
#[cfg(test)]
mod tests {
use super::*;
fn test_merger(d_vit: usize, d_hidden: usize, d_model: usize, merge: usize) -> MlpMerger {
let d_in = d_vit * merge * merge;
let w1 = vec![0.0f32; d_hidden * d_in];
let b1 = vec![0.0f32; d_hidden];
let w2 = vec![0.0f32; d_model * d_hidden];
let b2 = vec![0.0f32; d_model];
let weights = MlpMergerWeights { w1, b1, w2, b2 };
MlpMerger::new(weights, d_vit, d_hidden, d_model, merge).expect("test merger")
}
#[test]
fn merger_output_shape_4_patches_merge2() {
let d_vit = 8usize;
let d_hidden = 16usize;
let d_model = 4usize;
let merger = test_merger(d_vit, d_hidden, d_model, 2);
let vit_out = vec![1.0f32; 4 * d_vit];
let out = merger.merge_and_project(&vit_out, 4).expect("merge");
assert_eq!(out.len(), 1 * d_model);
}
#[test]
fn merger_output_shape_784_patches_merge2() {
let d_vit = 4usize;
let d_hidden = 8usize;
let d_model = 6usize;
let merger = test_merger(d_vit, d_hidden, d_model, 2);
let vit_out = vec![0.5f32; 784 * d_vit];
let out = merger.merge_and_project(&vit_out, 784).expect("merge");
assert_eq!(out.len(), 196 * d_model);
}
#[test]
fn merger_rejects_non_square_patch_count() {
let merger = test_merger(4, 8, 6, 2);
let vit_out = vec![0.0f32; 6 * 4];
let result = merger.merge_and_project(&vit_out, 6);
assert!(result.is_err());
}
#[test]
fn merger_rejects_non_divisible_patches() {
let merger = test_merger(4, 8, 6, 2);
let vit_out = vec![0.0f32; 9 * 4];
let result = merger.merge_and_project(&vit_out, 9);
assert!(result.is_err());
}
#[test]
fn merger_rejects_wrong_vit_out_len() {
let merger = test_merger(4, 8, 6, 2);
let vit_out = vec![0.0f32; 3 * 4]; let result = merger.merge_and_project(&vit_out, 4);
assert!(result.is_err());
}
#[test]
fn merger_output_is_finite() {
let d_vit = 8usize;
let d_hidden = 16usize;
let d_model = 6usize;
let merger = test_merger(d_vit, d_hidden, d_model, 2);
let vit_out: Vec<f32> = (0..4 * d_vit).map(|i| (i as f32) * 0.01).collect();
let out = merger.merge_and_project(&vit_out, 4).expect("merge");
for &v in &out {
assert!(v.is_finite(), "merger output non-finite: {v}");
}
}
#[test]
fn merger_construction_wrong_w1_size() {
let d_vit = 4usize;
let d_hidden = 8usize;
let d_model = 6usize;
let merge = 2usize;
let d_in = d_vit * merge * merge;
let weights = MlpMergerWeights {
w1: vec![0.0f32; d_hidden * d_in + 1], b1: vec![0.0f32; d_hidden],
w2: vec![0.0f32; d_model * d_hidden],
b2: vec![0.0f32; d_model],
};
let result = MlpMerger::new(weights, d_vit, d_hidden, d_model, merge);
assert!(result.is_err());
}
}