use super::{VisionError, config::VisionConfig};
#[derive(Debug, Clone)]
pub struct AttentionWeights {
pub q_proj: Vec<f32>,
pub k_proj: Vec<f32>,
pub v_proj: Vec<f32>,
pub o_proj: Vec<f32>,
pub q_norm_weight: Vec<f32>,
pub k_norm_weight: Vec<f32>,
}
#[derive(Debug, Clone)]
pub struct MlpWeights {
pub gate_proj: Vec<f32>,
pub up_proj: Vec<f32>,
pub down_proj: Vec<f32>,
}
#[derive(Debug, Clone)]
pub struct ViTBlockWeights {
pub ln1_weight: Vec<f32>, pub ln1_bias: Vec<f32>,
pub attn: AttentionWeights,
pub ln2_weight: Vec<f32>, pub ln2_bias: Vec<f32>,
pub mlp: MlpWeights,
}
#[derive(Debug, Clone)]
pub struct VisionWeights {
pub patch_embed_weight: Vec<f32>,
pub patch_embed_bias: Vec<f32>,
pub norm_weight: Vec<f32>,
pub norm_bias: Vec<f32>,
pub blocks: Vec<ViTBlockWeights>,
}
pub struct ViT {
pub(crate) config: VisionConfig,
pub(crate) weights: VisionWeights,
}
pub(crate) fn matvec(a: &[f32], x: &[f32], rows: usize, cols: usize) -> Vec<f32> {
crate::forward::cpu::validate_gemm_nn(a.len(), x.len(), rows, rows, cols, 1, "vit_matvec");
let mut y = vec![0.0f32; rows];
for r in 0..rows {
let row = &a[r * cols..(r + 1) * cols];
let mut acc = 0.0f32;
for (a_val, x_val) in row.iter().zip(x.iter()) {
acc += a_val * x_val;
}
y[r] = acc;
}
y
}
pub(crate) fn batch_matvec(a: &[f32], x: &[f32], n: usize, rows: usize, cols: usize) -> Vec<f32> {
assert!(
n.checked_mul(rows).is_some(),
"vit_batch_matvec: output shape overflow: n*rows"
);
let mut out = vec![0.0f32; n * rows];
crate::forward::cpu::validate_gemm_bt(
x.len(),
a.len(),
out.len(),
n,
cols,
rows,
"vit_batch_matvec",
);
for i in 0..n {
let xi = &x[i * cols..(i + 1) * cols];
let yi = &mut out[i * rows..(i + 1) * rows];
for r in 0..rows {
let row = &a[r * cols..(r + 1) * cols];
let mut acc = 0.0f32;
for (a_val, x_val) in row.iter().zip(xi.iter()) {
acc += a_val * x_val;
}
yi[r] = acc;
}
}
out
}
pub(crate) fn layer_norm(x: &mut [f32], weight: &[f32], bias: &[f32], eps: f32) {
let n = x.len();
assert_eq!(weight.len(), n);
assert_eq!(bias.len(), n);
let mean = x.iter().sum::<f32>() / n as f32;
let var = x.iter().map(|v| (v - mean).powi(2)).sum::<f32>() / n as f32;
let inv_std = 1.0 / (var + eps).sqrt();
for (i, v) in x.iter_mut().enumerate() {
*v = (*v - mean) * inv_std * weight[i] + bias[i];
}
}
#[inline(always)]
pub(crate) 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())
}
fn swiglu(gate: &[f32], up: &[f32]) -> Vec<f32> {
gate.iter()
.zip(up.iter())
.map(|(&g, &u)| gelu(g) * u)
.collect()
}
pub(crate) fn softmax_inplace(x: &mut [f32]) {
let (max, any_nan) = crate::attention::softmax_row::row_max_and_any_nan(x);
if crate::attention::softmax_row::row_fails_closed_pre_exp(max, any_nan) {
x.fill(0.0);
return;
}
let mut sum = 0.0f32;
for v in x.iter_mut() {
*v = (*v - max).exp();
sum += *v;
}
crate::attention::softmax_row::finalize_row(x, sum);
}
fn apply_2d_rope(
qkv: &mut [f32], n_patches: usize,
d_model: usize,
n_heads: usize,
patches_per_side: usize,
) {
let head_dim = d_model / n_heads;
let rope_half = head_dim / 4; let theta_base = 10_000.0_f32;
for patch_idx in 0..n_patches {
let row = patch_idx / patches_per_side;
let col = patch_idx % patches_per_side;
let pos_row = row as f32;
let pos_col = col as f32;
for qk in 0..2usize {
let base = patch_idx * 3 * d_model + qk * d_model;
for h in 0..n_heads {
let head_base = base + h * head_dim;
for i in 0..rope_half {
let freq = 1.0 / theta_base.powf((2 * i) as f32 / head_dim as f32);
let angle = pos_row * freq;
let (sin, cos) = angle.sin_cos();
let x0 = qkv[head_base + 2 * i];
let x1 = qkv[head_base + 2 * i + 1];
qkv[head_base + 2 * i] = x0 * cos - x1 * sin;
qkv[head_base + 2 * i + 1] = x0 * sin + x1 * cos;
}
let col_offset = rope_half * 2;
for i in 0..rope_half {
let freq = 1.0 / theta_base.powf((2 * i) as f32 / head_dim as f32);
let angle = pos_col * freq;
let (sin, cos) = angle.sin_cos();
let x0 = qkv[head_base + col_offset + 2 * i];
let x1 = qkv[head_base + col_offset + 2 * i + 1];
qkv[head_base + col_offset + 2 * i] = x0 * cos - x1 * sin;
qkv[head_base + col_offset + 2 * i + 1] = x0 * sin + x1 * cos;
}
}
}
}
}
#[allow(clippy::ptr_arg)] fn vit_block_forward(
hidden: &mut Vec<f32>,
w: &ViTBlockWeights,
n_patches: usize,
d_model: usize,
d_mlp: usize,
n_heads: usize,
head_dim: usize,
_block_idx: usize,
global_attn_every: usize,
_window_size: usize,
patches_per_side: usize,
) {
let eps = 1e-6_f32;
let scale = 1.0_f32 / (head_dim as f32).sqrt();
let residual_attn = hidden.clone();
let mut normed = hidden.clone();
for i in 0..n_patches {
let slice = &mut normed[i * d_model..(i + 1) * d_model];
layer_norm(slice, &w.ln1_weight, &w.ln1_bias, eps);
}
let mut qkv = vec![0.0f32; n_patches * 3 * d_model];
for i in 0..n_patches {
let xi = &normed[i * d_model..(i + 1) * d_model];
let q = matvec(&w.attn.q_proj, xi, d_model, d_model);
let k = matvec(&w.attn.k_proj, xi, d_model, d_model);
let v = matvec(&w.attn.v_proj, xi, d_model, d_model);
let base = i * 3 * d_model;
qkv[base..base + d_model].copy_from_slice(&q);
qkv[base + d_model..base + 2 * d_model].copy_from_slice(&k);
qkv[base + 2 * d_model..base + 3 * d_model].copy_from_slice(&v);
}
for i in 0..n_patches {
let q_slice = &mut qkv[i * 3 * d_model..i * 3 * d_model + d_model];
apply_qk_norm(q_slice, &w.attn.q_norm_weight, eps);
let k_slice = &mut qkv[i * 3 * d_model + d_model..i * 3 * d_model + 2 * d_model];
apply_qk_norm(k_slice, &w.attn.k_norm_weight, eps);
}
apply_2d_rope(&mut qkv, n_patches, d_model, n_heads, patches_per_side);
let _ = global_attn_every; let attn_out = multihead_attention(&qkv, n_patches, d_model, n_heads, head_dim, scale);
let o_out = batch_matvec(&w.attn.o_proj, &attn_out, n_patches, d_model, d_model);
for i in 0..n_patches * d_model {
hidden[i] = residual_attn[i] + o_out[i];
}
let residual_mlp = hidden.clone();
let mut normed_mlp = hidden.clone();
for i in 0..n_patches {
let slice = &mut normed_mlp[i * d_model..(i + 1) * d_model];
layer_norm(slice, &w.ln2_weight, &w.ln2_bias, eps);
}
for i in 0..n_patches {
let xi = &normed_mlp[i * d_model..(i + 1) * d_model];
let gate = matvec(&w.mlp.gate_proj, xi, d_mlp, d_model);
let up = matvec(&w.mlp.up_proj, xi, d_mlp, d_model);
let activated = swiglu(&gate, &up);
let down = matvec(&w.mlp.down_proj, &activated, d_model, d_mlp);
let base = i * d_model;
for j in 0..d_model {
hidden[base + j] = residual_mlp[base + j] + down[j];
}
}
}
fn apply_qk_norm(x: &mut [f32], weight: &[f32], eps: f32) {
let n = x.len();
let rms_sq = x.iter().map(|v| v * v).sum::<f32>() / n as f32;
let inv_rms = 1.0 / (rms_sq + eps).sqrt();
for (v, &w) in x.iter_mut().zip(weight.iter()) {
*v = *v * inv_rms * w;
}
}
fn multihead_attention(
qkv: &[f32],
n: usize,
d_model: usize,
n_heads: usize,
head_dim: usize,
scale: f32,
) -> Vec<f32> {
let mut out = vec![0.0f32; n * d_model];
for h in 0..n_heads {
let mut q_h = vec![0.0f32; n * head_dim];
let mut k_h = vec![0.0f32; n * head_dim];
let mut v_h = vec![0.0f32; n * head_dim];
for i in 0..n {
let base = i * 3 * d_model;
let q_src = &qkv[base + h * head_dim..base + (h + 1) * head_dim];
let k_src = &qkv[base + d_model + h * head_dim..base + d_model + (h + 1) * head_dim];
let v_src =
&qkv[base + 2 * d_model + h * head_dim..base + 2 * d_model + (h + 1) * head_dim];
q_h[i * head_dim..(i + 1) * head_dim].copy_from_slice(q_src);
k_h[i * head_dim..(i + 1) * head_dim].copy_from_slice(k_src);
v_h[i * head_dim..(i + 1) * head_dim].copy_from_slice(v_src);
}
let mut scores = vec![0.0f32; n * n];
for i in 0..n {
for j in 0..n {
let qi = &q_h[i * head_dim..(i + 1) * head_dim];
let kj = &k_h[j * head_dim..(j + 1) * head_dim];
let dot: f32 = qi.iter().zip(kj.iter()).map(|(a, b)| a * b).sum();
scores[i * n + j] = dot * scale;
}
}
for i in 0..n {
let row = &mut scores[i * n..(i + 1) * n];
softmax_inplace(row);
}
for i in 0..n {
let attn_row = &scores[i * n..(i + 1) * n];
for j in 0..head_dim {
let mut acc = 0.0f32;
for k in 0..n {
acc += attn_row[k] * v_h[k * head_dim + j];
}
out[i * d_model + h * head_dim + j] += acc;
}
}
}
out
}
impl ViT {
pub fn new(weights: VisionWeights, config: VisionConfig) -> Result<Self, VisionError> {
config.validate()?;
if weights.blocks.len() != config.n_layers {
return Err(VisionError::ShapeMismatch {
expected: config.n_layers,
actual: weights.blocks.len(),
context: "number of ViT block weight sets must equal n_layers".into(),
});
}
if weights.patch_embed_weight.len()
!= config.d_model * (config.patch_size as usize).pow(2) * 3
{
let expected = config.d_model * (config.patch_size as usize).pow(2) * 3;
return Err(VisionError::ShapeMismatch {
expected,
actual: weights.patch_embed_weight.len(),
context: "patch_embed_weight size".into(),
});
}
Ok(Self { config, weights })
}
pub fn forward(&self, img: &super::preprocess::ImageTensor) -> Result<Vec<f32>, VisionError> {
if img.n_patches != self.config.n_patches {
return Err(VisionError::ShapeMismatch {
expected: self.config.n_patches,
actual: img.n_patches,
context: "ImageTensor n_patches mismatch".into(),
});
}
let cfg = &self.config;
let n = cfg.n_patches;
let d = cfg.d_model;
let patch_len = (cfg.patch_size as usize).pow(2) * 3;
let patches_per_side = cfg.image_size as usize / cfg.patch_size as usize;
let mut hidden = batch_matvec(
&self.weights.patch_embed_weight,
&img.patches,
n,
d,
patch_len,
);
for i in 0..n {
for j in 0..d {
hidden[i * d + j] += self.weights.patch_embed_bias[j];
}
}
for (block_idx, block_w) in self.weights.blocks.iter().enumerate() {
vit_block_forward(
&mut hidden,
block_w,
n,
d,
cfg.d_mlp,
cfg.n_heads,
cfg.head_dim(),
block_idx,
cfg.global_attn_every,
cfg.window_size,
patches_per_side,
);
}
for i in 0..n {
let slice = &mut hidden[i * d..(i + 1) * d];
layer_norm(
slice,
&self.weights.norm_weight,
&self.weights.norm_bias,
1e-6,
);
}
Ok(hidden)
}
}
#[cfg(test)]
pub(crate) fn make_test_vit(cfg: &VisionConfig) -> ViT {
let d = cfg.d_model;
let patch_len = (cfg.patch_size as usize).pow(2) * 3;
let d_mlp = cfg.d_mlp;
let make_block = |_i: usize| ViTBlockWeights {
ln1_weight: vec![1.0f32; d],
ln1_bias: vec![0.0f32; d],
attn: AttentionWeights {
q_proj: identity_weight(d),
k_proj: identity_weight(d),
v_proj: identity_weight(d),
o_proj: identity_weight(d),
q_norm_weight: vec![1.0f32; d],
k_norm_weight: vec![1.0f32; d],
},
ln2_weight: vec![1.0f32; d],
ln2_bias: vec![0.0f32; d],
mlp: MlpWeights {
gate_proj: vec![0.0f32; d_mlp * d], up_proj: vec![0.0f32; d_mlp * d],
down_proj: vec![0.0f32; d * d_mlp],
},
};
let blocks = (0..cfg.n_layers).map(make_block).collect();
let weights = VisionWeights {
patch_embed_weight: identity_weight_mn(d, patch_len),
patch_embed_bias: vec![0.0f32; d],
norm_weight: vec![1.0f32; d],
norm_bias: vec![0.0f32; d],
blocks,
};
ViT::new(weights, cfg.clone()).expect("test ViT construction")
}
#[cfg(test)]
fn identity_weight(n: usize) -> Vec<f32> {
identity_weight_mn(n, n)
}
#[cfg(test)]
fn identity_weight_mn(rows: usize, cols: usize) -> Vec<f32> {
let min_dim = rows.min(cols);
let mut w = vec![0.0f32; rows * cols];
for i in 0..min_dim {
w[i * cols + i] = 1.0;
}
w
}
#[cfg(test)]
mod tests {
use super::*;
use crate::vision::config::VisionConfig;
fn tiny_cfg() -> VisionConfig {
let image_size = 8u32;
let patch_size = 4u32;
let n_patches = ((image_size / patch_size) as usize).pow(2); let d_model = 8usize;
let mlp_ratio = 2usize;
VisionConfig {
image_size,
patch_size,
n_patches,
d_model,
n_heads: 2,
n_layers: 1,
spatial_merge_size: 2,
global_attn_every: 1,
window_size: 2,
mlp_ratio,
use_gelu: true,
d_decoder: 16,
d_mlp: d_model * mlp_ratio,
}
}
#[test]
fn vit_forward_output_shape() {
let cfg = tiny_cfg();
let vit = make_test_vit(&cfg);
let n_patches = cfg.n_patches;
let patch_len = (cfg.patch_size as usize).pow(2) * 3;
let img = crate::vision::preprocess::ImageTensor {
patches: vec![0.1f32; n_patches * patch_len],
n_patches,
patch_hw: cfg.patch_size as usize,
};
let out = vit.forward(&img).expect("ViT forward");
assert_eq!(out.len(), n_patches * cfg.d_model);
}
#[test]
fn vit_forward_values_are_finite() {
let cfg = tiny_cfg();
let vit = make_test_vit(&cfg);
let n_patches = cfg.n_patches;
let patch_len = (cfg.patch_size as usize).pow(2) * 3;
let img = crate::vision::preprocess::ImageTensor {
patches: vec![0.5f32; n_patches * patch_len],
n_patches,
patch_hw: cfg.patch_size as usize,
};
let out = vit.forward(&img).expect("ViT forward");
for &v in &out {
assert!(v.is_finite(), "ViT output contained non-finite: {v}");
}
}
#[test]
fn softmax_inplace_nan_fails_closed() {
let mut x = vec![1.0f32, f32::NAN, 2.0, 0.5];
softmax_inplace(&mut x);
assert!(x.iter().all(|&v| v == 0.0), "expected all-zero row: {x:?}");
}
#[test]
fn softmax_inplace_pos_inf_fails_closed() {
let mut x = vec![1.0f32, f32::INFINITY, 2.0];
softmax_inplace(&mut x);
assert!(x.iter().all(|&v| v == 0.0), "expected all-zero row: {x:?}");
}
#[test]
fn softmax_inplace_all_neg_inf_fails_closed() {
let mut x = vec![f32::NEG_INFINITY; 4];
softmax_inplace(&mut x);
assert!(x.iter().all(|&v| v == 0.0), "expected all-zero row: {x:?}");
}
#[test]
fn softmax_inplace_normal_row_still_normalizes() {
let mut x = vec![1.0f32, 2.0, 3.0];
softmax_inplace(&mut x);
let sum: f32 = x.iter().sum();
assert!((sum - 1.0).abs() < 1e-6, "row must sum to ~1: {x:?}");
assert!(x.iter().all(|v| v.is_finite() && *v >= 0.0));
}
#[test]
fn vit_construction_wrong_block_count() {
let cfg = tiny_cfg(); let d = cfg.d_model;
let patch_len = (cfg.patch_size as usize).pow(2) * 3;
let d_mlp = cfg.d_mlp;
let dummy_block = ViTBlockWeights {
ln1_weight: vec![1.0f32; d],
ln1_bias: vec![0.0f32; d],
attn: AttentionWeights {
q_proj: vec![0.0f32; d * d],
k_proj: vec![0.0f32; d * d],
v_proj: vec![0.0f32; d * d],
o_proj: vec![0.0f32; d * d],
q_norm_weight: vec![1.0f32; d],
k_norm_weight: vec![1.0f32; d],
},
ln2_weight: vec![1.0f32; d],
ln2_bias: vec![0.0f32; d],
mlp: MlpWeights {
gate_proj: vec![0.0f32; d_mlp * d],
up_proj: vec![0.0f32; d_mlp * d],
down_proj: vec![0.0f32; d * d_mlp],
},
};
let weights = VisionWeights {
patch_embed_weight: vec![0.0f32; d * patch_len],
patch_embed_bias: vec![0.0f32; d],
norm_weight: vec![1.0f32; d],
norm_bias: vec![0.0f32; d],
blocks: vec![dummy_block.clone(), dummy_block],
};
let err = ViT::new(weights, cfg);
assert!(err.is_err());
}
#[test]
fn vit_forward_shape_mismatch_rejected() {
let cfg = tiny_cfg();
let vit = make_test_vit(&cfg);
let patch_len = (cfg.patch_size as usize).pow(2) * 3;
let img = crate::vision::preprocess::ImageTensor {
patches: vec![0.0f32; 9 * patch_len],
n_patches: 9,
patch_hw: cfg.patch_size as usize,
};
assert!(vit.forward(&img).is_err());
}
#[test]
#[should_panic(expected = "a too short for m*k")]
fn vit_matvec_rejects_short_a() {
let a = [0.0f32; 3]; let x = [0.0f32; 2];
matvec(&a, &x, 2, 2);
}
#[test]
#[should_panic(expected = "b too short for k*n")]
fn vit_matvec_rejects_short_x() {
let a = [0.0f32; 4];
let x = [0.0f32; 1]; matvec(&a, &x, 2, 2);
}
#[test]
fn vit_matvec_accepts_oversized_x() {
let a = [1.0f32, 1.0, 1.0, 1.0];
let x = [1.0f32, 1.0, 1.0, 99.0]; let y = matvec(&a, &x, 2, 2);
assert_eq!(y, vec![2.0, 2.0]);
}
#[test]
#[should_panic(expected = "a too short for m*k")]
fn vit_batch_matvec_rejects_short_x() {
let a = [0.0f32; 4]; let x = [0.0f32; 3]; batch_matvec(&a, &x, 2, 2, 2);
}
#[test]
#[should_panic(expected = "b too short for n*k")]
fn vit_batch_matvec_rejects_short_a() {
let a = [0.0f32; 3]; let x = [0.0f32; 4];
batch_matvec(&a, &x, 2, 2, 2);
}
#[test]
#[should_panic(expected = "vit_batch_matvec: output shape overflow: n*rows")]
fn vit_batch_matvec_rejects_output_overflow() {
let a = [0.0f32; 4];
let x = [0.0f32; 4];
batch_matvec(&a, &x, usize::MAX, 2, 2);
}
}