use ferrox_core::matmul::rms_norm;
use ferrox_core::weight_matrix::WeightMatrix;
#[derive(Debug, Clone, Copy)]
pub struct VisionConfig {
pub in_dim: usize,
pub patch_size: usize,
pub grid_h: usize,
pub grid_w: usize,
pub hidden_dim: usize,
pub num_heads: usize,
pub qkv_hidden: usize,
pub mlp_dim: usize,
pub rms_norm_eps: f32,
pub theta_base: f32,
pub merge_kh: usize,
pub merge_kw: usize,
pub projector_ln_eps: f32,
}
impl VisionConfig {
fn head_dim(&self) -> usize {
self.qkv_hidden / self.num_heads
}
fn n_patches(&self) -> usize {
self.grid_h * self.grid_w
}
}
pub struct VisionEncoderLayerWeights {
pub norm0_weight: Vec<f32>,
pub wqkv: WeightMatrix, pub wo: WeightMatrix, pub norm1_weight: Vec<f32>,
pub fc0: WeightMatrix, pub fc1: WeightMatrix, }
pub struct VisionEncoderWeights {
pub patch_embed: WeightMatrix, pub pos_emb: Vec<f32>,
pub layers: Vec<VisionEncoderLayerWeights>,
pub final_norm_weight: Vec<f32>,
}
pub struct VisionMergerWeights {
pub proj0: WeightMatrix, pub proj1: WeightMatrix, pub post_norm_weight: Vec<f32>, }
fn gelu_tanh(x: f32) -> f32 {
0.5 * x * (1.0 + (0.797_884_6 * (x + 0.044715 * x * x * x)).tanh())
}
fn erf(x: f32) -> f32 {
let sign = if x < 0.0 { -1.0 } else { 1.0 };
let x = x.abs();
let t = 1.0 / (1.0 + 0.3275911 * x);
let poly = ((((1.061_405_4 * t - 1.453_152_1) * t) + 1.421_413_8) * t - 0.284_496_72) * t
+ 0.254_829_6;
sign * (1.0 - poly * t * (-x * x).exp())
}
fn gelu_erf(x: f32) -> f32 {
0.5 * x * (1.0 + erf(x / std::f32::consts::SQRT_2))
}
fn precompute_rope_2d(cfg: &VisionConfig) -> (Vec<f32>, Vec<f32>) {
let head_dim = cfg.head_dim();
assert_eq!(head_dim % 4, 0, "vision head_dim must be divisible by 4");
let num_pairs = head_dim / 2;
let n = cfg.n_patches();
let mut cos_t = vec![0f32; n * num_pairs];
let mut sin_t = vec![0f32; n * num_pairs];
for p in 0..n {
let x_pos = (p % cfg.grid_w) as f32;
let y_pos = (p / cfg.grid_w) as f32;
for j in 0..num_pairs {
let i = j / 2;
let freq = cfg.theta_base.powf(-4.0 * i as f32 / head_dim as f32);
let angle = if j % 2 == 0 {
x_pos * freq
} else {
y_pos * freq
};
cos_t[p * num_pairs + j] = angle.cos();
sin_t[p * num_pairs + j] = angle.sin();
}
}
(cos_t, sin_t)
}
fn apply_rope_2d(x: &mut [f32], cfg: &VisionConfig, cos_t: &[f32], sin_t: &[f32]) {
let head_dim = cfg.head_dim();
let num_pairs = head_dim / 2;
for p in 0..cfg.n_patches() {
for h in 0..cfg.num_heads {
let base = (p * cfg.num_heads + h) * head_dim;
for j in 0..num_pairs {
let c = cos_t[p * num_pairs + j];
let s = sin_t[p * num_pairs + j];
let a = x[base + 2 * j];
let b = x[base + 2 * j + 1];
x[base + 2 * j] = a * c - b * s;
x[base + 2 * j + 1] = a * s + b * c;
}
}
}
}
pub fn embed_patches(
weights: &VisionEncoderWeights,
cfg: &VisionConfig,
patches: &[Vec<f32>],
) -> Vec<f32> {
assert_eq!(patches.len(), cfg.n_patches());
let mut out = Vec::with_capacity(cfg.n_patches() * cfg.hidden_dim);
for (p, patch) in patches.iter().enumerate() {
let embedded = weights.patch_embed.apply(patch);
let pos = &weights.pos_emb[p * cfg.hidden_dim..(p + 1) * cfg.hidden_dim];
for (e, po) in embedded.iter().zip(pos.iter()) {
out.push(e + po);
}
}
out
}
fn full_self_attention(
q: &[f32],
k: &[f32],
v: &[f32],
n_patches: usize,
num_heads: usize,
head_dim: usize,
) -> Vec<f32> {
let scale = 1.0 / (head_dim as f32).sqrt();
let mut out = vec![0f32; n_patches * num_heads * head_dim];
for h in 0..num_heads {
for i in 0..n_patches {
let q_i = &q[(i * num_heads + h) * head_dim..(i * num_heads + h + 1) * head_dim];
let mut scores = vec![0f32; n_patches];
for t in 0..n_patches {
let k_t = &k[(t * num_heads + h) * head_dim..(t * num_heads + h + 1) * head_dim];
let dot: f32 = q_i.iter().zip(k_t.iter()).map(|(a, b)| a * b).sum();
scores[t] = dot * scale;
}
let max = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0f32;
for s in scores.iter_mut() {
*s = (*s - max).exp();
sum += *s;
}
for s in scores.iter_mut() {
*s /= sum;
}
let out_base = (i * num_heads + h) * head_dim;
for t in 0..n_patches {
let v_t = &v[(t * num_heads + h) * head_dim..(t * num_heads + h + 1) * head_dim];
let w = scores[t];
for d in 0..head_dim {
out[out_base + d] += w * v_t[d];
}
}
}
}
out
}
fn encoder_layer_forward(
layer: &VisionEncoderLayerWeights,
cfg: &VisionConfig,
x: &[f32],
cos_t: &[f32],
sin_t: &[f32],
) -> Vec<f32> {
let n = cfg.n_patches();
let head_dim = cfg.head_dim();
let mut residual = x.to_vec();
let mut normed = Vec::with_capacity(n * cfg.hidden_dim);
for p in 0..n {
normed.extend(rms_norm(
&x[p * cfg.hidden_dim..(p + 1) * cfg.hidden_dim],
&layer.norm0_weight,
cfg.rms_norm_eps,
));
}
let mut q = vec![0f32; n * cfg.qkv_hidden];
let mut k = vec![0f32; n * cfg.qkv_hidden];
let mut v = vec![0f32; n * cfg.qkv_hidden];
for p in 0..n {
let qkv = layer
.wqkv
.apply(&normed[p * cfg.hidden_dim..(p + 1) * cfg.hidden_dim]);
q[p * cfg.qkv_hidden..(p + 1) * cfg.qkv_hidden].copy_from_slice(&qkv[0..cfg.qkv_hidden]);
k[p * cfg.qkv_hidden..(p + 1) * cfg.qkv_hidden]
.copy_from_slice(&qkv[cfg.qkv_hidden..2 * cfg.qkv_hidden]);
v[p * cfg.qkv_hidden..(p + 1) * cfg.qkv_hidden]
.copy_from_slice(&qkv[2 * cfg.qkv_hidden..3 * cfg.qkv_hidden]);
}
apply_rope_2d(&mut q, cfg, cos_t, sin_t);
apply_rope_2d(&mut k, cfg, cos_t, sin_t);
let attn_out = full_self_attention(&q, &k, &v, n, cfg.num_heads, head_dim);
for p in 0..n {
let attn_flat = &attn_out[p * cfg.qkv_hidden..(p + 1) * cfg.qkv_hidden];
let projected = layer.wo.apply(attn_flat);
let res = &mut residual[p * cfg.hidden_dim..(p + 1) * cfg.hidden_dim];
for (r, pr) in res.iter_mut().zip(projected.iter()) {
*r += pr;
}
}
let mut out = residual.clone();
for p in 0..n {
let res_p = &residual[p * cfg.hidden_dim..(p + 1) * cfg.hidden_dim];
let normed1 = rms_norm(res_p, &layer.norm1_weight, cfg.rms_norm_eps);
let mut hidden = layer.fc0.apply(&normed1);
for h in hidden.iter_mut() {
*h = gelu_tanh(*h);
}
let mlp_out = layer.fc1.apply(&hidden);
let out_p = &mut out[p * cfg.hidden_dim..(p + 1) * cfg.hidden_dim];
for (o, m) in out_p.iter_mut().zip(mlp_out.iter()) {
*o += m;
}
}
out
}
pub fn encoder_forward(
weights: &VisionEncoderWeights,
cfg: &VisionConfig,
patches: &[Vec<f32>],
) -> Vec<f32> {
let (cos_t, sin_t) = precompute_rope_2d(cfg);
let mut x = embed_patches(weights, cfg, patches);
for layer in &weights.layers {
x = encoder_layer_forward(layer, cfg, &x, &cos_t, &sin_t);
}
let n = cfg.n_patches();
let mut out = Vec::with_capacity(n * cfg.hidden_dim);
for p in 0..n {
out.extend(rms_norm(
&x[p * cfg.hidden_dim..(p + 1) * cfg.hidden_dim],
&weights.final_norm_weight,
cfg.rms_norm_eps,
));
}
out
}
pub fn patch_merge(encoder_out: &[f32], cfg: &VisionConfig) -> Vec<f32> {
assert_eq!(cfg.grid_h % cfg.merge_kh, 0);
assert_eq!(cfg.grid_w % cfg.merge_kw, 0);
let new_h = cfg.grid_h / cfg.merge_kh;
let new_w = cfg.grid_w / cfg.merge_kw;
let merge_block = cfg.merge_kh * cfg.merge_kw;
let mut out = vec![0f32; new_h * new_w * merge_block * cfg.hidden_dim];
for nh in 0..new_h {
for nw in 0..new_w {
for kh in 0..cfg.merge_kh {
for kw in 0..cfg.merge_kw {
let orig_h = nh * cfg.merge_kh + kh;
let orig_w = nw * cfg.merge_kw + kw;
let orig_patch = orig_h * cfg.grid_w + orig_w;
let merged_idx = nh * new_w + nw;
let sub_idx = kh * cfg.merge_kw + kw;
let src = &encoder_out
[orig_patch * cfg.hidden_dim..(orig_patch + 1) * cfg.hidden_dim];
let dst_base = (merged_idx * merge_block + sub_idx) * cfg.hidden_dim;
out[dst_base..dst_base + cfg.hidden_dim].copy_from_slice(src);
}
}
}
}
out
}
pub fn project_merged_patches(
weights: &VisionMergerWeights,
cfg: &VisionConfig,
merged: &[f32],
num_merged: usize,
) -> Vec<f32> {
let merge_hidden = cfg.merge_kh * cfg.merge_kw * cfg.hidden_dim;
let text_hidden = weights.post_norm_weight.len();
let mut out = Vec::with_capacity(num_merged * text_hidden);
for m in 0..num_merged {
let input = &merged[m * merge_hidden..(m + 1) * merge_hidden];
let mut hidden = weights.proj0.apply(input);
for h in hidden.iter_mut() {
*h = gelu_erf(*h);
}
let projected = weights.proj1.apply(&hidden);
out.extend(rms_norm(
&projected,
&weights.post_norm_weight,
cfg.projector_ln_eps,
));
}
out
}
#[cfg(test)]
mod tests {
use super::*;
use ferrox_core::tensor::Tensor;
const IN_DIM: usize = 3;
const PATCH_SIZE: usize = 2;
const GRID_H: usize = 2;
const GRID_W: usize = 2;
const HIDDEN_DIM: usize = 8;
const NUM_HEADS: usize = 2;
const QKV_HIDDEN: usize = 8;
const MLP_DIM: usize = 12;
const NORM_EPS: f32 = 1e-5;
const THETA_BASE: f32 = 10000.0;
const MERGE_KH: usize = 2;
const MERGE_KW: usize = 2;
const TEXT_HIDDEN: usize = 6;
const PROJECTOR_LN_EPS: f32 = 1e-5;
const VIS_PATCH_0: [f32; 12] = [
-0.396561, 0.120286, -0.948163, 0.697886, 0.319147, -0.146024, -0.155975, 0.151918,
-0.13383, -0.112954, 0.360034, 0.257353,
];
const VIS_PATCH_1: [f32; 12] = [
-0.032064, -0.0427383, 0.0804582, -0.307009, -0.201875, 0.27413, -0.0652414, -0.687213,
-0.238639, 0.328311, -0.116141, -0.0743664,
];
const VIS_PATCH_2: [f32; 12] = [
0.320918, 0.912305, -0.356594, 0.674103, -0.615006, 0.0874888, -0.584765, 0.675729,
0.416961, 0.568858, -0.442767, 0.342278,
];
const VIS_PATCH_3: [f32; 12] = [
-0.259507, -0.228693, 0.253269, 0.438359, 0.10221, -0.313994, -0.412908, 0.722158,
0.296973, 0.359864, 1.09174, -0.407931,
];
const VIS_PATCH_EMBED_W: [f32; 96] = [
0.767837,
0.945273,
0.485524,
0.248132,
-0.199147,
0.298346,
-0.132808,
-0.00649524,
-0.08713,
0.085149,
0.386423,
-0.166675,
-0.295621,
-0.300887,
-0.290483,
-0.429331,
-0.27388,
0.38798,
-0.177994,
0.0771323,
-0.365068,
0.0508952,
-0.522234,
-0.209627,
0.676362,
-0.174891,
0.335993,
0.136516,
-0.0458957,
-0.195632,
0.386073,
-0.053215,
0.458227,
-0.215724,
0.0172017,
0.13965,
0.111948,
-0.37014,
-0.199219,
-0.0587941,
-0.25611,
0.203198,
0.176406,
-0.587125,
-0.541576,
-0.384469,
0.0351779,
0.609952,
-0.114707,
0.0751952,
-0.318934,
-0.314051,
-0.587168,
-0.00850327,
0.284165,
-0.107014,
0.418936,
0.0593565,
-0.0109237,
0.155176,
0.146213,
0.344342,
-0.240586,
-0.686415,
0.034424,
-0.183503,
-0.00815316,
0.49929,
-0.330853,
0.229439,
0.28376,
0.138221,
0.335552,
-0.137509,
-0.204453,
0.311695,
0.21609,
0.417113,
0.0636869,
0.487069,
-0.0846241,
-0.318099,
-0.611127,
-0.33042,
0.245551,
-0.439479,
-0.135709,
0.631975,
0.254039,
0.537295,
-0.297469,
-0.737993,
0.454914,
-0.425083,
0.0272771,
0.065128,
];
const VIS_POS_EMB_W: [f32; 32] = [
-0.189779,
0.310527,
0.309868,
0.110289,
-0.00589813,
0.0382874,
-0.219182,
-0.0478128,
-0.0209866,
-0.170447,
0.169843,
-0.285296,
-0.0906838,
-0.450576,
0.0993637,
0.148521,
0.124258,
0.583617,
0.187958,
-0.188398,
0.532708,
-0.184064,
0.157748,
0.148626,
-0.175945,
0.27907,
0.0709187,
0.114191,
-0.135147,
0.260527,
-0.172188,
-0.355832,
];
const VIS_L0_NORM0_W: [f32; 8] = [
1.14251, 0.957327, 1.00434, 1.19228, 1.11365, 1.06327, 1.10699, 0.927249,
];
const VIS_L0_WQKV: [f32; 192] = [
0.225927, 0.332981, 0.226627, -0.0612404, -1.12422, 0.167191, 0.0344629, -0.110866,
0.23343, 0.0797109, 0.0611867, -0.445077, -0.437797, 0.111161, 0.0228171, 0.0416223,
0.0498789, 0.123437, 0.0108197, 0.0473313, 0.0768343, 0.278823, -0.154407, 0.235919,
0.190915, 0.0316206, -0.142534, 0.110742, 0.259192, -0.322493, 0.0186437, 0.0194232,
-0.198361, -0.185786, -0.465344, -0.330938, -0.461243, -0.184972, -0.167996, -0.162386,
0.111553, -0.05627, -0.218674, -0.216565, -0.339453, -0.0588504, -0.0642164, 0.486887,
0.460246, 0.443371, 0.623773, -0.36512, 0.201895, 0.228316, 0.142865, 0.638001, 0.761363,
0.29381, 0.195944, -0.336719, -0.336108, 0.210067, -0.0226898, -0.399168, -0.0372073,
0.108002, 0.280778, -0.0323038, -0.0807652, 0.00186235, 0.051895, -0.34894, 0.249993,
0.657098, 0.24426, -0.434459, 0.0283526, -0.529688, -0.548764, 0.197775, 0.0232947,
-0.206133, -0.626234, -0.38068, 0.0486588, 0.00956812, 0.20691, 0.457018, -0.101777,
0.14069, -0.0413395, -0.148962, -0.0734372, 0.288226, -0.308526, 0.342312, 0.443186,
0.0107199, 0.0764755, -0.42354, 0.70816, 0.293666, 0.0516828, 0.0172313, 0.135292,
-0.195371, 0.0849785, -0.277073, 0.149196, -0.203464, 0.268656, -0.0361107, 0.0806238,
0.888267, 0.24127, 0.0803401, 0.0133166, -0.311749, 0.325266, 0.480702, 0.0193065,
0.503025, -0.0336833, -0.531916, -0.195972, -0.317098, -0.082362, 0.188094, -0.278692,
-0.0240735, 0.0598056, -0.219908, -0.272796, 0.0261548, 0.261711, -0.180968, -0.225418,
-0.042376, 0.0869846, -0.162905, -0.100244, -0.0693806, -0.325737, -0.0532611, -0.168014,
0.475017, -0.155197, 0.137133, 0.113856, -0.33541, -0.374685, -0.133646, 0.198772,
-0.171972, -0.196846, -0.870066, -0.119052, 0.0937437, 0.133406, 0.128758, -0.0983533,
-0.0681474, 0.249683, -0.0407133, -0.690913, 0.306003, -0.323304, 0.412587, 0.508244,
0.228732, 0.271235, 0.55734, -0.330637, 0.0761556, 0.0717198, -0.290133, 0.603544,
-0.174208, -0.219374, 0.0803143, 0.309335, -0.180314, -0.17987, 0.151233, -0.427684,
-0.231776, -0.251471, -0.1303, -0.0426799, 0.553263, 0.272208, 0.00671946,
];
const VIS_L0_WO: [f32; 64] = [
-0.0674106,
0.262305,
0.12852,
0.267363,
0.414293,
-0.302455,
-0.675857,
-0.00654075,
0.0238384,
0.0860337,
-0.179647,
-0.0697861,
0.167621,
0.183768,
0.399294,
-0.649884,
-0.494598,
0.438582,
0.18246,
0.23803,
-0.0854503,
-0.142247,
0.0599411,
-0.281936,
-0.459049,
-0.333171,
0.0349579,
0.191157,
0.495378,
0.111546,
0.141755,
0.506442,
-0.232997,
-0.287787,
-0.16138,
-0.0288672,
0.606971,
-0.104652,
0.392134,
-0.141798,
0.147425,
0.204395,
0.159886,
0.225314,
-0.206712,
0.436844,
0.294298,
-0.210564,
0.151729,
-0.321948,
0.0284701,
-0.351326,
0.21551,
0.889311,
0.441252,
0.219234,
-0.270861,
0.114077,
0.355367,
-0.0265149,
0.0387836,
-0.508303,
-0.256626,
-0.554506,
];
const VIS_L0_NORM1_W: [f32; 8] = [
1.19145, 1.04726, 1.07247, 0.97973, 0.853933, 1.0176, 0.936351, 0.903017,
];
const VIS_L0_FC0: [f32; 96] = [
0.0166233, 0.0329455, -0.0541478, -0.232281, -0.29576, 0.00189617, 0.366746, -0.481293,
-0.197155, -0.0487075, -0.0476214, -0.177829, 0.578552, 0.11994, -0.295615, -0.564639,
-0.051645, 0.443841, 0.0578301, -0.694797, -0.413083, 0.0368751, 0.320283, 0.0423223,
0.044788, 0.321899, 0.376876, -0.202959, 0.131414, -0.287215, -0.322971, 0.176502,
-0.427032, 0.0509592, 0.0942597, 0.258386, -0.345129, 0.114245, -0.0376285, 0.445818,
-0.222826, -0.333006, 0.150066, -0.256744, -0.388289, 0.273022, 0.129116, -0.933427,
-0.276295, -0.332411, 0.193304, 0.441769, 0.617435, 0.211838, -0.328804, -0.339382,
-0.285821, 0.146822, 0.379626, -0.575327, 0.182773, 0.278357, 0.0947908, 0.97087, 0.160575,
0.00974755, 0.183156, -0.163253, -0.32978, 0.126654, 0.346088, -0.0704812, 0.277716,
-0.112732, 0.489417, 0.0644369, 0.160294, 0.22985, -0.405623, -0.0546382, 0.552158,
0.0109478, 0.420782, 0.265227, -0.548319, 0.212211, -0.326486, 0.0783693, -0.12468,
0.00652852, -0.0212635, 0.0728004, -0.134494, -0.405292, -0.302986, 0.274872,
];
const VIS_L0_FC1: [f32; 96] = [
0.517348,
0.0697715,
0.252689,
0.293056,
-0.409747,
-0.561795,
-0.0153769,
-0.567693,
-0.0817456,
-0.443467,
-0.187211,
-0.225356,
0.155999,
-0.194173,
-0.344326,
0.0146415,
0.890155,
0.147833,
-0.583478,
0.0273942,
0.771329,
0.0507425,
0.3842,
0.333034,
0.356512,
0.315238,
0.0847949,
-0.0552949,
-0.0835648,
0.0360814,
0.12811,
-0.382702,
-0.0841374,
-0.0121686,
0.0323577,
-0.0929058,
0.241921,
-0.0311971,
-0.112632,
-0.371058,
-0.23322,
0.0688837,
0.0357527,
-0.611612,
0.277291,
-0.210783,
-0.270577,
0.334413,
0.349005,
-0.213423,
-0.106685,
-0.367216,
0.244268,
0.271403,
-0.131644,
0.417505,
0.275575,
-0.0371131,
0.216475,
-0.496402,
-0.0788385,
0.075497,
-0.00312696,
0.614985,
0.0931599,
0.109938,
-0.398398,
-0.0562463,
-0.235057,
0.70463,
-0.124689,
0.0357173,
0.364742,
0.197,
-0.423393,
0.191556,
0.213094,
-0.0544942,
0.139289,
0.0314193,
-0.415614,
-0.16905,
0.00621485,
-0.175932,
0.161714,
0.300591,
0.082026,
-0.111803,
0.408852,
0.041258,
0.106057,
-0.229034,
-0.0652944,
0.299431,
0.223119,
-0.0561207,
];
const VIS_L1_NORM0_W: [f32; 8] = [
1.00082, 1.0199, 0.905881, 0.981468, 0.948844, 0.861423, 1.01522, 0.98562,
];
const VIS_L1_WQKV: [f32; 192] = [
-0.439142,
-0.238502,
-0.0261269,
-0.867316,
0.055201,
0.080017,
0.19467,
-0.178583,
-0.426824,
0.188486,
-0.412493,
0.162284,
-0.202711,
-0.303652,
0.0883011,
0.336673,
-0.047082,
-0.0737218,
-0.506555,
0.249362,
0.0218581,
-0.0815053,
0.392514,
-0.470229,
-0.151384,
0.00145506,
-0.0886644,
-0.73097,
0.104801,
-0.00915,
0.30323,
0.61196,
-0.0563681,
0.115342,
0.075918,
0.0509361,
-0.0414925,
0.115973,
0.338449,
-0.254719,
0.570279,
0.151889,
0.0908936,
-0.674739,
-0.00318789,
-0.292769,
0.172313,
0.286671,
0.344964,
-0.049921,
-0.0223788,
0.131137,
-0.62942,
-0.130129,
-0.162834,
0.156245,
-0.137213,
0.223168,
-0.0578504,
0.256585,
0.665501,
0.467405,
0.284134,
0.343418,
-0.109782,
-0.231188,
-0.253781,
0.0587258,
0.44202,
-0.218814,
0.1406,
-0.374162,
-0.364904,
0.257287,
-0.206018,
-0.214628,
0.592636,
0.0206818,
-0.0133439,
-0.0688817,
0.124602,
-0.212348,
0.017195,
0.326011,
-0.0961918,
-0.713652,
-0.196473,
0.230423,
0.336269,
0.109072,
0.428225,
-0.311886,
0.383659,
0.0629309,
0.20279,
-0.858124,
0.267745,
0.0698875,
0.572496,
0.0564592,
-0.291992,
-0.240974,
0.121683,
0.0923392,
-0.190241,
-0.320404,
0.454806,
-0.0716024,
-0.020555,
-0.160004,
0.419512,
0.334638,
-0.306493,
-0.0229742,
-0.276633,
0.556304,
0.132983,
-0.0840649,
-0.101249,
0.163251,
-0.0738091,
0.462798,
-0.27786,
-0.370396,
0.423145,
-0.127708,
0.425961,
0.263961,
-0.407428,
-0.182851,
0.014795,
-0.0661325,
0.0908633,
-0.183748,
-0.0254787,
0.300446,
0.399875,
-0.237013,
-0.0695918,
-0.0380361,
-0.404956,
0.281397,
0.169411,
0.156416,
-0.0330841,
0.2611,
0.289024,
0.348009,
-0.0510249,
0.104812,
-0.116636,
-0.0968812,
0.294442,
0.0161214,
-0.0461076,
0.180963,
0.199651,
-0.226287,
-0.19431,
0.21369,
0.157819,
-0.0628762,
0.358757,
-0.0297353,
-0.579409,
-0.238768,
-0.442426,
0.343682,
0.151564,
0.0356123,
-0.695277,
-0.443175,
-0.0452226,
-0.534062,
0.282538,
-0.0308858,
-0.0429336,
-0.019023,
-0.285045,
-0.312946,
-0.569596,
0.00704258,
0.121878,
-0.0712807,
-0.0847591,
-0.282027,
0.0752057,
-0.735423,
-0.331334,
-0.197524,
-0.152546,
0.0653733,
];
const VIS_L1_WO: [f32; 64] = [
-0.0577218, 0.0466404, 0.373732, 0.411631, 0.194282, 0.37458, 0.518893, 0.111126,
-0.419491, 0.301584, -0.298475, 0.558556, 0.517514, -0.155468, -0.329935, -0.206597,
-0.0195243, 0.384379, -0.370628, 0.26457, -0.620897, 0.364109, -0.018589, 0.0599231,
-0.0318378, 0.280722, 0.111394, 0.50638, -0.532875, -0.271522, -0.403808, 0.215335,
-0.506274, 0.729307, 0.0275459, 0.366571, -0.289899, 0.0588975, 0.187275, -0.211902,
-0.0683283, 0.467214, 0.379918, 0.173689, 0.591409, -0.58374, 0.08737, -0.114411, 0.427121,
0.332962, 0.350832, -0.179064, 0.16847, 0.324073, 0.389196, -0.277921, -0.204972, 0.073562,
0.418459, 0.186085, 0.412741, -0.58574, -0.0814907, 0.149653,
];
const VIS_L1_NORM1_W: [f32; 8] = [
1.02696, 0.810868, 1.11009, 1.03453, 0.953349, 1.11978, 0.827073, 0.948868,
];
const VIS_L1_FC0: [f32; 96] = [
0.265211,
0.112798,
-0.0674063,
-0.0461329,
-0.070063,
0.260172,
-0.0912916,
0.203377,
-0.626122,
0.368826,
0.732467,
0.149894,
-0.0259573,
-0.160522,
0.20993,
-0.195052,
-0.24562,
-0.638671,
0.0727824,
-0.0872477,
-0.0707903,
-0.000692914,
-0.212435,
0.0562055,
-0.323977,
-0.034617,
-0.0844252,
-0.512489,
-0.363591,
-0.177747,
-0.168726,
0.0777597,
-0.121498,
-0.0153486,
-0.122644,
-0.777041,
-0.0370188,
-0.252651,
-0.0479536,
-0.229824,
0.196703,
-0.0137741,
0.212481,
0.234418,
-0.169992,
0.244582,
0.0686207,
-0.317385,
-0.0347921,
0.0282208,
0.380975,
-0.10457,
0.471259,
0.230964,
-0.634277,
-0.310673,
0.175964,
-0.4366,
0.142855,
0.414785,
-0.0804484,
0.422308,
0.39718,
-0.0169027,
-0.0100078,
0.0543716,
0.133468,
-0.467238,
-0.0434751,
-0.501976,
-0.10354,
-0.0658216,
-0.545229,
-0.379061,
-0.119427,
0.648104,
-0.416021,
0.0373657,
-0.268065,
-0.112513,
0.295553,
0.0931424,
0.127734,
0.146628,
0.17265,
-0.64847,
0.115814,
-0.0255255,
0.119235,
0.250913,
0.124699,
0.207121,
-0.351437,
0.134342,
-0.0327808,
-0.0259479,
];
const VIS_L1_FC1: [f32; 96] = [
0.39004,
-0.146601,
-0.199982,
-0.39407,
-0.0886803,
0.321206,
-0.500337,
0.0753654,
-0.759829,
-0.207078,
0.131248,
-0.161568,
-0.360611,
-0.11047,
0.296733,
0.447581,
0.453663,
0.204687,
0.108174,
0.295752,
-0.346912,
-0.290761,
0.387467,
0.256882,
0.262525,
0.0665477,
-0.247498,
-0.324031,
0.281414,
0.324848,
-0.336142,
-0.221854,
-0.0308679,
0.262391,
-0.102162,
-0.255668,
0.296526,
0.179182,
-0.208277,
-0.197806,
0.332993,
-0.384262,
-0.0878628,
0.134792,
0.0493159,
0.278824,
-0.513304,
0.00528268,
0.0903486,
0.100079,
0.173074,
-0.35006,
0.433543,
0.136213,
0.168189,
-0.17306,
0.054225,
0.399138,
-0.00455385,
0.0586072,
0.0950148,
-0.489918,
-0.0517679,
0.207905,
-0.146564,
-0.0126508,
-0.00982426,
-0.473983,
0.241292,
-0.406213,
0.1268,
-0.040079,
-0.102337,
0.0200342,
0.226104,
0.244038,
-0.129102,
-0.0493391,
-0.257829,
-0.43845,
0.285282,
-0.329176,
-0.22059,
0.113821,
-0.404307,
0.144029,
-0.253124,
-0.340929,
-0.1742,
-0.255316,
0.0708169,
0.425441,
0.202217,
0.139482,
-0.184594,
-0.0723397,
];
const VIS_FINAL_NORM_W: [f32; 8] = [
1.00924, 0.992412, 1.01246, 0.877821, 1.17394, 0.862922, 0.981503, 1.07865,
];
const VIS_MERGE_PROJ0: [f32; 1024] = [
-0.236349,
0.302076,
-0.436699,
-0.236767,
-0.719464,
0.295538,
-0.373172,
-0.22321,
0.213612,
0.0864719,
0.302082,
0.176717,
0.249573,
-0.162176,
-0.0439175,
0.203247,
0.207236,
0.265128,
-0.248399,
-0.347564,
-0.395201,
-0.160791,
-0.0491745,
-0.376832,
-0.0740985,
0.170485,
0.277722,
-0.19514,
0.15138,
-0.401887,
-0.217022,
0.0129443,
-0.212729,
-0.252641,
0.0614228,
-0.316426,
0.0141244,
-0.0513168,
0.133343,
0.207095,
0.0205796,
-0.168545,
0.282756,
0.0393554,
-0.283789,
0.0319159,
-0.203635,
-0.777423,
-0.307626,
-0.133649,
-0.157371,
-0.481473,
-0.124941,
0.107979,
-0.362024,
0.0569102,
-0.628088,
0.131413,
-0.190082,
-0.422225,
0.312183,
0.118983,
0.264134,
0.0722628,
0.0832335,
-0.135963,
0.105563,
-0.0461713,
-0.212584,
-0.0778148,
-0.403216,
-0.0291645,
0.110195,
0.0157109,
0.148185,
-0.114513,
-0.0117906,
-0.013211,
0.5547,
-0.410289,
0.581179,
-0.239112,
-0.121669,
0.0614686,
0.151549,
-0.0293001,
0.0450231,
0.393947,
0.753717,
0.101535,
0.187023,
-0.0497537,
0.473786,
-0.0537569,
0.131424,
-0.330517,
-0.81053,
0.655177,
0.297821,
0.253245,
-0.295869,
-0.503227,
-0.0895497,
0.0160632,
0.195476,
-0.0153669,
0.361626,
0.08452,
0.0733994,
0.326634,
-0.216635,
0.0958955,
-0.183129,
-0.21866,
0.430652,
0.306151,
0.140998,
0.370526,
0.0936386,
0.383891,
-0.268531,
-0.0606306,
0.0965548,
-0.315549,
-0.216088,
-0.38473,
0.184519,
-0.153703,
0.0636753,
0.139368,
-0.0627465,
-0.00900365,
0.024932,
0.110283,
0.38189,
-0.341072,
0.0110382,
-0.370518,
0.413105,
-0.557165,
0.0337617,
-0.138768,
-0.185464,
0.19559,
0.141414,
-0.298639,
0.121156,
0.0692044,
-0.236338,
-0.446907,
0.217787,
0.189244,
0.136685,
-0.172439,
0.34132,
1.00006,
0.0300311,
0.13925,
-0.259658,
-0.257112,
0.501067,
0.653685,
-0.0926841,
-0.213994,
-0.29322,
-0.201359,
-0.00520445,
-0.391296,
-0.232173,
0.25891,
-0.0725001,
-0.25552,
-0.372763,
0.293815,
-0.336134,
-0.0673453,
0.272661,
0.1354,
0.0615359,
0.0644006,
-0.307396,
-0.239653,
-0.163541,
-0.88038,
-0.149103,
-0.0547648,
-0.0497764,
-0.0617094,
0.0359132,
0.055568,
0.462535,
0.0217722,
0.230232,
-0.201187,
0.0736835,
0.0634726,
-0.0907999,
-0.532513,
-0.423686,
0.112162,
-0.037038,
-0.0948945,
-0.223742,
0.00604771,
0.150508,
0.535644,
-0.0232303,
-0.0708197,
0.0633776,
0.33081,
0.672843,
0.151136,
-0.445781,
-0.00767956,
-0.354326,
-0.30255,
0.527207,
-0.117645,
-0.226993,
0.397529,
0.0112774,
-0.350571,
0.00682257,
-0.428325,
-0.267424,
0.0460792,
-0.772427,
0.204795,
0.0671937,
0.243194,
-0.125164,
0.11855,
-0.303156,
-0.283589,
0.0239112,
0.00140681,
-0.00935182,
0.398183,
-0.0882406,
0.12846,
0.0970286,
0.0414345,
0.117666,
-0.524796,
-0.171809,
0.324681,
-0.0936075,
0.203437,
0.135693,
-0.917639,
-0.398252,
-0.205078,
-0.209732,
-0.140251,
-0.135229,
-0.250041,
-0.205651,
-0.591604,
0.574865,
0.410002,
0.255556,
-0.0320876,
-0.280266,
0.723274,
-0.290018,
0.421423,
0.0849775,
-0.281182,
-0.252831,
-0.134178,
-0.195518,
-0.432479,
0.118937,
-0.283856,
0.220364,
-0.0620359,
0.0341665,
0.0261671,
-0.124931,
-0.0379213,
0.0432028,
0.68423,
-0.39457,
-0.219883,
-0.125563,
-0.387435,
-0.410432,
-0.476937,
-0.0312687,
-0.175228,
0.0353051,
-0.444102,
-0.283561,
0.257007,
0.208592,
-0.484891,
0.230308,
-0.182565,
-0.350742,
0.199333,
0.250345,
-0.767212,
0.238475,
-0.17154,
-0.180079,
-0.135758,
0.296082,
0.350977,
0.072524,
0.204048,
0.607167,
-0.0381462,
-0.3366,
0.124063,
0.177895,
-0.0799869,
-0.0899268,
0.610806,
-0.109059,
-0.299177,
-0.835396,
0.214294,
0.350627,
-0.380885,
0.144448,
-0.131053,
-0.00930346,
-0.204848,
-0.0314523,
0.231507,
-0.179883,
-0.230056,
-0.328797,
0.310598,
-0.276022,
0.160202,
-0.39869,
-0.664147,
-0.194,
0.0443349,
0.298699,
0.212516,
0.0706104,
-0.771505,
0.118546,
-0.467788,
-0.23413,
0.509708,
0.274984,
-0.347402,
-0.127842,
0.21163,
-0.0675676,
-0.0878888,
-0.358893,
0.137875,
0.177817,
-0.261019,
-0.628945,
0.518781,
-0.231445,
0.103466,
-0.276795,
-0.24368,
0.265653,
-0.115171,
0.210652,
0.217841,
-0.173374,
0.0684118,
0.212281,
0.35257,
-0.0986456,
-0.0154741,
-0.299062,
0.216102,
0.0126834,
-0.415971,
-0.121538,
-0.0510272,
-0.192594,
-0.231441,
-0.504684,
0.0135955,
-0.142346,
0.258746,
-0.0460905,
0.235684,
-0.371481,
0.0983083,
0.202997,
0.207863,
-0.0960896,
-0.278115,
0.485785,
-0.0703955,
-0.36832,
-0.0743555,
0.0874494,
-0.223167,
0.112969,
0.258647,
-0.0219521,
0.293628,
-0.0430369,
-0.148116,
0.465595,
0.180931,
0.0256682,
-0.544123,
-0.419134,
-0.504943,
-0.0196961,
-0.0160084,
0.116454,
0.301819,
-0.463058,
-0.409316,
-0.280946,
0.129024,
0.0789985,
-0.317849,
-0.217392,
0.10596,
0.0127895,
0.295834,
-0.250482,
-0.0489015,
-0.303991,
0.335735,
0.298507,
0.137767,
0.283669,
0.404645,
0.42822,
0.501337,
-0.218792,
-0.242324,
0.380491,
-0.129582,
-0.530074,
-0.13798,
0.303014,
-0.14071,
0.226784,
0.0899447,
-0.403096,
0.134022,
0.00218902,
-0.164883,
-0.428667,
0.47806,
-0.273605,
-0.610717,
0.248782,
-0.134394,
0.353971,
-0.0321702,
-0.0053093,
0.174744,
0.034076,
0.0328364,
0.323228,
-0.271353,
0.165744,
0.0854048,
-0.481041,
0.331567,
0.229811,
0.460195,
-0.130766,
-0.32927,
0.919488,
-0.304845,
-0.237677,
0.166027,
-0.129779,
0.123418,
-0.397317,
-0.509057,
-0.164184,
0.201893,
-0.00989215,
0.358239,
-0.0787409,
0.0564069,
-0.272082,
-0.437823,
-0.363646,
-0.238196,
0.588519,
-0.385448,
-0.264242,
0.410487,
-0.199585,
0.176318,
-0.682871,
-0.291284,
0.15313,
-0.509376,
0.0532867,
0.00989985,
-0.169128,
-0.0829413,
0.210541,
-0.0165126,
0.293903,
0.232969,
0.46007,
-0.128089,
0.0495355,
-0.317411,
0.184258,
-0.214564,
0.183357,
-0.0849151,
-0.0527055,
0.0672831,
-0.223927,
-0.202076,
-0.188411,
-0.16416,
0.184064,
-0.118991,
-0.174886,
0.242811,
0.203539,
-0.0508292,
0.734019,
-0.317723,
0.477404,
0.0731702,
0.296654,
-0.0947508,
0.345185,
-0.300144,
0.234292,
0.0730285,
0.265356,
-0.149729,
0.121972,
-0.239329,
-0.111945,
-0.0787418,
-0.185964,
-0.151469,
-0.370867,
-0.0418644,
0.430025,
0.846542,
-0.254989,
-0.231013,
0.0964694,
0.146091,
-0.39068,
-0.693897,
0.0743391,
0.396942,
0.422646,
-0.00549475,
0.413812,
-0.0535897,
-0.421364,
-0.0428651,
0.307348,
0.0362761,
-0.172384,
0.0135348,
0.0981934,
0.0389974,
0.0883734,
0.110775,
0.317814,
0.0373875,
-0.0478011,
0.105298,
0.175046,
-0.524262,
0.341699,
0.506601,
0.00578107,
0.182254,
-0.223231,
0.200005,
-0.256655,
0.0524954,
0.17099,
-0.237381,
-0.333608,
-0.105021,
-0.170642,
-0.0984785,
-0.33122,
0.323198,
-0.0827793,
0.130607,
-0.489473,
-0.177912,
0.476636,
-0.264911,
0.0834568,
0.624801,
-0.00254831,
0.44937,
-0.270769,
-0.145531,
-0.0201544,
-0.16437,
0.100481,
-0.534917,
-0.320538,
0.389989,
0.0333177,
0.143457,
0.501674,
-0.52644,
-0.898354,
0.527995,
-0.414444,
0.0201641,
-0.492257,
-0.771534,
-0.0286519,
0.133742,
-0.253637,
-0.768402,
0.085024,
0.562326,
-0.526449,
0.221449,
0.0751945,
-0.563033,
0.295024,
0.0474323,
-0.0501224,
0.375451,
0.00226611,
0.376463,
0.778869,
0.0582743,
0.12641,
-0.176428,
0.456384,
-0.371866,
-0.156221,
-0.0189524,
-0.0548232,
-0.214016,
0.263157,
-0.0931925,
-0.645682,
0.348074,
0.519181,
-0.167478,
0.407343,
-0.160051,
-0.298173,
-0.276582,
0.491172,
0.410276,
0.208572,
-0.303514,
0.239872,
-0.325825,
0.187611,
-0.16367,
-0.170864,
-0.156856,
-0.621923,
-0.526686,
-0.0450273,
-0.322389,
0.246644,
0.352417,
-0.467928,
-0.452286,
-0.0509965,
0.583273,
-0.236968,
0.148131,
-0.497616,
-0.275279,
0.141518,
-0.0770699,
0.145007,
0.0319243,
-0.00993607,
-0.016638,
0.405023,
-0.0263409,
0.251191,
-0.0873389,
0.0196543,
-0.330143,
-0.638085,
-0.136588,
0.019404,
-0.153683,
0.152007,
-0.251054,
-0.4297,
-0.454367,
0.183937,
-0.2615,
-0.0117493,
-0.438665,
-0.276094,
-0.0236997,
0.305428,
-0.339538,
-0.455488,
-0.0246618,
0.0245187,
-0.15223,
0.343269,
-0.538411,
-0.270215,
-0.0232314,
-0.286599,
0.0768719,
-0.0225739,
-0.0989169,
0.699144,
-0.321706,
-0.103017,
0.122571,
-0.249993,
0.0163995,
0.232155,
0.0407288,
0.198615,
-0.13675,
-0.10254,
0.552905,
0.11824,
0.266287,
-0.219535,
-0.0890321,
0.32599,
0.00448736,
-0.395851,
-0.292809,
-0.0430488,
0.448302,
0.1378,
-0.251007,
-0.0270933,
-0.00214663,
-0.523369,
-0.0309944,
-0.125962,
-0.0659885,
-0.289194,
0.0919669,
-0.144388,
-0.252316,
0.535763,
-0.132497,
-0.429882,
0.346153,
0.403493,
-0.309738,
-0.1503,
-0.0742551,
-0.141614,
-0.436618,
-0.345322,
0.620119,
-0.204794,
-0.181742,
0.209699,
-0.166686,
-0.223907,
-0.337174,
0.136383,
0.0377023,
0.312557,
-0.151656,
0.366969,
0.0219035,
-0.182791,
0.113654,
0.0807236,
0.0594,
-0.249578,
-0.0533519,
-0.081953,
0.196487,
0.362187,
0.226492,
0.201226,
0.308508,
0.347177,
-0.139138,
-0.371969,
0.0258503,
0.430667,
-0.398186,
-0.309352,
-0.00336896,
-0.187937,
0.276279,
0.00202034,
0.231139,
0.0684979,
0.100552,
-0.374401,
-0.214862,
-0.610311,
-0.252147,
-0.185271,
0.0818731,
-0.366164,
0.512052,
0.345181,
-0.325809,
-0.572269,
-0.500868,
0.0718371,
-0.391286,
-0.107487,
0.148359,
0.0972818,
0.200506,
-0.406185,
0.177848,
0.187755,
0.265422,
0.228593,
-0.401342,
-0.751146,
0.0909925,
0.460462,
-0.249553,
-0.252755,
0.0340475,
0.306479,
0.202305,
-0.372425,
-0.147982,
-0.186426,
0.154744,
0.209422,
-0.281605,
-0.259755,
0.0696664,
-0.0819251,
-0.247159,
-0.151843,
0.164756,
-0.121739,
-0.052189,
-0.248309,
0.162564,
-0.171285,
0.0632468,
0.0353644,
-0.132106,
-0.266058,
-0.0386287,
-0.10764,
0.288536,
-0.269559,
0.212119,
0.347515,
0.182925,
0.209263,
-0.0891225,
0.123076,
0.125977,
-0.493741,
-0.107514,
-0.371816,
-0.567552,
-0.0149517,
0.12388,
-0.206517,
0.148488,
0.165108,
0.0233112,
-0.15338,
0.197413,
0.52503,
-0.61807,
0.215751,
-0.349172,
0.251406,
0.296309,
-0.123011,
-0.130789,
-0.00651585,
0.035497,
-0.0609056,
0.341434,
-0.470598,
-0.150903,
0.195183,
-0.0862897,
-0.0746215,
-0.0906219,
0.285372,
0.208254,
-0.397821,
-0.341474,
-0.124454,
-0.0414253,
0.271967,
-0.0351484,
0.396491,
0.559566,
-0.260349,
0.0515127,
0.309838,
0.0710535,
0.104538,
-0.135632,
-0.163842,
0.063358,
0.300885,
-0.31289,
-0.0402696,
0.542997,
-0.0556195,
0.0532273,
0.243242,
0.175341,
-0.0434567,
-0.0617974,
0.461337,
-0.279505,
-0.402233,
0.338011,
0.383038,
-0.284724,
0.269527,
0.115576,
-0.0259366,
0.041392,
0.563885,
0.23412,
0.188308,
0.168665,
-0.166081,
-0.0822718,
-0.0302363,
-0.250662,
0.172037,
0.358585,
-0.0150482,
-0.232162,
-0.22336,
0.492373,
-0.162311,
-0.352932,
-0.432669,
-0.128367,
0.378123,
0.388758,
0.191049,
-0.208064,
0.261632,
-0.0438233,
0.197892,
0.668999,
-0.0518101,
-0.392415,
0.589557,
-0.0304808,
-0.0361237,
-0.473126,
0.243661,
0.336332,
0.0973627,
0.349706,
-0.096878,
0.427568,
0.255798,
0.553208,
-0.477473,
-0.351159,
0.291254,
0.137889,
-0.24987,
0.0544092,
0.305517,
0.369363,
0.108008,
0.0571747,
0.170541,
0.0154032,
0.603159,
-0.0985079,
0.0345475,
-0.0276445,
0.495077,
-0.130628,
0.29438,
0.0118292,
-0.107588,
-0.598261,
0.00946725,
0.530881,
-0.181808,
-0.851401,
0.113912,
-0.335427,
-0.511092,
-0.136116,
-0.382416,
0.332571,
0.0797383,
0.200826,
-0.175843,
-0.0182147,
-0.156549,
0.0881938,
-0.665213,
-0.00780874,
0.310998,
0.511727,
0.250097,
0.0436686,
-0.511367,
0.0453479,
];
const VIS_MERGE_PROJ1: [f32; 192] = [
-0.589584,
-0.270176,
-0.0642117,
-0.0839581,
0.263636,
-0.0866331,
-0.38092,
-0.14856,
-0.230682,
0.322097,
-0.444679,
-0.150674,
0.270976,
-0.261266,
0.110946,
0.00495127,
-0.106307,
-0.531284,
0.00520624,
-0.15978,
0.0808836,
-0.440902,
0.153633,
0.194342,
0.417523,
-0.148402,
-0.212159,
0.268179,
-0.317149,
-0.260497,
-0.09544,
-0.343195,
0.429478,
0.285979,
-0.654128,
-0.645155,
-0.196368,
0.148705,
0.203804,
0.422358,
0.0848344,
-0.321419,
0.125365,
-0.34804,
-0.251555,
-0.482738,
-0.271232,
-0.18782,
-0.172891,
0.530903,
0.0288639,
0.104308,
0.0468265,
0.475097,
-0.285449,
0.225047,
-0.256494,
-0.0416107,
0.148517,
-0.405077,
-0.115967,
0.00170611,
0.151352,
-0.0672187,
0.846001,
0.564592,
-0.660677,
0.0353704,
0.166182,
0.102616,
0.417172,
-0.858214,
-0.322372,
0.462712,
0.364483,
0.685643,
0.496457,
-0.38512,
0.310468,
-0.0848315,
0.0626254,
0.0168127,
-0.140008,
-0.151035,
0.232496,
-0.163294,
0.0647018,
0.145547,
-0.159983,
-0.0856119,
0.345313,
-0.509969,
-0.201798,
-0.401557,
0.0651107,
0.811157,
-0.368641,
0.368018,
0.426114,
-0.64137,
-0.402644,
-0.00634645,
-0.296265,
0.0834438,
0.409867,
0.211353,
-0.0590335,
0.170205,
0.105546,
-0.380194,
-0.180618,
0.140349,
-0.0205445,
0.0967417,
0.267025,
-0.725734,
0.444512,
-0.572309,
0.00349554,
0.0459292,
-0.0438933,
0.0563175,
-0.168883,
-0.219364,
-0.0316683,
0.051398,
-0.207059,
0.0949521,
-0.260774,
-0.177944,
-0.0742896,
0.131637,
-0.0958914,
0.114644,
0.480587,
-0.146184,
0.387162,
-0.400214,
0.190505,
-0.0616221,
0.628918,
-0.156545,
-0.253828,
0.311851,
0.119144,
0.0184293,
-0.729126,
0.113713,
-0.0374038,
0.0200473,
-0.239673,
0.385356,
-0.0108929,
-0.0806783,
-0.47495,
-0.0672919,
-0.336395,
0.401548,
0.0655212,
-0.190752,
-0.0640861,
0.425706,
0.0323832,
-0.0136082,
0.54734,
0.10145,
0.261653,
-0.115325,
-0.124108,
0.243261,
-0.0241873,
0.0977518,
-0.0293961,
-0.00992239,
-0.0499667,
-0.404479,
0.381245,
-0.189778,
-0.368962,
0.0141507,
-0.206907,
0.368309,
0.18214,
-0.186088,
-0.275483,
0.183006,
-0.218805,
-0.421012,
-0.508387,
0.201931,
-0.478372,
-0.364119,
];
const VIS_POST_NORM_W: [f32; 6] = [1.06551, 1.07693, 0.879476, 0.935329, 1.09384, 1.01741];
const VIS_GOLDEN_ENCODER_OUT: [f32; 32] = [
-0.791224, 0.0928006, -1.20306, 0.668509, -1.51065, 0.051747, -1.10594, 1.69022, 0.539921,
1.01363, -1.06878, -0.240584, -1.83794, 0.908773, -0.935426, 1.08535, -0.259518, 1.01078,
-0.0224143, -1.81579, 1.19159, 1.00525, -0.470024, -0.0242557, -1.95959, 0.713632,
-0.145732, -1.19647, -0.051718, 0.905252, -0.80985, -0.243886,
];
const VIS_GOLDEN_OUTPUT: [f32; 6] = [
-2.50858, 0.0255713, -0.376901, -0.370726, -0.301048, -0.203254,
];
fn wm(data: &[f32], rows: usize, cols: usize) -> WeightMatrix {
assert_eq!(data.len(), rows * cols);
WeightMatrix::F32(Tensor::new(data.to_vec(), vec![rows, cols]))
}
fn cfg() -> VisionConfig {
VisionConfig {
in_dim: IN_DIM,
patch_size: PATCH_SIZE,
grid_h: GRID_H,
grid_w: GRID_W,
hidden_dim: HIDDEN_DIM,
num_heads: NUM_HEADS,
qkv_hidden: QKV_HIDDEN,
mlp_dim: MLP_DIM,
rms_norm_eps: NORM_EPS,
theta_base: THETA_BASE,
merge_kh: MERGE_KH,
merge_kw: MERGE_KW,
projector_ln_eps: PROJECTOR_LN_EPS,
}
}
fn make_encoder_weights() -> VisionEncoderWeights {
let patch_dim = IN_DIM * PATCH_SIZE * PATCH_SIZE;
VisionEncoderWeights {
patch_embed: wm(&VIS_PATCH_EMBED_W, HIDDEN_DIM, patch_dim),
pos_emb: VIS_POS_EMB_W.to_vec(),
layers: vec![
VisionEncoderLayerWeights {
norm0_weight: VIS_L0_NORM0_W.to_vec(),
wqkv: wm(&VIS_L0_WQKV, 3 * QKV_HIDDEN, HIDDEN_DIM),
wo: wm(&VIS_L0_WO, HIDDEN_DIM, QKV_HIDDEN),
norm1_weight: VIS_L0_NORM1_W.to_vec(),
fc0: wm(&VIS_L0_FC0, MLP_DIM, HIDDEN_DIM),
fc1: wm(&VIS_L0_FC1, HIDDEN_DIM, MLP_DIM),
},
VisionEncoderLayerWeights {
norm0_weight: VIS_L1_NORM0_W.to_vec(),
wqkv: wm(&VIS_L1_WQKV, 3 * QKV_HIDDEN, HIDDEN_DIM),
wo: wm(&VIS_L1_WO, HIDDEN_DIM, QKV_HIDDEN),
norm1_weight: VIS_L1_NORM1_W.to_vec(),
fc0: wm(&VIS_L1_FC0, MLP_DIM, HIDDEN_DIM),
fc1: wm(&VIS_L1_FC1, HIDDEN_DIM, MLP_DIM),
},
],
final_norm_weight: VIS_FINAL_NORM_W.to_vec(),
}
}
fn make_merger_weights() -> VisionMergerWeights {
let merge_hidden = MERGE_KH * MERGE_KW * HIDDEN_DIM;
VisionMergerWeights {
proj0: wm(&VIS_MERGE_PROJ0, merge_hidden, merge_hidden),
proj1: wm(&VIS_MERGE_PROJ1, TEXT_HIDDEN, merge_hidden),
post_norm_weight: VIS_POST_NORM_W.to_vec(),
}
}
#[test]
fn encoder_matches_independent_python_reference() {
let cfg = cfg();
let weights = make_encoder_weights();
let patches = vec![
VIS_PATCH_0.to_vec(),
VIS_PATCH_1.to_vec(),
VIS_PATCH_2.to_vec(),
VIS_PATCH_3.to_vec(),
];
let out = encoder_forward(&weights, &cfg, &patches);
assert_eq!(out.len(), VIS_GOLDEN_ENCODER_OUT.len());
for (i, (a, b)) in out.iter().zip(VIS_GOLDEN_ENCODER_OUT.iter()).enumerate() {
assert!((a - b).abs() < 1e-3, "element {i}: rust={a} python={b}");
}
}
#[test]
fn full_pipeline_matches_independent_python_reference() {
let cfg = cfg();
let encoder_weights = make_encoder_weights();
let merger_weights = make_merger_weights();
let patches = vec![
VIS_PATCH_0.to_vec(),
VIS_PATCH_1.to_vec(),
VIS_PATCH_2.to_vec(),
VIS_PATCH_3.to_vec(),
];
let encoder_out = encoder_forward(&encoder_weights, &cfg, &patches);
let merged = patch_merge(&encoder_out, &cfg);
assert_eq!(merged.len(), MERGE_KH * MERGE_KW * HIDDEN_DIM);
let projected = project_merged_patches(&merger_weights, &cfg, &merged, 1);
assert_eq!(projected.len(), VIS_GOLDEN_OUTPUT.len());
for (i, (a, b)) in projected.iter().zip(VIS_GOLDEN_OUTPUT.iter()).enumerate() {
assert!((a - b).abs() < 1e-3, "element {i}: rust={a} python={b}");
}
}
#[test]
fn patch_merge_groups_spatially_adjacent_patches_row_major() {
let hidden_dim = 1;
let cfg = VisionConfig {
in_dim: 1,
patch_size: 1,
grid_h: 2,
grid_w: 2,
hidden_dim,
num_heads: 1,
qkv_hidden: 1,
mlp_dim: 1,
rms_norm_eps: 1e-5,
theta_base: 10000.0,
merge_kh: 2,
merge_kw: 2,
projector_ln_eps: 1e-5,
};
let encoder_out = vec![10.0, 20.0, 30.0, 40.0];
let merged = patch_merge(&encoder_out, &cfg);
assert_eq!(merged, vec![10.0, 20.0, 30.0, 40.0]);
}
#[test]
fn erf_matches_known_values() {
assert!((erf(0.0)).abs() < 1e-6);
assert!((erf(1.0) - 0.8427008).abs() < 1e-4);
assert!((erf(-1.0) + 0.8427008).abs() < 1e-4);
}
}