use super::nn;
use super::tensor::Mat;
use super::weights::Weights;
use crate::error::{FocrError, FocrResult};
use rayon::prelude::*;
#[cfg(not(target_arch = "wasm32"))]
use std::time::Instant;
#[cfg(target_arch = "wasm32")]
use web_time::Instant;
fn checked_shape_mul(context: &str, lhs: usize, rhs: usize, expression: &str) -> FocrResult<usize> {
lhs.checked_mul(rhs).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"{context}: usize overflow computing {expression} ({lhs} * {rhs})"
))
})
}
fn checked_shape_sub(context: &str, lhs: usize, rhs: usize, expression: &str) -> FocrResult<usize> {
lhs.checked_sub(rhs).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"{context}: usize underflow computing {expression} ({lhs} - {rhs})"
))
})
}
fn checked_shape_add(context: &str, lhs: usize, rhs: usize, expression: &str) -> FocrResult<usize> {
lhs.checked_add(rhs).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"{context}: usize overflow computing {expression} ({lhs} + {rhs})"
))
})
}
fn ensure_mat_data_len(mat: &Mat, context: &str) -> FocrResult<()> {
let expected_len = checked_shape_mul(context, mat.rows, mat.cols, "rows*cols")?;
if mat.data.len() == expected_len {
return Ok(());
}
Err(FocrError::Other(anyhow::anyhow!(
"{context}: data len {} != rows*cols {}",
mat.data.len(),
expected_len
)))
}
#[derive(Debug, Clone, Copy)]
pub struct ClipConfig {
pub num_layers: usize,
pub hidden_size: usize,
pub num_heads: usize,
pub ffn_hidden_size: usize,
pub patch_size: usize,
pub layernorm_eps: f32,
pub pre_layernorm_eps: f32,
}
impl Default for ClipConfig {
fn default() -> Self {
Self {
num_layers: 24,
hidden_size: 1024,
num_heads: 16,
ffn_hidden_size: 4096,
patch_size: 14,
layernorm_eps: 1e-5,
pre_layernorm_eps: 1e-5,
}
}
}
impl ClipConfig {
#[must_use]
pub fn head_dim(&self) -> usize {
self.hidden_size / self.num_heads
}
}
#[derive(Debug, Clone)]
pub struct LayerNormParams {
pub weight: Vec<f32>,
pub bias: Vec<f32>,
}
#[derive(Debug, Clone)]
pub struct LinearParams {
pub weight_t: Mat,
pub bias: Option<Vec<f32>>,
pub out_features: usize,
pub in_features: usize,
}
impl LinearParams {
pub fn from_row_major(
weight: &[f32],
bias: Option<Vec<f32>>,
out_features: usize,
in_features: usize,
) -> FocrResult<Self> {
let wt = transpose(weight, out_features, in_features)?;
if let Some(b) = &bias
&& b.len() != out_features
{
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip linear: bias len {} != out_features {}",
b.len(),
out_features
)));
}
Ok(Self {
weight_t: Mat::from_vec(in_features, out_features, wt),
bias,
out_features,
in_features,
})
}
}
#[derive(Debug, Clone)]
pub struct ClipBlockWeights {
pub layer_norm1: LayerNormParams,
pub qkv_proj: LinearParams,
pub out_proj: LinearParams,
pub layer_norm2: LayerNormParams,
pub fc1: LinearParams,
pub fc2: LinearParams,
}
#[derive(Debug, Clone)]
pub struct ClipWeights {
pub class_embedding: Vec<f32>,
pub position_embedding: Vec<f32>,
pub num_positions: usize,
pub pre_layernorm: LayerNormParams,
pub blocks: Vec<ClipBlockWeights>,
}
pub fn forward(weights: &Weights, _image: &Mat, sam_features: &Mat) -> FocrResult<Mat> {
ensure_mat_data_len(sam_features, "vision_clip forward sam_features")?;
let th = Instant::now();
let cw = clip_weights_from(weights)?;
super::timing_log(&format!(
" clip.hydrate {:.2}s",
th.elapsed().as_secs_f64()
));
forward_from_sam(&ClipConfig::default(), &cw, sam_features)
}
pub(crate) fn forward_from_sam(
cfg: &ClipConfig,
cw: &ClipWeights,
sam_features: &Mat,
) -> FocrResult<Mat> {
ensure_mat_data_len(sam_features, "vision_clip forward sam_features")?;
let sam_t = transpose(&sam_features.data, sam_features.rows, sam_features.cols)?;
let sam_mat = Mat::from_vec(sam_features.cols, sam_features.rows, sam_t);
let tb = Instant::now();
let out = forward_with(cfg, cw, &sam_mat);
super::timing_log(&format!(
" clip.blocks {:.2}s",
tb.elapsed().as_secs_f64()
));
out
}
pub(crate) fn forward_from_sam_streamed(
cfg: &ClipConfig,
weights: &Weights,
sam_features: &Mat,
) -> FocrResult<Mat> {
ensure_mat_data_len(sam_features, "vision_clip forward sam_features")?;
let sam_t = transpose(&sam_features.data, sam_features.rows, sam_features.cols)?;
let sam_mat = Mat::from_vec(sam_features.cols, sam_features.rows, sam_t);
let p = "model.vision_model";
let pos_name = format!("{p}.embeddings.position_embedding.weight");
let (num_positions, _pos_dim) = tensor_rank2_shape(weights, &pos_name)?;
let class_embedding = weights.vec(&format!("{p}.embeddings.class_embedding"))?;
let position_embedding = weights.vec(&pos_name)?;
let pre_layernorm = LayerNormParams {
weight: weights.vec(&format!("{p}.pre_layrnorm.weight"))?,
bias: weights.vec(&format!("{p}.pre_layrnorm.bias"))?,
};
let tb = Instant::now();
let out = forward_with_supplied(
cfg,
&class_embedding,
&position_embedding,
num_positions,
&pre_layernorm,
&sam_mat,
&mut |l| Ok(std::borrow::Cow::Owned(clip_block_from(weights, l)?)),
);
super::timing_log(&format!(
" clip.blocks(streamed) {:.2}s",
tb.elapsed().as_secs_f64()
));
out
}
pub(crate) fn forward_from_sam_streamed_views(
cfg: &ClipConfig,
weights: &Weights,
sam_features_per_view: &[&Mat],
) -> FocrResult<Vec<Mat>> {
if sam_features_per_view.is_empty() {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip::forward_from_sam_streamed_views: empty view list"
)));
}
let mut sam_mats = Vec::with_capacity(sam_features_per_view.len());
for sam_features in sam_features_per_view {
ensure_mat_data_len(sam_features, "vision_clip forward sam_features")?;
let sam_t = transpose(&sam_features.data, sam_features.rows, sam_features.cols)?;
sam_mats.push(Mat::from_vec(sam_features.cols, sam_features.rows, sam_t));
}
let p = "model.vision_model";
let pos_name = format!("{p}.embeddings.position_embedding.weight");
let (num_positions, _pos_dim) = tensor_rank2_shape(weights, &pos_name)?;
let class_embedding = weights.vec(&format!("{p}.embeddings.class_embedding"))?;
let position_embedding = weights.vec(&pos_name)?;
let pre_layernorm = LayerNormParams {
weight: weights.vec(&format!("{p}.pre_layrnorm.weight"))?,
bias: weights.vec(&format!("{p}.pre_layrnorm.bias"))?,
};
let sam_refs: Vec<&Mat> = sam_mats.iter().collect();
let tb = Instant::now();
let out = forward_with_supplied_views(
cfg,
&class_embedding,
&position_embedding,
num_positions,
&pre_layernorm,
&sam_refs,
&mut |l| Ok(std::borrow::Cow::Owned(clip_block_from(weights, l)?)),
);
super::timing_log(&format!(
" clip.blocks(streamed, {} views) {:.2}s",
sam_refs.len(),
tb.elapsed().as_secs_f64()
));
out
}
pub(crate) fn clip_block_from(weights: &Weights, l: usize) -> FocrResult<ClipBlockWeights> {
let b = format!("model.vision_model.transformer.layers.{l}");
let ln = |n: &str| -> FocrResult<LayerNormParams> {
Ok(LayerNormParams {
weight: weights.vec(&format!("{n}.weight"))?,
bias: weights.vec(&format!("{n}.bias"))?,
})
};
let lin = |n: &str| -> FocrResult<LinearParams> {
let weight_name = format!("{n}.weight");
let (out_features, in_features) = tensor_rank2_shape(weights, &weight_name)?;
LinearParams::from_row_major(
&weights.vec(&weight_name)?,
Some(weights.vec(&format!("{n}.bias"))?),
out_features,
in_features,
)
};
Ok(ClipBlockWeights {
layer_norm1: ln(&format!("{b}.layer_norm1"))?,
qkv_proj: lin(&format!("{b}.self_attn.qkv_proj"))?,
out_proj: lin(&format!("{b}.self_attn.out_proj"))?,
layer_norm2: ln(&format!("{b}.layer_norm2"))?,
fc1: lin(&format!("{b}.mlp.fc1"))?,
fc2: lin(&format!("{b}.mlp.fc2"))?,
})
}
pub(crate) fn clip_weights_from(weights: &Weights) -> FocrResult<ClipWeights> {
clip_weights_from_with(&ClipConfig::default(), weights)
}
pub(crate) fn clip_weights_from_with(
cfg: &ClipConfig,
weights: &Weights,
) -> FocrResult<ClipWeights> {
let p = "model.vision_model";
let ln = |n: &str| -> FocrResult<LayerNormParams> {
Ok(LayerNormParams {
weight: weights.vec(&format!("{n}.weight"))?,
bias: weights.vec(&format!("{n}.bias"))?,
})
};
struct RawLinear {
weight: Vec<f32>,
bias: Vec<f32>,
out_features: usize,
in_features: usize,
}
struct RawBlock {
layer_norm1: LayerNormParams,
qkv_proj: RawLinear,
out_proj: RawLinear,
layer_norm2: LayerNormParams,
fc1: RawLinear,
fc2: RawLinear,
}
let lin = |n: &str| -> FocrResult<RawLinear> {
let weight_name = format!("{n}.weight");
let (out_features, in_features) = tensor_rank2_shape(weights, &weight_name)?;
Ok(RawLinear {
weight: weights.vec(&weight_name)?,
bias: weights.vec(&format!("{n}.bias"))?,
out_features,
in_features,
})
};
let pos_name = format!("{p}.embeddings.position_embedding.weight");
let (num_positions, _pos_dim) = tensor_rank2_shape(weights, &pos_name)?;
let mut raw_blocks = Vec::with_capacity(cfg.num_layers);
for l in 0..cfg.num_layers {
let b = format!("{p}.transformer.layers.{l}");
raw_blocks.push(RawBlock {
layer_norm1: ln(&format!("{b}.layer_norm1"))?,
qkv_proj: lin(&format!("{b}.self_attn.qkv_proj"))?,
out_proj: lin(&format!("{b}.self_attn.out_proj"))?,
layer_norm2: ln(&format!("{b}.layer_norm2"))?,
fc1: lin(&format!("{b}.mlp.fc1"))?,
fc2: lin(&format!("{b}.mlp.fc2"))?,
});
}
let finish = |r: RawLinear| -> FocrResult<LinearParams> {
LinearParams::from_row_major(&r.weight, Some(r.bias), r.out_features, r.in_features)
};
let blocks = raw_blocks
.into_par_iter()
.map(|r| {
Ok(ClipBlockWeights {
layer_norm1: r.layer_norm1,
qkv_proj: finish(r.qkv_proj)?,
out_proj: finish(r.out_proj)?,
layer_norm2: r.layer_norm2,
fc1: finish(r.fc1)?,
fc2: finish(r.fc2)?,
})
})
.collect::<FocrResult<Vec<_>>>()?;
Ok(ClipWeights {
class_embedding: weights.vec(&format!("{p}.embeddings.class_embedding"))?,
position_embedding: weights.vec(&pos_name)?,
num_positions,
pre_layernorm: ln(&format!("{p}.pre_layrnorm"))?,
blocks,
})
}
fn tensor_rank2_shape(weights: &Weights, name: &str) -> FocrResult<(usize, usize)> {
let view = weights.tensor(name)?;
let [rows, cols] = view.shape else {
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?} has rank {}; expected 2 ([rows, cols])",
view.shape.len()
)));
};
Ok((*rows, *cols))
}
pub fn forward_with(
cfg: &ClipConfig,
weights: &ClipWeights,
sam_features: &Mat,
) -> FocrResult<Mat> {
let dim = cfg.hidden_size;
if sam_features.cols != dim {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip: sam_features width {} != hidden_size {}",
sam_features.cols,
dim
)));
}
ensure_mat_data_len(sam_features, "vision_clip forward_with sam_features")?;
if weights.class_embedding.len() != dim {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip: class_embedding len {} != hidden_size {}",
weights.class_embedding.len(),
dim
)));
}
if weights.blocks.len() != cfg.num_layers {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip: {} blocks != num_layers {}",
weights.blocks.len(),
cfg.num_layers
)));
}
forward_with_supplied(
cfg,
&weights.class_embedding,
&weights.position_embedding,
weights.num_positions,
&weights.pre_layernorm,
sam_features,
&mut |l| Ok(std::borrow::Cow::Borrowed(&weights.blocks[l])),
)
}
#[allow(clippy::too_many_arguments)]
fn forward_with_supplied<'w>(
cfg: &ClipConfig,
class_embedding: &[f32],
position_embedding: &[f32],
num_positions: usize,
pre_layernorm: &LayerNormParams,
sam_features: &Mat,
block_at: &mut dyn FnMut(usize) -> FocrResult<std::borrow::Cow<'w, ClipBlockWeights>>,
) -> FocrResult<Mat> {
let mut out = forward_with_supplied_views(
cfg,
class_embedding,
position_embedding,
num_positions,
pre_layernorm,
std::slice::from_ref(&sam_features),
block_at,
)?;
Ok(out.remove(0))
}
#[allow(clippy::too_many_arguments)]
fn forward_with_supplied_views<'w>(
cfg: &ClipConfig,
class_embedding: &[f32],
position_embedding: &[f32],
num_positions: usize,
pre_layernorm: &LayerNormParams,
sam_features_per_view: &[&Mat],
block_at: &mut dyn FnMut(usize) -> FocrResult<std::borrow::Cow<'w, ClipBlockWeights>>,
) -> FocrResult<Vec<Mat>> {
let dim = cfg.hidden_size;
let mut xs = Vec::with_capacity(sam_features_per_view.len());
for sam_features in sam_features_per_view {
let mut x = prepend_class_token(class_embedding, sam_features)?;
let seq = x.rows; let pos = abs_pos_for_len(position_embedding, num_positions, dim, seq)?;
add_in_place(&mut x, &pos)?;
x = nn::layer_norm(
&x,
Some(&pre_layernorm.weight),
Some(&pre_layernorm.bias),
cfg.pre_layernorm_eps,
)?;
xs.push(x);
}
for l in 0..cfg.num_layers {
let block = block_at(l)?;
for x in &mut xs {
*x = transformer_block(cfg, &block, x)?;
super::progress::vision_step();
}
}
Ok(xs)
}
pub fn forward_with_batched(
cfg: &ClipConfig,
weights: &ClipWeights,
sam_features_per_view: &[&Mat],
) -> FocrResult<Vec<Mat>> {
let v = sam_features_per_view.len();
if v == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip forward_with_batched: empty view batch"
)));
}
let dim = cfg.hidden_size;
let n = sam_features_per_view[0].rows;
for (i, sf) in sam_features_per_view.iter().enumerate() {
if sf.cols != dim {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip forward_with_batched: view {i} width {} != hidden_size {dim}",
sf.cols
)));
}
if sf.rows != n {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip forward_with_batched: view {i} rows {} != {n} (ragged batch)",
sf.rows
)));
}
ensure_mat_data_len(sf, "vision_clip forward_with_batched sam_features")?;
}
if weights.class_embedding.len() != dim {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip forward_with_batched: class_embedding len {} != hidden_size {dim}",
weights.class_embedding.len()
)));
}
if weights.blocks.len() != cfg.num_layers {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip forward_with_batched: {} blocks != num_layers {}",
weights.blocks.len(),
cfg.num_layers
)));
}
let seq = checked_shape_add("vision_clip forward_with_batched", n, 1, "N+1")?;
let pos = abs_pos_for_len(&weights.position_embedding, weights.num_positions, dim, seq)?;
ensure_mat_shape(&pos, seq, dim, "vision_clip forward_with_batched pos")?;
let total_rows = checked_shape_mul("vision_clip forward_with_batched", v, seq, "V*seq")?;
let stacked_len = checked_shape_mul(
"vision_clip forward_with_batched",
total_rows,
dim,
"V*seq*dim",
)?;
let mut data = vec![0.0f32; stacked_len];
for (vv, sf) in sam_features_per_view.iter().enumerate() {
let base = vv * seq * dim;
data[base..base + dim].copy_from_slice(&weights.class_embedding);
data[base + dim..base + seq * dim].copy_from_slice(&sf.data);
for s in 0..seq {
let off = base + s * dim;
let prow = &pos.data[s * dim..(s + 1) * dim];
for (d, pv) in prow.iter().enumerate() {
data[off + d] += *pv;
}
}
}
let mut x = Mat::from_vec(total_rows, dim, data);
x = nn::layer_norm(
&x,
Some(&weights.pre_layernorm.weight),
Some(&weights.pre_layernorm.bias),
cfg.pre_layernorm_eps,
)?;
for block in &weights.blocks {
x = transformer_block_batched(cfg, block, &x, v, seq)?;
}
let mut out = Vec::with_capacity(v);
for vv in 0..v {
let base = vv * seq * dim;
out.push(Mat::from_vec(
seq,
dim,
x.data[base..base + seq * dim].to_vec(),
));
}
Ok(out)
}
pub(crate) fn forward_batched_from_sam(
cfg: &ClipConfig,
weights: &ClipWeights,
sam_features_per_view: &[&Mat],
) -> FocrResult<Vec<Mat>> {
let mut transposed = Vec::with_capacity(sam_features_per_view.len());
for sf in sam_features_per_view {
ensure_mat_data_len(sf, "vision_clip forward_batched_from_sam sam_features")?;
let t = transpose(&sf.data, sf.rows, sf.cols)?;
transposed.push(Mat::from_vec(sf.cols, sf.rows, t));
}
let refs: Vec<&Mat> = transposed.iter().collect();
forward_with_batched(cfg, weights, &refs)
}
fn transformer_block_batched(
cfg: &ClipConfig,
w: &ClipBlockWeights,
x: &Mat,
v: usize,
seq: usize,
) -> FocrResult<Mat> {
let normed = nn::layer_norm(
x,
Some(&w.layer_norm1.weight),
Some(&w.layer_norm1.bias),
cfg.layernorm_eps,
)?;
let attn = self_attention_batched(cfg, w, &normed, v, seq)?;
let h = add(x, &attn)?;
let normed2 = nn::layer_norm(
&h,
Some(&w.layer_norm2.weight),
Some(&w.layer_norm2.bias),
cfg.layernorm_eps,
)?;
let mlp = feed_forward(w, &normed2)?;
add(&h, &mlp)
}
fn self_attention_batched(
cfg: &ClipConfig,
w: &ClipBlockWeights,
x: &Mat,
v: usize,
seq: usize,
) -> FocrResult<Mat> {
let dim = cfg.hidden_size;
let heads = cfg.num_heads;
if heads == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip self_attention_batched: num_heads must be non-zero"
)));
}
if dim == 0 || !dim.is_multiple_of(heads) {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip self_attention_batched: hidden_size {dim} must be non-zero and divisible by heads {heads}"
)));
}
let rows = checked_shape_mul("vision_clip self_attention_batched", v, seq, "V*seq")?;
if x.rows != rows {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip self_attention_batched: x.rows {} != V*seq {rows}",
x.rows
)));
}
if x.cols != dim {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip self_attention_batched: x.cols {} != hidden_size {dim}",
x.cols
)));
}
ensure_mat_data_len(x, "vision_clip self_attention_batched input")?;
let hd = dim / heads;
let three_dim = checked_shape_mul("vision_clip self_attention_batched", 3, dim, "3*dim")?;
let head_span = checked_shape_mul("vision_clip self_attention_batched", seq, hd, "seq*hd")?;
let num_bh = checked_shape_mul("vision_clip self_attention_batched", v, heads, "V*heads")?;
let buf_len = checked_shape_mul(
"vision_clip self_attention_batched",
num_bh,
head_span,
"V*heads*seq*hd",
)?;
let qkv = linear(&w.qkv_proj, x)?;
ensure_mat_shape(
&qkv,
rows,
three_dim,
"vision_clip self_attention_batched qkv",
)?;
let mut q = vec![0.0f32; buf_len];
let mut k = vec![0.0f32; buf_len];
let mut vbuf = vec![0.0f32; buf_len];
for view in 0..v {
for s in 0..seq {
let row = qkv.row(view * seq + s);
for hh in 0..heads {
let bh = view * heads + hh;
for d in 0..hd {
let q_src = hh * hd + d;
let k_src = dim + hh * hd + d;
let v_src = 2 * dim + hh * hd + d;
let dst = bh * head_span + s * hd + d;
q[dst] = row[q_src];
k[dst] = row[k_src];
vbuf[dst] = row[v_src];
}
}
}
}
let scale = 1.0f32 / (hd as f32).sqrt();
let ctx = nn::sdpa(&q, &k, &vbuf, num_bh, seq, seq, hd, hd, scale, false);
if ctx.len() != buf_len {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip self_attention_batched: sdpa context len {} != expected {buf_len}",
ctx.len()
)));
}
let merged_len =
checked_shape_mul("vision_clip self_attention_batched", rows, dim, "V*seq*dim")?;
let mut merged = Mat::from_vec(rows, dim, vec![0.0f32; merged_len]);
for view in 0..v {
for hh in 0..heads {
let bh = view * heads + hh;
for s in 0..seq {
for d in 0..hd {
let src = bh * head_span + s * hd + d;
let dst = (view * seq + s) * dim + hh * hd + d;
merged.data[dst] = ctx[src];
}
}
}
}
linear(&w.out_proj, &merged)
}
fn transformer_block(cfg: &ClipConfig, w: &ClipBlockWeights, x: &Mat) -> FocrResult<Mat> {
let normed = nn::layer_norm(
x,
Some(&w.layer_norm1.weight),
Some(&w.layer_norm1.bias),
cfg.layernorm_eps,
)?;
let attn = self_attention(cfg, w, &normed)?;
let h = add(x, &attn)?;
let normed2 = nn::layer_norm(
&h,
Some(&w.layer_norm2.weight),
Some(&w.layer_norm2.bias),
cfg.layernorm_eps,
)?;
let mlp = feed_forward(w, &normed2)?;
add(&h, &mlp)
}
fn self_attention(cfg: &ClipConfig, w: &ClipBlockWeights, x: &Mat) -> FocrResult<Mat> {
let dim = cfg.hidden_size;
let heads = cfg.num_heads;
let seq = x.rows;
if heads == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip self_attention: num_heads must be non-zero"
)));
}
if dim == 0 || !dim.is_multiple_of(heads) {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip self_attention: hidden_size {} must be non-zero and divisible by heads {}",
dim,
heads
)));
}
if x.cols != dim {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip self_attention: x.cols {} != hidden_size {}",
x.cols,
dim
)));
}
ensure_mat_data_len(x, "vision_clip self_attention input")?;
let hd = dim / heads;
let three_dim = checked_shape_mul("vision_clip self_attention", 3, dim, "3*hidden_size")?;
let head_span = checked_shape_mul("vision_clip self_attention", seq, hd, "seq*head_dim")?;
let qkv_buffer_len = checked_shape_mul(
"vision_clip self_attention",
heads,
head_span,
"heads*seq*head_dim",
)?;
let qkv = linear(&w.qkv_proj, x)?;
ensure_mat_shape(
&qkv,
seq,
three_dim,
"vision_clip self_attention qkv output",
)?;
let mut q = vec![0.0f32; qkv_buffer_len];
let mut k = vec![0.0f32; qkv_buffer_len];
let mut v = vec![0.0f32; qkv_buffer_len];
for s in 0..seq {
let row = qkv.row(s);
for hh in 0..heads {
for d in 0..hd {
let q_src = hh * hd + d;
let k_src = dim + hh * hd + d;
let v_src = 2 * dim + hh * hd + d;
let dst = hh * head_span + s * hd + d;
q[dst] = row[q_src];
k[dst] = row[k_src];
v[dst] = row[v_src];
}
}
}
let scale = 1.0f32 / (hd as f32).sqrt();
let ctx = nn::sdpa(&q, &k, &v, heads, seq, seq, hd, hd, scale, false);
let expected_ctx_len = qkv_buffer_len;
if ctx.len() != expected_ctx_len {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip self_attention: sdpa context len {} != expected {}",
ctx.len(),
expected_ctx_len
)));
}
let merged_len = checked_shape_mul("vision_clip self_attention", seq, dim, "seq*hidden_size")?;
let mut merged = Mat::from_vec(seq, dim, vec![0.0f32; merged_len]);
for hh in 0..heads {
for s in 0..seq {
for d in 0..hd {
let src = hh * head_span + s * hd + d;
let dst = s * dim + hh * hd + d;
merged.data[dst] = ctx[src];
}
}
}
let out = linear(&w.out_proj, &merged)?;
ensure_mat_shape(
&out,
seq,
dim,
"vision_clip self_attention projection output",
)?;
Ok(out)
}
fn ensure_mat_shape(mat: &Mat, rows: usize, cols: usize, context: &str) -> FocrResult<()> {
if mat.rows == rows && mat.cols == cols {
return Ok(());
}
Err(FocrError::Other(anyhow::anyhow!(
"{context}: shape {:?} != expected ({rows}, {cols})",
mat.shape()
)))
}
fn feed_forward(w: &ClipBlockWeights, x: &Mat) -> FocrResult<Mat> {
let mut hidden = linear(&w.fc1, x)?;
nn::quick_gelu(&mut hidden);
linear(&w.fc2, &hidden)
}
fn linear(w: &LinearParams, x: &Mat) -> FocrResult<Mat> {
if x.cols != w.in_features {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip linear: x.cols {} != in_features {}",
x.cols,
w.in_features
)));
}
ensure_mat_data_len(x, "vision_clip linear input")?;
if w.weight_t.rows != w.in_features || w.weight_t.cols != w.out_features {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip linear: weight_t shape {:?} != [in {}, out {}]",
w.weight_t.shape(),
w.in_features,
w.out_features
)));
}
let mut y = nn::matmul(x, &w.weight_t)?;
if let Some(b) = &w.bias {
if b.len() != w.out_features {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip linear: bias len {} != out_features {}",
b.len(),
w.out_features
)));
}
for r in 0..y.rows {
let row = y.row_mut(r);
for (c, bv) in b.iter().enumerate() {
row[c] += *bv;
}
}
}
Ok(y)
}
fn transpose(src: &[f32], rows: usize, cols: usize) -> FocrResult<Vec<f32>> {
let expected_len = checked_shape_mul("vision_clip transpose", rows, cols, "rows*cols")?;
if src.len() != expected_len {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip transpose: source len {} != rows*cols {}",
src.len(),
expected_len
)));
}
let mut out = vec![0.0f32; expected_len];
for c in 0..cols {
let dst = &mut out[c * rows..(c + 1) * rows];
for (r, slot) in dst.iter_mut().enumerate() {
*slot = src[r * cols + c];
}
}
Ok(out)
}
fn prepend_class_token(class_embedding: &[f32], patches: &Mat) -> FocrResult<Mat> {
let dim = patches.cols;
if class_embedding.len() != dim {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip class token: class_embedding len {} != patch dim {}",
class_embedding.len(),
dim
)));
}
ensure_mat_data_len(patches, "vision_clip class token patches")?;
let seq = checked_shape_add("vision_clip class token", patches.rows, 1, "patch_rows+1")?;
let out_len = checked_shape_mul("vision_clip class token", seq, dim, "seq*dim")?;
let mut data = Vec::with_capacity(out_len);
data.extend_from_slice(class_embedding);
data.extend_from_slice(&patches.data);
Ok(Mat::from_vec(seq, dim, data))
}
fn abs_pos_for_len(table: &[f32], num_positions: usize, dim: usize, seq: usize) -> FocrResult<Mat> {
let expected_table_len = checked_shape_mul(
"vision_clip abs_pos",
num_positions,
dim,
"num_positions*dim",
)?;
if table.len() != expected_table_len {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip abs_pos: table len {} != num_positions*dim {}",
table.len(),
expected_table_len
)));
}
if num_positions == 0 || seq == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip abs_pos: num_positions ({num_positions}) and seq ({seq}) must be non-zero"
)));
}
if seq == num_positions {
return Ok(Mat::from_vec(seq, dim, table.to_vec()));
}
let num_patches =
checked_shape_sub("vision_clip abs_pos", num_positions, 1, "num_positions-1")?;
let runtime_patches = checked_shape_sub("vision_clip abs_pos", seq, 1, "seq-1")?;
let src = isqrt(num_patches);
let tgt = isqrt(runtime_patches);
let src_square = checked_shape_mul("vision_clip abs_pos", src, src, "src*src")?;
let tgt_square = checked_shape_mul("vision_clip abs_pos", tgt, tgt, "tgt*tgt")?;
if src_square != num_patches || tgt_square != runtime_patches {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip abs_pos: non-square grids (src²={}, num_patches={}, tgt²={}, runtime_patches={})",
src_square,
num_patches,
tgt_square,
runtime_patches
)));
}
let out_len = checked_shape_mul("vision_clip abs_pos", seq, dim, "seq*dim")?;
let mut out = vec![0.0f32; out_len];
out[..dim].copy_from_slice(&table[..dim]);
if src == tgt {
out[dim..].copy_from_slice(&table[dim..]);
return Ok(Mat::from_vec(seq, dim, out));
}
let patch = &table[dim..]; interpolate_bicubic_into(patch, src, src, &mut out[dim..], tgt, tgt, dim);
Ok(Mat::from_vec(seq, dim, out))
}
fn interpolate_bicubic_into(
src: &[f32],
ih: usize,
iw: usize,
dst: &mut [f32],
oh: usize,
ow: usize,
dim: usize,
) {
let scale_h = ih as f32 / oh as f32;
let scale_w = iw as f32 / ow as f32;
for oy in 0..oh {
let sy = (oy as f32 + 0.5) * scale_h - 0.5;
let iy = sy.floor();
let fy = sy - iy;
let iy = iy as isize;
let wy = cubic_weights(fy);
for ox in 0..ow {
let sx = (ox as f32 + 0.5) * scale_w - 0.5;
let ix = sx.floor();
let fx = sx - ix;
let ix = ix as isize;
let wx = cubic_weights(fx);
for c in 0..dim {
let mut acc = 0.0f32;
for (m, &wym) in wy.iter().enumerate() {
let yy = clamp_index(iy - 1 + m as isize, ih);
for (n, &wxn) in wx.iter().enumerate() {
let xx = clamp_index(ix - 1 + n as isize, iw);
acc += wym * wxn * src[(yy * iw + xx) * dim + c];
}
}
dst[(oy * ow + ox) * dim + c] = acc;
}
}
}
}
fn cubic_weights(t: f32) -> [f32; 4] {
let a = -0.75f32;
let w0 = cubic_kernel(a, 1.0 + t);
let w1 = cubic_kernel(a, t);
let w2 = cubic_kernel(a, 1.0 - t);
let w3 = cubic_kernel(a, 2.0 - t);
[w0, w1, w2, w3]
}
fn cubic_kernel(a: f32, x: f32) -> f32 {
let x = x.abs();
if x <= 1.0 {
((a + 2.0) * x - (a + 3.0)) * x * x + 1.0
} else if x < 2.0 {
(((x - 5.0) * x + 8.0) * x - 4.0) * a
} else {
0.0
}
}
fn clamp_index(i: isize, n: usize) -> usize {
if i < 0 {
0
} else if i as usize >= n {
n - 1
} else {
i as usize
}
}
fn isqrt(v: usize) -> usize {
if v == 0 {
return 0;
}
let mut r = (v as f64).sqrt() as usize;
while r * r > v {
r -= 1;
}
while (r + 1) * (r + 1) <= v {
r += 1;
}
r
}
fn add(a: &Mat, b: &Mat) -> FocrResult<Mat> {
if a.rows != b.rows || a.cols != b.cols {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip add: shape mismatch [{},{}] vs [{},{}]",
a.rows,
a.cols,
b.rows,
b.cols
)));
}
ensure_mat_data_len(a, "vision_clip add lhs")?;
ensure_mat_data_len(b, "vision_clip add rhs")?;
let data = a
.data
.iter()
.zip(&b.data)
.map(|(x, y)| x + y)
.collect::<Vec<_>>();
Ok(Mat::from_vec(a.rows, a.cols, data))
}
fn add_in_place(a: &mut Mat, b: &Mat) -> FocrResult<()> {
if a.rows != b.rows || a.cols != b.cols {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_clip add_in_place: shape mismatch [{},{}] vs [{},{}]",
a.rows,
a.cols,
b.rows,
b.cols
)));
}
ensure_mat_data_len(a, "vision_clip add_in_place lhs")?;
ensure_mat_data_len(b, "vision_clip add_in_place rhs")?;
for (x, y) in a.data.iter_mut().zip(&b.data) {
*x += *y;
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::quant::focrq::{FocrqBuilder, WriteDType};
use half::bf16;
fn tiny_cfg() -> ClipConfig {
ClipConfig {
num_layers: 2,
hidden_size: 4,
num_heads: 2,
ffn_hidden_size: 8,
patch_size: 14,
layernorm_eps: 1e-5,
pre_layernorm_eps: 1e-5,
}
}
fn ln_identity(dim: usize) -> LayerNormParams {
LayerNormParams {
weight: vec![1.0; dim],
bias: vec![0.0; dim],
}
}
fn linear_identity(dim: usize) -> LinearParams {
let mut w = vec![0.0f32; dim * dim];
for i in 0..dim {
w[i * dim + i] = 1.0;
}
LinearParams::from_row_major(&w, Some(vec![0.0; dim]), dim, dim)
.expect("identity linear builds")
}
fn qkv_identity(dim: usize) -> LinearParams {
let mut w = vec![0.0f32; 3 * dim * dim];
for t in 0..3 {
for i in 0..dim {
let out_row = t * dim + i;
w[out_row * dim + i] = 1.0;
}
}
LinearParams::from_row_major(&w, Some(vec![0.0; 3 * dim]), 3 * dim, dim)
.expect("identity qkv builds")
}
fn fc1_identity_padded(dim: usize, ffn: usize) -> LinearParams {
let mut w = vec![0.0f32; ffn * dim];
for i in 0..dim {
w[i * dim + i] = 1.0;
}
LinearParams::from_row_major(&w, Some(vec![0.0; ffn]), ffn, dim)
.expect("identity fc1 builds")
}
fn fc2_identity_padded(dim: usize, ffn: usize) -> LinearParams {
let mut w = vec![0.0f32; dim * ffn];
for i in 0..dim {
w[i * ffn + i] = 1.0;
}
LinearParams::from_row_major(&w, Some(vec![0.0; dim]), dim, ffn)
.expect("identity fc2 builds")
}
fn block_identity(dim: usize, ffn: usize) -> ClipBlockWeights {
ClipBlockWeights {
layer_norm1: ln_identity(dim),
qkv_proj: qkv_identity(dim),
out_proj: linear_identity(dim),
layer_norm2: ln_identity(dim),
fc1: fc1_identity_padded(dim, ffn),
fc2: fc2_identity_padded(dim, ffn),
}
}
fn assert_err_contains<T>(res: FocrResult<T>, needle: &str) {
let message = match res {
Ok(_) => String::from("<ok>"),
Err(err) => err.to_string(),
};
assert!(
message.contains(needle),
"error {message:?} did not contain {needle:?}"
);
}
fn bf16_zeros(n: usize) -> Vec<u8> {
(0..n)
.flat_map(|_| bf16::from_f32(0.0).to_le_bytes())
.collect()
}
fn synth_values(len: usize, salt: u64) -> Vec<f32> {
(0..len)
.map(|i| {
let raw = ((i as u64)
.wrapping_mul(6364136223846793005)
.wrapping_add(salt)
>> 33) as u32;
(raw as f32 / u32::MAX as f32 - 0.5) * 0.6
})
.collect()
}
fn f32_bytes(values: &[f32]) -> Vec<u8> {
values.iter().flat_map(|v| v.to_le_bytes()).collect()
}
fn synth_clip_tower(cfg: &ClipConfig, num_patches: usize) -> Weights {
let dim = cfg.hidden_size;
let num_positions = num_patches + 1;
let p = "model.vision_model";
let mut b = FocrqBuilder::new();
let mut add = |name: String, shape: Vec<usize>, salt: u64| {
let len = shape.iter().product();
b.add_tensor(
name,
WriteDType::F32,
shape,
f32_bytes(&synth_values(len, salt)),
)
.expect("valid synthetic f32 tensor");
};
add(format!("{p}.embeddings.class_embedding"), vec![dim], 1);
add(
format!("{p}.embeddings.position_embedding.weight"),
vec![num_positions, dim],
2,
);
add(format!("{p}.pre_layrnorm.weight"), vec![dim], 3);
add(format!("{p}.pre_layrnorm.bias"), vec![dim], 4);
for l in 0..cfg.num_layers {
let base = format!("{p}.transformer.layers.{l}");
let salt = 100 * (l as u64 + 1);
add(format!("{base}.layer_norm1.weight"), vec![dim], salt + 1);
add(format!("{base}.layer_norm1.bias"), vec![dim], salt + 2);
add(
format!("{base}.self_attn.qkv_proj.weight"),
vec![3 * dim, dim],
salt + 3,
);
add(
format!("{base}.self_attn.qkv_proj.bias"),
vec![3 * dim],
salt + 4,
);
add(
format!("{base}.self_attn.out_proj.weight"),
vec![dim, dim],
salt + 5,
);
add(
format!("{base}.self_attn.out_proj.bias"),
vec![dim],
salt + 6,
);
add(format!("{base}.layer_norm2.weight"), vec![dim], salt + 7);
add(format!("{base}.layer_norm2.bias"), vec![dim], salt + 8);
add(
format!("{base}.mlp.fc1.weight"),
vec![cfg.ffn_hidden_size, dim],
salt + 9,
);
add(
format!("{base}.mlp.fc1.bias"),
vec![cfg.ffn_hidden_size],
salt + 10,
);
add(
format!("{base}.mlp.fc2.weight"),
vec![dim, cfg.ffn_hidden_size],
salt + 11,
);
add(format!("{base}.mlp.fc2.bias"), vec![dim], salt + 12);
}
Weights::from_bytes(b.build()).expect("synthetic tower parses")
}
#[test]
fn streamed_clip_forward_is_bit_identical_to_cached() {
let cfg = tiny_cfg();
let dim = cfg.hidden_size;
let num_patches = 4usize;
let weights = synth_clip_tower(&cfg, num_patches);
let sam = Mat::from_vec(dim, num_patches, synth_values(dim * num_patches, 999));
let cached_tower = clip_weights_from_with(&cfg, &weights).expect("cached tower hydrates");
let cached = forward_from_sam(&cfg, &cached_tower, &sam).expect("cached forward runs");
let streamed =
forward_from_sam_streamed(&cfg, &weights, &sam).expect("streamed forward runs");
assert_eq!(cached.shape(), streamed.shape());
assert_eq!(
cached
.data
.iter()
.map(|f| f.to_bits())
.collect::<Vec<u32>>(),
streamed
.data
.iter()
.map(|f| f.to_bits())
.collect::<Vec<u32>>(),
"streamed per-block CLIP must be bit-identical to the cached tower"
);
assert!(cached.data.iter().any(|&v| v != 0.0));
}
#[test]
fn streamed_clip_views_inner_is_bit_identical_to_views_outer() {
let cfg = tiny_cfg();
let dim = cfg.hidden_size;
let num_patches = 4usize;
let weights = synth_clip_tower(&cfg, num_patches);
let views: Vec<Mat> = (0..3)
.map(|v| {
Mat::from_vec(
dim,
num_patches,
synth_values(
dim * num_patches,
(v as u64 + 1).wrapping_mul(0x9E37_79B9_7F4A_7C15),
),
)
})
.collect();
let per_view: Vec<Mat> = views
.iter()
.map(|sam| {
forward_from_sam_streamed(&cfg, &weights, sam).expect("per-view streamed runs")
})
.collect();
let refs: Vec<&Mat> = views.iter().collect();
let hoisted = forward_from_sam_streamed_views(&cfg, &weights, &refs)
.expect("views-inner streamed runs");
assert_eq!(hoisted.len(), per_view.len());
for (v, (a, b)) in per_view.iter().zip(hoisted.iter()).enumerate() {
assert_eq!(a.shape(), b.shape(), "view {v} shape");
assert_eq!(
a.data.iter().map(|f| f.to_bits()).collect::<Vec<u32>>(),
b.data.iter().map(|f| f.to_bits()).collect::<Vec<u32>>(),
"view {v}: hydration-hoisted CLIP must be bit-identical to per-view"
);
}
assert!(per_view[0].data.iter().any(|&x| x != 0.0));
assert_ne!(views[0].data, views[1].data);
assert_ne!(per_view[0].data, per_view[1].data);
}
#[test]
fn clip_weights_from_rejects_rank1_linear_weight_without_panic() {
let p = "model.vision_model";
let layer0 = format!("{p}.transformer.layers.0");
let mut b = FocrqBuilder::new();
b.add_tensor(
format!("{p}.embeddings.position_embedding.weight"),
WriteDType::Bf16,
vec![1, 1],
bf16_zeros(1),
)
.unwrap();
b.add_tensor(
format!("{layer0}.layer_norm1.weight"),
WriteDType::Bf16,
vec![1],
bf16_zeros(1),
)
.unwrap();
b.add_tensor(
format!("{layer0}.layer_norm1.bias"),
WriteDType::Bf16,
vec![1],
bf16_zeros(1),
)
.unwrap();
b.add_tensor(
format!("{layer0}.self_attn.qkv_proj.weight"),
WriteDType::Bf16,
vec![4],
bf16_zeros(4),
)
.unwrap();
let weights = Weights::from_bytes(b.build()).unwrap();
assert_err_contains(clip_weights_from(&weights), "rank 1");
}
#[test]
fn clip_weights_from_rejects_rank1_position_embedding_without_panic() {
let p = "model.vision_model";
let mut b = FocrqBuilder::new();
b.add_tensor(
format!("{p}.embeddings.position_embedding.weight"),
WriteDType::Bf16,
vec![4],
bf16_zeros(4),
)
.unwrap();
let weights = Weights::from_bytes(b.build()).unwrap();
assert_err_contains(clip_weights_from(&weights), "rank 1");
}
#[test]
fn transpose_roundtrips() -> FocrResult<()> {
let src = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let t = transpose(&src, 2, 3)?;
assert_eq!(t, vec![1.0, 4.0, 2.0, 5.0, 3.0, 6.0]);
assert_eq!(transpose(&t, 3, 2)?, src);
Ok(())
}
#[test]
fn transpose_rejects_malformed_len_without_panic() {
assert_err_contains(transpose(&[1.0], 2, 2), "source len");
}
#[test]
fn transpose_rejects_shape_product_overflow_without_panic() {
assert_err_contains(transpose(&[], usize::MAX, 2), "rows*cols");
}
#[test]
fn linear_applies_weight_and_bias() -> FocrResult<()> {
let w = LinearParams::from_row_major(
&[1.0, 0.0, 0.0, 0.0, 1.0, 0.0],
Some(vec![10.0, 20.0]),
2,
3,
)?;
let x = Mat::from_vec(1, 3, vec![1.0, 2.0, 3.0]);
let y = linear(&w, &x)?;
assert_eq!(y.shape(), (1, 2));
assert!((y.data[0] - 11.0).abs() < 1e-6);
assert!((y.data[1] - 22.0).abs() < 1e-6);
Ok(())
}
#[test]
fn linear_rejects_dim_mismatch() {
let w = linear_identity(4);
let x = Mat::zeros(2, 3); assert!(linear(&w, &x).is_err());
}
#[test]
fn from_row_major_stores_the_exact_transpose() {
let (out, inn) = (3usize, 2usize);
let w: Vec<f32> = (0..out * inn).map(|v| v as f32 * 0.5 - 1.0).collect();
let p = LinearParams::from_row_major(&w, None, out, inn).expect("builds");
assert_eq!(p.weight_t.shape(), (inn, out));
for o in 0..out {
for i in 0..inn {
assert_eq!(
p.weight_t.data[i * out + o].to_bits(),
w[o * inn + i].to_bits(),
"weight_t[{i}][{o}] must be bit-equal to weight[{o}][{i}]"
);
}
}
}
#[test]
fn from_row_major_rejects_shape_product_overflow_without_panic() {
assert_err_contains(
LinearParams::from_row_major(&[], None, usize::MAX, 2),
"rows*cols",
);
}
#[test]
fn linear_rejects_mismatched_pretransposed_weight_shape() {
let w = LinearParams {
weight_t: Mat::zeros(1, 1),
bias: None,
out_features: 3,
in_features: 2,
};
let x = Mat::zeros(1, 2);
assert_err_contains(linear(&w, &x), "weight_t shape");
}
#[test]
fn linear_rejects_malformed_input_mat_without_panic() {
let w = LinearParams::from_row_major(&[1.0, 1.0], None, 1, 2).expect("tiny linear builds");
let x = Mat {
rows: 2,
cols: 2,
data: vec![1.0, 2.0],
};
assert_err_contains(linear(&w, &x), "data len");
}
#[test]
fn class_token_prepended_as_row0() -> FocrResult<()> {
let patches = Mat::from_vec(2, 3, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
let out = prepend_class_token(&[7.0, 8.0, 9.0], &patches)?;
assert_eq!(out.shape(), (3, 3));
assert_eq!(out.row(0), &[7.0, 8.0, 9.0]);
assert_eq!(out.row(1), &[1.0, 2.0, 3.0]);
assert_eq!(out.row(2), &[4.0, 5.0, 6.0]);
Ok(())
}
#[test]
fn prepend_class_token_rejects_bad_class_len_without_panic() {
let patches = Mat::zeros(2, 3);
assert_err_contains(
prepend_class_token(&[1.0, 2.0], &patches),
"class_embedding len",
);
}
#[test]
fn prepend_class_token_rejects_malformed_patch_data_without_panic() {
let patches = Mat {
rows: 2,
cols: 3,
data: vec![1.0, 2.0, 3.0],
};
assert_err_contains(prepend_class_token(&[7.0, 8.0, 9.0], &patches), "data len");
}
#[test]
fn prepend_class_token_rejects_seq_overflow_without_panic() {
let patches = Mat {
rows: usize::MAX,
cols: 0,
data: Vec::new(),
};
assert_err_contains(prepend_class_token(&[], &patches), "patch_rows+1");
}
#[test]
fn abs_pos_passthrough_when_lengths_match() -> FocrResult<()> {
let table: Vec<f32> = (0..10).map(|i| i as f32).collect();
let pos = abs_pos_for_len(&table, 5, 2, 5)?;
assert_eq!(pos.shape(), (5, 2));
assert_eq!(pos.data, table);
Ok(())
}
#[test]
fn abs_pos_same_grid_size_is_identity_even_via_resize_branch() -> FocrResult<()> {
let np = 10; let dim = 1;
let mut table = vec![0.0f32; np * dim];
table[0] = 99.0; for (i, slot) in table.iter_mut().enumerate().skip(1) {
*slot = i as f32;
}
let pos = abs_pos_for_len(&table, np, dim, 5)?;
assert_eq!(pos.shape(), (5, 1));
assert!((pos.data[0] - 99.0).abs() < 1e-6);
for &v in &pos.data[1..] {
assert!(
(0.5..=10.0).contains(&v),
"interp out of expected range: {v}"
);
}
Ok(())
}
#[test]
fn abs_pos_rejects_non_square_grid() {
let table = vec![0.0f32; 4 * 2];
assert!(abs_pos_for_len(&table, 4, 2, 5).is_err());
}
#[test]
fn abs_pos_rejects_shape_product_overflow_without_panic() {
assert_err_contains(abs_pos_for_len(&[], usize::MAX, 2, 1), "num_positions*dim");
}
#[test]
fn abs_pos_rejects_zero_lengths_without_underflow() {
assert_err_contains(abs_pos_for_len(&[], 0, 2, 1), "non-zero");
let table = vec![0.0f32; 2];
assert_err_contains(abs_pos_for_len(&table, 1, 2, 0), "non-zero");
}
#[test]
fn cubic_weights_partition_of_unity() {
for &t in &[0.0f32, 0.25, 0.5, 0.75, 0.999] {
let w = cubic_weights(t);
let s: f32 = w.iter().sum();
assert!((s - 1.0).abs() < 1e-5, "weights sum {s} != 1 at t={t}");
}
}
#[test]
fn cubic_weights_at_zero_is_unit_impulse() {
let w = cubic_weights(0.0);
assert!((w[0]).abs() < 1e-6);
assert!((w[1] - 1.0).abs() < 1e-6);
assert!((w[2]).abs() < 1e-6);
assert!((w[3]).abs() < 1e-6);
}
#[test]
fn self_attention_preserves_shape() -> FocrResult<()> {
let cfg = tiny_cfg();
let block = block_identity(cfg.hidden_size, cfg.ffn_hidden_size);
let x = Mat::from_vec(
3,
4,
vec![
0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0, 1.1, 1.2,
],
);
let out = self_attention(&cfg, &block, &x)?;
assert_eq!(out.shape(), (3, 4));
Ok(())
}
#[test]
fn self_attention_rejects_bad_head_config_without_panic() {
let mut cfg = tiny_cfg();
let block = block_identity(cfg.hidden_size, cfg.ffn_hidden_size);
let x = Mat::zeros(2, cfg.hidden_size);
cfg.num_heads = 0;
assert_err_contains(self_attention(&cfg, &block, &x), "num_heads");
cfg.num_heads = 3;
assert_err_contains(self_attention(&cfg, &block, &x), "divisible");
}
#[test]
fn self_attention_rejects_malformed_qkv_and_projection_shapes() {
let cfg = tiny_cfg();
let dim = cfg.hidden_size;
let x = Mat::zeros(3, dim);
let mut bad_qkv = block_identity(dim, cfg.ffn_hidden_size);
let wrong_out = 3 * dim - 1;
bad_qkv.qkv_proj =
LinearParams::from_row_major(&vec![0.0; wrong_out * dim], None, wrong_out, dim)
.expect("wrong-shaped qkv builds");
assert_err_contains(
self_attention(&cfg, &bad_qkv, &x),
"self_attention qkv output",
);
let mut bad_proj = block_identity(dim, cfg.ffn_hidden_size);
let wrong_out = dim - 1;
bad_proj.out_proj =
LinearParams::from_row_major(&vec![0.0; wrong_out * dim], None, wrong_out, dim)
.expect("wrong-shaped out_proj builds");
assert_err_contains(
self_attention(&cfg, &bad_proj, &x),
"self_attention projection output",
);
}
#[test]
fn self_attention_rejects_malformed_input_mat_without_panic() {
let cfg = tiny_cfg();
let block = block_identity(cfg.hidden_size, cfg.ffn_hidden_size);
let x = Mat {
rows: 2,
cols: cfg.hidden_size,
data: vec![0.0; cfg.hidden_size],
};
assert_err_contains(self_attention(&cfg, &block, &x), "data len");
}
#[test]
fn self_attention_is_convex_average_of_values() -> FocrResult<()> {
let cfg = tiny_cfg();
let block = block_identity(cfg.hidden_size, cfg.ffn_hidden_size);
let x = Mat::from_vec(
3,
4,
vec![
0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0,
],
);
let out = self_attention(&cfg, &block, &x)?;
for c in 0..4 {
let col_min = (0..3).map(|r| x.get(r, c)).fold(f32::INFINITY, f32::min);
let col_max = (0..3)
.map(|r| x.get(r, c))
.fold(f32::NEG_INFINITY, f32::max);
for r in 0..3 {
let v = out.get(r, c);
assert!(
v >= col_min - 1e-4 && v <= col_max + 1e-4,
"out[{r},{c}]={v} outside [{col_min},{col_max}]"
);
}
}
Ok(())
}
#[test]
fn feed_forward_applies_quick_gelu() -> FocrResult<()> {
let cfg = tiny_cfg();
let block = block_identity(cfg.hidden_size, cfg.ffn_hidden_size);
let x = Mat::from_vec(1, 4, vec![-1.0, 0.0, 1.0, 2.0]);
let out = feed_forward(&block, &x)?;
for (i, &xi) in x.data.iter().enumerate() {
let want = nn::quick_gelu_scalar(xi);
assert!(
(out.data[i] - want).abs() < 1e-5,
"ff[{i}]={} != quick_gelu({xi})={want}",
out.data[i]
);
}
Ok(())
}
#[test]
fn add_rejects_malformed_backing_data_without_panic() {
let malformed_lhs = Mat {
rows: 2,
cols: 2,
data: vec![1.0, 2.0],
};
let good = Mat::zeros(2, 2);
assert_err_contains(add(&malformed_lhs, &good), "add lhs");
let malformed_rhs = Mat {
rows: 2,
cols: 2,
data: vec![1.0, 2.0],
};
assert_err_contains(add(&good, &malformed_rhs), "add rhs");
}
#[test]
fn add_in_place_rejects_malformed_backing_data_before_mutating() {
let mut malformed_lhs = Mat {
rows: 2,
cols: 2,
data: vec![1.0, 2.0],
};
let good = Mat::zeros(2, 2);
let lhs_before = malformed_lhs.data.clone();
assert_err_contains(add_in_place(&mut malformed_lhs, &good), "add_in_place lhs");
assert_eq!(malformed_lhs.data, lhs_before);
let mut good_lhs = Mat::from_vec(2, 2, vec![1.0, 2.0, 3.0, 4.0]);
let malformed_rhs = Mat {
rows: 2,
cols: 2,
data: vec![10.0, 20.0],
};
let good_before = good_lhs.data.clone();
assert_err_contains(
add_in_place(&mut good_lhs, &malformed_rhs),
"add_in_place rhs",
);
assert_eq!(good_lhs.data, good_before);
}
fn tiny_weights(cfg: &ClipConfig, num_patches: usize) -> ClipWeights {
let dim = cfg.hidden_size;
let num_positions = num_patches + 1;
ClipWeights {
class_embedding: vec![0.0; dim],
position_embedding: vec![0.0; num_positions * dim],
num_positions,
pre_layernorm: ln_identity(dim),
blocks: (0..cfg.num_layers)
.map(|_| block_identity(dim, cfg.ffn_hidden_size))
.collect(),
}
}
fn det(i: usize, salt: usize) -> f32 {
let x = (i as u64)
.wrapping_mul(2_654_435_761)
.wrapping_add((salt as u64).wrapping_mul(40_503))
% 1000;
(x as f32) / 1000.0 - 0.5
}
fn rand_lin(out: usize, inn: usize, salt: usize) -> LinearParams {
let w: Vec<f32> = (0..out * inn).map(|i| det(i, salt)).collect();
LinearParams::from_row_major(
&w,
Some((0..out).map(|i| det(i, salt + 1)).collect()),
out,
inn,
)
.expect("rand linear builds")
}
fn rand_ln(dim: usize, salt: usize) -> LayerNormParams {
LayerNormParams {
weight: (0..dim).map(|i| 1.0 + det(i, salt)).collect(),
bias: (0..dim).map(|i| det(i, salt + 2)).collect(),
}
}
fn rand_weights(cfg: &ClipConfig, num_patches: usize) -> ClipWeights {
let dim = cfg.hidden_size;
let np = num_patches + 1;
ClipWeights {
class_embedding: (0..dim).map(|i| det(i, 11)).collect(),
position_embedding: (0..np * dim).map(|i| det(i, 13)).collect(),
num_positions: np,
pre_layernorm: rand_ln(dim, 17),
blocks: (0..cfg.num_layers)
.map(|l| ClipBlockWeights {
layer_norm1: rand_ln(dim, 100 + l * 10),
qkv_proj: rand_lin(3 * dim, dim, 200 + l * 10),
out_proj: rand_lin(dim, dim, 300 + l * 10),
layer_norm2: rand_ln(dim, 400 + l * 10),
fc1: rand_lin(cfg.ffn_hidden_size, dim, 500 + l * 10),
fc2: rand_lin(dim, cfg.ffn_hidden_size, 600 + l * 10),
})
.collect(),
}
}
fn parity_cfg() -> ClipConfig {
ClipConfig {
num_layers: 3,
hidden_size: 8,
num_heads: 2,
ffn_hidden_size: 16,
patch_size: 14,
layernorm_eps: 1e-5,
pre_layernorm_eps: 1e-5,
}
}
fn parity_view(num_patches: usize, dim: usize, salt: usize) -> Mat {
Mat::from_vec(
num_patches,
dim,
(0..num_patches * dim)
.map(|i| (((i + salt * 7) as f32) * 0.013).sin())
.collect(),
)
}
#[test]
fn batched_clip_equals_per_view_byte_for_byte() -> FocrResult<()> {
let cfg = parity_cfg();
let num_patches = 5;
let w = rand_weights(&cfg, num_patches);
let views: Vec<Mat> = (0..4)
.map(|vi| parity_view(num_patches, cfg.hidden_size, vi))
.collect();
let refs: Vec<&Mat> = views.iter().collect();
let batched = forward_with_batched(&cfg, &w, &refs)?;
assert_eq!(batched.len(), views.len());
for (vi, view) in views.iter().enumerate() {
let seq_out = forward_with(&cfg, &w, view)?;
assert_eq!(batched[vi].shape(), seq_out.shape(), "view {vi} shape");
assert_eq!(
batched[vi].data, seq_out.data,
"view {vi}: batched CLIP != per-view sequential (cross-view leak or M-dependence)"
);
}
Ok(())
}
#[test]
fn batched_clip_single_view_equals_forward_with() -> FocrResult<()> {
let cfg = parity_cfg();
let num_patches = 6;
let w = rand_weights(&cfg, num_patches);
let view = parity_view(num_patches, cfg.hidden_size, 3);
let batched = forward_with_batched(&cfg, &w, &[&view])?;
let seq_out = forward_with(&cfg, &w, &view)?;
assert_eq!(batched.len(), 1);
assert_eq!(batched[0].data, seq_out.data);
Ok(())
}
#[test]
fn batched_from_sam_equals_per_view_forward_path() -> FocrResult<()> {
let cfg = parity_cfg();
let num_patches = 5;
let w = rand_weights(&cfg, num_patches);
let sams: Vec<Mat> = (0..3)
.map(|vi| parity_view(cfg.hidden_size, num_patches, vi))
.collect();
let refs: Vec<&Mat> = sams.iter().collect();
let batched = forward_batched_from_sam(&cfg, &w, &refs)?;
assert_eq!(batched.len(), sams.len());
for (vi, sam) in sams.iter().enumerate() {
let t = transpose(&sam.data, sam.rows, sam.cols)?;
let per_view = forward_with(&cfg, &w, &Mat::from_vec(sam.cols, sam.rows, t))?;
assert_eq!(
batched[vi].data, per_view.data,
"view {vi}: from-sam batched CLIP != per-view transpose+forward_with"
);
}
Ok(())
}
#[test]
fn batched_clip_rejects_ragged_and_empty() {
let cfg = parity_cfg();
let w = rand_weights(&cfg, 4);
assert!(forward_with_batched(&cfg, &w, &[]).is_err());
let a = parity_view(4, cfg.hidden_size, 1);
let b = parity_view(5, cfg.hidden_size, 2); assert!(forward_with_batched(&cfg, &w, &[&a, &b]).is_err());
}
#[test]
fn forward_with_produces_class_plus_patches_rows() -> FocrResult<()> {
let cfg = tiny_cfg();
let num_patches = 4; let w = tiny_weights(&cfg, num_patches);
let sam = Mat::from_vec(
num_patches,
cfg.hidden_size,
(0..num_patches * cfg.hidden_size)
.map(|i| (i as f32) * 0.01)
.collect(),
);
let out = forward_with(&cfg, &w, &sam)?;
assert_eq!(out.shape(), (num_patches + 1, cfg.hidden_size));
assert!(out.data.iter().all(|v| v.is_finite()));
Ok(())
}
#[test]
fn forward_with_rejects_wrong_width() {
let cfg = tiny_cfg();
let w = tiny_weights(&cfg, 4);
let sam = Mat::zeros(4, cfg.hidden_size + 1); assert!(forward_with(&cfg, &w, &sam).is_err());
}
#[test]
fn forward_with_rejects_wrong_block_count() {
let cfg = tiny_cfg();
let mut w = tiny_weights(&cfg, 4);
w.blocks.pop(); let sam = Mat::zeros(4, cfg.hidden_size);
assert!(forward_with(&cfg, &w, &sam).is_err());
}
#[test]
fn forward_rejects_malformed_sam_features_before_weight_hydration() -> FocrResult<()> {
let weights = Weights::from_bytes(FocrqBuilder::new().build())?;
let image = Mat::zeros(1, 1);
let sam = Mat {
rows: 2,
cols: 2,
data: vec![0.0],
};
assert_err_contains(forward(&weights, &image, &sam), "sam_features");
Ok(())
}
#[test]
fn forward_with_rejects_malformed_sam_features_without_panic() {
let cfg = tiny_cfg();
let w = tiny_weights(&cfg, 2);
let sam = Mat {
rows: 2,
cols: cfg.hidden_size,
data: vec![0.0; cfg.hidden_size],
};
assert_err_contains(forward_with(&cfg, &w, &sam), "data len");
}
#[test]
fn zeroed_sublayers_make_block_identity() -> FocrResult<()> {
let cfg = tiny_cfg();
let dim = cfg.hidden_size;
let mut block = block_identity(dim, cfg.ffn_hidden_size);
block
.out_proj
.weight_t
.data
.iter_mut()
.for_each(|v| *v = 0.0);
block.out_proj.bias = Some(vec![0.0; dim]);
block.fc2.weight_t.data.iter_mut().for_each(|v| *v = 0.0);
block.fc2.bias = Some(vec![0.0; dim]);
let x = Mat::from_vec(2, 4, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0]);
let out = transformer_block(&cfg, &block, &x)?;
for (o, i) in out.data.iter().zip(&x.data) {
assert!((o - i).abs() < 1e-5, "block not identity: {o} vs {i}");
}
Ok(())
}
#[test]
fn isqrt_is_floor_sqrt() {
assert_eq!(isqrt(0), 0);
assert_eq!(isqrt(1), 1);
assert_eq!(isqrt(3), 1);
assert_eq!(isqrt(4), 2);
assert_eq!(isqrt(255), 15);
assert_eq!(isqrt(256), 16);
assert_eq!(isqrt(257), 16);
}
}