use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct VisionConfig {
pub hidden_size: usize,
#[serde(rename = "num_heads")]
pub num_attention_heads: usize,
#[serde(rename = "depth")]
pub num_hidden_layers: usize,
pub patch_size: usize,
pub intermediate_size: usize,
pub image_size: usize,
#[serde(default = "default_out_hidden_size")]
pub out_hidden_size: usize,
#[serde(default = "default_spatial_merge_size")]
pub spatial_merge_size: usize,
#[serde(default = "default_temporal_patch_size")]
pub temporal_patch_size: usize,
#[serde(default = "default_num_channels")]
pub num_channels: usize,
pub rms_norm_eps: f64,
#[serde(default = "default_hidden_act")]
pub hidden_act: String,
#[serde(default = "default_attention_bias")]
pub attention_bias: bool,
}
fn default_hidden_act() -> String {
"silu".to_string()
}
fn default_attention_bias() -> bool {
true
}
fn default_out_hidden_size() -> usize {
1536
}
fn default_spatial_merge_size() -> usize {
2
}
fn default_temporal_patch_size() -> usize {
2
}
fn default_num_channels() -> usize {
3
}
impl Default for VisionConfig {
fn default() -> Self {
Self {
hidden_size: 1024,
num_attention_heads: 16,
num_hidden_layers: 24,
patch_size: 14,
intermediate_size: 4096,
image_size: 336,
out_hidden_size: default_out_hidden_size(),
spatial_merge_size: default_spatial_merge_size(),
temporal_patch_size: default_temporal_patch_size(),
num_channels: default_num_channels(),
rms_norm_eps: 1e-5,
hidden_act: default_hidden_act(),
attention_bias: default_attention_bias(),
}
}
}
#[cfg(not(target_arch = "wasm32"))]
mod imp {
use candle_core::{D, DType, Device, IndexOp, Module, Result as CandleResult, Tensor};
use candle_nn::VarBuilder;
use super::VisionConfig;
use crate::CandleOcrError;
use crate::error::Result;
const ROPE_THETA: f64 = 10000.0;
#[derive(Debug, Clone)]
struct RmsNorm {
weight: Tensor,
eps: f64,
}
impl RmsNorm {
fn new(size: usize, eps: f64, vb: VarBuilder) -> CandleResult<Self> {
let weight = vb.get(size, "weight")?;
Ok(Self { weight, eps })
}
fn forward(&self, xs: &Tensor) -> CandleResult<Tensor> {
let input_dtype = xs.dtype();
let xs_f32 = xs.to_dtype(DType::F32)?;
let variance = xs_f32.sqr()?.mean_keepdim(D::Minus1)?;
let normed = xs_f32.broadcast_div(&(variance + self.eps)?.sqrt()?)?;
normed.to_dtype(input_dtype)?.broadcast_mul(&self.weight)
}
}
struct Vision2dRope {
inv_freq: Tensor,
half_dim: usize,
device: Device,
}
impl Vision2dRope {
fn new(head_dim: usize, device: Device) -> CandleResult<Self> {
let half_dim = head_dim / 2;
let freq_dim = half_dim / 2;
let inv_freq: Vec<f32> = (0..freq_dim)
.map(|i| (1.0 / ROPE_THETA.powf((i as f64 * 2.0) / half_dim as f64)) as f32)
.collect();
let inv_freq = Tensor::from_vec(inv_freq, freq_dim, &device)?;
Ok(Self {
inv_freq,
half_dim,
device,
})
}
fn cos_sin(&self, grid_h: usize, grid_w: usize, dtype: DType) -> CandleResult<(Tensor, Tensor)> {
let positions_h: Vec<f32> = (0..grid_h)
.flat_map(|h| std::iter::repeat_n(h as f32, grid_w))
.collect();
let positions_w: Vec<f32> = (0..grid_h).flat_map(|_| (0..grid_w).map(|w| w as f32)).collect();
let seq_len = grid_h * grid_w;
let pos_h = Tensor::from_vec(positions_h, (seq_len, 1), &self.device)?.to_dtype(DType::F32)?;
let pos_w = Tensor::from_vec(positions_w, (seq_len, 1), &self.device)?.to_dtype(DType::F32)?;
let inv_freq = self.inv_freq.reshape((1, ()))?;
let freqs_h = pos_h.broadcast_mul(&inv_freq)?;
let freqs_w = pos_w.broadcast_mul(&inv_freq)?;
let freqs = Tensor::cat(&[&freqs_h, &freqs_w], D::Minus1)?;
let emb = Tensor::cat(&[&freqs, &freqs], D::Minus1)?;
let cos = emb.cos()?.to_dtype(dtype)?;
let sin = emb.sin()?.to_dtype(dtype)?;
Ok((cos, sin))
}
fn apply(&self, q: &Tensor, k: &Tensor, grid_h: usize, grid_w: usize) -> CandleResult<(Tensor, Tensor)> {
debug_assert_eq!(q.dim(D::Minus1)?, 2 * self.half_dim);
let (cos, sin) = self.cos_sin(grid_h, grid_w, q.dtype())?;
let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
let q_rot = q.broadcast_mul(&cos)? + rotate_half(q)?.broadcast_mul(&sin)?;
let k_rot = k.broadcast_mul(&cos)? + rotate_half(k)?.broadcast_mul(&sin)?;
Ok((q_rot?, k_rot?))
}
}
fn rotate_half(x: &Tensor) -> CandleResult<Tensor> {
let last = x.dim(D::Minus1)?;
let h = last / 2;
let x1 = x.narrow(D::Minus1, 0, h)?;
let x2 = x.narrow(D::Minus1, h, last - h)?;
Tensor::cat(&[&x2.neg()?, &x1], D::Minus1)
}
struct Attention {
qkv: candle_nn::Linear,
q_norm: RmsNorm,
k_norm: RmsNorm,
proj: candle_nn::Linear,
num_heads: usize,
head_dim: usize,
scale: f64,
}
impl Attention {
fn new(config: &VisionConfig, vb: VarBuilder) -> CandleResult<Self> {
let hidden = config.hidden_size;
let head_dim = hidden / config.num_attention_heads;
if !head_dim.is_multiple_of(2) {
return Err(candle_core::Error::Msg(format!(
"head_dim must be even for rotate_half; got {head_dim}"
)));
}
let scale = 1.0 / (head_dim as f64).sqrt();
let qkv = if config.attention_bias {
candle_nn::linear(hidden, 3 * hidden, vb.pp("qkv"))?
} else {
candle_nn::linear_no_bias(hidden, 3 * hidden, vb.pp("qkv"))?
};
let q_norm = RmsNorm::new(head_dim, config.rms_norm_eps, vb.pp("q_norm"))?;
let k_norm = RmsNorm::new(head_dim, config.rms_norm_eps, vb.pp("k_norm"))?;
let proj = candle_nn::linear(hidden, hidden, vb.pp("proj"))?;
Ok(Self {
qkv,
q_norm,
k_norm,
proj,
num_heads: config.num_attention_heads,
head_dim,
scale,
})
}
fn forward(&self, xs: &Tensor, rope: &Vision2dRope, grid_h: usize, grid_w: usize) -> CandleResult<Tensor> {
let (batch, seq_len, _) = xs.dims3()?;
let qkv = self.qkv.forward(xs)?;
let qkv = qkv
.reshape((batch, seq_len, 3, self.num_heads, self.head_dim))?
.permute([2, 0, 3, 1, 4])?;
let q = qkv.i(0)?.contiguous()?;
let k = qkv.i(1)?.contiguous()?;
let v = qkv.i(2)?.contiguous()?;
let q = self.q_norm.forward(&q)?;
let k = self.k_norm.forward(&k)?;
let (q, k) = rope.apply(&q, &k, grid_h, grid_w)?;
let scores = (q.matmul(&k.transpose(D::Minus2, D::Minus1)?)? * self.scale)?;
let attn = candle_nn::ops::softmax_last_dim(&scores)?;
let out = attn.matmul(&v)?;
let out = out
.permute([0, 2, 1, 3])?
.reshape((batch, seq_len, self.num_heads * self.head_dim))?;
self.proj.forward(&out)
}
}
struct Mlp {
gate_proj: candle_nn::Linear,
up_proj: candle_nn::Linear,
down_proj: candle_nn::Linear,
}
impl Mlp {
fn new(config: &VisionConfig, vb: VarBuilder) -> CandleResult<Self> {
let hidden = config.hidden_size;
let inter = config.intermediate_size;
let gate_proj = candle_nn::linear(hidden, inter, vb.pp("gate_proj"))?;
let up_proj = candle_nn::linear(hidden, inter, vb.pp("up_proj"))?;
let down_proj = candle_nn::linear(inter, hidden, vb.pp("down_proj"))?;
Ok(Self {
gate_proj,
up_proj,
down_proj,
})
}
fn forward(&self, xs: &Tensor) -> CandleResult<Tensor> {
let gate = self.gate_proj.forward(xs)?.silu()?;
let up = self.up_proj.forward(xs)?;
self.down_proj.forward(&(gate * up)?)
}
}
struct TransformerBlock {
norm1: RmsNorm,
attn: Attention,
norm2: RmsNorm,
mlp: Mlp,
}
impl TransformerBlock {
fn new(config: &VisionConfig, vb: VarBuilder) -> CandleResult<Self> {
let norm1 = RmsNorm::new(config.hidden_size, config.rms_norm_eps, vb.pp("norm1"))?;
let attn = Attention::new(config, vb.pp("attn"))?;
let norm2 = RmsNorm::new(config.hidden_size, config.rms_norm_eps, vb.pp("norm2"))?;
let mlp = Mlp::new(config, vb.pp("mlp"))?;
Ok(Self {
norm1,
attn,
norm2,
mlp,
})
}
fn forward(&self, xs: &Tensor, rope: &Vision2dRope, grid_h: usize, grid_w: usize) -> CandleResult<Tensor> {
let xs = (xs + self.attn.forward(&self.norm1.forward(xs)?, rope, grid_h, grid_w)?)?;
&xs + self.mlp.forward(&self.norm2.forward(&xs)?)?
}
}
struct PatchEmbedding {
conv: candle_nn::Conv2d,
}
impl PatchEmbedding {
fn new(config: &VisionConfig, vb: VarBuilder) -> CandleResult<Self> {
let w5d = vb.get(
(
config.hidden_size,
config.num_channels,
config.temporal_patch_size,
config.patch_size,
config.patch_size,
),
"proj.weight",
)?;
let w2d = w5d.sum(2)?.contiguous()?;
let bias = vb.get(config.hidden_size, "proj.bias")?;
let cfg = candle_nn::Conv2dConfig {
stride: config.patch_size,
padding: 0,
..Default::default()
};
let conv = candle_nn::Conv2d::new(w2d, Some(bias), cfg);
Ok(Self { conv })
}
fn forward(&self, x: &Tensor) -> CandleResult<Tensor> {
let x = self.conv.forward(x)?;
let (b, c, h, w) = x.dims4()?;
x.reshape((b, c, h * w))?.permute([0, 2, 1])
}
}
pub struct CogVit {
patch_embed: PatchEmbedding,
blocks: Vec<TransformerBlock>,
post_layernorm: RmsNorm,
rope: Vision2dRope,
config: VisionConfig,
#[allow(dead_code)]
device: Device,
}
impl CogVit {
pub fn new(config: &VisionConfig, vb: VarBuilder, device: Device) -> Result<Self> {
let patch_embed = PatchEmbedding::new(config, vb.pp("patch_embed"))
.map_err(|e| CandleOcrError::ModelLoadFailed(format!("Patch embedding init: {}", e)))?;
let mut blocks = Vec::with_capacity(config.num_hidden_layers);
for i in 0..config.num_hidden_layers {
let block = TransformerBlock::new(config, vb.pp(format!("blocks.{}", i)))
.map_err(|e| CandleOcrError::ModelLoadFailed(format!("Block {} init: {}", i, e)))?;
blocks.push(block);
}
let post_layernorm = RmsNorm::new(config.hidden_size, config.rms_norm_eps, vb.pp("post_layernorm"))
.map_err(|e| CandleOcrError::ModelLoadFailed(format!("post_layernorm init: {}", e)))?;
let head_dim = config.hidden_size / config.num_attention_heads;
let rope = Vision2dRope::new(head_dim, device.clone())
.map_err(|e| CandleOcrError::ModelLoadFailed(format!("Vision RoPE init: {}", e)))?;
Ok(Self {
patch_embed,
blocks,
post_layernorm,
rope,
config: config.clone(),
device,
})
}
pub fn forward(&self, pixel_values: &Tensor) -> Result<Tensor> {
let dims = pixel_values
.dims4()
.map_err(|e| CandleOcrError::InferenceFailed(format!("pixel_values must be 4-D: {}", e)))?;
let (_, _, height, width) = dims;
let patch = self.config.patch_size;
if height % patch != 0 || width % patch != 0 {
return Err(CandleOcrError::InferenceFailed(format!(
"pixel_values H/W ({}, {}) must be multiples of patch_size {}",
height, width, patch
)));
}
let grid_h = height / patch;
let grid_w = width / patch;
let mut x = self
.patch_embed
.forward(pixel_values)
.map_err(|e| CandleOcrError::InferenceFailed(format!("Patch embedding forward: {}", e)))?;
for (i, block) in self.blocks.iter().enumerate() {
x = block
.forward(&x, &self.rope, grid_h, grid_w)
.map_err(|e| CandleOcrError::InferenceFailed(format!("Block {} forward: {}", i, e)))?;
}
self.post_layernorm
.forward(&x)
.map_err(|e| CandleOcrError::InferenceFailed(format!("post_layernorm forward: {}", e)))
}
}
}
#[cfg(not(target_arch = "wasm32"))]
pub use imp::CogVit;