use super::bert::{ConfigJson, Shape, TokenEmbeddings, padding_bias};
use candle_core::{D, Device, Module, Result, Tensor};
use candle_nn::{LayerNorm, Linear, VarBuilder};
#[derive(Debug, Clone, PartialEq)]
pub struct Config {
pub shape: Shape,
pub gelu_gate: bool,
pub qk_norm: bool,
}
impl Config {
pub fn from_json(text: &str) -> std::result::Result<Self, String> {
let v = ConfigJson::parse(text)?;
if v.text("position_embedding_type") != "alibi" {
return Err("config.json: not a JinaBERT v2 model (no ALiBi)".to_string());
}
let gelu_gate = match v.text("feed_forward_type") {
"geglu" => true,
"reglu" => false,
other => {
return Err(format!(
"config.json: unsupported feed_forward_type '{other}'"
));
}
};
Ok(Self {
shape: Shape {
vocab_size: v.int("vocab_size")?,
hidden_size: v.int("hidden_size")?,
num_hidden_layers: v.int("num_hidden_layers")?,
num_attention_heads: v.int("num_attention_heads")?,
intermediate_size: v.int("intermediate_size")?,
layer_norm_eps: v.number("layer_norm_eps").unwrap_or(1e-12),
type_vocab_size: v.int_or("type_vocab_size", 2)?,
},
gelu_gate,
qk_norm: v.text("_name_or_path").contains("qk-post-norm"),
})
}
}
pub fn alibi_slopes(heads: usize) -> Vec<f32> {
fn power_of_two(n: usize) -> Vec<f64> {
let start = 2f64.powf(-(2f64.powf(-((n as f64).log2() - 3.0))));
(0..n).map(|i| start * start.powi(i as i32)).collect()
}
let slopes: Vec<f64> = if heads.is_power_of_two() {
power_of_two(heads)
} else {
let closest = 1usize << (usize::BITS - 1 - heads.leading_zeros());
let mut s = power_of_two(closest);
s.extend(
power_of_two(2 * closest)
.into_iter()
.step_by(2)
.take(heads - closest),
);
s
};
slopes.into_iter().map(|s| s as f32).collect()
}
fn alibi_bias(slopes: &[f32], len: usize, device: &Device) -> Result<Tensor> {
let mut data = Vec::with_capacity(slopes.len() * len * len);
for &slope in slopes {
for i in 0..len {
for j in 0..len {
data.push(-slope * i.abs_diff(j) as f32);
}
}
}
Tensor::from_vec(data, (1, slopes.len(), len, len), device)
}
struct SelfAttention {
query: Linear,
key: Linear,
value: Linear,
norm_q: Option<LayerNorm>,
norm_k: Option<LayerNorm>,
dense: Linear,
norm_out: LayerNorm,
heads: usize,
head_dim: usize,
}
impl SelfAttention {
fn load(vb: VarBuilder, cfg: &Config) -> Result<Self> {
let h = cfg.shape.hidden_size;
let eps = cfg.shape.layer_norm_eps;
let own = vb.pp("self");
let norm = |name: &str| {
cfg.qk_norm
.then(|| candle_nn::layer_norm(h, eps, own.pp(name)))
.transpose()
};
Ok(Self {
query: candle_nn::linear(h, h, own.pp("query"))?,
key: candle_nn::linear(h, h, own.pp("key"))?,
value: candle_nn::linear(h, h, own.pp("value"))?,
norm_q: norm("layer_norm_q")?,
norm_k: norm("layer_norm_k")?,
dense: candle_nn::linear(h, h, vb.pp("output").pp("dense"))?,
norm_out: candle_nn::layer_norm(h, eps, vb.pp("output").pp("LayerNorm"))?,
heads: cfg.shape.num_attention_heads,
head_dim: h / cfg.shape.num_attention_heads,
})
}
fn forward(&self, x: &Tensor, alibi: &Tensor, mask: &Tensor) -> Result<Tensor> {
let (b, len, hidden) = x.dims3()?;
let project = |linear: &Linear, norm: &Option<LayerNorm>| -> Result<Tensor> {
let t = linear.forward(x)?;
let t = match norm {
Some(norm) => norm.forward(&t)?,
None => t,
};
t.reshape((b, len, self.heads, self.head_dim))?
.transpose(1, 2)?
.contiguous()
};
let q = project(&self.query, &self.norm_q)?;
let k = project(&self.key, &self.norm_k)?;
let v = project(&self.value, &None)?;
let scores = (q.matmul(&k.t()?)? * (1.0 / (self.head_dim as f64).sqrt()))?;
let scores = scores.broadcast_add(alibi)?.broadcast_add(mask)?;
let probs = candle_nn::ops::softmax_last_dim(&scores)?;
let context = probs
.matmul(&v)?
.transpose(1, 2)?
.reshape((b, len, hidden))?;
self.norm_out.forward(&(self.dense.forward(&context)? + x)?)
}
}
struct Layer {
attention: SelfAttention,
norm_1: LayerNorm,
up_gated: Linear,
down: Linear,
norm_2: LayerNorm,
intermediate: usize,
gelu_gate: bool,
}
impl Layer {
fn load(vb: VarBuilder, cfg: &Config) -> Result<Self> {
let Shape {
hidden_size: h,
intermediate_size: i,
layer_norm_eps: eps,
..
} = cfg.shape;
Ok(Self {
attention: SelfAttention::load(vb.pp("attention"), cfg)?,
norm_1: candle_nn::layer_norm(h, eps, vb.pp("layer_norm_1"))?,
up_gated: candle_nn::linear_no_bias(h, 2 * i, vb.pp("mlp").pp("up_gated_layer"))?,
down: candle_nn::linear(i, h, vb.pp("mlp").pp("down_layer"))?,
norm_2: candle_nn::layer_norm(h, eps, vb.pp("layer_norm_2"))?,
intermediate: i,
gelu_gate: cfg.gelu_gate,
})
}
fn forward(&self, x: &Tensor, alibi: &Tensor, mask: &Tensor) -> Result<Tensor> {
let attended = self.attention.forward(x, alibi, mask)?;
let x = self.norm_1.forward(&(x + attended)?)?;
let up_gated = self.up_gated.forward(&x)?;
let up = up_gated.narrow(D::Minus1, 0, self.intermediate)?;
let gate = up_gated.narrow(D::Minus1, self.intermediate, self.intermediate)?;
let gate = match self.gelu_gate {
true => gate.gelu_erf()?,
false => gate.relu()?,
};
let mlp = self.down.forward(&(up * gate)?)?;
self.norm_2.forward(&(x + mlp)?)
}
}
pub struct JinaBert {
embeddings: TokenEmbeddings,
layers: Vec<Layer>,
slopes: Vec<f32>,
device: Device,
}
impl JinaBert {
pub fn load(vb: VarBuilder, cfg: &Config) -> Result<Self> {
Ok(Self {
embeddings: TokenEmbeddings::load(
&vb,
vb.pp("embeddings").pp("LayerNorm"),
&cfg.shape,
)?,
layers: (0..cfg.shape.num_hidden_layers)
.map(|n| Layer::load(vb.pp("encoder").pp("layer").pp(n), cfg))
.collect::<Result<_>>()?,
slopes: alibi_slopes(cfg.shape.num_attention_heads),
device: vb.device().clone(),
})
}
pub fn embed(&self, ids: &Tensor, mask: &Tensor) -> Result<Tensor> {
let (_, len) = ids.dims2()?;
let mut x = self.embeddings.forward(ids)?;
let alibi = alibi_bias(&self.slopes, len, &self.device)?;
let padding = padding_bias(mask)?;
for layer in &self.layers {
x = layer.forward(&x, &alibi, &padding)?;
}
let weights = mask.unsqueeze(2)?;
let summed = x.broadcast_mul(&weights)?.sum(1)?;
let counts = weights.sum(1)?.clamp(1e-9, f64::MAX)?;
summed.broadcast_div(&counts)
}
}
#[cfg(test)]
pub(crate) fn tiny_weights() -> (Config, super::bert::Tensors) {
let cfg = Config {
shape: super::bert::TINY_SHAPE,
gelu_gate: true,
qk_norm: true,
};
let mut weights = super::bert::TestWeights::new();
let mut put = |name: String, shape: &[usize]| weights.put(name, shape);
let (h, i) = (cfg.shape.hidden_size, cfg.shape.intermediate_size);
put(
"embeddings.word_embeddings.weight".into(),
&[cfg.shape.vocab_size, h],
);
put("embeddings.token_type_embeddings.weight".into(), &[2, h]);
for part in ["weight", "bias"] {
put(format!("embeddings.LayerNorm.{part}"), &[h]);
}
for n in 0..cfg.shape.num_hidden_layers {
let p = format!("encoder.layer.{n}");
for proj in ["query", "key", "value"] {
put(format!("{p}.attention.self.{proj}.weight"), &[h, h]);
put(format!("{p}.attention.self.{proj}.bias"), &[h]);
}
for norm in [
"attention.self.layer_norm_q",
"attention.self.layer_norm_k",
"attention.output.LayerNorm",
"layer_norm_1",
"layer_norm_2",
] {
put(format!("{p}.{norm}.weight"), &[h]);
put(format!("{p}.{norm}.bias"), &[h]);
}
put(format!("{p}.attention.output.dense.weight"), &[h, h]);
put(format!("{p}.attention.output.dense.bias"), &[h]);
put(format!("{p}.mlp.up_gated_layer.weight"), &[2 * i, h]);
put(format!("{p}.mlp.down_layer.weight"), &[h, i]);
put(format!("{p}.mlp.down_layer.bias"), &[h]);
}
(cfg, weights.tensors)
}
#[cfg(test)]
mod tests {
use super::*;
use candle_core::DType;
#[test]
fn slopes_match_the_reference_for_twelve_heads() {
let s = alibi_slopes(12);
let expected: Vec<f32> = (1..=8)
.map(|k| 2f32.powi(-k))
.chain([0.5f32, 1.5, 2.5, 3.5].map(|e| 2f32.powf(-e)))
.collect();
assert_eq!(s.len(), 12);
for (got, want) in s.iter().zip(&expected) {
assert!((got - want).abs() < 1e-7, "{s:?}");
}
assert_eq!(alibi_slopes(8), expected[..8].to_vec());
}
#[test]
fn config_reads_the_code_model_and_refuses_others() {
let text = r#"{"_name_or_path": "jinaai/jina-bert-v2-qk-post-norm", "position_embedding_type": "alibi",
"feed_forward_type": "geglu", "vocab_size": 61056, "hidden_size": 768, "num_hidden_layers": 12,
"num_attention_heads": 12, "intermediate_size": 3072, "layer_norm_eps": 1e-12}"#;
let cfg = Config::from_json(text).unwrap();
assert!(cfg.qk_norm && cfg.gelu_gate);
assert_eq!(
(cfg.shape.hidden_size, cfg.shape.num_hidden_layers),
(768, 12)
);
let absolute = text.replace("alibi", "absolute");
assert!(Config::from_json(&absolute).unwrap_err().contains("ALiBi"));
}
fn tiny() -> (JinaBert, Config) {
let (cfg, tensors) = tiny_weights();
let vb = VarBuilder::from_tensors(tensors, DType::F32, &Device::Cpu);
(JinaBert::load(vb, &cfg).unwrap(), cfg)
}
#[test]
fn padding_does_not_change_an_embedding() {
let (model, cfg) = tiny();
let alone = model
.embed(
&Tensor::new(&[[0u32, 5, 7, 2]], &Device::Cpu).unwrap(),
&Tensor::new(&[[1f32, 1., 1., 1.]], &Device::Cpu).unwrap(),
)
.unwrap()
.to_vec2::<f32>()
.unwrap();
let padded = model
.embed(
&Tensor::new(&[[0u32, 5, 7, 2, 1, 1], [0, 3, 2, 1, 1, 1]], &Device::Cpu).unwrap(),
&Tensor::new(
&[[1f32, 1., 1., 1., 0., 0.], [1., 1., 1., 0., 0., 0.]],
&Device::Cpu,
)
.unwrap(),
)
.unwrap()
.to_vec2::<f32>()
.unwrap();
assert_eq!(alone[0].len(), cfg.shape.hidden_size);
for (a, b) in alone[0].iter().zip(&padded[0]) {
assert!((a - b).abs() < 1e-5, "{alone:?} vs {padded:?}");
}
assert_ne!(padded[0], padded[1]);
}
}