use anyhow::{Context, Result};
use crate::gguf::GgufFile;
use std::sync::Arc;
use crate::model::weights::MmapWeight;
const KEY_HAS_VISION: &str = "clip.has_vision_encoder";
const KEY_N_LAYER: &str = "clip.vision.block_count";
const KEY_N_EMBD: &str = "clip.vision.embedding_length";
const KEY_N_FF: &str = "clip.vision.feed_forward_length";
const KEY_N_HEAD: &str = "clip.vision.attention.head_count";
const KEY_LN_EPS: &str = "clip.vision.attention.layer_norm_epsilon";
const KEY_IMAGE_SIZE: &str = "clip.vision.image_size";
const KEY_PATCH_SIZE: &str = "clip.vision.patch_size";
const KEY_PROJECTION_DIM: &str = "clip.vision.projection_dim";
const KEY_SCALE_FACTOR: &str = "clip.vision.projector.scale_factor";
const KEY_IMAGE_MEAN: &str = "clip.vision.image_mean";
const KEY_IMAGE_STD: &str = "clip.vision.image_std";
#[derive(Debug, Clone)]
pub struct VisionEncoderConfig {
pub n_layer: usize,
pub n_embd: usize,
pub n_ff: usize,
pub n_head: usize,
pub eps: f32,
pub image_size: usize,
pub patch_size: usize,
pub n_trained_patches: usize,
pub projection_dim: usize,
pub scale_factor: usize,
pub image_mean: [f32; 3],
pub image_std: [f32; 3],
pub image_min_pixels: usize,
pub image_max_pixels: usize,
}
impl VisionEncoderConfig {
pub fn from_gguf(gguf: &Arc<GgufFile>) -> Result<Self> {
let has_vision = gguf.get_bool(KEY_HAS_VISION).unwrap_or(false);
anyhow::ensure!(
has_vision,
"mmproj GGUF missing or false `{KEY_HAS_VISION}`; \
not a vision encoder"
);
let n_layer = gguf
.get_u32(KEY_N_LAYER)
.with_context(|| format!("missing `{KEY_N_LAYER}`"))? as usize;
let n_embd = gguf
.get_u32(KEY_N_EMBD)
.with_context(|| format!("missing `{KEY_N_EMBD}`"))? as usize;
let n_ff = gguf
.get_u32(KEY_N_FF)
.with_context(|| format!("missing `{KEY_N_FF}`"))? as usize;
let n_head = gguf
.get_u32(KEY_N_HEAD)
.with_context(|| format!("missing `{KEY_N_HEAD}`"))? as usize;
let eps = gguf
.get_f32(KEY_LN_EPS)
.with_context(|| format!("missing `{KEY_LN_EPS}`"))?;
let image_size =
gguf.get_u32(KEY_IMAGE_SIZE)
.with_context(|| format!("missing `{KEY_IMAGE_SIZE}`"))? as usize;
let patch_size =
gguf.get_u32(KEY_PATCH_SIZE)
.with_context(|| format!("missing `{KEY_PATCH_SIZE}`"))? as usize;
anyhow::ensure!(
patch_size > 0 && image_size % patch_size == 0,
"image_size ({image_size}) must be a positive multiple of patch_size ({patch_size})"
);
let n_trained_patches = (image_size / patch_size).pow(2);
let projection_dim =
gguf.get_u32(KEY_PROJECTION_DIM)
.with_context(|| format!("missing `{KEY_PROJECTION_DIM}`"))? as usize;
let scale_factor =
gguf.get_u32(KEY_SCALE_FACTOR)
.with_context(|| format!("missing `{KEY_SCALE_FACTOR}`"))? as usize;
anyhow::ensure!(
scale_factor > 0,
"scale_factor ({scale_factor}) must be > 0; a zero or missing value would \
collapse the pixel band and trip a divide-by-zero in `pixel_shuffle`",
);
let image_mean = read_rgb_array(gguf, KEY_IMAGE_MEAN)?;
let image_std = read_rgb_array(gguf, KEY_IMAGE_STD)?;
let n_tokens_min: usize = 64;
let n_tokens_max: usize = 256;
let pixels_per_token = (patch_size * scale_factor).pow(2);
let image_min_pixels = n_tokens_min * pixels_per_token;
let image_max_pixels = n_tokens_max * pixels_per_token;
Ok(Self {
n_layer,
n_embd,
n_ff,
n_head,
eps,
image_size,
patch_size,
n_trained_patches,
projection_dim,
scale_factor,
image_mean,
image_std,
image_min_pixels,
image_max_pixels,
})
}
}
pub struct PatchEmbedWeights {
pub conv_w: Vec<f32>,
pub conv_b: Vec<f32>,
}
pub struct VitBlockWeights {
pub ln1_w: Vec<f32>,
pub ln1_b: Vec<f32>,
pub q_w: MmapWeight,
pub q_b: Vec<f32>,
pub k_w: MmapWeight,
pub k_b: Vec<f32>,
pub v_w: MmapWeight,
pub v_b: Vec<f32>,
pub o_w: MmapWeight,
pub o_b: Vec<f32>,
pub ln2_w: Vec<f32>,
pub ln2_b: Vec<f32>,
pub ffn_up_w: MmapWeight,
pub ffn_up_b: Vec<f32>,
pub ffn_down_w: MmapWeight,
pub ffn_down_b: Vec<f32>,
}
pub struct ProjectorWeights {
pub mm1_w: MmapWeight,
pub mm1_b: Vec<f32>,
pub mm2_w: MmapWeight,
pub mm2_b: Vec<f32>,
}
pub struct VisionEncoderWeights {
pub config: VisionEncoderConfig,
pub patch_embed: PatchEmbedWeights,
pub position_embed: Vec<f32>,
pub blocks: Vec<VitBlockWeights>,
pub post_ln_w: Vec<f32>,
pub post_ln_b: Vec<f32>,
pub projector: ProjectorWeights,
}
impl VisionEncoderWeights {
pub fn from_gguf(gguf: &Arc<GgufFile>) -> Result<Self> {
let config = VisionEncoderConfig::from_gguf(gguf)?;
let patch_t = gguf
.get_tensor("v.patch_embd.weight")
.context("loading v.patch_embd.weight")?;
let patch_shape = patch_t.shape().to_vec();
anyhow::ensure!(
patch_shape.len() == 4
&& patch_shape[0] == config.patch_size
&& patch_shape[1] == config.patch_size
&& patch_shape[2] == 3
&& patch_shape[3] == config.n_embd,
"v.patch_embd.weight shape {patch_shape:?} != [patch_size={}, patch_size={}, 3, n_embd={}]",
config.patch_size,
config.patch_size,
config.n_embd,
);
let raw = patch_t.to_f32_vec();
let p = config.patch_size;
let in_dim = 3 * p * p;
let out_dim = config.n_embd;
let mut conv_w = vec![0f32; in_dim * out_dim];
for oc in 0..out_dim {
for c in 0..3 {
for kh in 0..p {
for kw in 0..p {
let src = kw + p * kh + p * p * c + p * p * 3 * oc;
let in_idx = c * p * p + kh * p + kw;
conv_w[in_idx * out_dim + oc] = raw[src];
}
}
}
}
let conv_b = load_vec_f32(gguf, "v.patch_embd.bias")?;
anyhow::ensure!(
conv_b.len() == config.n_embd,
"v.patch_embd.bias len ({}) != n_embd ({})",
conv_b.len(),
config.n_embd,
);
let patch_embed = PatchEmbedWeights { conv_w, conv_b };
let pos_t = gguf
.get_tensor("v.position_embd.weight")
.context("loading v.position_embd.weight")?;
let pos_shape = pos_t.shape();
anyhow::ensure!(
pos_shape.len() == 2
&& pos_shape[0] == config.n_embd
&& pos_shape[1] == config.n_trained_patches,
"v.position_embd.weight shape {pos_shape:?} != [n_embd={}, n_trained_patches={}]",
config.n_embd,
config.n_trained_patches,
);
let position_embed = pos_t.to_f32_vec();
let mut blocks = Vec::with_capacity(config.n_layer);
for il in 0..config.n_layer {
blocks.push(load_vit_block(gguf, il, &config)?);
}
let post_ln_w = load_vec_f32(gguf, "v.post_ln.weight")?;
let post_ln_b = load_vec_f32(gguf, "v.post_ln.bias")?;
anyhow::ensure!(
post_ln_w.len() == config.n_embd && post_ln_b.len() == config.n_embd,
"v.post_ln {{weight,bias}} len ({}, {}) != n_embd ({})",
post_ln_w.len(),
post_ln_b.len(),
config.n_embd,
);
let mm1_w = wt_f32(gguf, "mm.1.weight")?;
let mm1_b = load_vec_f32(gguf, "mm.1.bias")?;
let mm2_w = wt_f32(gguf, "mm.2.weight")?;
let mm2_b = load_vec_f32(gguf, "mm.2.bias")?;
let mm1_in_dim = config.n_embd * config.scale_factor.pow(2);
let intermediate_dim = mm1_w.rows;
anyhow::ensure!(
mm1_w.cols == mm1_in_dim,
"mm.1.weight cols ({}) != n_embd*sf² ({mm1_in_dim})",
mm1_w.cols,
);
anyhow::ensure!(
mm1_b.len() == intermediate_dim,
"mm.1.bias len ({}) != mm.1.weight.rows ({intermediate_dim})",
mm1_b.len(),
);
anyhow::ensure!(
mm2_w.cols == intermediate_dim,
"mm.2.weight cols ({}) != mm.1.weight.rows ({intermediate_dim}) — \
projector mm.1→mm.2 dimensions don't line up",
mm2_w.cols,
);
anyhow::ensure!(
mm2_w.rows == config.projection_dim,
"mm.2.weight rows ({}) != projection_dim ({})",
mm2_w.rows,
config.projection_dim,
);
anyhow::ensure!(
mm2_b.len() == config.projection_dim,
"mm.2.bias len ({}) != projection_dim ({})",
mm2_b.len(),
config.projection_dim,
);
let projector = ProjectorWeights {
mm1_w,
mm1_b,
mm2_w,
mm2_b,
};
Ok(Self {
config,
patch_embed,
position_embed,
blocks,
post_ln_w,
post_ln_b,
projector,
})
}
pub fn encode_image(&self, pixels: &[f32], grid_w: usize, grid_h: usize) -> Result<Vec<f32>> {
let cfg = &self.config;
anyhow::ensure!(grid_w > 0 && grid_h > 0, "grid dims must be > 0");
anyhow::ensure!(
cfg.scale_factor > 0,
"vision encoder config has scale_factor=0; refusing to divide by zero"
);
anyhow::ensure!(
grid_w % cfg.scale_factor == 0 && grid_h % cfg.scale_factor == 0,
"grid {grid_w}×{grid_h} not divisible by scale_factor ({})",
cfg.scale_factor,
);
let target_w = grid_w
.checked_mul(cfg.patch_size)
.ok_or_else(|| anyhow::anyhow!("grid_w·patch_size overflow"))?;
let target_h = grid_h
.checked_mul(cfg.patch_size)
.ok_or_else(|| anyhow::anyhow!("grid_h·patch_size overflow"))?;
let n_pix = target_w
.checked_mul(target_h)
.and_then(|x| x.checked_mul(3))
.ok_or_else(|| anyhow::anyhow!("3·target_w·target_h overflow"))?;
anyhow::ensure!(
pixels.len() == n_pix,
"encode_image: pixels.len() {} != 3·target_w·target_h ({n_pix})",
pixels.len()
);
let n_patches = grid_w
.checked_mul(grid_h)
.ok_or_else(|| anyhow::anyhow!("grid_w·grid_h overflow"))?;
let mut tokens = patch_embed_compute(pixels, &self.patch_embed, cfg, grid_w, grid_h);
let pos = self.resolved_position_embed(grid_w, grid_h);
debug_assert_eq!(pos.len(), n_patches * cfg.n_embd);
for (t, p) in tokens.iter_mut().zip(pos.iter()) {
*t += *p;
}
let mut scratch = VitScratch::new(cfg, n_patches);
for block in &self.blocks {
self.vit_block_forward(&mut tokens, block, &mut scratch, n_patches);
}
for t in 0..n_patches {
let row = &mut tokens[t * cfg.n_embd..(t + 1) * cfg.n_embd];
crate::backend::cpu::layer_norm_inplace(row, &self.post_ln_w, &self.post_ln_b, cfg.eps);
}
let pooled = pixel_shuffle(&tokens, cfg, grid_w, grid_h);
Ok(self.projector_forward(&pooled, cfg))
}
fn resolved_position_embed(&self, grid_w: usize, grid_h: usize) -> std::borrow::Cow<'_, [f32]> {
let cfg = &self.config;
let trained_side = (cfg.n_trained_patches as f64).sqrt().round() as usize;
debug_assert_eq!(
trained_side * trained_side,
cfg.n_trained_patches,
"non-square trained pos-embed grid is not currently supported"
);
if grid_w == trained_side && grid_h == trained_side {
std::borrow::Cow::Borrowed(self.position_embed.as_slice())
} else {
std::borrow::Cow::Owned(interpolate_pos_embed_2d(
&self.position_embed,
trained_side,
trained_side,
grid_h,
grid_w,
cfg.n_embd,
))
}
}
fn vit_block_forward(
&self,
tokens: &mut [f32],
block: &VitBlockWeights,
scratch: &mut VitScratch,
n_tokens: usize,
) {
let cfg = &self.config;
let n_embd = cfg.n_embd;
let n_head = cfg.n_head;
let head_dim = n_embd / n_head;
let scale = 1.0f32 / (head_dim as f32).sqrt();
scratch.pre_norm.copy_from_slice(tokens);
for t in 0..n_tokens {
let row = &mut scratch.pre_norm[t * n_embd..(t + 1) * n_embd];
crate::backend::cpu::layer_norm_inplace(row, &block.ln1_w, &block.ln1_b, cfg.eps);
}
for t in 0..n_tokens {
let pre_row = &scratch.pre_norm[t * n_embd..(t + 1) * n_embd];
let q_row = &mut scratch.q[t * n_embd..(t + 1) * n_embd];
let k_row = &mut scratch.k[t * n_embd..(t + 1) * n_embd];
let v_row = &mut scratch.v[t * n_embd..(t + 1) * n_embd];
block.q_w.gemv(pre_row, q_row);
block.k_w.gemv(pre_row, k_row);
block.v_w.gemv(pre_row, v_row);
for (q, b) in q_row.iter_mut().zip(block.q_b.iter()) {
*q += *b;
}
for (k, b) in k_row.iter_mut().zip(block.k_b.iter()) {
*k += *b;
}
for (v, b) in v_row.iter_mut().zip(block.v_b.iter()) {
*v += *b;
}
}
for h in 0..n_head {
for q_idx in 0..n_tokens {
let q_off = q_idx * n_embd + h * head_dim;
let q = &scratch.q[q_off..q_off + head_dim];
for (k_idx, score) in scratch.scores.iter_mut().enumerate() {
let k_off = k_idx * n_embd + h * head_dim;
let k = &scratch.k[k_off..k_off + head_dim];
let dot: f32 = q.iter().zip(k).map(|(a, b)| a * b).sum();
*score = dot * scale;
}
crate::backend::cpu::softmax_inplace(&mut scratch.scores);
let out_off = q_idx * n_embd + h * head_dim;
let out_slice = &mut scratch.attn_out[out_off..out_off + head_dim];
out_slice.iter_mut().for_each(|v| *v = 0.0);
for (k_idx, &s) in scratch.scores.iter().enumerate() {
let v_off = k_idx * n_embd + h * head_dim;
let v = &scratch.v[v_off..v_off + head_dim];
for (o, vv) in out_slice.iter_mut().zip(v) {
*o += s * vv;
}
}
}
}
for t in 0..n_tokens {
let attn_row = &scratch.attn_out[t * n_embd..(t + 1) * n_embd];
let proj_row = &mut scratch.attn_proj[t * n_embd..(t + 1) * n_embd];
block.o_w.gemv(attn_row, proj_row);
for (o, b) in proj_row.iter_mut().zip(block.o_b.iter()) {
*o += *b;
}
let tok_row = &mut tokens[t * n_embd..(t + 1) * n_embd];
for (tk, p) in tok_row.iter_mut().zip(proj_row.iter()) {
*tk += *p;
}
}
scratch.pre_norm.copy_from_slice(tokens);
for t in 0..n_tokens {
let row = &mut scratch.pre_norm[t * n_embd..(t + 1) * n_embd];
crate::backend::cpu::layer_norm_inplace(row, &block.ln2_w, &block.ln2_b, cfg.eps);
}
let n_ff = cfg.n_ff;
for t in 0..n_tokens {
let pre_row = &scratch.pre_norm[t * n_embd..(t + 1) * n_embd];
let ff_row = &mut scratch.ffn_mid[t * n_ff..(t + 1) * n_ff];
block.ffn_up_w.gemv(pre_row, ff_row);
for (f, b) in ff_row.iter_mut().zip(block.ffn_up_b.iter()) {
*f += *b;
}
crate::backend::cpu::gelu_inplace(ff_row);
let down_row = &mut scratch.ffn_out[t * n_embd..(t + 1) * n_embd];
block.ffn_down_w.gemv(ff_row, down_row);
for (d, b) in down_row.iter_mut().zip(block.ffn_down_b.iter()) {
*d += *b;
}
let tok_row = &mut tokens[t * n_embd..(t + 1) * n_embd];
for (tk, d) in tok_row.iter_mut().zip(down_row.iter()) {
*tk += *d;
}
}
}
fn projector_forward(&self, pooled: &[f32], cfg: &VisionEncoderConfig) -> Vec<f32> {
let p = &self.projector;
let in_dim = p.mm1_w.cols;
let mid_dim = p.mm1_w.rows; let out_dim = cfg.projection_dim;
let n_tokens = pooled.len() / in_dim;
let mut mid = vec![0f32; n_tokens * mid_dim];
let mut out = vec![0f32; n_tokens * out_dim];
for t in 0..n_tokens {
let in_row = &pooled[t * in_dim..(t + 1) * in_dim];
let mid_row = &mut mid[t * mid_dim..(t + 1) * mid_dim];
p.mm1_w.gemv(in_row, mid_row);
for (m, b) in mid_row.iter_mut().zip(p.mm1_b.iter()) {
*m += *b;
}
crate::backend::cpu::gelu_inplace(mid_row);
let out_row = &mut out[t * out_dim..(t + 1) * out_dim];
p.mm2_w.gemv(mid_row, out_row);
for (o, b) in out_row.iter_mut().zip(p.mm2_b.iter()) {
*o += *b;
}
}
out
}
}
struct VitScratch {
pre_norm: Vec<f32>,
q: Vec<f32>,
k: Vec<f32>,
v: Vec<f32>,
attn_out: Vec<f32>,
scores: Vec<f32>,
attn_proj: Vec<f32>,
ffn_mid: Vec<f32>,
ffn_out: Vec<f32>,
}
impl VitScratch {
fn new(cfg: &VisionEncoderConfig, n_tokens: usize) -> Self {
let n_pe = n_tokens * cfg.n_embd;
let n_pf = n_tokens * cfg.n_ff;
Self {
pre_norm: vec![0.0; n_pe],
q: vec![0.0; n_pe],
k: vec![0.0; n_pe],
v: vec![0.0; n_pe],
attn_out: vec![0.0; n_pe],
scores: vec![0.0; n_tokens],
attn_proj: vec![0.0; n_pe],
ffn_mid: vec![0.0; n_pf],
ffn_out: vec![0.0; n_pe],
}
}
}
fn patch_embed_compute(
image: &[f32],
patch_embed: &PatchEmbedWeights,
cfg: &VisionEncoderConfig,
grid_w: usize,
grid_h: usize,
) -> Vec<f32> {
let p = cfg.patch_size;
let in_dim = 3 * p * p;
let out_dim = cfg.n_embd;
let target_w = grid_w * p;
let target_h = grid_h * p;
debug_assert_eq!(
patch_embed.conv_w.len(),
in_dim * out_dim,
"patch_embed.conv_w length should be in_dim*out_dim after load-time transpose"
);
debug_assert_eq!(image.len(), 3 * target_h * target_w);
let h_stride = target_w;
let c_stride = target_h * target_w;
let n_patches = grid_w * grid_h;
let conv_w = &patch_embed.conv_w;
let conv_b = &patch_embed.conv_b;
let compute_patch = |patch: &mut [f32], patch_idx: usize, out_row: &mut [f32]| {
let gr = patch_idx / grid_w;
let gc = patch_idx % grid_w;
for c in 0..3 {
for kh in 0..p {
for kw in 0..p {
let pixel_r = gr * p + kh;
let pixel_c = gc * p + kw;
let in_idx = c * p * p + kh * p + kw;
let img_idx = c * c_stride + pixel_r * h_stride + pixel_c;
patch[in_idx] = image[img_idx];
}
}
}
out_row.copy_from_slice(conv_b);
crate::backend::cpu::matmul_f32(patch, conv_w, out_row, 1, out_dim, in_dim);
};
let mut out = vec![0f32; n_patches * out_dim];
#[cfg(feature = "parallel")]
{
use rayon::prelude::*;
out.par_chunks_mut(out_dim).enumerate().for_each_init(
|| vec![0f32; in_dim],
|patch, (patch_idx, out_row)| compute_patch(patch, patch_idx, out_row),
);
}
#[cfg(not(feature = "parallel"))]
{
let mut patch = vec![0f32; in_dim];
for (patch_idx, out_row) in out.chunks_mut(out_dim).enumerate() {
compute_patch(&mut patch, patch_idx, out_row);
}
}
out
}
fn interpolate_pos_embed_2d(
pos: &[f32],
in_h: usize,
in_w: usize,
out_h: usize,
out_w: usize,
n_embd: usize,
) -> Vec<f32> {
debug_assert_eq!(pos.len(), in_h * in_w * n_embd);
let mut out = vec![0f32; out_h * out_w * n_embd];
let scale_y = in_h as f32 / out_h as f32;
let scale_x = in_w as f32 / out_w as f32;
let in_h_max = (in_h - 1) as f32;
let in_w_max = (in_w - 1) as f32;
for or in 0..out_h {
for oc in 0..out_w {
let sy = ((or as f32 + 0.5) * scale_y - 0.5).clamp(0.0, in_h_max);
let sx = ((oc as f32 + 0.5) * scale_x - 0.5).clamp(0.0, in_w_max);
let y0 = sy.floor() as usize;
let x0 = sx.floor() as usize;
let y1 = (y0 + 1).min(in_h - 1);
let x1 = (x0 + 1).min(in_w - 1);
let dy = sy - y0 as f32;
let dx = sx - x0 as f32;
let w00 = (1.0 - dy) * (1.0 - dx);
let w01 = (1.0 - dy) * dx;
let w10 = dy * (1.0 - dx);
let w11 = dy * dx;
let dst = (or * out_w + oc) * n_embd;
let src00 = (y0 * in_w + x0) * n_embd;
let src01 = (y0 * in_w + x1) * n_embd;
let src10 = (y1 * in_w + x0) * n_embd;
let src11 = (y1 * in_w + x1) * n_embd;
for k in 0..n_embd {
out[dst + k] = w00 * pos[src00 + k]
+ w01 * pos[src01 + k]
+ w10 * pos[src10 + k]
+ w11 * pos[src11 + k];
}
}
}
out
}
fn pixel_shuffle(
tokens: &[f32],
cfg: &VisionEncoderConfig,
grid_w: usize,
grid_h: usize,
) -> Vec<f32> {
let sf = cfg.scale_factor;
debug_assert!(
sf > 0 && grid_w % sf == 0 && grid_h % sf == 0,
"patch grid {grid_w}×{grid_h} must be divisible by scale_factor ({sf})"
);
let new_w = grid_w / sf;
let new_h = grid_h / sf;
let n_embd = cfg.n_embd;
let new_dim = n_embd * sf * sf;
let n_out = new_w * new_h;
let mut out = vec![0f32; n_out * new_dim];
for or in 0..new_h {
for oc in 0..new_w {
let dst_off = (or * new_w + oc) * new_dim;
for sr in 0..sf {
for sc in 0..sf {
let in_r = or * sf + sr;
let in_c = oc * sf + sc;
let src_off = (in_r * grid_w + in_c) * n_embd;
let chan_base = (sr * sf + sc) * n_embd;
out[dst_off + chan_base..dst_off + chan_base + n_embd]
.copy_from_slice(&tokens[src_off..src_off + n_embd]);
}
}
}
}
out
}
fn read_rgb_array(gguf: &Arc<GgufFile>, key: &str) -> Result<[f32; 3]> {
let arr = gguf
.get_f32_array(key)
.with_context(|| format!("missing `{key}`"))?;
anyhow::ensure!(
arr.len() == 3,
"`{key}` length {} != 3 (RGB triple expected)",
arr.len()
);
Ok([arr[0], arr[1], arr[2]])
}
fn load_vit_block(
gguf: &Arc<GgufFile>,
il: usize,
cfg: &VisionEncoderConfig,
) -> Result<VitBlockWeights> {
let pfx = format!("v.blk.{il}");
let vec_f32 = |suffix: &str| load_vec_f32(gguf, &format!("{pfx}.{suffix}"));
let weight_f32 = |suffix: &str| -> Result<MmapWeight> {
let name = format!("{pfx}.{suffix}");
MmapWeight::from_gguf(gguf, &name).with_context(|| format!("loading {name}"))
};
let ln1_w = vec_f32("ln1.weight")?;
let ln1_b = vec_f32("ln1.bias")?;
let q_w = weight_f32("attn_q.weight")?;
let q_b = vec_f32("attn_q.bias")?;
let k_w = weight_f32("attn_k.weight")?;
let k_b = vec_f32("attn_k.bias")?;
let v_w = weight_f32("attn_v.weight")?;
let v_b = vec_f32("attn_v.bias")?;
let o_w = weight_f32("attn_out.weight")?;
let o_b = vec_f32("attn_out.bias")?;
let ln2_w = vec_f32("ln2.weight")?;
let ln2_b = vec_f32("ln2.bias")?;
let ffn_up_w = weight_f32("ffn_up.weight")?;
let ffn_up_b = vec_f32("ffn_up.bias")?;
let ffn_down_w = weight_f32("ffn_down.weight")?;
let ffn_down_b = vec_f32("ffn_down.bias")?;
let n_embd = cfg.n_embd;
let n_ff = cfg.n_ff;
for (name, w) in [
("attn_q.weight", &q_w),
("attn_k.weight", &k_w),
("attn_v.weight", &v_w),
("attn_out.weight", &o_w),
] {
anyhow::ensure!(
w.rows == n_embd && w.cols == n_embd,
"block {il} {name} shape ({}, {}) != ({n_embd}, {n_embd})",
w.rows,
w.cols,
);
}
anyhow::ensure!(
ffn_up_w.rows == n_ff && ffn_up_w.cols == n_embd,
"block {il} ffn_up.weight shape ({}, {}) != ({n_ff}, {n_embd})",
ffn_up_w.rows,
ffn_up_w.cols,
);
anyhow::ensure!(
ffn_down_w.rows == n_embd && ffn_down_w.cols == n_ff,
"block {il} ffn_down.weight shape ({}, {}) != ({n_embd}, {n_ff})",
ffn_down_w.rows,
ffn_down_w.cols,
);
for (name, v) in [
("ln1.weight", &ln1_w),
("ln1.bias", &ln1_b),
("attn_q.bias", &q_b),
("attn_k.bias", &k_b),
("attn_v.bias", &v_b),
("attn_out.bias", &o_b),
("ln2.weight", &ln2_w),
("ln2.bias", &ln2_b),
("ffn_down.bias", &ffn_down_b),
] {
anyhow::ensure!(
v.len() == n_embd,
"block {il} {name} len ({}) != n_embd ({n_embd})",
v.len(),
);
}
anyhow::ensure!(
ffn_up_b.len() == n_ff,
"block {il} ffn_up.bias len ({}) != n_ff ({n_ff})",
ffn_up_b.len(),
);
anyhow::ensure!(
cfg.n_head > 0 && n_embd % cfg.n_head == 0,
"n_embd ({n_embd}) not divisible by n_head ({})",
cfg.n_head,
);
Ok(VitBlockWeights {
ln1_w,
ln1_b,
q_w,
q_b,
k_w,
k_b,
v_w,
v_b,
o_w,
o_b,
ln2_w,
ln2_b,
ffn_up_w,
ffn_up_b,
ffn_down_w,
ffn_down_b,
})
}
fn wt_f32(gguf: &Arc<GgufFile>, name: &str) -> Result<MmapWeight> {
MmapWeight::from_gguf(gguf, name).with_context(|| format!("loading {name}"))
}
fn load_vec_f32(gguf: &Arc<GgufFile>, name: &str) -> Result<Vec<f32>> {
let tensor = gguf
.get_tensor(name)
.with_context(|| format!("loading {name}"))?;
anyhow::ensure!(
tensor.shape().len() == 1,
"tensor {name} must be 1D, got rank {}",
tensor.shape().len()
);
Ok(tensor.to_f32_vec())
}
#[cfg(test)]
mod tests {
use super::*;
fn synth_cfg(grid: usize, n_embd: usize, scale_factor: usize) -> VisionEncoderConfig {
let patch_size = 16;
let image_size = grid * patch_size;
let pixels_per_token = (patch_size * scale_factor).pow(2);
VisionEncoderConfig {
n_layer: 1,
n_embd,
n_ff: n_embd * 4,
n_head: 1,
eps: 1e-6,
image_size,
patch_size,
n_trained_patches: grid * grid,
projection_dim: n_embd,
scale_factor,
image_mean: [0.0; 3],
image_std: [1.0; 3],
image_min_pixels: 64 * pixels_per_token,
image_max_pixels: 256 * pixels_per_token,
}
}
#[test]
fn pixel_shuffle_2x2_reshapes_grid_in_row_major_order() {
let cfg = synth_cfg(4, 1, 2);
let n_in = cfg.n_trained_patches * cfg.n_embd; let mut tokens = vec![0f32; n_in];
for (i, t) in tokens.iter_mut().enumerate() {
*t = i as f32; }
let pooled = pixel_shuffle(&tokens, &cfg, 4, 4);
assert_eq!(pooled.len(), 4 * 4);
assert_eq!(&pooled[0..4], &[0.0, 1.0, 4.0, 5.0]);
assert_eq!(&pooled[4..8], &[2.0, 3.0, 6.0, 7.0]);
assert_eq!(&pooled[8..12], &[8.0, 9.0, 12.0, 13.0]);
assert_eq!(&pooled[12..16], &[10.0, 11.0, 14.0, 15.0]);
}
#[test]
fn pixel_shuffle_2x2_preserves_per_patch_channels() {
let cfg = synth_cfg(2, 3, 2);
let mut tokens = vec![0f32; 4 * 3];
for i in 0..4 {
for c in 0..3 {
tokens[i * 3 + c] = (i * 10 + c) as f32;
}
}
let pooled = pixel_shuffle(&tokens, &cfg, 2, 2);
assert_eq!(pooled.len(), 12);
assert_eq!(
pooled,
vec![
0.0, 1.0, 2.0, 10.0, 11.0, 12.0, 20.0, 21.0, 22.0, 30.0, 31.0, 32.0, ]
);
}
#[test]
fn patch_embed_loader_transpose_round_trip() {
let p = 2usize; let in_dim = 3 * p * p; let out_dim = 5; let mut raw = vec![0f32; p * p * 3 * out_dim];
let tag = |oc: usize, c: usize, kh: usize, kw: usize| -> f32 {
(oc * 1000 + c * 100 + kh * 10 + kw) as f32
};
for oc in 0..out_dim {
for c in 0..3 {
for kh in 0..p {
for kw in 0..p {
let src = kw + p * kh + p * p * c + p * p * 3 * oc;
raw[src] = tag(oc, c, kh, kw);
}
}
}
}
let mut conv_w = vec![0f32; in_dim * out_dim];
for oc in 0..out_dim {
for c in 0..3 {
for kh in 0..p {
for kw in 0..p {
let src = kw + p * kh + p * p * c + p * p * 3 * oc;
let in_idx = c * p * p + kh * p + kw;
conv_w[in_idx * out_dim + oc] = raw[src];
}
}
}
}
for oc in 0..out_dim {
for c in 0..3 {
for kh in 0..p {
for kw in 0..p {
let in_idx = c * p * p + kh * p + kw;
let got = conv_w[in_idx * out_dim + oc];
let expected = tag(oc, c, kh, kw);
assert_eq!(
got, expected,
"conv_w[in_idx={in_idx}, oc={oc}] = {got} != {expected} \
(oc={oc}, c={c}, kh={kh}, kw={kw})"
);
}
}
}
}
}
#[test]
fn patch_embed_compute_matches_explicit_dot_product() {
let grid = 2usize;
let cfg = synth_cfg(grid, 3, 1);
let p = cfg.patch_size;
let in_dim = 3 * p * p; let out_dim = cfg.n_embd; let n_patches = cfg.n_trained_patches;
let mut conv_w = vec![0f32; in_dim * out_dim];
for in_idx in 0..in_dim {
for oc in 0..out_dim {
conv_w[in_idx * out_dim + oc] = (in_idx as f32) * 0.01 + (oc as f32) * 0.1;
}
}
let conv_b = vec![0.5, 1.0, 1.5];
let patch_embed = PatchEmbedWeights {
conv_w: conv_w.clone(),
conv_b: conv_b.clone(),
};
let n_pix = 3 * cfg.image_size * cfg.image_size;
let mut image = vec![0f32; n_pix];
for c in 0..3 {
for y in 0..cfg.image_size {
for x in 0..cfg.image_size {
image[c * cfg.image_size * cfg.image_size + y * cfg.image_size + x] =
(c as f32) * 0.3 + (y as f32) * 0.05 + (x as f32) * 0.02;
}
}
}
let actual = patch_embed_compute(&image, &patch_embed, &cfg, grid, grid);
assert_eq!(actual.len(), n_patches * out_dim);
for gr in 0..grid {
for gc in 0..grid {
let patch_idx = gr * grid + gc;
for oc in 0..out_dim {
let mut expected = conv_b[oc];
for c in 0..3 {
for kh in 0..p {
for kw in 0..p {
let in_idx = c * p * p + kh * p + kw;
let pixel_r = gr * p + kh;
let pixel_c = gc * p + kw;
let img_v = image[c * cfg.image_size * cfg.image_size
+ pixel_r * cfg.image_size
+ pixel_c];
expected += img_v * conv_w[in_idx * out_dim + oc];
}
}
}
let got = actual[patch_idx * out_dim + oc];
assert!(
(got - expected).abs() < 1e-5,
"patch_idx={patch_idx} oc={oc}: got {got}, expected {expected}"
);
}
}
}
}
#[test]
fn interpolate_pos_embed_2d_identity() {
let n_embd = 3;
let mut input = vec![0f32; 4 * 4 * n_embd];
for (i, v) in input.iter_mut().enumerate() {
*v = i as f32;
}
let out = interpolate_pos_embed_2d(&input, 4, 4, 4, 4, n_embd);
for (i, (a, b)) in out.iter().zip(input.iter()).enumerate() {
assert!((a - b).abs() < 1e-5, "diff at {i}: {a} vs {b}");
}
}
#[test]
fn interpolate_pos_embed_2d_corners_match() {
let n_embd = 1;
let input: Vec<f32> = vec![1.0, 2.0, 3.0, 4.0];
let out = interpolate_pos_embed_2d(&input, 2, 2, 4, 6, n_embd);
assert!((out[0] - 1.0).abs() < 1e-5);
assert!((out[5] - 2.0).abs() < 1e-5);
assert!((out[3 * 6] - 3.0).abs() < 1e-5);
assert!((out[3 * 6 + 5] - 4.0).abs() < 1e-5);
}
#[test]
fn pixel_shuffle_non_square() {
let cfg = synth_cfg(
4, 1, 2,
);
let grid_w = 6;
let grid_h = 4;
let n_in = grid_w * grid_h;
let mut tokens = vec![0f32; n_in];
for (i, t) in tokens.iter_mut().enumerate() {
*t = i as f32;
}
let pooled = pixel_shuffle(&tokens, &cfg, grid_w, grid_h);
assert_eq!(pooled.len(), 6 * 4);
assert_eq!(&pooled[0..4], &[0.0, 1.0, 6.0, 7.0]);
assert_eq!(&pooled[4..8], &[2.0, 3.0, 8.0, 9.0]);
assert_eq!(&pooled[8..12], &[4.0, 5.0, 10.0, 11.0]);
assert_eq!(&pooled[12..16], &[12.0, 13.0, 18.0, 19.0]);
}
}