use rayon::prelude::*;
use super::nn;
use super::tensor::Mat;
use super::weights::Weights;
use crate::error::{FocrError, FocrResult};
#[cfg(not(target_arch = "wasm32"))]
use std::time::Instant;
#[cfg(target_arch = "wasm32")]
use web_time::Instant;
pub const EMBED_DIM: usize = 768;
pub const DEPTH: usize = 12;
pub const NUM_HEADS: usize = 12;
pub const HEAD_DIM: usize = EMBED_DIM / NUM_HEADS;
pub const PATCH: usize = 16;
pub const WINDOW: usize = 14;
pub const GLOBAL_BLOCKS: [usize; 4] = [2, 5, 8, 11];
pub const NECK_CH: usize = 256;
pub const NET2_CH: usize = 512;
pub const OUT_CH: usize = 1024;
pub const LN_EPS: f32 = 1e-6;
pub const MLP_HIDDEN: usize = EMBED_DIM * 4;
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_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 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_nchw_len(
context: &str,
ch: usize,
h: usize,
w: usize,
expression: &str,
) -> FocrResult<usize> {
let hw = checked_shape_mul(context, h, w, "h*w")?;
checked_shape_mul(context, ch, hw, expression)
}
fn checked_conv_weight_len(
context: &str,
out_ch: usize,
in_ch: usize,
kh: usize,
kw: usize,
) -> FocrResult<usize> {
let out_in = checked_shape_mul(context, out_ch, in_ch, "out_ch*in_ch")?;
let kernel = checked_shape_mul(context, kh, kw, "kh*kw")?;
checked_shape_mul(context, out_in, kernel, "out_ch*in_ch*kh*kw")
}
#[derive(Debug, Clone)]
pub struct Linear {
wt: Mat,
pub b: Vec<f32>,
pub out: usize,
pub in_: usize,
}
impl Linear {
pub fn from_row_major(w: &[f32], b: Vec<f32>, out: usize, in_: usize) -> FocrResult<Self> {
let expected_weight_len = checked_shape_mul("vision_sam linear", out, in_, "out*in")?;
if w.len() != expected_weight_len {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam linear: weight len {} != out*in {}",
w.len(),
expected_weight_len
)));
}
if !b.is_empty() && b.len() != out {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam linear: bias len {} != out_features {}",
b.len(),
out
)));
}
let wt = Mat::from_vec(in_, out, transpose(w, out, in_));
Ok(Self { wt, b, out, in_ })
}
pub fn apply(&self, x: &Mat) -> FocrResult<Mat> {
if x.cols != self.in_ {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam linear: input cols {} != expected in_features {}",
x.cols,
self.in_
)));
}
if self.wt.rows != self.in_ || self.wt.cols != self.out {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam linear: pretransposed weight shape [{}, {}] != [in,out] [{}, {}]",
self.wt.rows,
self.wt.cols,
self.in_,
self.out
)));
}
if !self.b.is_empty() && self.b.len() != self.out {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam linear: bias len {} != out_features {}",
self.b.len(),
self.out
)));
}
let mut y = nn::matmul(x, &self.wt)?;
if !self.b.is_empty() {
for r in 0..y.rows {
let row = y.row_mut(r);
for (c, v) in row.iter_mut().enumerate() {
*v += self.b[c];
}
}
}
Ok(y)
}
}
#[derive(Debug, Clone)]
pub struct LayerNormP {
pub w: Vec<f32>,
pub b: Vec<f32>,
}
#[derive(Debug, Clone)]
pub struct Conv {
pub w: Vec<f32>,
pub b: Option<Vec<f32>>,
pub out_ch: usize,
pub in_ch: usize,
pub kh: usize,
pub kw: usize,
}
#[derive(Debug, Clone)]
pub struct AttnP {
pub qkv: Linear,
pub proj: Linear,
pub rel_pos_h: Vec<f32>,
pub rel_pos_w: Vec<f32>,
pub size_h: usize,
pub size_w: usize,
}
#[derive(Debug, Clone)]
pub struct BlockP {
pub norm1: LayerNormP,
pub attn: AttnP,
pub norm2: LayerNormP,
pub lin1: Linear,
pub lin2: Linear,
pub window: usize,
}
#[derive(Debug, Clone)]
pub struct SamWeights {
pub patch_embed: Conv,
pub pos_embed: Vec<f32>,
pub pos_grid_h: usize,
pub pos_grid_w: usize,
pub blocks: Vec<BlockP>,
pub neck_conv1: Conv,
pub neck_ln1: LayerNormP,
pub neck_conv2: Conv,
pub neck_ln2: LayerNormP,
pub net2: Conv,
pub net3: Conv,
}
pub fn forward(weights: &Weights, image: &Mat) -> FocrResult<Mat> {
forward_prefix(weights, image, "model.sam_model")
}
pub fn forward_prefix(weights: &Weights, image: &Mat, prefix: &str) -> FocrResult<Mat> {
if image.rows != 3 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam::forward: expected 3 input channels, got {}",
image.rows
)));
}
let side = (image.cols as f64).sqrt() as usize;
if side * side != image.cols {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam::forward: image.cols {} is not a perfect square",
image.cols
)));
}
let th = Instant::now();
let w = sam_weights_from(weights, prefix)?;
super::timing_log(&format!(
" sam.hydrate {:.2}s",
th.elapsed().as_secs_f64()
));
let tf = Instant::now();
let out = forward_with(&w, image, side, side);
super::timing_log(&format!(
" sam.forward {:.2}s",
tf.elapsed().as_secs_f64()
));
out
}
pub(crate) fn sam_weights_from(weights: &Weights, prefix: &str) -> FocrResult<SamWeights> {
let (patch_embed, pos_embed, pos_grid_h, pos_grid_w) = sam_embed_from(weights, prefix)?;
let mut blocks = Vec::with_capacity(DEPTH);
for i in 0..DEPTH {
blocks.push(sam_block_from(weights, prefix, i)?);
}
let mut w = sam_neck_shell(
weights,
prefix,
patch_embed,
pos_embed,
pos_grid_h,
pos_grid_w,
)?;
w.blocks = blocks;
Ok(w)
}
pub(crate) fn sam_block_from(weights: &Weights, prefix: &str, i: usize) -> FocrResult<BlockP> {
let flat = |n: &str| weights.vec(n);
let ln = |n: &str| -> FocrResult<LayerNormP> {
Ok(LayerNormP {
w: flat(&format!("{n}.weight"))?,
b: flat(&format!("{n}.bias"))?,
})
};
let b = format!("{prefix}.blocks.{i}");
let window = if GLOBAL_BLOCKS.contains(&i) {
0
} else {
WINDOW
};
let qkv_name = format!("{b}.attn.qkv.weight");
let proj_name = format!("{b}.attn.proj.weight");
let rph_name = format!("{b}.attn.rel_pos_h");
let rpw_name = format!("{b}.attn.rel_pos_w");
let lin1_name = format!("{b}.mlp.lin1.weight");
let lin2_name = format!("{b}.mlp.lin2.weight");
let (q_out, q_in) = tensor_rank2_shape(weights, &qkv_name)?;
let (proj_out, proj_in) = tensor_rank2_shape(weights, &proj_name)?;
let (rph_rows, _rph_cols) = tensor_rank2_shape(weights, &rph_name)?;
let (rpw_rows, _rpw_cols) = tensor_rank2_shape(weights, &rpw_name)?;
let (lin1_out, lin1_in) = tensor_rank2_shape(weights, &lin1_name)?;
let (lin2_out, lin2_in) = tensor_rank2_shape(weights, &lin2_name)?;
Ok(BlockP {
norm1: ln(&format!("{b}.norm1"))?,
attn: AttnP {
qkv: Linear::from_row_major(
&flat(&qkv_name)?,
flat(&format!("{b}.attn.qkv.bias"))?,
q_out,
q_in,
)?,
proj: Linear::from_row_major(
&flat(&proj_name)?,
flat(&format!("{b}.attn.proj.bias"))?,
proj_out,
proj_in,
)?,
rel_pos_h: flat(&rph_name)?,
rel_pos_w: flat(&rpw_name)?,
size_h: rph_rows.div_ceil(2),
size_w: rpw_rows.div_ceil(2),
},
norm2: ln(&format!("{b}.norm2"))?,
lin1: Linear::from_row_major(
&flat(&lin1_name)?,
flat(&format!("{b}.mlp.lin1.bias"))?,
lin1_out,
lin1_in,
)?,
lin2: Linear::from_row_major(
&flat(&lin2_name)?,
flat(&format!("{b}.mlp.lin2.bias"))?,
lin2_out,
lin2_in,
)?,
window,
})
}
fn sam_embed_from(weights: &Weights, prefix: &str) -> FocrResult<(Conv, Vec<f32>, usize, usize)> {
let p = prefix;
let flat = |n: &str| weights.vec(n);
let pe = format!("{p}.patch_embed.proj.weight");
let (pe_out, pe_in, pe_kh, pe_kw) = tensor_rank4_shape(weights, &pe)?;
let patch_embed = Conv {
w: flat(&pe)?,
b: Some(flat(&format!("{p}.patch_embed.proj.bias"))?),
out_ch: pe_out,
in_ch: pe_in,
kh: pe_kh,
kw: pe_kw,
};
let pos_name = format!("{p}.pos_embed");
let (pgh, pgw) = tensor_pos_grid_shape(weights, &pos_name)?;
let pos_embed = flat(&pos_name)?;
Ok((patch_embed, pos_embed, pgh, pgw))
}
fn sam_neck_shell(
weights: &Weights,
prefix: &str,
patch_embed: Conv,
pos_embed: Vec<f32>,
pos_grid_h: usize,
pos_grid_w: usize,
) -> FocrResult<SamWeights> {
let p = prefix;
let flat = |n: &str| weights.vec(n);
let conv = |n: &str, bias: bool| -> FocrResult<Conv> {
let d = tensor_min_rank_shape(weights, n, 2)?;
Ok(Conv {
w: flat(n)?,
b: if bias {
Some(flat(&n.replace(".weight", ".bias"))?)
} else {
None
},
out_ch: d[0],
in_ch: d[1],
kh: d.get(2).copied().unwrap_or(1),
kw: d.get(3).copied().unwrap_or(1),
})
};
let ln = |n: &str| -> FocrResult<LayerNormP> {
Ok(LayerNormP {
w: flat(&format!("{n}.weight"))?,
b: flat(&format!("{n}.bias"))?,
})
};
Ok(SamWeights {
patch_embed,
pos_embed,
pos_grid_h,
pos_grid_w,
blocks: Vec::new(),
neck_conv1: conv(&format!("{p}.neck.0.weight"), false)?,
neck_ln1: ln(&format!("{p}.neck.1"))?,
neck_conv2: conv(&format!("{p}.neck.2.weight"), false)?,
neck_ln2: ln(&format!("{p}.neck.3"))?,
net2: conv(&format!("{p}.net_2.weight"), false)?,
net3: conv(&format!("{p}.net_3.weight"), false)?,
})
}
pub(crate) fn sam_head_from(weights: &Weights, prefix: &str) -> FocrResult<SamWeights> {
let (patch_embed, pos_embed, pgh, pgw) = sam_embed_from(weights, prefix)?;
sam_neck_shell(weights, prefix, patch_embed, pos_embed, pgh, pgw)
}
fn tensor_min_rank_shape(weights: &Weights, name: &str, min_rank: usize) -> FocrResult<Vec<usize>> {
let view = weights.tensor(name)?;
if view.shape.len() < min_rank {
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?} has rank {}; expected at least {min_rank}",
view.shape.len()
)));
}
Ok(view.shape.to_vec())
}
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))
}
fn tensor_rank4_shape(weights: &Weights, name: &str) -> FocrResult<(usize, usize, usize, usize)> {
let view = weights.tensor(name)?;
let [out_ch, in_ch, kh, kw] = view.shape else {
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?} has rank {}; expected 4 ([out_ch, in_ch, kh, kw])",
view.shape.len()
)));
};
Ok((*out_ch, *in_ch, *kh, *kw))
}
fn tensor_pos_grid_shape(weights: &Weights, name: &str) -> FocrResult<(usize, usize)> {
let view = weights.tensor(name)?;
match view.shape {
[_h0, h, w, _c] => Ok((*h, *w)),
[h, w, _c] => Ok((*h, *w)),
shape => Err(FocrError::FormatMismatch(format!(
"tensor {name:?} has rank {}; expected 3 ([h, w, c]) or 4 ([1, h, w, c]), got {shape:?}",
shape.len()
))),
}
}
pub fn forward_with(w: &SamWeights, image: &Mat, h: usize, win: usize) -> FocrResult<Mat> {
forward_core(
w,
image,
h,
win,
w.blocks.len(),
&mut |i| Ok(std::borrow::Cow::Borrowed(&w.blocks[i])),
false,
)
}
pub(crate) fn forward_streamed(weights: &Weights, image: &Mat, prefix: &str) -> FocrResult<Mat> {
if image.rows != 3 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam::forward: expected 3 input channels, got {}",
image.rows
)));
}
let side = (image.cols as f64).sqrt() as usize;
if side * side != image.cols {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam::forward: image.cols {} is not a perfect square",
image.cols
)));
}
let th = Instant::now();
let head = sam_head_from(weights, prefix)?;
super::timing_log(&format!(
" sam.hydrate(head) {:.2}s",
th.elapsed().as_secs_f64()
));
let tf = Instant::now();
let out = forward_core(
&head,
image,
side,
side,
DEPTH,
&mut |i| Ok(std::borrow::Cow::Owned(sam_block_from(weights, prefix, i)?)),
true,
);
super::timing_log(&format!(
" sam.forward(streamed) {:.2}s",
tf.elapsed().as_secs_f64()
));
out
}
pub(crate) fn forward_streamed_views(
weights: &Weights,
images: &[&Mat],
prefix: &str,
) -> FocrResult<Vec<Mat>> {
if images.is_empty() {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam::forward_streamed_views: empty view list"
)));
}
let mut dims = Vec::with_capacity(images.len());
for image in images {
if image.rows != 3 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam::forward: expected 3 input channels, got {}",
image.rows
)));
}
let side = (image.cols as f64).sqrt() as usize;
if side * side != image.cols {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam::forward: image.cols {} is not a perfect square",
image.cols
)));
}
dims.push((side, side));
}
let th = Instant::now();
let head = sam_head_from(weights, prefix)?;
super::timing_log(&format!(
" sam.hydrate(head) {:.2}s",
th.elapsed().as_secs_f64()
));
let tf = Instant::now();
let out = forward_core_views(
&head,
images,
&dims,
DEPTH,
&mut |i| Ok(std::borrow::Cow::Owned(sam_block_from(weights, prefix, i)?)),
true,
);
super::timing_log(&format!(
" sam.forward(streamed, {} views) {:.2}s",
images.len(),
tf.elapsed().as_secs_f64()
));
out
}
#[allow(clippy::too_many_arguments)]
fn forward_core<'w>(
w: &SamWeights,
image: &Mat,
h: usize,
win: usize,
depth: usize,
block_at: &mut dyn FnMut(usize) -> FocrResult<std::borrow::Cow<'w, BlockP>>,
low_mem: bool,
) -> FocrResult<Mat> {
let mut out = forward_core_views(
w,
std::slice::from_ref(&image),
&[(h, win)],
depth,
block_at,
low_mem,
)?;
Ok(out.remove(0))
}
fn forward_core_views<'w>(
w: &SamWeights,
images: &[&Mat],
dims: &[(usize, usize)],
depth: usize,
block_at: &mut dyn FnMut(usize) -> FocrResult<std::borrow::Cow<'w, BlockP>>,
low_mem: bool,
) -> FocrResult<Vec<Mat>> {
if images.len() != dims.len() {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam: {} views but {} geometries",
images.len(),
dims.len()
)));
}
let dim = w.patch_embed.out_ch;
let mut grids = Vec::with_capacity(images.len());
let mut xs = Vec::with_capacity(images.len());
for (image, &(h, win)) in images.iter().zip(dims.iter()) {
if image.rows != 3 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam: expected 3 input channels, got {}",
image.rows
)));
}
if h == 0 || win == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam: spatial dims ({h},{win}) must be non-zero"
)));
}
let expected_cols = checked_shape_mul("vision_sam", h, win, "H*W")?;
if image.cols != expected_cols {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam: image.cols {} != H*W {}*{} ({})",
image.cols,
h,
win,
expected_cols
)));
}
if !h.is_multiple_of(PATCH) || !win.is_multiple_of(PATCH) {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam: spatial dims ({h},{win}) must be multiples of patch {PATCH}"
)));
}
let gh = h / PATCH;
let gw = win / PATCH;
let conv_out = conv_apply(&w.patch_embed, &image.data, h, win, 0, PATCH)?;
ensure_flat_len("vision_sam patch_embed output", &conv_out, dim, gh, gw)?;
let mut x = nchw_to_nhwc_rows(&conv_out, dim, gh, gw);
let pos = abs_pos(&w.pos_embed, w.pos_grid_h, w.pos_grid_w, dim, gh, gw)?;
for (xv, pv) in x.data.iter_mut().zip(pos.iter()) {
*xv += *pv;
}
grids.push((gh, gw));
xs.push(x);
}
let tb = Instant::now();
for i in 0..depth {
let blk = block_at(i)?;
for (x, &(gh, gw)) in xs.iter_mut().zip(grids.iter()) {
*x = block_forward_impl(&blk, x, gh, gw, low_mem)?;
super::progress::vision_step();
}
}
super::timing_log(&format!(
" sam.blocks {:.2}s",
tb.elapsed().as_secs_f64()
));
let mut out = Vec::with_capacity(xs.len());
for (x, &(gh, gw)) in xs.iter().zip(grids.iter()) {
let x_nchw = nhwc_rows_to_nchw(x, dim, gh, gw);
let nc1 = conv_apply(&w.neck_conv1, &x_nchw, gh, gw, 0, 1)?;
ensure_flat_len("vision_sam neck_conv1 output", &nc1, NECK_CH, gh, gw)?;
let nc1 = layer_norm_2d(&nc1, &w.neck_ln1, NECK_CH, gh, gw)?;
let nc2 = conv_apply(&w.neck_conv2, &nc1, gh, gw, 1, 1)?;
ensure_flat_len("vision_sam neck_conv2 output", &nc2, NECK_CH, gh, gw)?;
let neck = layer_norm_2d(&nc2, &w.neck_ln2, NECK_CH, gh, gw)?;
let (gh2, gw2) = (gh.div_ceil(2), gw.div_ceil(2));
let x2 = conv_apply(&w.net2, &neck, gh, gw, 1, 2)?;
ensure_flat_len("vision_sam net2 output", &x2, NET2_CH, gh2, gw2)?;
let (gh3, gw3) = (gh2.div_ceil(2), gw2.div_ceil(2));
let x3 = conv_apply(&w.net3, &x2, gh2, gw2, 1, 2)?;
ensure_flat_len("vision_sam net3 output", &x3, OUT_CH, gh3, gw3)?;
out.push(Mat::from_vec(OUT_CH, gh3 * gw3, x3));
}
Ok(out)
}
pub fn forward_with_batched(
w: &SamWeights,
images: &[&Mat],
h: usize,
win: usize,
) -> FocrResult<Vec<Mat>> {
let v = images.len();
if v == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam forward_with_batched: empty view batch"
)));
}
if h == 0 || win == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam forward_with_batched: spatial dims ({h},{win}) must be non-zero"
)));
}
let expected_cols = checked_shape_mul("vision_sam forward_with_batched", h, win, "H*W")?;
if !h.is_multiple_of(PATCH) || !win.is_multiple_of(PATCH) {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam forward_with_batched: spatial dims ({h},{win}) must be multiples of patch {PATCH}"
)));
}
for (i, img) in images.iter().enumerate() {
if img.rows != 3 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam forward_with_batched: view {i} expected 3 input channels, got {}",
img.rows
)));
}
if img.cols != expected_cols {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam forward_with_batched: view {i} cols {} != H*W {}*{} ({expected_cols}) (ragged batch)",
img.cols,
h,
win
)));
}
}
let gh = h / PATCH;
let gw = win / PATCH;
let n = checked_shape_mul("vision_sam forward_with_batched", gh, gw, "gh*gw")?;
let dim = w.patch_embed.out_ch;
let pos = abs_pos(&w.pos_embed, w.pos_grid_h, w.pos_grid_w, dim, gh, gw)?;
let row_span = checked_shape_mul("vision_sam forward_with_batched", n, dim, "n*dim")?;
let total_rows = checked_shape_mul("vision_sam forward_with_batched", v, n, "V*n")?;
let stacked_len = checked_shape_mul(
"vision_sam forward_with_batched",
total_rows,
dim,
"V*n*dim",
)?;
let mut stacked = vec![0.0f32; stacked_len];
for (vv, img) in images.iter().enumerate() {
let conv_out = conv_apply(&w.patch_embed, &img.data, h, win, 0, PATCH)?;
ensure_flat_len(
"vision_sam forward_with_batched patch_embed output",
&conv_out,
dim,
gh,
gw,
)?;
let x_view = nchw_to_nhwc_rows(&conv_out, dim, gh, gw);
let base = vv * row_span;
let dst = &mut stacked[base..base + row_span];
for (xv, pv) in dst.iter_mut().zip(x_view.data.iter().zip(pos.iter())) {
*xv = *pv.0 + *pv.1;
}
}
let mut x = Mat::from_vec(total_rows, dim, stacked);
for blk in &w.blocks {
x = block_forward_batched(blk, &x, gh, gw, v)?;
}
let mut out = Vec::with_capacity(v);
for vv in 0..v {
let base = vv * row_span;
let x_view = Mat::from_vec(n, dim, x.data[base..base + row_span].to_vec());
let x_nchw = nhwc_rows_to_nchw(&x_view, dim, gh, gw);
let nc1 = conv_apply(&w.neck_conv1, &x_nchw, gh, gw, 0, 1)?;
ensure_flat_len(
"vision_sam forward_with_batched neck_conv1 output",
&nc1,
NECK_CH,
gh,
gw,
)?;
let nc1 = layer_norm_2d(&nc1, &w.neck_ln1, NECK_CH, gh, gw)?;
let nc2 = conv_apply(&w.neck_conv2, &nc1, gh, gw, 1, 1)?;
ensure_flat_len(
"vision_sam forward_with_batched neck_conv2 output",
&nc2,
NECK_CH,
gh,
gw,
)?;
let neck = layer_norm_2d(&nc2, &w.neck_ln2, NECK_CH, gh, gw)?;
let (gh2, gw2) = (gh.div_ceil(2), gw.div_ceil(2));
let x2 = conv_apply(&w.net2, &neck, gh, gw, 1, 2)?;
ensure_flat_len(
"vision_sam forward_with_batched net2 output",
&x2,
NET2_CH,
gh2,
gw2,
)?;
let (gh3, gw3) = (gh2.div_ceil(2), gw2.div_ceil(2));
let x3 = conv_apply(&w.net3, &x2, gh2, gw2, 1, 2)?;
ensure_flat_len(
"vision_sam forward_with_batched net3 output",
&x3,
OUT_CH,
gh3,
gw3,
)?;
out.push(Mat::from_vec(OUT_CH, gh3 * gw3, x3));
}
Ok(out)
}
#[cfg_attr(not(test), allow(dead_code))]
fn block_forward(blk: &BlockP, x: &Mat, gh: usize, gw: usize) -> FocrResult<Mat> {
block_forward_impl(blk, x, gh, gw, false)
}
fn block_forward_impl(
blk: &BlockP,
x: &Mat,
gh: usize,
gw: usize,
low_mem: bool,
) -> FocrResult<Mat> {
let normed = layer_norm_rows(x, &blk.norm1)?;
let ta = Instant::now();
let attn_out = if blk.window > 0 {
attention_windowed(&blk.attn, &normed, gh, gw, blk.window)?
} else if low_mem {
attention_global_bounded(&blk.attn, &normed, gh, gw)?
} else {
attention(&blk.attn, &normed, gh, gw, None)?
};
super::timing_log(&format!(
" sam.block attn({}) {:.3}s",
if blk.window > 0 { "win" } else { "GLOBAL" },
ta.elapsed().as_secs_f64()
));
ensure_same_shape("vision_sam block attention residual", &attn_out, x)?;
let mut h1 = x.clone();
for (a, b) in h1.data.iter_mut().zip(attn_out.data.iter()) {
*a += *b;
}
let tm = Instant::now();
let normed2 = layer_norm_rows(&h1, &blk.norm2)?;
let mlp = if low_mem {
mlp_row_chunked(blk, &normed2)?
} else {
let mut hidden = blk.lin1.apply(&normed2)?;
nn::gelu(&mut hidden);
blk.lin2.apply(&hidden)?
};
super::timing_log(&format!(
" sam.block mlp {:.3}s",
tm.elapsed().as_secs_f64()
));
ensure_same_shape("vision_sam block mlp residual", &mlp, &h1)?;
for (a, b) in h1.data.iter_mut().zip(mlp.data.iter()) {
*a += *b;
}
Ok(h1)
}
const MLP_ROW_CHUNK: usize = 512;
fn mlp_row_chunked(blk: &BlockP, normed2: &Mat) -> FocrResult<Mat> {
mlp_row_chunked_with(blk, normed2, MLP_ROW_CHUNK)
}
fn mlp_row_chunked_with(blk: &BlockP, normed2: &Mat, chunk: usize) -> FocrResult<Mat> {
if chunk == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam mlp_row_chunked: row chunk must be non-zero"
)));
}
let (rows, cols) = (normed2.rows, normed2.cols);
let mut out: Vec<f32> = Vec::with_capacity(rows * blk.lin2.out);
let mut start = 0usize;
while start < rows {
let take = chunk.min(rows - start);
let chunk = Mat::from_vec(
take,
cols,
normed2.data[start * cols..(start + take) * cols].to_vec(),
);
let mut hidden = blk.lin1.apply(&chunk)?;
nn::gelu(&mut hidden);
let done = blk.lin2.apply(&hidden)?;
out.extend_from_slice(&done.data);
start += take;
}
Ok(Mat::from_vec(rows, blk.lin2.out, out))
}
fn block_forward_batched(blk: &BlockP, x: &Mat, gh: usize, gw: usize, v: usize) -> FocrResult<Mat> {
let dim = x.cols;
let n = checked_shape_mul("vision_sam block_forward_batched", gh, gw, "gh*gw")?;
let row_span = checked_shape_mul("vision_sam block_forward_batched", n, dim, "n*dim")?;
let total_rows = checked_shape_mul("vision_sam block_forward_batched", v, n, "V*n")?;
if x.rows != total_rows {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam block_forward_batched: x.rows {} != V*n {total_rows}",
x.rows
)));
}
ensure_flat_len("vision_sam block_forward_batched input", &x.data, v, n, dim)?;
let normed = layer_norm_rows(x, &blk.norm1)?;
let mut attn_data = vec![0.0f32; x.data.len()];
for view in 0..v {
let base = view * row_span;
let normed_view = Mat::from_vec(n, dim, normed.data[base..base + row_span].to_vec());
let av = if blk.window > 0 {
attention_windowed(&blk.attn, &normed_view, gh, gw, blk.window)?
} else {
attention(&blk.attn, &normed_view, gh, gw, None)?
};
ensure_mat_shape(
&av,
n,
dim,
"vision_sam block_forward_batched attention view",
)?;
attn_data[base..base + row_span].copy_from_slice(&av.data);
}
let mut h1 = x.clone();
for (a, b) in h1.data.iter_mut().zip(attn_data.iter()) {
*a += *b;
}
let normed2 = layer_norm_rows(&h1, &blk.norm2)?;
let mut mlp = blk.lin1.apply(&normed2)?;
nn::gelu(&mut mlp);
let mlp = blk.lin2.apply(&mlp)?;
ensure_same_shape("vision_sam block_forward_batched mlp residual", &mlp, &h1)?;
for (a, b) in h1.data.iter_mut().zip(mlp.data.iter()) {
*a += *b;
}
Ok(h1)
}
fn attention_windowed(p: &AttnP, normed: &Mat, gh: usize, gw: usize, ws: usize) -> FocrResult<Mat> {
let dim = normed.cols;
let pad_h = (ws - gh % ws) % ws;
let pad_w = (ws - gw % ws) % ws;
let hp = gh + pad_h;
let wp = gw + pad_w;
let nwin_h = hp / ws;
let nwin_w = wp / ws;
let nwin = nwin_h * nwin_w;
let mut windows = vec![0.0f32; nwin * ws * ws * dim];
for wy in 0..nwin_h {
for wx in 0..nwin_w {
let widx = wy * nwin_w + wx;
for ly in 0..ws {
for lx in 0..ws {
let gy = wy * ws + ly;
let gxx = wx * ws + lx;
let dst = ((widx * ws + ly) * ws + lx) * dim;
if gy < gh && gxx < gw {
let src = (gy * gw + gxx) * dim;
windows[dst..dst + dim].copy_from_slice(&normed.data[src..src + dim]);
}
}
}
}
}
let nh = NUM_HEADS;
let hd = dim / nh;
let rh = get_rel_pos(ws, ws, &p.rel_pos_h, p.size_h, hd);
let rw = get_rel_pos(ws, ws, &p.rel_pos_w, p.size_w, hd);
let win_span = ws * ws * dim;
let mut out_windows = vec![0.0f32; windows.len()];
out_windows
.par_chunks_mut(win_span)
.enumerate()
.try_for_each(|(widx, out_chunk)| -> FocrResult<()> {
let base = widx * win_span;
let win_in = Mat::from_vec(ws * ws, dim, windows[base..base + win_span].to_vec());
let win_out = attention(p, &win_in, ws, ws, Some((&rh, &rw)))?;
out_chunk.copy_from_slice(&win_out.data);
Ok(())
})?;
let mut merged = vec![0.0f32; gh * gw * dim];
for wy in 0..nwin_h {
for wx in 0..nwin_w {
let widx = wy * nwin_w + wx;
for ly in 0..ws {
for lx in 0..ws {
let gy = wy * ws + ly;
let gxx = wx * ws + lx;
if gy < gh && gxx < gw {
let src = ((widx * ws + ly) * ws + lx) * dim;
let dst = (gy * gw + gxx) * dim;
merged[dst..dst + dim].copy_from_slice(&out_windows[src..src + dim]);
}
}
}
}
}
Ok(Mat::from_vec(gh * gw, dim, merged))
}
fn attention(
p: &AttnP,
x: &Mat,
gh: usize,
gw: usize,
relpos: Option<(&[f32], &[f32])>,
) -> FocrResult<Mat> {
let n = gh * gw;
if x.rows != n {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam attention: input rows {} != gh*gw {}*{}",
x.rows,
gh,
gw
)));
}
let dim = x.cols;
let nh = NUM_HEADS;
if dim == 0 || !dim.is_multiple_of(nh) {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam attention: dim {} must be non-zero and divisible by heads {}",
dim,
nh
)));
}
if p.qkv.in_ != dim || p.qkv.out != 3 * dim {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam attention: qkv shape out/in {}x{} incompatible with dim {}",
p.qkv.out,
p.qkv.in_,
dim
)));
}
if p.proj.in_ != dim || p.proj.out != dim {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam attention: proj shape out/in {}x{} incompatible with dim {}",
p.proj.out,
p.proj.in_,
dim
)));
}
let hd = dim / nh;
ensure_rel_pos_len("rel_pos_h", p.size_h, hd, p.rel_pos_h.len())?;
ensure_rel_pos_len("rel_pos_w", p.size_w, hd, p.rel_pos_w.len())?;
let scale = (hd as f32).powf(-0.5);
let qkv = p.qkv.apply(x)?;
ensure_mat_shape(&qkv, n, 3 * dim, "vision_sam attention qkv output")?;
let mut q = vec![0.0f32; nh * n * hd];
let mut k = vec![0.0f32; nh * n * hd];
let mut v = vec![0.0f32; nh * n * hd];
for r in 0..n {
let row = qkv.row(r);
for head in 0..nh {
let dst = (head * n + r) * hd;
q[dst..dst + hd].copy_from_slice(&row[head * hd..(head + 1) * hd]);
k[dst..dst + hd].copy_from_slice(&row[(nh + head) * hd..(nh + head + 1) * hd]);
v[dst..dst + hd].copy_from_slice(&row[(2 * nh + head) * hd..(2 * nh + head + 1) * hd]);
}
}
let (rh_owned, rw_owned);
let (rh, rw): (&[f32], &[f32]) = match relpos {
Some((rh, rw)) => (rh, rw),
None => {
rh_owned = get_rel_pos(gh, gh, &p.rel_pos_h, p.size_h, hd);
rw_owned = get_rel_pos(gw, gw, &p.rel_pos_w, p.size_w, hd);
(&rh_owned, &rw_owned)
}
};
let mut out = vec![0.0f32; nh * n * hd]; out.par_chunks_mut(n * hd)
.enumerate()
.try_for_each(|(head, out_chunk)| -> FocrResult<()> {
let qh = &q[head * n * hd..(head + 1) * n * hd];
let kh = &k[head * n * hd..(head + 1) * n * hd];
let vh = &v[head * n * hd..(head + 1) * n * hd];
let (rel_h_bias, rel_w_bias) = decomposed_rel_pos_bias(qh, rh, rw, gh, gw, hd);
let qh_mat = Mat::from_vec(n, hd, qh.to_vec());
let kt_mat = Mat::from_vec(hd, n, transpose_contiguous_stores(kh, n, hd));
let mut lm = nn::matmul(&qh_mat, &kt_mat)?;
for i in 0..n {
let lrow = &mut lm.data[i * n..(i + 1) * n];
let brow_w = &rel_w_bias[i * gw..(i + 1) * gw];
for ky in 0..gh {
let bh = rel_h_bias[i * gh + ky];
let seg = &mut lrow[ky * gw..(ky + 1) * gw];
for (l, &bw) in seg.iter_mut().zip(brow_w) {
*l = scale * *l + bh + bw;
}
}
}
nn::softmax_rows(&mut lm)?;
let vh_mat = Mat::from_vec(n, hd, vh.to_vec());
let head_out = nn::matmul(&lm, &vh_mat)?;
out_chunk.copy_from_slice(&head_out.data);
Ok(())
})?;
let mut ctx = vec![0.0f32; n * dim];
for head in 0..nh {
for r in 0..n {
let src = (head * n + r) * hd;
let dst = r * dim + head * hd;
ctx[dst..dst + hd].copy_from_slice(&out[src..src + hd]);
}
}
let ctx_mat = Mat::from_vec(n, dim, ctx);
let y = p.proj.apply(&ctx_mat)?;
ensure_mat_shape(&y, n, dim, "vision_sam attention projection output")?;
Ok(y)
}
const GLOBAL_ATTN_QUERY_SLAB: usize = 128;
fn attention_global_bounded(p: &AttnP, x: &Mat, gh: usize, gw: usize) -> FocrResult<Mat> {
attention_global_bounded_with(p, x, gh, gw, GLOBAL_ATTN_QUERY_SLAB)
}
fn attention_global_bounded_with(
p: &AttnP,
x: &Mat,
gh: usize,
gw: usize,
slab: usize,
) -> FocrResult<Mat> {
let n = gh * gw;
if x.rows != n {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam attention: input rows {} != gh*gw {}*{}",
x.rows,
gh,
gw
)));
}
let dim = x.cols;
let nh = NUM_HEADS;
if dim == 0 || !dim.is_multiple_of(nh) {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam attention: dim {} must be non-zero and divisible by heads {}",
dim,
nh
)));
}
if p.qkv.in_ != dim || p.qkv.out != 3 * dim {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam attention: qkv shape out/in {}x{} incompatible with dim {}",
p.qkv.out,
p.qkv.in_,
dim
)));
}
if p.proj.in_ != dim || p.proj.out != dim {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam attention: proj shape out/in {}x{} incompatible with dim {}",
p.proj.out,
p.proj.in_,
dim
)));
}
if slab == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam attention: query slab must be non-zero"
)));
}
let hd = dim / nh;
ensure_rel_pos_len("rel_pos_h", p.size_h, hd, p.rel_pos_h.len())?;
ensure_rel_pos_len("rel_pos_w", p.size_w, hd, p.rel_pos_w.len())?;
let scale = (hd as f32).powf(-0.5);
let qkv = p.qkv.apply(x)?;
ensure_mat_shape(&qkv, n, 3 * dim, "vision_sam attention qkv output")?;
let mut q = vec![0.0f32; nh * n * hd];
let mut k = vec![0.0f32; nh * n * hd];
let mut v = vec![0.0f32; nh * n * hd];
for r in 0..n {
let row = qkv.row(r);
for head in 0..nh {
let dst = (head * n + r) * hd;
q[dst..dst + hd].copy_from_slice(&row[head * hd..(head + 1) * hd]);
k[dst..dst + hd].copy_from_slice(&row[(nh + head) * hd..(nh + head + 1) * hd]);
v[dst..dst + hd].copy_from_slice(&row[(2 * nh + head) * hd..(2 * nh + head + 1) * hd]);
}
}
drop(qkv);
let rh = get_rel_pos(gh, gh, &p.rel_pos_h, p.size_h, hd);
let rw = get_rel_pos(gw, gw, &p.rel_pos_w, p.size_w, hd);
let mut out = vec![0.0f32; nh * n * hd]; out.par_chunks_mut(n * hd)
.enumerate()
.try_for_each(|(head, out_chunk)| -> FocrResult<()> {
let qh = &q[head * n * hd..(head + 1) * n * hd];
let kh = &k[head * n * hd..(head + 1) * n * hd];
let vh = &v[head * n * hd..(head + 1) * n * hd];
let (rel_h_bias, rel_w_bias) = decomposed_rel_pos_bias(qh, &rh, &rw, gh, gw, hd);
let kt_mat = Mat::from_vec(hd, n, transpose_contiguous_stores(kh, n, hd));
let vh_mat = Mat::from_vec(n, hd, vh.to_vec());
let mut q0 = 0usize;
while q0 < n {
let qc = slab.min(n - q0);
let qb_mat = Mat::from_vec(qc, hd, qh[q0 * hd..(q0 + qc) * hd].to_vec());
let mut lm = nn::matmul(&qb_mat, &kt_mat)?;
for li in 0..qc {
let i = q0 + li;
let lrow = &mut lm.data[li * n..(li + 1) * n];
let brow_w = &rel_w_bias[i * gw..(i + 1) * gw];
for ky in 0..gh {
let bh = rel_h_bias[i * gh + ky];
let seg = &mut lrow[ky * gw..(ky + 1) * gw];
for (l, &bw) in seg.iter_mut().zip(brow_w) {
*l = scale * *l + bh + bw;
}
}
}
nn::softmax_rows(&mut lm)?;
let slab_out = nn::matmul(&lm, &vh_mat)?;
out_chunk[q0 * hd..(q0 + qc) * hd].copy_from_slice(&slab_out.data);
q0 += qc;
}
Ok(())
})?;
let mut ctx = vec![0.0f32; n * dim];
for head in 0..nh {
for r in 0..n {
let src = (head * n + r) * hd;
let dst = r * dim + head * hd;
ctx[dst..dst + hd].copy_from_slice(&out[src..src + hd]);
}
}
let ctx_mat = Mat::from_vec(n, dim, ctx);
let y = p.proj.apply(&ctx_mat)?;
ensure_mat_shape(&y, n, dim, "vision_sam attention projection output")?;
Ok(y)
}
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 ensure_flat_len(context: &str, data: &[f32], ch: usize, h: usize, w: usize) -> FocrResult<()> {
let expected = checked_nchw_len(context, ch, h, w, "ch*h*w")?;
if data.len() == expected {
return Ok(());
}
Err(FocrError::Other(anyhow::anyhow!(
"{context}: len {} != ch*h*w {}*{}*{} ({expected})",
data.len(),
ch,
h,
w
)))
}
fn ensure_same_shape(context: &str, actual: &Mat, expected: &Mat) -> FocrResult<()> {
if actual.shape() == expected.shape() {
return Ok(());
}
Err(FocrError::Other(anyhow::anyhow!(
"{context}: shape {:?} != expected {:?}",
actual.shape(),
expected.shape()
)))
}
fn ensure_rel_pos_len(name: &str, size: usize, hd: usize, actual_len: usize) -> FocrResult<()> {
let rows = size
.checked_mul(2)
.and_then(|n| n.checked_sub(1))
.ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"vision_sam attention: {name} size {size} is invalid"
))
})?;
let expected_len = rows.checked_mul(hd).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"vision_sam attention: {name} size {size} overflows for head_dim {hd}"
))
})?;
if actual_len == expected_len {
return Ok(());
}
Err(FocrError::Other(anyhow::anyhow!(
"vision_sam attention: {name} len {actual_len} != expected {expected_len} \
for size {size} and head_dim {hd}"
)))
}
fn decomposed_rel_pos_bias(
qh: &[f32],
rh: &[f32],
rw: &[f32],
gh: usize,
gw: usize,
hd: usize,
) -> (Vec<f32>, Vec<f32>) {
let n = gh * gw;
debug_assert_eq!(qh.len(), n * hd);
debug_assert_eq!(rh.len(), gh * gh * hd);
debug_assert_eq!(rw.len(), gw * gw * hd);
let mut rel_h_bias = vec![0.0f32; n * gh];
let mut rel_w_bias = vec![0.0f32; n * gw];
for i in 0..n {
let qy = i / gw;
let qx = i % gw;
let qi = &qh[i * hd..(i + 1) * hd];
for ky in 0..gh {
let rh_base = (qy * gh + ky) * hd;
let mut bh = 0.0f32;
for c in 0..hd {
bh += qi[c] * rh[rh_base + c];
}
rel_h_bias[i * gh + ky] = bh;
}
for kx in 0..gw {
let rw_base = (qx * gw + kx) * hd;
let mut bw = 0.0f32;
for c in 0..hd {
bw += qi[c] * rw[rw_base + c];
}
rel_w_bias[i * gw + kx] = bw;
}
}
(rel_h_bias, rel_w_bias)
}
fn get_rel_pos(q_size: usize, k_size: usize, rel_pos: &[f32], size: usize, hd: usize) -> Vec<f32> {
let max_rel = 2 * q_size.max(k_size) - 1;
let table_rows = rel_pos.len() / hd;
debug_assert_eq!(table_rows, 2 * size - 1);
let resized: Vec<f32> = if table_rows != max_rel {
interp_linear_rows(rel_pos, table_rows, hd, max_rel)
} else {
rel_pos.to_vec()
};
let qf = q_size as f32;
let kf = k_size as f32;
let ratio_qk = (kf / qf).max(1.0);
let ratio_kq = (qf / kf).max(1.0);
let mut out = vec![0.0f32; q_size * k_size * hd];
for qi in 0..q_size {
for ki in 0..k_size {
let qc = qi as f32 * ratio_qk;
let kc = ki as f32 * ratio_kq;
let rc = (qc - kc) + (k_size as f32 - 1.0) * ratio_kq;
let idx = rc as usize; let src = idx * hd;
let dst = (qi * k_size + ki) * hd;
out[dst..dst + hd].copy_from_slice(&resized[src..src + hd]);
}
}
out
}
fn interp_linear_rows(src: &[f32], rows: usize, hd: usize, new_rows: usize) -> Vec<f32> {
if new_rows == rows {
return src.to_vec();
}
let mut out = vec![0.0f32; new_rows * hd];
let scale = rows as f32 / new_rows as f32;
for i in 0..new_rows {
let s = (i as f32 + 0.5) * scale - 0.5;
let s_clamped = s.clamp(0.0, (rows - 1) as f32);
let lo = s_clamped.floor() as usize;
let hi = (lo + 1).min(rows - 1);
let frac = s_clamped - lo as f32;
for c in 0..hd {
let a = src[lo * hd + c];
let b = src[hi * hd + c];
out[i * hd + c] = a + (b - a) * frac;
}
}
out
}
fn layer_norm_rows(x: &Mat, ln: &LayerNormP) -> FocrResult<Mat> {
nn::layer_norm(x, Some(&ln.w), Some(&ln.b), LN_EPS)
}
fn layer_norm_2d(
x: &[f32],
ln: &LayerNormP,
ch: usize,
gh: usize,
gw: usize,
) -> FocrResult<Vec<f32>> {
if ch == 0 || gh == 0 || gw == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam layer_norm_2d: channels and grid dims must be non-zero (ch={ch}, gh={gh}, gw={gw})"
)));
}
ensure_flat_len("vision_sam layer_norm_2d input", x, ch, gh, gw)?;
if ln.w.len() != ch {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam layer_norm_2d: weight len {} != channels {}",
ln.w.len(),
ch
)));
}
if ln.b.len() != ch {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam layer_norm_2d: bias len {} != channels {}",
ln.b.len(),
ch
)));
}
let hw = checked_shape_mul("vision_sam layer_norm_2d", gh, gw, "gh*gw")?;
let out_len = checked_shape_mul("vision_sam layer_norm_2d", ch, hw, "ch*gh*gw")?;
let mut out = vec![0.0f32; out_len];
for s in 0..hw {
let mut mean = 0.0f32;
for c in 0..ch {
mean += x[c * hw + s];
}
mean /= ch as f32;
let mut var = 0.0f32;
for c in 0..ch {
let d = x[c * hw + s] - mean;
var += d * d;
}
var /= ch as f32;
let inv = 1.0 / (var + LN_EPS).sqrt();
for c in 0..ch {
let norm = (x[c * hw + s] - mean) * inv;
out[c * hw + s] = ln.w[c] * norm + ln.b[c];
}
}
Ok(out)
}
pub(crate) fn conv_apply(
conv: &Conv,
input: &[f32],
gh: usize,
gw: usize,
pad: usize,
stride: usize,
) -> FocrResult<Vec<f32>> {
if gh == 0 || gw == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam conv: input grid ({gh},{gw}) must be non-zero"
)));
}
if stride == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam conv: stride must be non-zero"
)));
}
if conv.in_ch == 0 || conv.out_ch == 0 || conv.kh == 0 || conv.kw == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam conv: channels and kernel dims must be non-zero (out={}, in={}, kh={}, kw={})",
conv.out_ch,
conv.in_ch,
conv.kh,
conv.kw
)));
}
let expected_input =
checked_nchw_len("vision_sam conv input", conv.in_ch, gh, gw, "in_ch*gh*gw")?;
if input.len() != expected_input {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam conv: input len {} != in_ch*gh*gw {}",
input.len(),
expected_input
)));
}
let expected_weight = checked_conv_weight_len(
"vision_sam conv weight",
conv.out_ch,
conv.in_ch,
conv.kh,
conv.kw,
)?;
if conv.w.len() != expected_weight {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam conv: weight len {} != out_ch*in_ch*kh*kw {}",
conv.w.len(),
expected_weight
)));
}
if let Some(bias) = &conv.b
&& bias.len() != conv.out_ch
{
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam conv: bias len {} != out_ch {}",
bias.len(),
conv.out_ch
)));
}
let two_pad = checked_shape_mul("vision_sam conv", 2, pad, "2*pad")?;
let ph = checked_shape_add("vision_sam conv", gh, two_pad, "gh+2*pad")?;
let pw = checked_shape_add("vision_sam conv", gw, two_pad, "gw+2*pad")?;
let oh_base = checked_shape_sub("vision_sam conv", ph, conv.kh, "padded_h-kh")?;
let ow_base = checked_shape_sub("vision_sam conv", pw, conv.kw, "padded_w-kw")?;
let oh = checked_shape_add("vision_sam conv", oh_base / stride, 1, "output_h+1")?;
let ow = checked_shape_add("vision_sam conv", ow_base / stride, 1, "output_w+1")?;
let expected_out = checked_nchw_len(
"vision_sam conv output",
conv.out_ch,
oh,
ow,
"out_ch*oh*ow",
)?;
let padded = pad_nchw(input, conv.in_ch, gh, gw, pad)?;
let out = nn::conv2d(
&padded,
&conv.w,
conv.b.as_deref(),
1,
conv.in_ch,
ph,
pw,
conv.kh,
conv.kw,
oh,
ow,
stride,
stride,
conv.out_ch,
);
if out.len() != expected_out {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam conv: kernel output len {} != out_ch*oh*ow {}",
out.len(),
expected_out
)));
}
Ok(out)
}
fn pad_nchw(input: &[f32], ch: usize, gh: usize, gw: usize, pad: usize) -> FocrResult<Vec<f32>> {
let expected_input = checked_nchw_len("vision_sam pad_nchw input", ch, gh, gw, "ch*gh*gw")?;
if input.len() != expected_input {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam pad_nchw: input len {} != ch*gh*gw {}",
input.len(),
expected_input
)));
}
if pad == 0 {
return Ok(input.to_vec());
}
let two_pad = checked_shape_mul("vision_sam pad_nchw", 2, pad, "2*pad")?;
let ph = checked_shape_add("vision_sam pad_nchw", gh, two_pad, "gh+2*pad")?;
let pw = checked_shape_add("vision_sam pad_nchw", gw, two_pad, "gw+2*pad")?;
let out_len = checked_nchw_len("vision_sam pad_nchw output", ch, ph, pw, "ch*ph*pw")?;
let mut out = vec![0.0f32; out_len];
for c in 0..ch {
for y in 0..gh {
for x in 0..gw {
let src = c * gh * gw + y * gw + x;
let dst = c * ph * pw + (y + pad) * pw + (x + pad);
out[dst] = input[src];
}
}
}
Ok(out)
}
pub(crate) fn nchw_to_nhwc_rows(nchw: &[f32], ch: usize, gh: usize, gw: usize) -> Mat {
let n = gh * gw;
let mut data = vec![0.0f32; n * ch];
for s in 0..n {
let dst = &mut data[s * ch..(s + 1) * ch];
for c in 0..ch {
dst[c] = nchw[c * n + s];
}
}
Mat::from_vec(n, ch, data)
}
fn nhwc_rows_to_nchw(x: &Mat, ch: usize, gh: usize, gw: usize) -> Vec<f32> {
let n = gh * gw;
debug_assert_eq!(x.rows, n);
debug_assert_eq!(x.cols, ch);
let mut out = vec![0.0f32; ch * n];
for c in 0..ch {
let dst = &mut out[c * n..(c + 1) * n];
for (s, slot) in dst.iter_mut().enumerate() {
*slot = x.data[s * ch + c];
}
}
out
}
fn transpose(m: &[f32], rows: usize, cols: usize) -> Vec<f32> {
let mut out = vec![0.0f32; rows * cols];
for c in 0..cols {
let dst = &mut out[c * rows..(c + 1) * rows];
for (r, slot) in dst.iter_mut().enumerate() {
*slot = m[r * cols + c];
}
}
out
}
fn transpose_contiguous_stores(m: &[f32], rows: usize, cols: usize) -> Vec<f32> {
let mut out = vec![0.0f32; rows * cols];
for c in 0..cols {
let dst = &mut out[c * rows..(c + 1) * rows];
for r in 0..rows {
dst[r] = m[r * cols + c];
}
}
out
}
fn abs_pos(
pos: &[f32],
src_h: usize,
src_w: usize,
dim: usize,
gh: usize,
gw: usize,
) -> FocrResult<Vec<f32>> {
if src_h == 0 || src_w == 0 || dim == 0 || gh == 0 || gw == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam abs_pos: source ({src_h},{src_w}), target ({gh},{gw}), and dim {dim} must be non-zero"
)));
}
let src_hw = checked_shape_mul("vision_sam abs_pos", src_h, src_w, "src_h*src_w")?;
let expected_pos_len = checked_shape_mul("vision_sam abs_pos", src_hw, dim, "src_h*src_w*dim")?;
if pos.len() != expected_pos_len {
return Err(FocrError::Other(anyhow::anyhow!(
"vision_sam abs_pos: pos_embed len {} != src_h*src_w*dim {}*{}*{} ({expected_pos_len})",
pos.len(),
src_h,
src_w,
dim
)));
}
let tgt_hw = checked_shape_mul("vision_sam abs_pos", gh, gw, "gh*gw")?;
let out_len = checked_shape_mul("vision_sam abs_pos", tgt_hw, dim, "gh*gw*dim")?;
if src_h == gh && src_w == gw {
return Ok(pos.to_vec()); }
let mut out = vec![0.0f32; out_len];
let scale_y = src_h as f32 / gh as f32;
let scale_x = src_w as f32 / gw as f32;
for oy in 0..gh {
let sy = (oy as f32 + 0.5) * scale_y - 0.5;
for ox in 0..gw {
let sx = (ox as f32 + 0.5) * scale_x - 0.5;
let dst = (oy * gw + ox) * dim;
for c in 0..dim {
out[dst + c] = bicubic_sample(pos, src_h, src_w, dim, c, sy, sx);
}
}
}
Ok(out)
}
fn bicubic_sample(
pos: &[f32],
src_h: usize,
src_w: usize,
dim: usize,
c: usize,
sy: f32,
sx: f32,
) -> f32 {
let iy = sy.floor();
let ix = sx.floor();
let fy = sy - iy;
let fx = sx - ix;
let wy = cubic_weights(fy);
let wx = cubic_weights(fx);
let mut acc = 0.0f32;
#[allow(clippy::needless_range_loop)]
for m in 0..4 {
let yy = clamp_idx(iy as isize - 1 + m as isize, src_h);
#[allow(clippy::needless_range_loop)]
for n in 0..4 {
let xx = clamp_idx(ix as isize - 1 + n as isize, src_w);
let val = pos[(yy * src_w + xx) * dim + c];
acc += val * wy[m] * wx[n];
}
}
acc
}
fn cubic_weights(t: f32) -> [f32; 4] {
let a = -0.75f32;
let d0 = 1.0 + t;
let d1 = t;
let d2 = 1.0 - t;
let d3 = 2.0 - t;
[
cubic_k(d0, a),
cubic_k(d1, a),
cubic_k(d2, a),
cubic_k(d3, a),
]
}
fn cubic_k(x: f32, a: f32) -> f32 {
let x = x.abs();
if x <= 1.0 {
(a + 2.0) * x * x * x - (a + 3.0) * x * x + 1.0
} else if x < 2.0 {
a * x * x * x - 5.0 * a * x * x + 8.0 * a * x - 4.0 * a
} else {
0.0
}
}
fn clamp_idx(i: isize, n: usize) -> usize {
if i < 0 {
0
} else if i as usize >= n {
n - 1
} else {
i as usize
}
}
#[cfg(test)]
pub(crate) mod test_support {
use super::{DEPTH, GLOBAL_BLOCKS, NUM_HEADS, PATCH, WINDOW};
use crate::quant::focrq::{FocrqBuilder, WriteDType};
pub(crate) fn synth_values(len: usize, salt: u64) -> Vec<f32> {
(0..len)
.map(|i| {
let raw = ((i as u64)
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(salt)
>> 33) as u32;
(raw as f32 / u32::MAX as f32 - 0.5) * 0.3
})
.collect()
}
pub(crate) fn add_synth_sam_tower(b: &mut FocrqBuilder, prefix: &str, dim: usize) {
let f32_bytes =
|values: &[f32]| -> Vec<u8> { values.iter().flat_map(|v| v.to_le_bytes()).collect() };
let mut add = |name: String, shape: Vec<usize>, salt: u64| {
let len: usize = shape.iter().product();
b.add_tensor(
name,
WriteDType::F32,
shape,
f32_bytes(&synth_values(len, salt)),
)
.expect("valid synthetic f32 tensor");
};
add(
format!("{prefix}.patch_embed.proj.weight"),
vec![dim, 3, PATCH, PATCH],
1,
);
add(format!("{prefix}.patch_embed.proj.bias"), vec![dim], 2);
add(format!("{prefix}.pos_embed"), vec![1, 4, 4, dim], 3);
for (idx, name) in ["neck.0", "neck.2"].iter().enumerate() {
let ch = 8usize;
let in_ch = if idx == 0 { dim } else { ch };
add(
format!("{prefix}.{name}.weight"),
vec![ch, in_ch, 1, 1],
10 + idx as u64,
);
}
for (idx, name) in ["neck.1", "neck.3"].iter().enumerate() {
add(format!("{prefix}.{name}.weight"), vec![8], 20 + idx as u64);
add(format!("{prefix}.{name}.bias"), vec![8], 30 + idx as u64);
}
add(format!("{prefix}.net_2.weight"), vec![8, 8, 3, 3], 40);
add(format!("{prefix}.net_3.weight"), vec![8, 8, 3, 3], 41);
let hd = dim / NUM_HEADS;
for i in 0..DEPTH {
let bb = format!("{prefix}.blocks.{i}");
let salt = 1000 * (i as u64 + 1);
let rel_rows = if GLOBAL_BLOCKS.contains(&i) {
2 * 4 - 1 } else {
2 * WINDOW - 1
};
add(format!("{bb}.norm1.weight"), vec![dim], salt + 1);
add(format!("{bb}.norm1.bias"), vec![dim], salt + 2);
add(
format!("{bb}.attn.qkv.weight"),
vec![3 * dim, dim],
salt + 3,
);
add(format!("{bb}.attn.qkv.bias"), vec![3 * dim], salt + 4);
add(format!("{bb}.attn.proj.weight"), vec![dim, dim], salt + 5);
add(format!("{bb}.attn.proj.bias"), vec![dim], salt + 6);
add(format!("{bb}.attn.rel_pos_h"), vec![rel_rows, hd], salt + 7);
add(format!("{bb}.attn.rel_pos_w"), vec![rel_rows, hd], salt + 8);
add(format!("{bb}.norm2.weight"), vec![dim], salt + 9);
add(format!("{bb}.norm2.bias"), vec![dim], salt + 10);
add(
format!("{bb}.mlp.lin1.weight"),
vec![4 * dim, dim],
salt + 11,
);
add(format!("{bb}.mlp.lin1.bias"), vec![4 * dim], salt + 12);
add(
format!("{bb}.mlp.lin2.weight"),
vec![dim, 4 * dim],
salt + 13,
);
add(format!("{bb}.mlp.lin2.bias"), vec![dim], salt + 14);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::quant::focrq::{FocrqBuilder, WriteDType};
use half::bf16;
use serde_json::json;
use std::time::Instant;
fn identity_conv1(ch: usize) -> Conv {
let mut w = vec![0.0f32; ch * ch];
for c in 0..ch {
w[c * ch + c] = 1.0;
}
Conv {
w,
b: None,
out_ch: ch,
in_ch: ch,
kh: 1,
kw: 1,
}
}
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 add_minimal_patch_embed(b: &mut FocrqBuilder) {
let p = "model.sam_model";
b.add_tensor(
format!("{p}.patch_embed.proj.weight"),
WriteDType::Bf16,
vec![1, 1, 1, 1],
bf16_zeros(1),
)
.unwrap();
b.add_tensor(
format!("{p}.patch_embed.proj.bias"),
WriteDType::Bf16,
vec![1],
bf16_zeros(1),
)
.unwrap();
}
#[test]
fn sam_weights_from_rejects_rank1_patch_embed_without_panic() {
let p = "model.sam_model";
let mut b = FocrqBuilder::new();
b.add_tensor(
format!("{p}.patch_embed.proj.weight"),
WriteDType::Bf16,
vec![4],
bf16_zeros(4),
)
.unwrap();
let weights = Weights::from_bytes(b.build()).unwrap();
assert_err_contains(sam_weights_from(&weights, "model.sam_model"), "rank 1");
}
#[test]
fn sam_weights_from_rejects_rank1_pos_embed_without_panic() {
let p = "model.sam_model";
let mut b = FocrqBuilder::new();
add_minimal_patch_embed(&mut b);
b.add_tensor(
format!("{p}.pos_embed"),
WriteDType::Bf16,
vec![4],
bf16_zeros(4),
)
.unwrap();
let weights = Weights::from_bytes(b.build()).unwrap();
assert_err_contains(sam_weights_from(&weights, "model.sam_model"), "rank 1");
}
#[test]
fn sam_weights_from_rejects_rank1_block_qkv_without_panic() {
let p = "model.sam_model";
let mut b = FocrqBuilder::new();
add_minimal_patch_embed(&mut b);
b.add_tensor(
format!("{p}.pos_embed"),
WriteDType::Bf16,
vec![1, 1, 1, 1],
bf16_zeros(1),
)
.unwrap();
b.add_tensor(
format!("{p}.blocks.0.attn.qkv.weight"),
WriteDType::Bf16,
vec![4],
bf16_zeros(4),
)
.unwrap();
let weights = Weights::from_bytes(b.build()).unwrap();
assert_err_contains(sam_weights_from(&weights, "model.sam_model"), "rank 1");
}
#[test]
fn transpose_roundtrips() {
let m = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]; let t = transpose(&m, 2, 3); assert_eq!(t, vec![1.0, 4.0, 2.0, 5.0, 3.0, 6.0]);
assert_eq!(transpose(&t, 3, 2), m);
}
fn test_linear(w: Vec<f32>, b: Vec<f32>, out: usize, in_: usize) -> Linear {
Linear::from_row_major(&w, b, out, in_).expect("test linear shape is valid")
}
#[test]
fn linear_applies_weight_and_bias() -> FocrResult<()> {
let lin = test_linear(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![10.0, 20.0], 2, 3);
let x = Mat::from_vec(1, 3, vec![1.0, 1.0, 1.0]);
let y = lin.apply(&x)?;
assert_eq!(y.shape(), (1, 2));
assert!((y.data[0] - 16.0).abs() < 1e-5);
assert!((y.data[1] - 35.0).abs() < 1e-5);
Ok(())
}
#[test]
fn linear_rejects_malformed_shapes_without_panic() {
assert_err_contains(
Linear::from_row_major(&[1.0; 5], vec![], 2, 3),
"weight len",
);
assert_err_contains(
Linear::from_row_major(&[1.0; 6], vec![0.0], 2, 3),
"bias len",
);
let bad_x = Mat::from_vec(1, 2, vec![1.0, 2.0]);
let lin = test_linear(vec![1.0; 6], vec![], 2, 3);
assert_err_contains(lin.apply(&bad_x), "input cols");
}
#[test]
fn pad_nchw_zeros_border() -> FocrResult<()> {
let input = vec![1.0, 2.0, 3.0, 4.0];
let out = pad_nchw(&input, 1, 2, 2, 1)?;
assert_eq!(out.len(), 16);
assert_eq!(out[4 + 1], 1.0);
assert_eq!(out[4 + 2], 2.0);
assert_eq!(out[8 + 1], 3.0);
assert_eq!(out[8 + 2], 4.0);
assert_eq!(out[0], 0.0);
assert_eq!(out[15], 0.0);
Ok(())
}
#[test]
fn nchw_nhwc_roundtrip() {
let nchw = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let rows = nchw_to_nhwc_rows(&nchw, 2, 2, 2);
assert_eq!(rows.row(0), &[1.0, 5.0]);
assert_eq!(rows.row(3), &[4.0, 8.0]);
let back = nhwc_rows_to_nchw(&rows, 2, 2, 2);
assert_eq!(back, nchw);
}
#[test]
fn nchw_nhwc_roundtrip_nonsquare_grid() {
let nchw = vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 20.0, 21.0, 22.0, 23.0, 24.0, 25.0, ];
let rows = nchw_to_nhwc_rows(&nchw, 3, 2, 3);
assert_eq!(rows.row(0), &[1.0, 10.0, 20.0]);
assert_eq!(rows.row(5), &[6.0, 15.0, 25.0]);
let back = nhwc_rows_to_nchw(&rows, 3, 2, 3);
assert_eq!(back, nchw);
}
#[test]
fn layer_norm_2d_normalizes_channels() -> FocrResult<()> {
let x = vec![1.0, 9.0, 3.0, 11.0]; let ln = LayerNormP {
w: vec![1.0, 1.0],
b: vec![0.0, 0.0],
};
let out = layer_norm_2d(&x, &ln, 2, 1, 2)?;
assert!((out[0] - (-1.0)).abs() < 1e-3); assert!((out[2] - 1.0).abs() < 1e-3); assert!((out[1] - (-1.0)).abs() < 1e-3); assert!((out[3] - 1.0).abs() < 1e-3); Ok(())
}
#[test]
fn layer_norm_2d_rejects_malformed_affine_without_panic() {
let x = vec![1.0, 2.0, 3.0, 4.0];
assert_err_contains(
layer_norm_2d(
&x,
&LayerNormP {
w: vec![1.0],
b: vec![0.0, 0.0],
},
2,
1,
2,
),
"weight len",
);
assert_err_contains(
layer_norm_2d(
&x,
&LayerNormP {
w: vec![1.0, 1.0],
b: vec![0.0],
},
2,
1,
2,
),
"bias len",
);
}
#[test]
fn layer_norm_2d_rejects_malformed_input_without_panic() {
let ln = LayerNormP {
w: vec![1.0, 1.0],
b: vec![0.0, 0.0],
};
assert_err_contains(layer_norm_2d(&[1.0, 2.0, 3.0], &ln, 2, 1, 2), "input");
assert_err_contains(layer_norm_2d(&[], &ln, 0, 1, 2), "non-zero");
}
#[test]
fn conv_apply_identity_1x1_preserves() -> FocrResult<()> {
let input = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let conv = identity_conv1(2);
let out = conv_apply(&conv, &input, 2, 2, 0, 1)?;
assert_eq!(out, input);
Ok(())
}
#[test]
fn conv_apply_rejects_malformed_geometry_without_panic() {
let input = vec![1.0, 2.0, 3.0, 4.0];
let conv = identity_conv1(1);
assert_err_contains(conv_apply(&conv, &input, 2, 2, 0, 0), "stride");
let oversized_kernel = Conv {
w: vec![1.0; 25],
b: None,
out_ch: 1,
in_ch: 1,
kh: 5,
kw: 5,
};
assert_err_contains(
conv_apply(&oversized_kernel, &input, 2, 2, 0, 1),
"padded_h-kh",
);
}
#[test]
fn conv_apply_rejects_buffer_mismatches_without_panic() {
let input = vec![1.0, 2.0, 3.0, 4.0];
let short_weight = Conv {
w: vec![1.0],
b: None,
out_ch: 2,
in_ch: 1,
kh: 1,
kw: 1,
};
assert_err_contains(conv_apply(&short_weight, &input, 2, 2, 0, 1), "weight len");
let bad_bias = Conv {
w: vec![1.0],
b: Some(vec![0.0, 1.0]),
out_ch: 1,
in_ch: 1,
kh: 1,
kw: 1,
};
assert_err_contains(conv_apply(&bad_bias, &input, 2, 2, 0, 1), "bias len");
assert_err_contains(
conv_apply(&identity_conv1(1), &input[..3], 2, 2, 0, 1),
"input len",
);
}
#[test]
fn conv_apply_rejects_padding_overflow_before_allocating() {
let input = vec![1.0];
assert_err_contains(
conv_apply(&identity_conv1(1), &input, 1, 1, usize::MAX / 2 + 1, 1),
"2*pad",
);
}
#[test]
fn cubic_weights_sum_to_one() {
for &t in &[0.0f32, 0.25, 0.5, 0.75, 0.99] {
let w = cubic_weights(t);
let s: f32 = w.iter().sum();
assert!((s - 1.0).abs() < 1e-5, "t={t} sum={s}");
}
}
#[test]
fn abs_pos_identity_when_grid_matches() -> FocrResult<()> {
let pos = vec![1.0, 2.0, 3.0, 4.0]; let out = abs_pos(&pos, 2, 2, 1, 2, 2)?;
assert_eq!(out, pos);
Ok(())
}
#[test]
fn abs_pos_bicubic_constant_field_is_constant() -> FocrResult<()> {
let dim = 1;
let pos = vec![7.0f32; 4 * 4 * dim];
let out = abs_pos(&pos, 4, 4, dim, 6, 6)?;
assert_eq!(out.len(), 6 * 6);
for &v in &out {
assert!((v - 7.0).abs() < 1e-3, "got {v}");
}
Ok(())
}
#[test]
fn abs_pos_rejects_malformed_source_len_without_panic() {
assert_err_contains(abs_pos(&[1.0, 2.0, 3.0], 2, 2, 1, 2, 2), "pos_embed len");
assert_err_contains(abs_pos(&[1.0, 2.0, 3.0], 2, 2, 1, 3, 3), "pos_embed len");
}
#[test]
fn abs_pos_rejects_invalid_geometry_without_panic() {
assert_err_contains(abs_pos(&[], 0, 2, 1, 2, 2), "must be non-zero");
assert_err_contains(abs_pos(&[], usize::MAX, 2, 1, 2, 2), "src_h*src_w");
}
#[test]
fn get_rel_pos_indexes_table_directly() {
let table = vec![10.0, 20.0, 30.0];
let r = get_rel_pos(2, 2, &table, 2, 1);
assert_eq!(r[0], 20.0); assert_eq!(r[1], 10.0); assert_eq!(r[2], 30.0); assert_eq!(r[3], 20.0); }
#[test]
fn decomposed_rel_pos_bias_matches_direct_inner_loop_formula() {
let (gh, gw, hd) = (3, 2, 5);
let n = gh * gw;
let qh: Vec<f32> = (0..n * hd)
.map(|i| ((i % 11) as f32 - 5.0) * 0.013)
.collect();
let rh: Vec<f32> = (0..gh * gh * hd)
.map(|i| ((i % 7) as f32 - 3.0) * 0.017)
.collect();
let rw: Vec<f32> = (0..gw * gw * hd)
.map(|i| ((i % 5) as f32 - 2.0) * 0.019)
.collect();
let (rel_h_bias, rel_w_bias) = decomposed_rel_pos_bias(&qh, &rh, &rw, gh, gw, hd);
for i in 0..n {
let qy = i / gw;
let qx = i % gw;
let qi = &qh[i * hd..(i + 1) * hd];
for j in 0..n {
let ky = j / gw;
let kx = j % gw;
let rh_base = (qy * gh + ky) * hd;
let rw_base = (qx * gw + kx) * hd;
let mut expected_h = 0.0f32;
let mut expected_w = 0.0f32;
for c in 0..hd {
expected_h += qi[c] * rh[rh_base + c];
expected_w += qi[c] * rw[rw_base + c];
}
assert_eq!(rel_h_bias[i * gh + ky], expected_h);
assert_eq!(rel_w_bias[i * gw + kx], expected_w);
}
}
}
fn tiny_block(window: usize) -> BlockP {
let dim = EMBED_DIM;
let hd = HEAD_DIM;
let size = 2;
let rel_rows = 2 * size - 1;
BlockP {
norm1: LayerNormP {
w: vec![1.0; dim],
b: vec![0.0; dim],
},
attn: AttnP {
qkv: test_linear(identity_block_3(dim), vec![0.0; 3 * dim], 3 * dim, dim),
proj: test_linear(identity_mat(dim), vec![0.0; dim], dim, dim),
rel_pos_h: vec![0.0; rel_rows * hd],
rel_pos_w: vec![0.0; rel_rows * hd],
size_h: size,
size_w: size,
},
norm2: LayerNormP {
w: vec![1.0; dim],
b: vec![0.0; dim],
},
lin1: test_linear(
vec![0.0; MLP_HIDDEN * dim],
vec![0.0; MLP_HIDDEN],
MLP_HIDDEN,
dim,
),
lin2: test_linear(vec![0.0; dim * MLP_HIDDEN], vec![0.0; dim], dim, MLP_HIDDEN),
window,
}
}
fn identity_mat(dim: usize) -> Vec<f32> {
let mut w = vec![0.0f32; dim * dim];
for i in 0..dim {
w[i * dim + i] = 1.0;
}
w
}
#[test]
fn row_chunked_block_mlp_is_bit_identical() -> FocrResult<()> {
let dim = EMBED_DIM;
let mut blk = tiny_block(0);
blk.lin1 = test_linear(
(0..MLP_HIDDEN * dim)
.map(|i| (((i * 31 + 7) % 101) as f32 - 50.0) * 3.0e-4)
.collect(),
(0..MLP_HIDDEN)
.map(|i| ((i % 9) as f32 - 4.0) * 0.01)
.collect(),
MLP_HIDDEN,
dim,
);
blk.lin2 = test_linear(
(0..dim * MLP_HIDDEN)
.map(|i| (((i * 17 + 3) % 83) as f32 - 41.0) * 4.0e-4)
.collect(),
(0..dim).map(|i| ((i % 6) as f32 - 3.0) * 0.02).collect(),
dim,
MLP_HIDDEN,
);
let rows = 13usize;
let x = Mat::from_vec(
rows,
dim,
(0..rows * dim)
.map(|i| ((i % 23) as f32 - 11.0) * 0.004)
.collect(),
);
let mut whole = blk.lin1.apply(&x)?;
nn::gelu(&mut whole);
let whole = blk.lin2.apply(&whole)?;
for chunk in [1usize, 2, 5, 13, 512] {
let chunked = mlp_row_chunked_with(&blk, &x, chunk)?;
assert_eq!(chunked.shape(), whole.shape());
assert_eq!(
chunked
.data
.iter()
.map(|f| f.to_bits())
.collect::<Vec<u32>>(),
whole.data.iter().map(|f| f.to_bits()).collect::<Vec<u32>>(),
"chunk {chunk}: row-chunked MLP must be bit-identical"
);
}
assert!(
mlp_row_chunked_with(&blk, &x, 0).is_err(),
"a zero chunk height must be rejected, not loop forever"
);
Ok(())
}
#[test]
fn bounded_global_attention_is_bit_identical() -> FocrResult<()> {
let dim = EMBED_DIM;
let hd = HEAD_DIM;
for &(grid, slabs) in &[(6usize, [5usize, 16, 36, 512]), (8, [3, 16, 64, 512])] {
let n = grid * grid;
let rel_rows = 2 * grid - 1;
let qkv_w: Vec<f32> = (0..3 * dim * dim)
.map(|i| (((i * 37 + 11) % 97) as f32 - 48.0) * 4.0e-4)
.collect();
let qkv_b: Vec<f32> = (0..3 * dim)
.map(|i| ((i % 7) as f32 - 3.0) * 0.01)
.collect();
let proj_w: Vec<f32> = (0..dim * dim)
.map(|i| (((i * 53 + 5) % 89) as f32 - 44.0) * 5.0e-4)
.collect();
let proj_b: Vec<f32> = (0..dim).map(|i| ((i % 5) as f32 - 2.0) * 0.02).collect();
let attn = AttnP {
qkv: test_linear(qkv_w, qkv_b, 3 * dim, dim),
proj: test_linear(proj_w, proj_b, dim, dim),
rel_pos_h: (0..rel_rows * hd)
.map(|i| ((i % 13) as f32 - 6.0) * 0.0011)
.collect(),
rel_pos_w: (0..rel_rows * hd)
.map(|i| ((i % 11) as f32 - 5.0) * 0.0009)
.collect(),
size_h: grid,
size_w: grid,
};
let x = Mat::from_vec(
n,
dim,
(0..n * dim)
.map(|i| ((i % 29) as f32 - 14.0) * 0.003)
.collect(),
);
let full = attention(&attn, &x, grid, grid, None)?;
for &slab in &slabs {
let bounded = attention_global_bounded_with(&attn, &x, grid, grid, slab)?;
assert_eq!(bounded.shape(), full.shape());
assert_eq!(
bounded
.data
.iter()
.map(|f| f.to_bits())
.collect::<Vec<u32>>(),
full.data.iter().map(|f| f.to_bits()).collect::<Vec<u32>>(),
"grid {grid} slab {slab}: bounded global attention must be bit-identical"
);
}
assert!(full.data.iter().any(|&v| v != 0.0));
}
Ok(())
}
#[test]
fn sam_block_from_matches_whole_tower_hydration() -> FocrResult<()> {
use crate::quant::focrq::FocrqBuilder;
let prefix = "model.sam_model";
let dim = 24usize; let mut b = FocrqBuilder::new();
test_support::add_synth_sam_tower(&mut b, prefix, dim);
let weights = Weights::from_bytes(b.build()).expect("synthetic SAM parses");
let whole = sam_weights_from(&weights, prefix)?;
assert_eq!(whole.blocks.len(), DEPTH);
for i in 0..DEPTH {
let solo = sam_block_from(&weights, prefix, i)?;
let cached = &whole.blocks[i];
assert_eq!(solo.window, cached.window, "block {i} window");
assert_eq!(solo.norm1.w, cached.norm1.w, "block {i} norm1.w");
assert_eq!(
solo.attn.qkv.wt.data, cached.attn.qkv.wt.data,
"block {i} qkv"
);
assert_eq!(solo.attn.qkv.b, cached.attn.qkv.b, "block {i} qkv bias");
assert_eq!(
solo.attn.proj.wt.data, cached.attn.proj.wt.data,
"block {i} proj"
);
assert_eq!(
solo.attn.rel_pos_h, cached.attn.rel_pos_h,
"block {i} rel_pos_h"
);
assert_eq!(solo.lin1.wt.data, cached.lin1.wt.data, "block {i} lin1");
assert_eq!(solo.lin2.wt.data, cached.lin2.wt.data, "block {i} lin2");
}
let head = sam_head_from(&weights, prefix)?;
assert!(head.blocks.is_empty());
assert_eq!(head.patch_embed.w, whole.patch_embed.w);
assert_eq!(head.pos_embed, whole.pos_embed);
assert_eq!(head.net3.w, whole.net3.w);
Ok(())
}
fn identity_block_3(dim: usize) -> Vec<f32> {
let mut w = vec![0.0f32; 3 * dim * dim];
for s in 0..3 {
for i in 0..dim {
let row = s * dim + i;
w[row * dim + i] = 1.0;
}
}
w
}
fn attention_scalar_reference(p: &AttnP, x: &Mat, gh: usize, gw: usize) -> FocrResult<Mat> {
let n = gh * gw;
let dim = x.cols;
let nh = NUM_HEADS;
let hd = dim / nh;
let scale = (hd as f32).powf(-0.5);
let qkv = p.qkv.apply(x)?;
let mut q = vec![0.0f32; nh * n * hd];
let mut k = vec![0.0f32; nh * n * hd];
let mut v = vec![0.0f32; nh * n * hd];
for r in 0..n {
let row = qkv.row(r);
for head in 0..nh {
for d in 0..hd {
q[(head * n + r) * hd + d] = row[head * hd + d];
k[(head * n + r) * hd + d] = row[(nh + head) * hd + d];
v[(head * n + r) * hd + d] = row[(2 * nh + head) * hd + d];
}
}
}
let rh = get_rel_pos(gh, gh, &p.rel_pos_h, p.size_h, hd);
let rw = get_rel_pos(gw, gw, &p.rel_pos_w, p.size_w, hd);
let mut out = vec![0.0f32; nh * n * hd];
for head in 0..nh {
let qh = &q[head * n * hd..(head + 1) * n * hd];
let kh = &k[head * n * hd..(head + 1) * n * hd];
let vh = &v[head * n * hd..(head + 1) * n * hd];
let (rel_h_bias, rel_w_bias) = decomposed_rel_pos_bias(qh, &rh, &rw, gh, gw, hd);
let mut logits = vec![0.0f32; n * n];
for i in 0..n {
let qi = &qh[i * hd..(i + 1) * hd];
for j in 0..n {
let ky = j / gw;
let kx = j % gw;
let kj = &kh[j * hd..(j + 1) * hd];
let mut dot = 0.0f32;
for c in 0..hd {
dot += qi[c] * kj[c];
}
logits[i * n + j] =
scale * dot + rel_h_bias[i * gh + ky] + rel_w_bias[i * gw + kx];
}
}
let mut lm = Mat::from_vec(n, n, logits);
nn::softmax_rows(&mut lm)?;
for i in 0..n {
let probs = lm.row(i);
let o = &mut out[(head * n + i) * hd..(head * n + i + 1) * hd];
for (j, &pj) in probs.iter().enumerate() {
let vj = &vh[j * hd..(j + 1) * hd];
for c in 0..hd {
o[c] += pj * vj[c];
}
}
}
}
let mut ctx = vec![0.0f32; n * dim];
for head in 0..nh {
for r in 0..n {
let src = (head * n + r) * hd;
let dst = r * dim + head * hd;
ctx[dst..dst + hd].copy_from_slice(&out[src..src + hd]);
}
}
p.proj.apply(&Mat::from_vec(n, dim, ctx))
}
#[test]
fn attention_gemm_matches_scalar_reference_with_relpos() -> FocrResult<()> {
let dim = EMBED_DIM;
let hd = HEAD_DIM;
let grid = 3usize;
let n = grid * grid;
let rel_rows = 2 * grid - 1;
let attn = AttnP {
qkv: test_linear(identity_block_3(dim), vec![0.0; 3 * dim], 3 * dim, dim),
proj: test_linear(identity_mat(dim), vec![0.0; dim], dim, dim),
rel_pos_h: (0..rel_rows * hd)
.map(|i| ((i % 13) as f32 - 6.0) * 0.0011)
.collect(),
rel_pos_w: (0..rel_rows * hd)
.map(|i| ((i % 11) as f32 - 5.0) * 0.0009)
.collect(),
size_h: grid,
size_w: grid,
};
let x = Mat::from_vec(
n,
dim,
(0..n * dim)
.map(|i| ((i % 29) as f32 - 14.0) * 0.003)
.collect(),
);
let got = attention(&attn, &x, grid, grid, None)?;
let expected = attention_scalar_reference(&attn, &x, grid, grid)?;
let max_abs = got
.data
.iter()
.zip(expected.data.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(max_abs <= 2.0e-6, "max_abs={max_abs}");
Ok(())
}
#[test]
fn attention_rejects_malformed_qkv_and_projection_shapes() {
let dim = EMBED_DIM;
let x = Mat::from_vec(4, dim, vec![0.0; 4 * dim]);
let mut bad_qkv = tiny_block(0).attn;
bad_qkv.qkv = test_linear(
vec![0.0; (3 * dim - 1) * dim],
vec![0.0; 3 * dim - 1],
3 * dim - 1,
dim,
);
assert_err_contains(attention(&bad_qkv, &x, 2, 2, None), "qkv shape");
let mut bad_proj = tiny_block(0).attn;
bad_proj.proj = test_linear(vec![0.0; (dim - 1) * dim], vec![0.0; dim - 1], dim - 1, dim);
assert_err_contains(attention(&bad_proj, &x, 2, 2, None), "proj shape");
}
#[test]
fn block_forward_preserves_shape_global() -> FocrResult<()> {
let blk = tiny_block(0);
let n = 4;
let dim = EMBED_DIM;
let mut data = vec![0.0f32; n * dim];
for (i, v) in data.iter_mut().enumerate() {
*v = ((i % 7) as f32) * 0.1 - 0.3;
}
let x = Mat::from_vec(n, dim, data);
let out = block_forward(&blk, &x, 2, 2)?;
assert_eq!(out.shape(), (n, dim));
Ok(())
}
#[test]
fn block_forward_windowed_pads_and_unpartitions() -> FocrResult<()> {
let blk = tiny_block(3);
let n = 4;
let dim = EMBED_DIM;
let data: Vec<f32> = (0..n * dim).map(|i| (i as f32 % 5.0) * 0.01).collect();
let x = Mat::from_vec(n, dim, data);
let out = block_forward(&blk, &x, 2, 2)?;
assert_eq!(out.shape(), (n, dim));
Ok(())
}
#[test]
fn block_forward_rejects_mlp_residual_shape_mismatch() {
let mut blk = tiny_block(0);
blk.lin2 = test_linear(
vec![0.0; (EMBED_DIM - 1) * MLP_HIDDEN],
vec![0.0; EMBED_DIM - 1],
EMBED_DIM - 1,
MLP_HIDDEN,
);
let n = 4;
let dim = EMBED_DIM;
let data: Vec<f32> = (0..n * dim).map(|i| (i as f32 % 5.0) * 0.01).collect();
let x = Mat::from_vec(n, dim, data);
assert_err_contains(
block_forward(&blk, &x, 2, 2),
"vision_sam block mlp residual",
);
}
#[test]
fn attention_zero_relpos_is_uniform_average_for_equal_q() -> FocrResult<()> {
let dim = EMBED_DIM;
let hd = HEAD_DIM;
let size = 2;
let rel_rows = 2 * size - 1;
let attn = AttnP {
qkv: test_linear(identity_block_3(dim), vec![0.0; 3 * dim], 3 * dim, dim),
proj: test_linear(identity_mat(dim), vec![0.0; dim], dim, dim),
rel_pos_h: vec![0.0; rel_rows * hd],
rel_pos_w: vec![0.0; rel_rows * hd],
size_h: size,
size_w: size,
};
let x = Mat::from_vec(4, dim, vec![1.0; 4 * dim]);
let out = attention(&attn, &x, 2, 2, None)?;
assert_eq!(out.shape(), (4, dim));
for &v in &out.data {
assert!((v - 1.0).abs() < 1e-4, "got {v}");
}
Ok(())
}
#[test]
#[ignore = "local perf probe; run explicitly with --ignored --nocapture"]
fn sam_attention_relpos_bias_local_probe() -> FocrResult<()> {
let dim = EMBED_DIM;
let hd = HEAD_DIM;
let grid = 14usize;
let n = grid * grid;
let rel_rows = 2 * grid - 1;
let attn = AttnP {
qkv: test_linear(identity_block_3(dim), vec![0.0; 3 * dim], 3 * dim, dim),
proj: test_linear(identity_mat(dim), vec![0.0; dim], dim, dim),
rel_pos_h: (0..rel_rows * hd)
.map(|i| ((i % 17) as f32 - 8.0) * 0.0007)
.collect(),
rel_pos_w: (0..rel_rows * hd)
.map(|i| ((i % 19) as f32 - 9.0) * 0.0005)
.collect(),
size_h: grid,
size_w: grid,
};
let x = Mat::from_vec(
n,
dim,
(0..n * dim)
.map(|i| ((i % 31) as f32 - 15.0) * 0.002)
.collect(),
);
let runs = std::env::var("FOCR_SAM_ATTN_PROBE_RUNS")
.ok()
.and_then(|s| s.parse::<usize>().ok())
.unwrap_or(3)
.max(1);
let warm = attention(&attn, &x, grid, grid, None)?;
let warm_checksum: f32 = warm.data.iter().step_by(97).copied().sum();
let start = Instant::now();
let mut checksum = 0.0f32;
for _ in 0..runs {
let out = attention(&attn, &x, grid, grid, None)?;
checksum += out.data.iter().step_by(97).copied().sum::<f32>();
}
let elapsed = start.elapsed();
let total_ms = elapsed.as_secs_f64() * 1000.0;
let avg_ms = total_ms / runs as f64;
assert!(checksum.is_finite());
println!(
"{}",
json!({
"probe": "sam_attention_relpos_bias_local_probe",
"grid": grid,
"tokens": n,
"dim": dim,
"heads": NUM_HEADS,
"head_dim": hd,
"runs": runs,
"total_ms": total_ms,
"avg_ms": avg_ms,
"warm_checksum": warm_checksum,
"checksum": checksum
})
);
Ok(())
}
#[test]
fn forward_with_end_to_end_shapes() -> FocrResult<()> {
let h = 32;
let gh = h / PATCH; let gw = gh;
let patch_embed = Conv {
w: vec![0.0; EMBED_DIM * 3 * PATCH * PATCH],
b: Some(vec![0.0; EMBED_DIM]),
out_ch: EMBED_DIM,
in_ch: 3,
kh: PATCH,
kw: PATCH,
};
let blocks: Vec<BlockP> = (0..DEPTH)
.map(|i| {
let window = if GLOBAL_BLOCKS.contains(&i) {
0
} else {
WINDOW
};
let size = if window == 0 { gh } else { window };
let rel_rows = 2 * size - 1;
let dim = EMBED_DIM;
let hd = HEAD_DIM;
BlockP {
norm1: LayerNormP {
w: vec![1.0; dim],
b: vec![0.0; dim],
},
attn: AttnP {
qkv: test_linear(
vec![0.0; 3 * dim * dim],
vec![0.0; 3 * dim],
3 * dim,
dim,
),
proj: test_linear(vec![0.0; dim * dim], vec![0.0; dim], dim, dim),
rel_pos_h: vec![0.0; rel_rows * hd],
rel_pos_w: vec![0.0; rel_rows * hd],
size_h: size,
size_w: size,
},
norm2: LayerNormP {
w: vec![1.0; dim],
b: vec![0.0; dim],
},
lin1: test_linear(
vec![0.0; MLP_HIDDEN * dim],
vec![0.0; MLP_HIDDEN],
MLP_HIDDEN,
dim,
),
lin2: test_linear(vec![0.0; dim * MLP_HIDDEN], vec![0.0; dim], dim, MLP_HIDDEN),
window,
}
})
.collect();
let w = SamWeights {
patch_embed,
pos_embed: vec![0.0; gh * gw * EMBED_DIM],
pos_grid_h: gh,
pos_grid_w: gw,
blocks,
neck_conv1: Conv {
w: vec![0.0; NECK_CH * EMBED_DIM],
b: None,
out_ch: NECK_CH,
in_ch: EMBED_DIM,
kh: 1,
kw: 1,
},
neck_ln1: LayerNormP {
w: vec![1.0; NECK_CH],
b: vec![0.0; NECK_CH],
},
neck_conv2: Conv {
w: vec![0.0; NECK_CH * NECK_CH * 9],
b: None,
out_ch: NECK_CH,
in_ch: NECK_CH,
kh: 3,
kw: 3,
},
neck_ln2: LayerNormP {
w: vec![1.0; NECK_CH],
b: vec![0.0; NECK_CH],
},
net2: Conv {
w: vec![0.0; NET2_CH * NECK_CH * 9],
b: None,
out_ch: NET2_CH,
in_ch: NECK_CH,
kh: 3,
kw: 3,
},
net3: Conv {
w: vec![0.0; OUT_CH * NET2_CH * 9],
b: None,
out_ch: OUT_CH,
in_ch: NET2_CH,
kh: 3,
kw: 3,
},
};
let image = Mat::from_vec(3, h * h, vec![0.5; 3 * h * h]);
let out = forward_with(&w, &image, h, h)?;
assert_eq!(out.shape(), (OUT_CH, 1));
assert!(out.data.iter().all(|&v| v.abs() < 1e-6));
Ok(())
}
#[test]
fn forward_with_rejects_bad_channels() {
let w = tiny_weights_minimal();
let bad = Mat::from_vec(2, 32 * 32, vec![0.0; 2 * 32 * 32]);
assert!(forward_with(&w, &bad, 32, 32).is_err());
}
#[test]
fn forward_with_rejects_non_patch_multiple() {
let w = tiny_weights_minimal();
let img = Mat::from_vec(3, 20 * 20, vec![0.0; 3 * 20 * 20]);
assert!(forward_with(&w, &img, 20, 20).is_err());
}
#[test]
fn forward_with_rejects_zero_spatial_dims_before_conv() {
let w = tiny_weights_minimal();
let img = Mat::from_vec(3, 0, Vec::new());
assert!(matches!(
forward_with(&w, &img, 0, 0),
Err(err) if err.to_string().contains("non-zero")
));
}
#[test]
fn forward_with_rejects_spatial_product_overflow_before_conv() {
let w = tiny_weights_minimal();
let img = Mat::from_vec(3, 0, Vec::new());
assert!(matches!(
forward_with(&w, &img, usize::MAX, 2),
Err(err) if err.to_string().contains("H*W")
));
}
fn tiny_weights_minimal() -> SamWeights {
let gh = 2;
let blocks: Vec<BlockP> = (0..DEPTH)
.map(|i| {
let window = if GLOBAL_BLOCKS.contains(&i) {
0
} else {
WINDOW
};
let size = if window == 0 { gh } else { window };
let rel_rows = 2 * size - 1;
let dim = EMBED_DIM;
let hd = HEAD_DIM;
BlockP {
norm1: LayerNormP {
w: vec![1.0; dim],
b: vec![0.0; dim],
},
attn: AttnP {
qkv: test_linear(
vec![0.0; 3 * dim * dim],
vec![0.0; 3 * dim],
3 * dim,
dim,
),
proj: test_linear(vec![0.0; dim * dim], vec![0.0; dim], dim, dim),
rel_pos_h: vec![0.0; rel_rows * hd],
rel_pos_w: vec![0.0; rel_rows * hd],
size_h: size,
size_w: size,
},
norm2: LayerNormP {
w: vec![1.0; dim],
b: vec![0.0; dim],
},
lin1: test_linear(
vec![0.0; MLP_HIDDEN * dim],
vec![0.0; MLP_HIDDEN],
MLP_HIDDEN,
dim,
),
lin2: test_linear(vec![0.0; dim * MLP_HIDDEN], vec![0.0; dim], dim, MLP_HIDDEN),
window,
}
})
.collect();
SamWeights {
patch_embed: Conv {
w: vec![0.0; EMBED_DIM * 3 * PATCH * PATCH],
b: Some(vec![0.0; EMBED_DIM]),
out_ch: EMBED_DIM,
in_ch: 3,
kh: PATCH,
kw: PATCH,
},
pos_embed: vec![0.0; gh * gh * EMBED_DIM],
pos_grid_h: gh,
pos_grid_w: gh,
blocks,
neck_conv1: Conv {
w: vec![0.0; NECK_CH * EMBED_DIM],
b: None,
out_ch: NECK_CH,
in_ch: EMBED_DIM,
kh: 1,
kw: 1,
},
neck_ln1: LayerNormP {
w: vec![1.0; NECK_CH],
b: vec![0.0; NECK_CH],
},
neck_conv2: Conv {
w: vec![0.0; NECK_CH * NECK_CH * 9],
b: None,
out_ch: NECK_CH,
in_ch: NECK_CH,
kh: 3,
kw: 3,
},
neck_ln2: LayerNormP {
w: vec![1.0; NECK_CH],
b: vec![0.0; NECK_CH],
},
net2: Conv {
w: vec![0.0; NET2_CH * NECK_CH * 9],
b: None,
out_ch: NET2_CH,
in_ch: NECK_CH,
kh: 3,
kw: 3,
},
net3: Conv {
w: vec![0.0; OUT_CH * NET2_CH * 9],
b: None,
out_ch: OUT_CH,
in_ch: NET2_CH,
kh: 3,
kw: 3,
},
}
}
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) -> Linear {
test_linear(
(0..out * inn).map(|i| det(i, salt)).collect(),
(0..out).map(|i| det(i, salt + 1)).collect(),
out,
inn,
)
}
fn rand_ln(dim: usize, salt: usize) -> LayerNormP {
LayerNormP {
w: (0..dim).map(|i| 1.0 + det(i, salt)).collect(),
b: (0..dim).map(|i| det(i, salt + 2)).collect(),
}
}
fn rand_conv(
out_ch: usize,
in_ch: usize,
kh: usize,
kw: usize,
bias: bool,
salt: usize,
) -> Conv {
Conv {
w: (0..out_ch * in_ch * kh * kw)
.map(|i| det(i, salt))
.collect(),
b: if bias {
Some((0..out_ch).map(|i| det(i, salt + 1)).collect())
} else {
None
},
out_ch,
in_ch,
kh,
kw,
}
}
fn tiny_weights_nontrivial(gh: usize, gw: usize) -> SamWeights {
let dim = EMBED_DIM;
let hd = HEAD_DIM;
let blocks: Vec<BlockP> = (0..DEPTH)
.map(|i| {
let window = if GLOBAL_BLOCKS.contains(&i) {
0
} else {
WINDOW
};
let size_h = if window == 0 { gh } else { window };
let size_w = if window == 0 { gw } else { window };
let rel_rows_h = 2 * size_h - 1;
let rel_rows_w = 2 * size_w - 1;
BlockP {
norm1: rand_ln(dim, 100 + i * 13),
attn: AttnP {
qkv: rand_lin(3 * dim, dim, 200 + i * 13),
proj: rand_lin(dim, dim, 300 + i * 13),
rel_pos_h: (0..rel_rows_h * hd).map(|j| det(j, 700 + i * 13)).collect(),
rel_pos_w: (0..rel_rows_w * hd).map(|j| det(j, 800 + i * 13)).collect(),
size_h,
size_w,
},
norm2: rand_ln(dim, 400 + i * 13),
lin1: rand_lin(MLP_HIDDEN, dim, 500 + i * 13),
lin2: rand_lin(dim, MLP_HIDDEN, 600 + i * 13),
window,
}
})
.collect();
SamWeights {
patch_embed: rand_conv(EMBED_DIM, 3, PATCH, PATCH, true, 11),
pos_embed: (0..gh * gw * EMBED_DIM).map(|i| det(i, 13)).collect(),
pos_grid_h: gh,
pos_grid_w: gw,
blocks,
neck_conv1: rand_conv(NECK_CH, EMBED_DIM, 1, 1, false, 21),
neck_ln1: rand_ln(NECK_CH, 23),
neck_conv2: rand_conv(NECK_CH, NECK_CH, 3, 3, false, 25),
neck_ln2: rand_ln(NECK_CH, 27),
net2: rand_conv(NET2_CH, NECK_CH, 3, 3, false, 29),
net3: rand_conv(OUT_CH, NET2_CH, 3, 3, false, 31),
}
}
fn nontrivial_image(h: usize, win: usize, salt: usize) -> Mat {
let cols = h * win;
Mat::from_vec(
3,
cols,
(0..3 * cols)
.map(|i| ((((i + salt * 7919) as f32) * 0.0007).sin()) * 0.5)
.collect(),
)
}
#[test]
fn batched_sam_equals_per_view_byte_for_byte() -> FocrResult<()> {
let h = 240;
let win = 240;
let gh = h / PATCH;
let gw = win / PATCH;
let w = tiny_weights_nontrivial(gh, gw);
let v0 = nontrivial_image(h, win, 1);
let v1 = nontrivial_image(h, win, 2);
let v2 = nontrivial_image(h, win, 3);
let batched = forward_with_batched(&w, &[&v0, &v1, &v2], h, win)?;
assert_eq!(batched.len(), 3);
for (i, view) in [&v0, &v1, &v2].iter().enumerate() {
let seq = forward_with(&w, view, h, win)?;
assert_eq!(batched[i].shape(), seq.shape(), "view {i} shape");
assert_eq!(
batched[i].data, seq.data,
"view {i}: batched SAM != per-view sequential (cross-view leak or M-dependence)"
);
}
Ok(())
}
#[test]
fn batched_sam_single_view_equals_forward_with() -> FocrResult<()> {
let h = 240;
let win = 240;
let gh = h / PATCH;
let gw = win / PATCH;
let w = tiny_weights_nontrivial(gh, gw);
let view = nontrivial_image(h, win, 5);
let batched = forward_with_batched(&w, &[&view], h, win)?;
let seq = forward_with(&w, &view, h, win)?;
assert_eq!(batched.len(), 1);
assert_eq!(batched[0].shape(), seq.shape());
assert_eq!(batched[0].data, seq.data);
Ok(())
}
#[test]
fn streamed_views_inner_equals_views_outer_byte_for_byte() -> FocrResult<()> {
let h = 240;
let win = 240;
let gh = h / PATCH;
let gw = win / PATCH;
let w = tiny_weights_nontrivial(gh, gw);
let mut head = tiny_weights_nontrivial(gh, gw);
head.blocks = Vec::new();
let blocks = w.blocks.clone();
let mut hydrate = |i: usize| -> FocrResult<std::borrow::Cow<'static, BlockP>> {
Ok(std::borrow::Cow::Owned(blocks[i].clone()))
};
let small = 128usize;
let views = [
nontrivial_image(h, win, 1),
nontrivial_image(h, win, 2),
nontrivial_image(small, small, 3),
];
let dims = [(h, win), (h, win), (small, small)];
let mut per_view = Vec::new();
for (view, &(vh, vw)) in views.iter().zip(dims.iter()) {
per_view.push(forward_core(
&head,
view,
vh,
vw,
DEPTH,
&mut hydrate,
true,
)?);
}
let refs: Vec<&Mat> = views.iter().collect();
let hoisted = forward_core_views(&head, &refs, &dims, DEPTH, &mut hydrate, true)?;
assert_eq!(hoisted.len(), per_view.len());
for (i, (outer, inner)) in per_view.iter().zip(hoisted.iter()).enumerate() {
assert_eq!(outer.shape(), inner.shape(), "view {i} shape");
assert_eq!(
outer.data.iter().map(|f| f.to_bits()).collect::<Vec<u32>>(),
inner.data.iter().map(|f| f.to_bits()).collect::<Vec<u32>>(),
"view {i}: hydration-hoisted SAM must be bit-identical to per-view"
);
}
assert!(per_view[0].data.iter().any(|&v| v != 0.0));
assert_ne!(per_view[0].data, per_view[1].data);
Ok(())
}
#[test]
fn batched_sam_rejects_ragged_and_empty() {
let h = 32;
let win = 32;
let w = tiny_weights_minimal(); assert!(forward_with_batched(&w, &[], h, win).is_err());
let a = nontrivial_image(h, win, 1);
let b = nontrivial_image(h, win + PATCH, 2); assert!(forward_with_batched(&w, &[&a, &b], h, win).is_err());
let bad = Mat::from_vec(2, h * win, vec![0.0; 2 * h * win]);
assert!(forward_with_batched(&w, &[&bad], h, win).is_err());
}
}