use anyhow::Result;
use rlx_core::weight_loader::WeightLoader;
use rlx_ir::infer::GraphExt;
use rlx_ir::{DType, Graph, NodeId, Shape};
use std::collections::HashMap;
#[derive(Debug, Clone, Copy)]
pub struct VisionConfig {
pub hidden: usize, pub layers: usize, pub heads: usize, pub head_dim: usize, pub intermediate: usize, pub eps: f32, pub rope_theta: f64, pub lm_hidden: usize, }
impl Default for VisionConfig {
fn default() -> Self {
Self {
hidden: 768,
layers: 16,
heads: 12,
head_dim: 64,
intermediate: 3072,
eps: 1e-6,
rope_theta: 100.0,
lm_hidden: 1536,
}
}
}
pub fn vision_inv_freq(cfg: &VisionConfig) -> Vec<f64> {
let spatial = cfg.head_dim / 2; (0..spatial)
.step_by(2)
.map(|i| 1.0 / cfg.rope_theta.powf(i as f64 / spatial as f64))
.collect()
}
pub fn vision_rope_tables(cfg: &VisionConfig, positions: &[(u32, u32)]) -> (Vec<f32>, Vec<f32>) {
let inv = vision_inv_freq(cfg); let hd = cfg.head_dim; let half = hd / 2; let nf = inv.len(); let p = positions.len();
let mut cos = vec![0f32; p * hd];
let mut sin = vec![0f32; p * hd];
for (pi, &(x, y)) in positions.iter().enumerate() {
let base = pi * hd;
for (axis, coord) in [x, y].into_iter().enumerate() {
let off = base + axis * half; for j in 0..nf {
let angle = coord as f64 * inv[j];
let (c, s) = (angle.cos() as f32, angle.sin() as f32);
cos[off + j] = c;
cos[off + nf + j] = c;
sin[off + j] = s;
sin[off + nf + j] = s;
}
}
}
(cos, sin)
}
fn apply_vision_rope_2d(
g: &mut Graph,
x: NodeId, cos: NodeId,
sin: NodeId,
batch: usize,
p: usize,
heads: usize,
hd: usize, f: DType,
) -> NodeId {
let half = hd / 2; let quarter = half / 2; let x4 = g.reshape_(x, vec![batch as i64, p as i64, heads as i64, hd as i64]);
let cos4 = g.reshape_(cos, vec![1, p as i64, 1, hd as i64]);
let sin4 = g.reshape_(sin, vec![1, p as i64, 1, hd as i64]);
let sh = Shape::new(&[batch, p, heads, quarter], f);
let nx = |g: &mut Graph, start: usize| g.narrow_(x4, 3, start, quarter);
let x_lo = nx(g, 0); let x_hi = nx(g, quarter); let y_lo = nx(g, half); let y_hi = nx(g, half + quarter); let neg = |g: &mut Graph, n: NodeId| {
g.add_node(
rlx_ir::op::Op::Activation(rlx_ir::op::Activation::Neg),
vec![n],
sh.clone(),
)
};
let nx_hi = neg(g, x_hi);
let ny_hi = neg(g, y_hi);
let rot = g.concat_(vec![nx_hi, x_lo, ny_hi, y_lo], 3);
let xc = g.mul(x4, cos4);
let rs = g.mul(rot, sin4);
let out = g.add(xc, rs);
g.reshape_(out, vec![batch as i64, p as i64, (heads * hd) as i64])
}
pub fn build_vision_encoder(
cfg: &VisionConfig,
weights: &mut dyn WeightLoader,
batch: usize,
num_patches: usize,
) -> Result<(Graph, HashMap<String, Vec<f32>>)> {
let mut g = Graph::new("gemma4_vision");
let mut params: HashMap<String, Vec<f32>> = HashMap::new();
let hs = build_encoder_core(&mut g, &mut params, cfg, weights, batch, num_patches)?;
g.set_outputs(vec![hs]);
Ok((g, params))
}
fn load_t(
g: &mut Graph,
params: &mut HashMap<String, Vec<f32>>,
w: &mut dyn WeightLoader,
key: &str,
) -> Result<NodeId> {
let (data, shape) = w.take_transposed(key)?;
let id = g.param(key, Shape::new(&shape, DType::F32));
params.insert(key.to_string(), data);
Ok(id)
}
fn load_v(
g: &mut Graph,
params: &mut HashMap<String, Vec<f32>>,
w: &mut dyn WeightLoader,
key: &str,
shape: &[usize],
) -> Result<NodeId> {
let (data, sh) = w.take(key)?;
debug_assert_eq!(&sh, shape, "{key} shape");
let id = g.param(key, Shape::new(&sh, DType::F32));
params.insert(key.to_string(), data);
Ok(id)
}
fn vrms(
g: &mut Graph,
params: &mut HashMap<String, Vec<f32>>,
x: NodeId,
key: &str,
w: &mut dyn WeightLoader,
dim: usize,
eps: f32,
) -> Result<NodeId> {
let wv = load_v(g, params, w, key, &[dim])?;
let ones = {
let id = g.param(format!("{key}.ones"), Shape::new(&[dim], DType::F32));
params.insert(format!("{key}.ones"), vec![1.0f32; dim]);
id
};
let gamma = g.add(ones, wv);
let beta = {
let id = g.param(format!("{key}.beta"), Shape::new(&[dim], DType::F32));
params.insert(format!("{key}.beta"), vec![0.0f32; dim]);
id
};
Ok(g.rms_norm(x, gamma, beta, eps))
}
fn vrms_noscale(
g: &mut Graph,
params: &mut HashMap<String, Vec<f32>>,
x: NodeId,
name: &str,
dim: usize,
eps: f32,
) -> NodeId {
let gamma = synth(
g,
params,
&format!("{name}.ones"),
vec![1.0f32; dim],
&[dim],
);
let beta = synth(
g,
params,
&format!("{name}.beta"),
vec![0.0f32; dim],
&[dim],
);
g.rms_norm(x, gamma, beta, eps)
}
fn build_encoder_core(
g: &mut Graph,
params: &mut HashMap<String, Vec<f32>>,
cfg: &VisionConfig,
weights: &mut dyn WeightLoader,
batch: usize,
num_patches: usize,
) -> Result<NodeId> {
let f = DType::F32;
let (h, nh, hd, eps) = (cfg.hidden, cfg.heads, cfg.head_dim, cfg.eps);
let p = num_patches;
let vis_tap = std::env::var("RLX_VIS_TAP").unwrap_or_default();
let pixels = g.input("vision_pixels", Shape::new(&[batch, p, 3 * 16 * 16], f));
let pos_embed = g.input("vision_pos_embed", Shape::new(&[batch, p, h], f));
let rope_cos = g.input("vision_rope_cos", Shape::new(&[1, p, hd], f));
let rope_sin = g.input("vision_rope_sin", Shape::new(&[1, p, hd], f));
let two = synth(g, params, "vis.two", vec![2.0], &[1]);
let one = synth(g, params, "vis.one", vec![1.0], &[1]);
let px2 = g.mul(pixels, two);
let scaled = g.sub(px2, one); if vis_tap == "scaled" {
return Ok(scaled);
}
let ip = load_t(
g,
params,
weights,
"model.vision_tower.patch_embedder.input_proj.weight",
)?;
let mut hs = g.mm(scaled, ip); if vis_tap == "mm" {
return Ok(hs);
}
hs = g.add(hs, pos_embed);
if vis_tap == "embed" {
return Ok(hs);
}
for layer in 0..cfg.layers {
let lp = format!("model.vision_tower.encoder.layers.{layer}");
let normed = vrms(
g,
params,
hs,
&format!("{lp}.input_layernorm.weight"),
weights,
h,
eps,
)?;
if layer == 0 && vis_tap == "norm" {
return Ok(normed);
}
let q = {
let w = load_t(
g,
params,
weights,
&format!("{lp}.self_attn.q_proj.linear.weight"),
)?;
g.mm(normed, w)
};
let k = {
let w = load_t(
g,
params,
weights,
&format!("{lp}.self_attn.k_proj.linear.weight"),
)?;
g.mm(normed, w)
};
let v = {
let w = load_t(
g,
params,
weights,
&format!("{lp}.self_attn.v_proj.linear.weight"),
)?;
g.mm(normed, w)
};
let q = per_head_norm(
g,
params,
q,
&format!("{lp}.self_attn.q_norm.weight"),
weights,
batch,
p,
nh,
hd,
eps,
true,
)?;
if layer == 0 && vis_tap == "qn" {
return Ok(q);
}
let k = per_head_norm(
g,
params,
k,
&format!("{lp}.self_attn.k_norm.weight"),
weights,
batch,
p,
nh,
hd,
eps,
true,
)?;
let v = per_head_norm(
g,
params,
v,
&format!("{lp}.self_attn.v_norm"),
weights,
batch,
p,
nh,
hd,
eps,
false,
)?;
let q = apply_vision_rope_2d(g, q, rope_cos, rope_sin, batch, p, nh, hd, f);
let k = apply_vision_rope_2d(g, k, rope_cos, rope_sin, batch, p, nh, hd, f);
if layer == 0 && vis_tap == "rope" {
return Ok(q);
}
let attn_shape = rlx_ir::shape::attention_shape(g.shape(q));
let attn = g.attention_kind_opts(
q,
k,
v,
nh,
hd,
rlx_ir::op::MaskKind::None,
attn_shape,
Some(1.0),
None,
);
if layer == 0 && vis_tap == "attn" {
return Ok(attn);
}
let o = {
let w = load_t(
g,
params,
weights,
&format!("{lp}.self_attn.o_proj.linear.weight"),
)?;
g.mm(attn, w)
};
let o = vrms(
g,
params,
o,
&format!("{lp}.post_attention_layernorm.weight"),
weights,
h,
eps,
)?;
hs = g.add(hs, o);
if layer == 0 && vis_tap == "res1" {
return Ok(hs);
}
let normed = vrms(
g,
params,
hs,
&format!("{lp}.pre_feedforward_layernorm.weight"),
weights,
h,
eps,
)?;
let gate = {
let w = load_t(
g,
params,
weights,
&format!("{lp}.mlp.gate_proj.linear.weight"),
)?;
g.mm(normed, w)
};
let up = {
let w = load_t(
g,
params,
weights,
&format!("{lp}.mlp.up_proj.linear.weight"),
)?;
g.mm(normed, w)
};
let gact = g.gelu_approx(gate);
let inner = g.mul(gact, up);
let down = {
let w = load_t(
g,
params,
weights,
&format!("{lp}.mlp.down_proj.linear.weight"),
)?;
g.mm(inner, w)
};
let down = vrms(
g,
params,
down,
&format!("{lp}.post_feedforward_layernorm.weight"),
weights,
h,
eps,
)?;
hs = g.add(hs, down);
}
Ok(hs)
}
pub fn vision_pool_weights(positions: &[(u32, u32)], k: usize) -> (Vec<f32>, usize) {
let p = positions.len();
let max_x = positions.iter().map(|&(x, _)| x).max().unwrap_or(0) as usize + 1;
let bx = max_x / k; let l = p / (k * k);
let inv = 1.0f32 / (k * k) as f32;
let mut w = vec![0f32; l * p];
for (i, &(x, y)) in positions.iter().enumerate() {
let block = (x as usize / k) + bx * (y as usize / k);
w[block * p + i] = inv;
}
(w, l)
}
pub fn build_vision_features(
cfg: &VisionConfig,
weights: &mut dyn WeightLoader,
positions: &[(u32, u32)],
pooling_kernel: usize,
) -> Result<(Graph, HashMap<String, Vec<f32>>)> {
let mut g = Graph::new("gemma4_vision_features");
let mut params: HashMap<String, Vec<f32>> = HashMap::new();
let p = positions.len();
let h = cfg.hidden;
let hs = build_encoder_core(&mut g, &mut params, cfg, weights, 1, p)?; let hs2d = g.reshape_(hs, vec![p as i64, h as i64]);
let (pw, l) = vision_pool_weights(positions, pooling_kernel);
let pool = synth(&mut g, &mut params, "vis.pool_w", pw, &[l, p]);
let pooled = g.mm(pool, hs2d);
let root = synth(
&mut g,
&mut params,
"vis.root_h",
vec![(h as f32).sqrt()],
&[1],
);
let scaled = g.mul(pooled, root);
let normed = vrms_noscale(
&mut g,
&mut params,
scaled,
"embed_vision.pre_norm",
h,
cfg.eps,
);
let proj = load_t(
&mut g,
&mut params,
weights,
"model.embed_vision.embedding_projection.weight",
)?;
let feats = g.mm(normed, proj); let out = g.reshape_(feats, vec![1, l as i64, cfg.lm_hidden as i64]);
g.set_outputs(vec![out]);
Ok((g, params))
}
fn synth(
g: &mut Graph,
params: &mut HashMap<String, Vec<f32>>,
name: &str,
data: Vec<f32>,
shape: &[usize],
) -> NodeId {
let id = g.param(name, Shape::new(shape, DType::F32));
params.insert(name.to_string(), data);
id
}
#[allow(clippy::too_many_arguments)]
fn per_head_norm(
g: &mut Graph,
params: &mut HashMap<String, Vec<f32>>,
x: NodeId,
key: &str,
weights: &mut dyn WeightLoader,
batch: usize,
p: usize,
heads: usize,
hd: usize,
eps: f32,
with_weight: bool,
) -> Result<NodeId> {
let x4 = g.reshape_(x, vec![batch as i64, p as i64, heads as i64, hd as i64]);
let gamma = if with_weight {
let (data, _) = weights.take(key)?;
let ones = vec![1.0f32; hd];
let g_id = g.param(key, Shape::new(&[hd], DType::F32));
params.insert(key.to_string(), data);
let ones_id = synth(g, params, &format!("{key}.ones"), ones, &[hd]);
g.add(ones_id, g_id)
} else {
synth(g, params, &format!("{key}.ones"), vec![1.0f32; hd], &[hd])
};
let beta = synth(g, params, &format!("{key}.vbeta"), vec![0.0f32; hd], &[hd]);
let normed = g.rms_norm(x4, gamma, beta, eps);
Ok(g.reshape_(normed, vec![batch as i64, p as i64, (heads * hd) as i64]))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rope_tables_shape_and_pos0_identity() {
let cfg = VisionConfig::default();
let positions = vec![(0u32, 0u32), (1, 0), (0, 1), (3, 5)];
let (cos, sin) = vision_rope_tables(&cfg, &positions);
assert_eq!(cos.len(), positions.len() * cfg.head_dim);
assert_eq!(sin.len(), cos.len());
for j in 0..cfg.head_dim {
assert!((cos[j] - 1.0).abs() < 1e-6);
assert!(sin[j].abs() < 1e-6);
}
let row = cfg.head_dim; let nf = 16;
assert!((cos[row] - cos[row + nf]).abs() < 1e-6);
}
}