use crate::pooling::Pooling;
use anyhow::{Context, Result};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum EncArch {
Bert,
XlmRoberta,
ModernBert,
NomicBert,
SiglipText,
SiglipVision,
Qwen3Embed,
Lfm2Colbert,
DebertaV2,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct RelAttn {
pub span: usize,
pub max_rel: usize,
pub c2p: bool,
pub p2c: bool,
}
impl RelAttn {
pub fn scale_factor(&self) -> usize {
1 + usize::from(self.c2p) + usize::from(self.p2c)
}
pub fn bucket(&self, rel: isize) -> isize {
let mid = (self.span / 2) as isize;
if rel.abs() <= mid {
return rel;
}
let sign = rel.signum() as f32;
let abs_pos = rel.abs() as f32;
let mid_f = mid as f32;
let log_pos = ((abs_pos / mid_f).ln() / ((self.max_rel as f32 - 1.0) / mid_f).ln()
* (mid_f - 1.0))
.ceil()
+ mid_f;
(sign * log_pos) as isize
}
pub fn index(&self, i: usize, j: usize) -> usize {
let span = self.span as isize;
(self.bucket(i as isize - j as isize) + span).clamp(0, 2 * span - 1) as usize
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum PosKind {
Learned {
offset: usize,
},
Rope {
theta: f32,
local_theta: f32,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum NormKind {
LayerNorm {
bias: bool,
},
RmsNorm,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Act {
GeluErf,
GeluTanh,
Silu,
Tanh,
Relu,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MlpKind {
Dense {
act: Act,
bias: bool,
},
Glu {
act: Act,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MaskKind {
Bidirectional,
Causal,
}
#[derive(Clone, Debug)]
pub struct EncoderConfig {
pub arch: EncArch,
pub hidden: usize,
pub n_layers: usize,
pub n_heads: usize,
pub n_kv_heads: usize,
pub head_dim: usize,
pub intermediate: usize,
pub vocab: usize,
pub max_pos: usize,
pub type_vocab: usize,
pub pos_kind: PosKind,
pub norm_kind: NormKind,
pub mlp_kind: MlpKind,
pub attn_mask: MaskKind,
pub layer_window: Vec<u32>,
pub prenorm: bool,
pub skip_first_attn_norm: bool,
pub qkv_bias: bool,
pub qk_norm: bool,
pub conv_l: usize,
pub layer_is_attn: Vec<bool>,
pub pooling: Pooling,
pub eps: f32,
pub rel_attn: Option<RelAttn>,
}
impl EncoderConfig {
pub fn from_json(bytes: &[u8]) -> Result<Self> {
let v: serde_json::Value = serde_json::from_slice(bytes)?;
let g = |k: &str| v.get(k).and_then(|x| x.as_u64()).map(|x| x as usize);
let f = |k: &str| v.get(k).and_then(|x| x.as_f64());
let arch_name = v
.get("architectures")
.and_then(|a| a.get(0))
.and_then(|x| x.as_str())
.context("architectures[0] missing")?;
let arch = encoder_arch(arch_name)
.with_context(|| format!("`{arch_name}` is not a supported encoder architecture"))?;
match arch {
EncArch::Bert | EncArch::XlmRoberta => {
let pos_type = v
.get("position_embedding_type")
.and_then(|x| x.as_str())
.unwrap_or("absolute");
anyhow::ensure!(
pos_type == "absolute",
"BERT-family position_embedding_type `{pos_type}` not supported (absolute only)"
);
let act = act_from(
v.get("hidden_act")
.and_then(|x| x.as_str())
.unwrap_or("gelu"),
)?;
let hidden = g("hidden_size").context("hidden_size")?;
let n_heads = g("num_attention_heads").context("num_attention_heads")?;
let n_layers = g("num_hidden_layers").context("num_hidden_layers")?;
let offset = if arch == EncArch::XlmRoberta {
g("pad_token_id").context("pad_token_id")? + 1
} else {
0
};
let pooling = if arch_name.ends_with("ForSequenceClassification") {
let n_labels = g("num_labels")
.or_else(|| {
v.get("id2label")
.and_then(|m| m.as_object())
.map(|m| m.len())
})
.unwrap_or(1)
.max(1);
Pooling::CrossEncoder { n_labels }
} else {
Pooling::Mean
};
Ok(Self {
arch,
hidden,
n_layers,
n_heads,
n_kv_heads: n_heads,
head_dim: hidden / n_heads,
intermediate: g("intermediate_size").context("intermediate_size")?,
vocab: g("vocab_size").context("vocab_size")?,
max_pos: g("max_position_embeddings").context("max_position_embeddings")?,
type_vocab: g("type_vocab_size").unwrap_or(0),
pos_kind: PosKind::Learned { offset },
norm_kind: NormKind::LayerNorm { bias: true },
mlp_kind: MlpKind::Dense { act, bias: true },
attn_mask: MaskKind::Bidirectional,
layer_window: vec![0; n_layers],
prenorm: false,
skip_first_attn_norm: false,
qkv_bias: true,
qk_norm: false,
conv_l: 0,
layer_is_attn: vec![true; n_layers],
pooling,
eps: f("layer_norm_eps").unwrap_or(1e-12) as f32,
rel_attn: None,
})
}
EncArch::ModernBert => {
let hidden = g("hidden_size").context("hidden_size")?;
let n_heads = g("num_attention_heads").context("num_attention_heads")?;
let n_layers = g("num_hidden_layers").context("num_hidden_layers")?;
let every = g("global_attn_every_n_layers").unwrap_or(3);
anyhow::ensure!(every > 0, "global_attn_every_n_layers must be > 0");
let window = g("local_attention").unwrap_or(128) as u32;
let act = act_from(
v.get("hidden_activation")
.and_then(|x| x.as_str())
.unwrap_or("gelu"),
)?;
let norm_bias = v
.get("norm_bias")
.and_then(|x| x.as_bool())
.unwrap_or(false);
Ok(Self {
arch,
hidden,
n_layers,
n_heads,
n_kv_heads: n_heads,
head_dim: hidden / n_heads,
intermediate: g("intermediate_size").context("intermediate_size")?,
vocab: g("vocab_size").context("vocab_size")?,
max_pos: g("max_position_embeddings").context("max_position_embeddings")?,
type_vocab: 0,
pos_kind: PosKind::Rope {
theta: f("global_rope_theta").unwrap_or(160_000.0) as f32,
local_theta: f("local_rope_theta").unwrap_or(10_000.0) as f32,
},
norm_kind: NormKind::LayerNorm { bias: norm_bias },
mlp_kind: MlpKind::Glu { act },
attn_mask: MaskKind::Bidirectional,
layer_window: (0..n_layers)
.map(|i| if i % every == 0 { 0 } else { window })
.collect(),
prenorm: true,
skip_first_attn_norm: true,
qkv_bias: v
.get("attention_bias")
.and_then(|x| x.as_bool())
.unwrap_or(false),
qk_norm: false,
conv_l: 0,
layer_is_attn: vec![true; n_layers],
pooling: Pooling::Mean,
eps: f("norm_eps").unwrap_or(1e-5) as f32,
rel_attn: None,
})
}
EncArch::NomicBert => {
let hidden = g("n_embd").or_else(|| g("hidden_size")).context("n_embd")?;
let n_heads = g("n_head")
.or_else(|| g("num_attention_heads"))
.context("n_head")?;
let n_layers = g("n_layer")
.or_else(|| g("num_hidden_layers"))
.context("n_layer")?;
let intermediate = g("n_inner")
.or_else(|| g("intermediate_size"))
.context("n_inner")?;
let frac = f("rotary_emb_fraction").unwrap_or(1.0);
anyhow::ensure!(
(frac - 1.0).abs() < 1e-6,
"nomic rotary_emb_fraction {frac} != 1.0 (partial rotary unsupported)"
);
anyhow::ensure!(
!v.get("rotary_emb_interleaved")
.and_then(|x| x.as_bool())
.unwrap_or(false),
"nomic interleaved rotary embeddings unsupported (rotate_half only)"
);
anyhow::ensure!(
v.get("rotary_emb_scale_base").map(|x| x.is_null()) != Some(false),
"nomic rotary_emb_scale_base scaling unsupported"
);
anyhow::ensure!(
v.get("rotary_scaling_factor").map(|x| x.is_null()) != Some(false),
"nomic rotary_scaling_factor unsupported"
);
anyhow::ensure!(
!v.get("use_rms_norm")
.and_then(|x| x.as_bool())
.unwrap_or(false),
"nomic use_rms_norm unsupported (biased LayerNorm only)"
);
anyhow::ensure!(
!v.get("causal").and_then(|x| x.as_bool()).unwrap_or(false),
"nomic causal attention unsupported (bidirectional embedder only)"
);
anyhow::ensure!(
!v.get("prenorm").and_then(|x| x.as_bool()).unwrap_or(false),
"nomic prenorm=true unsupported (nomic-embed is post-norm)"
);
let act = match v.get("activation_function").and_then(|x| x.as_str()) {
Some("swiglu") | Some("silu") | Some("swish") => Act::Silu,
Some(other) => anyhow::bail!("unsupported nomic activation `{other}`"),
None => act_from(
v.get("hidden_act")
.and_then(|x| x.as_str())
.unwrap_or("silu"),
)?,
};
let theta = f("rotary_emb_base").unwrap_or(1000.0) as f32;
Ok(Self {
arch,
hidden,
n_layers,
n_heads,
n_kv_heads: n_heads,
head_dim: hidden / n_heads,
intermediate,
vocab: g("vocab_size").context("vocab_size")?,
max_pos: g("max_position_embeddings")
.or_else(|| g("n_positions"))
.context("max_position_embeddings")?,
type_vocab: g("type_vocab_size").unwrap_or(0),
pos_kind: PosKind::Rope {
theta,
local_theta: theta,
},
norm_kind: NormKind::LayerNorm { bias: true },
mlp_kind: MlpKind::Glu { act },
attn_mask: MaskKind::Bidirectional,
layer_window: vec![0; n_layers],
prenorm: false,
skip_first_attn_norm: false,
qkv_bias: v
.get("qkv_proj_bias")
.and_then(|x| x.as_bool())
.unwrap_or(false),
qk_norm: false,
conv_l: 0,
layer_is_attn: vec![true; n_layers],
pooling: Pooling::Mean,
eps: f("layer_norm_epsilon")
.or_else(|| f("layer_norm_eps"))
.unwrap_or(1e-12) as f32,
rel_attn: None,
})
}
EncArch::DebertaV2 => {
let hidden = g("hidden_size").context("hidden_size")?;
let n_heads = g("num_attention_heads").context("num_attention_heads")?;
let n_layers = g("num_hidden_layers").context("num_hidden_layers")?;
let max_pos = g("max_position_embeddings").context("max_position_embeddings")?;
anyhow::ensure!(
v.get("relative_attention").and_then(|x| x.as_bool()) == Some(true),
"DeBERTa without relative_attention is not supported (nothing would position \
the tokens: position_biased_input is false)"
);
anyhow::ensure!(
!v.get("position_biased_input")
.and_then(|x| x.as_bool())
.unwrap_or(false),
"DeBERTa with position_biased_input=true (absolute position embeddings) is \
untested — refuse loudly"
);
let norm_rel = v
.get("norm_rel_ebd")
.and_then(|x| x.as_str())
.unwrap_or("none");
anyhow::ensure!(
norm_rel.contains("layer_norm"),
"DeBERTa norm_rel_ebd `{norm_rel}` not supported (layer_norm only)"
);
anyhow::ensure!(
v.get("share_att_key").and_then(|x| x.as_bool()) == Some(true),
"DeBERTa with share_att_key=false carries separate pos_key_proj/pos_query_proj \
weights this loader does not read — refuse loudly"
);
let pos_att: Vec<&str> = match v.get("pos_att_type") {
Some(serde_json::Value::Array(a)) => {
a.iter().filter_map(|x| x.as_str()).collect()
}
Some(serde_json::Value::String(s)) => s.split('|').collect(),
_ => anyhow::bail!("pos_att_type missing"),
};
let rel = RelAttn {
span: g("position_buckets").context("position_buckets")?,
max_rel: v
.get("max_relative_positions")
.and_then(|x| x.as_i64())
.filter(|x| *x > 0)
.map(|x| x as usize)
.unwrap_or(max_pos),
c2p: pos_att.contains(&"c2p"),
p2c: pos_att.contains(&"p2c"),
};
anyhow::ensure!(
rel.c2p || rel.p2c,
"DeBERTa pos_att_type has neither c2p nor p2c — nothing positions the tokens"
);
Ok(Self {
arch,
hidden,
n_layers,
n_heads,
n_kv_heads: n_heads,
head_dim: hidden / n_heads,
intermediate: g("intermediate_size").context("intermediate_size")?,
vocab: g("vocab_size").context("vocab_size")?,
max_pos,
type_vocab: g("type_vocab_size").unwrap_or(0),
pos_kind: PosKind::Learned { offset: 0 },
norm_kind: NormKind::LayerNorm { bias: true },
mlp_kind: MlpKind::Dense {
act: act_from(
v.get("hidden_act")
.and_then(|x| x.as_str())
.unwrap_or("gelu"),
)?,
bias: true,
},
attn_mask: MaskKind::Bidirectional,
layer_window: vec![0; n_layers],
prenorm: false,
skip_first_attn_norm: false,
qkv_bias: true,
qk_norm: false,
conv_l: 0,
layer_is_attn: vec![true; n_layers],
pooling: Pooling::Mean,
eps: f("layer_norm_eps").unwrap_or(1e-7) as f32,
rel_attn: Some(rel),
})
}
EncArch::SiglipText
| EncArch::SiglipVision
| EncArch::Qwen3Embed
| EncArch::Lfm2Colbert => anyhow::bail!(
"encoder architecture `{arch_name}` is recognized but not yet loadable by this build"
),
}
}
}
#[derive(Clone, Debug)]
pub struct EncBatch {
pub tokens: Vec<u32>,
pub seq_starts: Vec<u32>,
pub valid: Option<Vec<u32>>,
pub type_ids: Option<Vec<u32>>,
}
impl EncBatch {
pub fn single(tokens: Vec<u32>) -> Self {
let end = tokens.len() as u32;
Self {
tokens,
seq_starts: vec![0, end],
valid: None,
type_ids: None,
}
}
pub fn single_with_validity(tokens: Vec<u32>, valid: Vec<u32>) -> Self {
let end = tokens.len() as u32;
Self {
tokens,
seq_starts: vec![0, end],
valid: Some(valid),
type_ids: None,
}
}
pub fn single_pair(tokens: Vec<u32>, type_ids: Vec<u32>) -> Self {
let end = tokens.len() as u32;
Self {
tokens,
seq_starts: vec![0, end],
valid: None,
type_ids: Some(type_ids),
}
}
pub fn from_seqs<I: IntoIterator<Item = Vec<u32>>>(seqs: I) -> Self {
let mut tokens = Vec::new();
let mut seq_starts = vec![0u32];
for s in seqs {
tokens.extend_from_slice(&s);
seq_starts.push(tokens.len() as u32);
}
Self {
tokens,
seq_starts,
valid: None,
type_ids: None,
}
}
pub fn from_pairs<I: IntoIterator<Item = (Vec<u32>, Vec<u32>)>>(seqs: I) -> Self {
let mut tokens = Vec::new();
let mut type_ids = Vec::new();
let mut seq_starts = vec![0u32];
for (s, t) in seqs {
tokens.extend_from_slice(&s);
type_ids.extend_from_slice(&t);
seq_starts.push(tokens.len() as u32);
}
Self {
tokens,
seq_starts,
valid: None,
type_ids: Some(type_ids),
}
}
pub fn n_seqs(&self) -> usize {
self.seq_starts.len().saturating_sub(1)
}
}
pub fn encoder_arch(name: &str) -> Option<EncArch> {
if name.starts_with("Bert") {
Some(EncArch::Bert)
} else if name.starts_with("XLMRoberta") {
Some(EncArch::XlmRoberta)
} else if name.starts_with("ModernBert") {
Some(EncArch::ModernBert)
} else if name.starts_with("NomicBert") {
Some(EncArch::NomicBert)
} else if name.starts_with("Siglip") {
Some(EncArch::SiglipText)
} else if name.starts_with("DebertaV2") || name.starts_with("DebertaV3") {
Some(EncArch::DebertaV2)
} else {
None
}
}
fn act_from(name: &str) -> Result<Act> {
Ok(match name {
"relu" => Act::Relu,
"gelu" => Act::GeluErf,
"gelu_new" | "gelu_pytorch_tanh" => Act::GeluTanh,
"silu" | "swish" => Act::Silu,
"tanh" => Act::Tanh,
other => anyhow::bail!("unsupported activation `{other}`"),
})
}
pub fn pooling_marker(dir: &std::path::Path) -> Option<Pooling> {
let bytes = std::fs::read(dir.join("1_Pooling").join("config.json")).ok()?;
let v: serde_json::Value = serde_json::from_slice(&bytes).ok()?;
let flag = |k: &str| v.get(k).and_then(|x| x.as_bool()).unwrap_or(false);
if flag("pooling_mode_lasttoken") {
Some(Pooling::LastToken)
} else if flag("pooling_mode_cls_token") {
Some(Pooling::Cls)
} else if flag("pooling_mode_mean_tokens") {
Some(Pooling::Mean)
} else {
None
}
}
pub fn encoder_config_from_dir(dir: &std::path::Path) -> Result<EncoderConfig> {
let bytes = std::fs::read(dir.join("config.json"))
.with_context(|| format!("read {}/config.json", dir.display()))?;
let v: serde_json::Value = serde_json::from_slice(&bytes)?;
let arch_name = v
.get("architectures")
.and_then(|a| a.get(0))
.and_then(|x| x.as_str())
.context("architectures[0] missing")?;
let marker = pooling_marker(dir);
if encoder_arch(arch_name).is_some() {
let mut cfg = EncoderConfig::from_json(&bytes)?;
if let Some(p) = marker {
cfg.pooling = p;
}
return Ok(cfg);
}
if arch_name == "Lfm2Model" {
let dense_path = dir.join("1_Dense").join("config.json");
anyhow::ensure!(
dense_path.exists(),
"`Lfm2Model` is a bare backbone; embedding it requires the pylate Dense module \
(1_Dense/config.json) — LFM2-ColBERT checkpoints ship it, plain backbones are \
refused rather than silently mis-embedded"
);
let dense: serde_json::Value = serde_json::from_slice(&std::fs::read(&dense_path)?)?;
return lfm2_colbert_from_json(&v, &dense);
}
if arch_name == "Qwen3ForCausalLM" {
anyhow::ensure!(
marker == Some(Pooling::LastToken),
"`{arch_name}` is a decoder; embedding it requires the sentence-transformers \
last-token marker (1_Pooling/config.json with pooling_mode_lasttoken) — a plain \
generation checkpoint is refused rather than silently mis-embedded"
);
return qwen3_embedder_from_json(&v);
}
anyhow::bail!("`{arch_name}` is not a supported encoder architecture")
}
fn qwen3_embedder_from_json(v: &serde_json::Value) -> Result<EncoderConfig> {
let g = |k: &str| v.get(k).and_then(|x| x.as_u64()).map(|x| x as usize);
let f = |k: &str| v.get(k).and_then(|x| x.as_f64());
let hidden = g("hidden_size").context("hidden_size")?;
let n_heads = g("num_attention_heads").context("num_attention_heads")?;
let n_layers = g("num_hidden_layers").context("num_hidden_layers")?;
let head_dim = g("head_dim").unwrap_or(hidden / n_heads);
let n_kv_heads = g("num_key_value_heads").unwrap_or(n_heads);
anyhow::ensure!(
n_heads % n_kv_heads == 0,
"GQA requires n_heads ({n_heads}) divisible by n_kv_heads ({n_kv_heads})"
);
let theta = f("rope_theta").unwrap_or(1_000_000.0) as f32;
Ok(EncoderConfig {
arch: EncArch::Qwen3Embed,
hidden,
n_layers,
n_heads,
n_kv_heads,
head_dim,
intermediate: g("intermediate_size").context("intermediate_size")?,
vocab: g("vocab_size").context("vocab_size")?,
max_pos: g("max_position_embeddings").context("max_position_embeddings")?,
type_vocab: 0,
pos_kind: PosKind::Rope {
theta,
local_theta: theta,
},
norm_kind: NormKind::RmsNorm,
mlp_kind: MlpKind::Glu { act: Act::Silu },
attn_mask: MaskKind::Causal,
layer_window: vec![0; n_layers],
prenorm: true,
skip_first_attn_norm: false,
qkv_bias: false,
qk_norm: true,
conv_l: 0,
layer_is_attn: vec![true; n_layers],
pooling: Pooling::LastToken,
eps: f("rms_norm_eps").unwrap_or(1e-6) as f32,
rel_attn: None,
})
}
fn lfm2_colbert_from_json(
v: &serde_json::Value,
dense: &serde_json::Value,
) -> Result<EncoderConfig> {
let g = |k: &str| v.get(k).and_then(|x| x.as_u64()).map(|x| x as usize);
let f = |k: &str| v.get(k).and_then(|x| x.as_f64());
let hidden = g("hidden_size").context("hidden_size")?;
let n_heads = g("num_attention_heads").context("num_attention_heads")?;
let n_layers = g("num_hidden_layers").context("num_hidden_layers")?;
let n_kv_heads = g("num_key_value_heads").unwrap_or(n_heads);
anyhow::ensure!(n_heads % n_kv_heads == 0, "GQA head divisibility");
anyhow::ensure!(
!dense.get("bias").and_then(|x| x.as_bool()).unwrap_or(false),
"ColBERT Dense with bias is unverified — refused"
);
if let Some(act) = dense.get("activation_function").and_then(|x| x.as_str()) {
anyhow::ensure!(
act.ends_with("Identity"),
"ColBERT Dense activation `{act}` unverified (Identity only)"
);
}
let dim = dense
.get("out_features")
.and_then(|x| x.as_u64())
.context("1_Dense out_features")? as usize;
anyhow::ensure!(
dense.get("in_features").and_then(|x| x.as_u64()) == Some(hidden as u64),
"Dense in_features must equal hidden"
);
let raw_ff = g("intermediate_size").context("intermediate_size")?;
let intermediate = if v
.get("block_auto_adjust_ff_dim")
.and_then(|x| x.as_bool())
.unwrap_or(false)
{
let m = g("block_multiple_of").unwrap_or(256);
let ff = (2 * raw_ff) / 3;
m * ff.div_ceil(m)
} else {
raw_ff
};
let layer_is_attn: Vec<bool> = v
.get("layer_types")
.and_then(|x| x.as_array())
.context("layer_types")?
.iter()
.map(|x| x.as_str() == Some("full_attention"))
.collect();
anyhow::ensure!(layer_is_attn.len() == n_layers, "layer_types length");
let theta = f("rope_theta").unwrap_or(1_000_000.0) as f32;
Ok(EncoderConfig {
arch: EncArch::Lfm2Colbert,
hidden,
n_layers,
n_heads,
n_kv_heads,
head_dim: g("head_dim").unwrap_or(hidden / n_heads),
intermediate,
vocab: g("vocab_size").context("vocab_size")?,
max_pos: g("max_position_embeddings").context("max_position_embeddings")?,
type_vocab: 0,
pos_kind: PosKind::Rope {
theta,
local_theta: theta,
},
norm_kind: NormKind::RmsNorm,
mlp_kind: MlpKind::Glu { act: Act::Silu },
attn_mask: MaskKind::Causal,
layer_window: vec![0; n_layers],
prenorm: true,
skip_first_attn_norm: false,
qkv_bias: false,
qk_norm: true,
conv_l: g("conv_L_cache").unwrap_or(3),
layer_is_attn,
pooling: Pooling::PerToken { dim },
eps: f("norm_eps").unwrap_or(1e-5) as f32,
rel_attn: None,
})
}
pub fn siglip_text_config(v: &serde_json::Value) -> Result<EncoderConfig> {
let t = v.get("text_config").context("text_config")?;
let g = |k: &str| t.get(k).and_then(|x| x.as_u64()).map(|x| x as usize);
let f = |k: &str| t.get(k).and_then(|x| x.as_f64());
let hidden = g("hidden_size").unwrap_or(768);
let n_heads = g("num_attention_heads").unwrap_or(12);
let n_layers = g("num_hidden_layers").unwrap_or(12);
let act = act_from(
t.get("hidden_act")
.and_then(|x| x.as_str())
.unwrap_or("gelu_pytorch_tanh"),
)?;
Ok(EncoderConfig {
arch: EncArch::SiglipText,
hidden,
n_layers,
n_heads,
n_kv_heads: n_heads,
head_dim: hidden / n_heads,
intermediate: g("intermediate_size").unwrap_or(3072),
vocab: g("vocab_size").unwrap_or(256_000),
max_pos: g("max_position_embeddings").unwrap_or(64),
type_vocab: 0,
pos_kind: PosKind::Learned { offset: 0 },
norm_kind: NormKind::LayerNorm { bias: true },
mlp_kind: MlpKind::Dense { act, bias: true },
attn_mask: MaskKind::Bidirectional,
layer_window: vec![0; n_layers],
prenorm: true,
skip_first_attn_norm: false,
qkv_bias: true,
qk_norm: false,
conv_l: 0,
layer_is_attn: vec![true; n_layers],
pooling: Pooling::LastToken,
eps: f("layer_norm_eps").unwrap_or(1e-6) as f32,
rel_attn: None,
})
}
pub fn siglip_vision_config(v: &serde_json::Value) -> Result<SiglipVisionSpec> {
let c = v.get("vision_config").context("vision_config")?;
let g = |k: &str| c.get(k).and_then(|x| x.as_u64()).map(|x| x as usize);
let f = |k: &str| c.get(k).and_then(|x| x.as_f64());
let hidden = g("hidden_size").unwrap_or(768);
let n_heads = g("num_attention_heads").unwrap_or(12);
let n_layers = g("num_hidden_layers").unwrap_or(12);
let image_size = g("image_size").unwrap_or(224);
let patch_size = g("patch_size").unwrap_or(16);
anyhow::ensure!(
image_size % patch_size == 0,
"image_size {image_size} not divisible by patch_size {patch_size}"
);
let n_patches = (image_size / patch_size) * (image_size / patch_size);
let act = act_from(
c.get("hidden_act")
.and_then(|x| x.as_str())
.unwrap_or("gelu_pytorch_tanh"),
)?;
Ok(SiglipVisionSpec {
config: EncoderConfig {
arch: EncArch::SiglipVision,
hidden,
n_layers,
n_heads,
n_kv_heads: n_heads,
head_dim: hidden / n_heads,
intermediate: g("intermediate_size").unwrap_or(3072),
vocab: 0,
max_pos: n_patches,
type_vocab: 0,
pos_kind: PosKind::Learned { offset: 0 },
norm_kind: NormKind::LayerNorm { bias: true },
mlp_kind: MlpKind::Dense { act, bias: true },
attn_mask: MaskKind::Bidirectional,
layer_window: vec![0; n_layers],
prenorm: true,
skip_first_attn_norm: false,
qkv_bias: true,
qk_norm: false,
conv_l: 0,
layer_is_attn: vec![true; n_layers],
pooling: Pooling::MapHead,
eps: f("layer_norm_eps").unwrap_or(1e-6) as f32,
rel_attn: None,
},
image_size,
patch_size,
})
}
#[derive(Clone, Debug)]
pub struct SiglipVisionSpec {
pub config: EncoderConfig,
pub image_size: usize,
pub patch_size: usize,
}
pub fn siglip_configs_from_dir(dir: &std::path::Path) -> Result<(EncoderConfig, SiglipVisionSpec)> {
let bytes = std::fs::read(dir.join("config.json"))
.with_context(|| format!("read {}/config.json", dir.display()))?;
let v: serde_json::Value = serde_json::from_slice(&bytes)?;
let is_siglip = v
.get("model_type")
.and_then(|x| x.as_str())
.is_some_and(|t| t.starts_with("siglip"))
|| v.get("architectures")
.and_then(|a| a.get(0))
.and_then(|x| x.as_str())
.is_some_and(|a| a.starts_with("Siglip"));
anyhow::ensure!(is_siglip, "{} is not a SigLIP checkpoint", dir.display());
Ok((siglip_text_config(&v)?, siglip_vision_config(&v)?))
}