use std::path::Path;
use anyhow::{Context as _, Result, ensure};
use rayon::prelude::*;
use super::gguf::{Gguf, TensorHandle, TensorType};
use super::kernels::{
dequantize_row, dim_to_f32, dot_f32, matrix_matrix, matrix_matrix_triple, softmax, vector_add,
};
use super::wordpiece::WordPieceTokenizer;
const LAYERS: usize = 12;
const HIDDEN: usize = 384;
const FEED_FORWARD: usize = 1_536;
const HEADS: usize = 12;
const HEAD_DIMENSION: usize = HIDDEN / HEADS;
const CONTEXT: usize = 512;
const VOCABULARY: usize = 30_522;
const TOKEN_TYPES: usize = 2;
const EPSILON: f32 = 1.0e-12;
const CLS_POOLING: usize = 2;
pub const EMBEDDING_DIMENSIONS: usize = HIDDEN;
pub struct EmbeddingModel {
gguf: Gguf,
tokenizer: WordPieceTokenizer,
layers: Vec<LayerWeights>,
token_embeddings: TensorHandle,
token_types: TensorHandle,
position_embeddings: TensorHandle,
token_norm: TensorHandle,
token_norm_bias: TensorHandle,
}
struct LayerWeights {
query: TensorHandle,
query_bias: TensorHandle,
key: TensorHandle,
key_bias: TensorHandle,
value: TensorHandle,
value_bias: TensorHandle,
attention_output: TensorHandle,
attention_output_bias: TensorHandle,
attention_norm: TensorHandle,
attention_norm_bias: TensorHandle,
feed_forward_up: TensorHandle,
feed_forward_up_bias: TensorHandle,
feed_forward_down: TensorHandle,
feed_forward_down_bias: TensorHandle,
output_norm: TensorHandle,
output_norm_bias: TensorHandle,
}
struct LayerNames {
query: String,
query_bias: String,
key: String,
key_bias: String,
value: String,
value_bias: String,
attention_output: String,
attention_output_bias: String,
attention_norm: String,
attention_norm_bias: String,
feed_forward_up: String,
feed_forward_up_bias: String,
feed_forward_down: String,
feed_forward_down_bias: String,
output_norm: String,
output_norm_bias: String,
}
impl LayerNames {
fn new(layer: usize) -> Self {
let prefix = format!("blk.{layer}");
Self {
query: format!("{prefix}.attn_q.weight"),
query_bias: format!("{prefix}.attn_q.bias"),
key: format!("{prefix}.attn_k.weight"),
key_bias: format!("{prefix}.attn_k.bias"),
value: format!("{prefix}.attn_v.weight"),
value_bias: format!("{prefix}.attn_v.bias"),
attention_output: format!("{prefix}.attn_output.weight"),
attention_output_bias: format!("{prefix}.attn_output.bias"),
attention_norm: format!("{prefix}.attn_output_norm.weight"),
attention_norm_bias: format!("{prefix}.attn_output_norm.bias"),
feed_forward_up: format!("{prefix}.ffn_up.weight"),
feed_forward_up_bias: format!("{prefix}.ffn_up.bias"),
feed_forward_down: format!("{prefix}.ffn_down.weight"),
feed_forward_down_bias: format!("{prefix}.ffn_down.bias"),
output_norm: format!("{prefix}.layer_output_norm.weight"),
output_norm_bias: format!("{prefix}.layer_output_norm.bias"),
}
}
}
impl LayerWeights {
fn resolve(gguf: &Gguf, names: &LayerNames) -> Result<Self> {
Ok(Self {
query: gguf.resolve(&names.query)?,
query_bias: gguf.resolve(&names.query_bias)?,
key: gguf.resolve(&names.key)?,
key_bias: gguf.resolve(&names.key_bias)?,
value: gguf.resolve(&names.value)?,
value_bias: gguf.resolve(&names.value_bias)?,
attention_output: gguf.resolve(&names.attention_output)?,
attention_output_bias: gguf.resolve(&names.attention_output_bias)?,
attention_norm: gguf.resolve(&names.attention_norm)?,
attention_norm_bias: gguf.resolve(&names.attention_norm_bias)?,
feed_forward_up: gguf.resolve(&names.feed_forward_up)?,
feed_forward_up_bias: gguf.resolve(&names.feed_forward_up_bias)?,
feed_forward_down: gguf.resolve(&names.feed_forward_down)?,
feed_forward_down_bias: gguf.resolve(&names.feed_forward_down_bias)?,
output_norm: gguf.resolve(&names.output_norm)?,
output_norm_bias: gguf.resolve(&names.output_norm_bias)?,
})
}
}
impl EmbeddingModel {
pub fn load(path: impl AsRef<Path>) -> Result<Self> {
let gguf = Gguf::load(path).context("load BGE-small-en-v1.5 GGUF")?;
validate_metadata(&gguf).context("validate BGE-small-en-v1.5 metadata")?;
validate_tensors(&gguf).context("validate BGE-small-en-v1.5 tensors")?;
let tokenizer = WordPieceTokenizer::from_bert(&gguf).context("load BERT tokenizer")?;
let layers = (0..LAYERS)
.map(|layer| LayerWeights::resolve(&gguf, &LayerNames::new(layer)))
.collect::<Result<Vec<_>>>()
.context("resolve BERT layer weights")?;
let token_embeddings = gguf.resolve("token_embd.weight")?;
let token_types = gguf.resolve("token_types.weight")?;
let position_embeddings = gguf.resolve("position_embd.weight")?;
let token_norm = gguf.resolve("token_embd_norm.weight")?;
let token_norm_bias = gguf.resolve("token_embd_norm.bias")?;
Ok(Self {
gguf,
tokenizer,
layers,
token_embeddings,
token_types,
position_embeddings,
token_norm,
token_norm_bias,
})
}
pub fn embed(&self, text: &str) -> Result<Vec<f32>> {
let tokens = self
.tokenizer
.encode(text, CONTEXT)
.context("tokenize embedding input")?;
let input = self.embeddings(&tokens)?;
let hidden = self.forward(&input, tokens.len())?;
normalize_embedding(hidden[..HIDDEN].to_vec())
}
fn embeddings(&self, tokens: &[u32]) -> Result<Vec<f32>> {
let token_embeddings = self.gguf.tensor_from_handle(&self.token_embeddings);
let token_types = self.gguf.tensor_from_handle(&self.token_types);
let position_embeddings = self.gguf.tensor_from_handle(&self.position_embeddings);
let token_type = token_types.f32_row(0)?;
let mut output = vec![0.0; tokens.len() * HIDDEN];
for (position, (token, row)) in tokens
.iter()
.zip(output.chunks_exact_mut(HIDDEN))
.enumerate()
{
dequantize_row(token_embeddings.q8_row(usize::try_from(*token)?)?, row)?;
for ((value, type_value), position_value) in row
.iter_mut()
.zip(token_type)
.zip(position_embeddings.f32_row(position)?)
{
*value += type_value + position_value;
}
}
let weight = self.gguf.tensor_from_handle(&self.token_norm);
let bias = self.gguf.tensor_from_handle(&self.token_norm_bias);
layer_norm(&output, weight.f32_slice()?, bias.f32_slice()?, EPSILON)
}
fn forward(&self, input: &[f32], tokens: usize) -> Result<Vec<f32>> {
ensure!(
tokens != 0 && input.len() == tokens * HIDDEN,
"BERT input shape differs"
);
let mut hidden = input.to_vec();
for layer in &self.layers {
hidden = self.attention_layer(&hidden, tokens, layer)?;
hidden = self.feed_forward_layer(&hidden, tokens, layer)?;
}
Ok(hidden)
}
fn attention_layer(
&self,
input: &[f32],
tokens: usize,
layer: &LayerWeights,
) -> Result<Vec<f32>> {
let query_weight = self.gguf.tensor_from_handle(&layer.query);
let key_weight = self.gguf.tensor_from_handle(&layer.key);
let value_weight = self.gguf.tensor_from_handle(&layer.value);
let (mut query, mut key, mut value) =
matrix_matrix_triple(&query_weight, &key_weight, &value_weight, input, tokens)?;
add_bias_batch(
&mut query,
self.gguf
.tensor_from_handle(&layer.query_bias)
.f32_slice()?,
)?;
add_bias_batch(
&mut key,
self.gguf.tensor_from_handle(&layer.key_bias).f32_slice()?,
)?;
add_bias_batch(
&mut value,
self.gguf
.tensor_from_handle(&layer.value_bias)
.f32_slice()?,
)?;
let attended = bidirectional_attention(&query, &key, &value, tokens)?;
let output_weight = self.gguf.tensor_from_handle(&layer.attention_output);
let mut projected = matrix_matrix(&output_weight, &attended, tokens)?;
add_bias_batch(
&mut projected,
self.gguf
.tensor_from_handle(&layer.attention_output_bias)
.f32_slice()?,
)?;
vector_add(&mut projected, input)?;
let norm = self.gguf.tensor_from_handle(&layer.attention_norm);
let norm_bias = self.gguf.tensor_from_handle(&layer.attention_norm_bias);
layer_norm(
&projected,
norm.f32_slice()?,
norm_bias.f32_slice()?,
EPSILON,
)
}
fn feed_forward_layer(
&self,
input: &[f32],
tokens: usize,
layer: &LayerWeights,
) -> Result<Vec<f32>> {
let up_weight = self.gguf.tensor_from_handle(&layer.feed_forward_up);
let mut intermediate = matrix_matrix(&up_weight, input, tokens)?;
add_bias_batch(
&mut intermediate,
self.gguf
.tensor_from_handle(&layer.feed_forward_up_bias)
.f32_slice()?,
)?;
gelu_inplace(&mut intermediate);
let down_weight = self.gguf.tensor_from_handle(&layer.feed_forward_down);
let mut output = matrix_matrix(&down_weight, &intermediate, tokens)?;
add_bias_batch(
&mut output,
self.gguf
.tensor_from_handle(&layer.feed_forward_down_bias)
.f32_slice()?,
)?;
vector_add(&mut output, input)?;
let norm = self.gguf.tensor_from_handle(&layer.output_norm);
let norm_bias = self.gguf.tensor_from_handle(&layer.output_norm_bias);
layer_norm(&output, norm.f32_slice()?, norm_bias.f32_slice()?, EPSILON)
}
}
fn bidirectional_attention(
query: &[f32],
key: &[f32],
value: &[f32],
tokens: usize,
) -> Result<Vec<f32>> {
let expected = tokens
.checked_mul(HIDDEN)
.context("attention tensor size overflow")?;
ensure!(
query.len() == expected && key.len() == expected && value.len() == expected,
"attention input shape differs"
);
let scale = dim_to_f32(HEAD_DIMENSION).sqrt().recip();
let mut output = vec![0.0; expected];
output
.par_chunks_exact_mut(HIDDEN)
.enumerate()
.for_each(|(query_token, token_output)| {
let query_row = &query[query_token * HIDDEN..(query_token + 1) * HIDDEN];
let mut scores = vec![0.0; tokens];
for head in 0..HEADS {
let offset = head * HEAD_DIMENSION;
let query_head = &query_row[offset..offset + HEAD_DIMENSION];
for (key_token, score) in scores.iter_mut().enumerate() {
let start = key_token * HIDDEN + offset;
*score = dot_f32(query_head, &key[start..start + HEAD_DIMENSION]) * scale;
}
softmax(&mut scores);
let head_output = &mut token_output[offset..offset + HEAD_DIMENSION];
for (value_token, weight) in scores.iter().copied().enumerate() {
let start = value_token * HIDDEN + offset;
for (output_value, value) in head_output
.iter_mut()
.zip(&value[start..start + HEAD_DIMENSION])
{
*output_value += value * weight;
}
}
}
});
ensure!(
output.iter().all(|value| value.is_finite()),
"attention output is not finite"
);
Ok(output)
}
fn layer_norm(input: &[f32], weight: &[f32], bias: &[f32], epsilon: f32) -> Result<Vec<f32>> {
ensure!(
weight.len() == HIDDEN
&& bias.len() == HIDDEN
&& !input.is_empty()
&& input.len().is_multiple_of(HIDDEN),
"layer norm shape differs"
);
let mut output = input.to_vec();
output.par_chunks_exact_mut(HIDDEN).for_each(|row| {
let mean = row.iter().sum::<f32>() / dim_to_f32(HIDDEN);
let variance = row
.iter()
.map(|value| {
let centered = value - mean;
centered * centered
})
.sum::<f32>()
/ dim_to_f32(HIDDEN);
let inverse = (variance + epsilon).sqrt().recip();
for ((value, scale), offset) in row.iter_mut().zip(weight).zip(bias) {
*value = (*value - mean) * inverse * scale + offset;
}
});
ensure!(
output.iter().all(|value| value.is_finite()),
"layer norm output is not finite"
);
Ok(output)
}
fn add_bias_batch(values: &mut [f32], bias: &[f32]) -> Result<()> {
ensure!(
!bias.is_empty() && values.len().is_multiple_of(bias.len()),
"bias batch shape differs"
);
values.par_chunks_exact_mut(bias.len()).for_each(|row| {
row.iter_mut()
.zip(bias)
.for_each(|(value, bias)| *value += bias);
});
Ok(())
}
fn gelu_inplace(values: &mut [f32]) {
const SCALE: f32 = 0.797_884_6;
values.par_iter_mut().for_each(|value| {
let input = *value;
*value = 0.5 * input * (1.0 + (SCALE * (input + 0.044_715 * input.powi(3))).tanh());
});
}
fn normalize_embedding(mut embedding: Vec<f32>) -> Result<Vec<f32>> {
let magnitude = dot_f32(&embedding, &embedding).sqrt();
ensure!(
magnitude.is_finite() && magnitude > 0.0,
"embedding magnitude is invalid"
);
embedding
.par_iter_mut()
.for_each(|value| *value /= magnitude);
Ok(embedding)
}
fn validate_metadata(gguf: &Gguf) -> Result<()> {
ensure!(gguf.architecture() == "bert", "expected BERT architecture");
validate_u32(gguf, "bert.block_count", LAYERS)?;
validate_u32(gguf, "bert.context_length", CONTEXT)?;
validate_u32(gguf, "bert.embedding_length", HIDDEN)?;
validate_u32(gguf, "bert.feed_forward_length", FEED_FORWARD)?;
validate_u32(gguf, "bert.attention.head_count", HEADS)?;
validate_u32(gguf, "bert.pooling_type", CLS_POOLING)?;
validate_u32(gguf, "tokenizer.ggml.token_type_count", TOKEN_TYPES)?;
ensure!(
!gguf.bool("bert.attention.causal")?,
"BERT attention is causal"
);
let epsilon = gguf.f32("bert.attention.layer_norm_epsilon")?;
ensure!(
epsilon.to_bits() == EPSILON.to_bits(),
"layer norm epsilon is {epsilon}, expected {EPSILON}"
);
ensure!(
gguf.strings("tokenizer.ggml.tokens")?.len() == VOCABULARY,
"vocabulary size differs"
);
Ok(())
}
fn validate_u32(gguf: &Gguf, key: &str, expected: usize) -> Result<()> {
let value = gguf.u32(key)?;
ensure!(
usize::try_from(value)? == expected,
"{key} is {value}, expected {expected}"
);
Ok(())
}
fn validate_tensors(gguf: &Gguf) -> Result<()> {
validate_kind(
gguf,
"token_embd.weight",
&[HIDDEN, VOCABULARY],
TensorType::Q8_0,
)?;
validate_kind(
gguf,
"token_types.weight",
&[HIDDEN, TOKEN_TYPES],
TensorType::F32,
)?;
validate_kind(
gguf,
"position_embd.weight",
&[HIDDEN, CONTEXT],
TensorType::F32,
)?;
validate_kind(gguf, "token_embd_norm.weight", &[HIDDEN], TensorType::F32)?;
validate_kind(gguf, "token_embd_norm.bias", &[HIDDEN], TensorType::F32)?;
for layer in 0..LAYERS {
let names = LayerNames::new(layer);
for (name, dimensions, kind) in [
(&names.query, &[HIDDEN, HIDDEN][..], TensorType::Q8_0),
(&names.query_bias, &[HIDDEN][..], TensorType::F32),
(&names.key, &[HIDDEN, HIDDEN][..], TensorType::Q8_0),
(&names.key_bias, &[HIDDEN][..], TensorType::F32),
(&names.value, &[HIDDEN, HIDDEN][..], TensorType::Q8_0),
(&names.value_bias, &[HIDDEN][..], TensorType::F32),
(
&names.attention_output,
&[HIDDEN, HIDDEN][..],
TensorType::Q8_0,
),
(&names.attention_output_bias, &[HIDDEN][..], TensorType::F32),
(&names.attention_norm, &[HIDDEN][..], TensorType::F32),
(&names.attention_norm_bias, &[HIDDEN][..], TensorType::F32),
(
&names.feed_forward_up,
&[HIDDEN, FEED_FORWARD][..],
TensorType::Q8_0,
),
(
&names.feed_forward_up_bias,
&[FEED_FORWARD][..],
TensorType::F32,
),
(
&names.feed_forward_down,
&[FEED_FORWARD, HIDDEN][..],
TensorType::Q8_0,
),
(
&names.feed_forward_down_bias,
&[HIDDEN][..],
TensorType::F32,
),
(&names.output_norm, &[HIDDEN][..], TensorType::F32),
(&names.output_norm_bias, &[HIDDEN][..], TensorType::F32),
] {
validate_kind(gguf, name, dimensions, kind)?;
}
}
Ok(())
}
fn validate_kind(gguf: &Gguf, name: &str, dimensions: &[usize], kind: TensorType) -> Result<()> {
let tensor = gguf.tensor(name)?;
ensure!(
tensor.dimensions() == dimensions,
"tensor `{name}` has dimensions {:?}, expected {dimensions:?}",
tensor.dimensions()
);
ensure!(
tensor.tensor_type() == kind,
"tensor `{name}` is {:?}, expected {kind:?}",
tensor.tensor_type()
);
Ok(())
}