use serde::{Deserialize, Serialize};
use super::vision::VisionConfig;
const DEFAULT_MERGER_INTERMEDIATE_SIZE: usize = 4608;
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct ConnectorConfig {
pub vision_hidden_size: usize,
pub text_hidden_size: usize,
#[serde(default = "default_merger_intermediate_size")]
pub merger_intermediate_size: usize,
pub spatial_merge_size: usize,
}
fn default_merger_intermediate_size() -> usize {
DEFAULT_MERGER_INTERMEDIATE_SIZE
}
impl ConnectorConfig {
pub fn from_vision_config(vision: &VisionConfig) -> Self {
Self {
vision_hidden_size: vision.hidden_size,
text_hidden_size: vision.out_hidden_size,
merger_intermediate_size: DEFAULT_MERGER_INTERMEDIATE_SIZE,
spatial_merge_size: vision.spatial_merge_size,
}
}
}
impl Default for ConnectorConfig {
fn default() -> Self {
Self::from_vision_config(&VisionConfig::default())
}
}
#[cfg(not(target_arch = "wasm32"))]
mod imp {
use candle_core::Tensor;
use candle_nn::{Conv2d, Conv2dConfig, LayerNorm, Linear, Module, VarBuilder};
use super::ConnectorConfig;
use crate::CandleOcrError;
use crate::error::Result;
pub struct VisionConnector {
downsample: Conv2d,
gate_proj: Linear,
up_proj: Linear,
down_proj: Linear,
proj: Linear,
post_projection_norm: LayerNorm,
config: ConnectorConfig,
}
impl VisionConnector {
pub fn new(config: &ConnectorConfig, vb: VarBuilder) -> Result<Self> {
let downsample = candle_nn::conv2d(
config.vision_hidden_size,
config.text_hidden_size,
config.spatial_merge_size,
Conv2dConfig {
stride: config.spatial_merge_size,
padding: 0,
..Default::default()
},
vb.pp("downsample"),
)
.map_err(|e| {
CandleOcrError::ModelLoadFailed(format!("Failed to load connector downsample Conv2d: {}", e))
})?;
let merger_vb = vb.pp("merger");
let gate_proj = candle_nn::linear_no_bias(
config.text_hidden_size,
config.merger_intermediate_size,
merger_vb.pp("gate_proj"),
)
.map_err(|e| CandleOcrError::ModelLoadFailed(format!("Failed to load merger gate_proj: {}", e)))?;
let up_proj = candle_nn::linear_no_bias(
config.text_hidden_size,
config.merger_intermediate_size,
merger_vb.pp("up_proj"),
)
.map_err(|e| CandleOcrError::ModelLoadFailed(format!("Failed to load merger up_proj: {}", e)))?;
let down_proj = candle_nn::linear_no_bias(
config.merger_intermediate_size,
config.text_hidden_size,
merger_vb.pp("down_proj"),
)
.map_err(|e| CandleOcrError::ModelLoadFailed(format!("Failed to load merger down_proj: {}", e)))?;
let proj =
candle_nn::linear_no_bias(config.text_hidden_size, config.text_hidden_size, merger_vb.pp("proj"))
.map_err(|e| CandleOcrError::ModelLoadFailed(format!("Failed to load merger proj: {}", e)))?;
let post_projection_norm = candle_nn::layer_norm(
config.text_hidden_size,
candle_nn::LayerNormConfig::default(),
merger_vb.pp("post_projection_norm"),
)
.map_err(|e| {
CandleOcrError::ModelLoadFailed(format!("Failed to load merger post_projection_norm: {}", e))
})?;
Ok(Self {
downsample,
gate_proj,
up_proj,
down_proj,
proj,
post_projection_norm,
config: config.clone(),
})
}
pub fn forward(&self, vision_embeds: &Tensor, h_patches: usize, w_patches: usize) -> Result<Tensor> {
let (batch, num_tokens, vision_hidden) = vision_embeds
.dims3()
.map_err(|e| CandleOcrError::InferenceFailed(format!("Invalid input shape: {}", e)))?;
if vision_hidden != self.config.vision_hidden_size {
return Err(CandleOcrError::InferenceFailed(format!(
"Vision hidden mismatch: expected {}, got {}",
self.config.vision_hidden_size, vision_hidden
)));
}
if h_patches * w_patches != num_tokens {
return Err(CandleOcrError::InferenceFailed(format!(
"Patch grid mismatch: h_patches * w_patches = {} but N = {}",
h_patches * w_patches,
num_tokens
)));
}
let merge = self.config.spatial_merge_size;
if merge == 0 || !h_patches.is_multiple_of(merge) || !w_patches.is_multiple_of(merge) {
return Err(CandleOcrError::InferenceFailed(format!(
"Patch grid ({}, {}) not divisible by spatial_merge_size {}",
h_patches, w_patches, merge
)));
}
let x = vision_embeds
.reshape((batch, h_patches, w_patches, vision_hidden))
.and_then(|t| t.permute([0, 3, 1, 2]))
.and_then(|t| t.contiguous())
.map_err(|e| CandleOcrError::InferenceFailed(format!("Connector reshape to (B,C,H,W): {}", e)))?;
let x = self
.downsample
.forward(&x)
.map_err(|e| CandleOcrError::InferenceFailed(format!("Connector downsample Conv2d forward: {}", e)))?;
let (_b, c_out, h_merged, w_merged) = x
.dims4()
.map_err(|e| CandleOcrError::InferenceFailed(format!("Connector downsample output shape: {}", e)))?;
let new_num_tokens = h_merged * w_merged;
let x = x
.permute([0, 2, 3, 1])
.and_then(|t| t.contiguous())
.and_then(|t| t.reshape((batch, new_num_tokens, c_out)))
.map_err(|e| CandleOcrError::InferenceFailed(format!("Connector reshape to (B,N',C): {}", e)))?;
let x = x
.apply(&self.proj)
.map_err(|e| CandleOcrError::InferenceFailed(format!("Merger proj: {}", e)))?;
let x = x
.apply(&self.post_projection_norm)
.map_err(|e| CandleOcrError::InferenceFailed(format!("Merger post_projection_norm: {}", e)))?;
let x = x
.gelu_erf()
.map_err(|e| CandleOcrError::InferenceFailed(format!("Merger post_projection_norm gelu_erf: {}", e)))?;
let gate = x
.apply(&self.gate_proj)
.map_err(|e| CandleOcrError::InferenceFailed(format!("Merger gate_proj: {}", e)))?;
let up = x
.apply(&self.up_proj)
.map_err(|e| CandleOcrError::InferenceFailed(format!("Merger up_proj: {}", e)))?;
let hidden = gate
.silu()
.and_then(|g| g.mul(&up))
.map_err(|e| CandleOcrError::InferenceFailed(format!("Merger SwiGLU activation: {}", e)))?;
let x = hidden
.apply(&self.down_proj)
.map_err(|e| CandleOcrError::InferenceFailed(format!("Merger down_proj: {}", e)))?;
Ok(x)
}
#[deprecated(
note = "Use `forward(vision_embeds, h_patches, w_patches)`; the engine-wiring pass will remove this shim."
)]
pub fn forward_compat(&self, vision_embeds: &Tensor) -> Result<Tensor> {
let (_batch, num_tokens, _hidden) = vision_embeds
.dims3()
.map_err(|e| CandleOcrError::InferenceFailed(format!("Invalid input shape: {}", e)))?;
let side = (num_tokens as f64).sqrt() as usize;
if side * side != num_tokens {
return Err(CandleOcrError::InferenceFailed(format!(
"forward_compat requires a square patch grid; got N = {}",
num_tokens
)));
}
self.forward(vision_embeds, side, side)
}
}
}
#[cfg(not(target_arch = "wasm32"))]
pub use imp::VisionConnector;