use std::path::Path;
use anyhow::{Context as _, Result, ensure};
use rayon::prelude::*;
use super::gguf::{Gguf, TensorHandle, TensorType};
use super::kernels::{
dequantize_row, dot_f32, matrix_matrix_segmented, matrix_matrix_triple_segmented, 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;
const BATCH_TOKEN_BUDGET: usize = 1_024;
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())
}
pub fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
let tokenized = texts
.iter()
.map(|text| {
self.tokenizer
.encode(text, CONTEXT)
.context("tokenize embedding input")
})
.collect::<Result<Vec<_>>>()?;
let mut embeddings = Vec::with_capacity(tokenized.len());
let mut start = 0;
while start < tokenized.len() {
let end = embedding_batch_end(&tokenized, start)?;
embeddings.extend(self.embed_token_batch(&tokenized[start..end])?);
start = end;
}
Ok(embeddings)
}
fn embed_token_batch(&self, tokenized: &[Vec<u32>]) -> Result<Vec<Vec<f32>>> {
let sequence_lengths = tokenized.iter().map(Vec::len).collect::<Vec<_>>();
let input = self.embeddings_batch(tokenized)?;
let hidden = self.forward_batch(&input, &sequence_lengths)?;
let mut embeddings = Vec::with_capacity(sequence_lengths.len());
let mut token_offset = 0_usize;
for tokens in sequence_lengths {
let start = token_offset
.checked_mul(HIDDEN)
.context("embedding offset overflow")?;
let end = start
.checked_add(HIDDEN)
.context("embedding offset overflow")?;
embeddings.push(normalize_embedding(hidden[start..end].to_vec())?);
token_offset = token_offset
.checked_add(tokens)
.context("embedding token count overflow")?;
}
Ok(embeddings)
}
fn embeddings(&self, tokens: &[u32]) -> Result<Vec<f32>> {
self.embeddings_batch(std::slice::from_ref(&tokens))
}
fn embeddings_batch<T: AsRef<[u32]>>(&self, sequences: &[T]) -> Result<Vec<f32>> {
ensure!(
!sequences.is_empty()
&& sequences
.iter()
.all(|sequence| !sequence.as_ref().is_empty()),
"BERT sequence batch is empty"
);
let token_count = sequences.iter().try_fold(0_usize, |total, sequence| {
total
.checked_add(sequence.as_ref().len())
.context("BERT sequence token count overflow")
})?;
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 output_len = token_count
.checked_mul(HIDDEN)
.context("embedding tensor size overflow")?;
let mut output = vec![0.0; output_len];
let mut rows = output.chunks_exact_mut(HIDDEN);
for sequence in sequences {
for (position, token) in sequence.as_ref().iter().enumerate() {
let row = rows.next().context("embedding row count differs")?;
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;
}
}
}
ensure!(rows.next().is_none(), "embedding row count differs");
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>> {
self.forward_batch(input, std::slice::from_ref(&tokens))
}
fn forward_batch(&self, input: &[f32], sequence_lengths: &[usize]) -> Result<Vec<f32>> {
let tokens = sequence_token_count(sequence_lengths)?;
let expected = tokens
.checked_mul(HIDDEN)
.context("BERT input size overflow")?;
ensure!(input.len() == expected, "BERT input shape differs");
let mut hidden = input.to_vec();
for layer in &self.layers {
hidden = self.attention_layer(&hidden, sequence_lengths, layer)?;
hidden = self.feed_forward_layer(&hidden, sequence_lengths, layer)?;
}
Ok(hidden)
}
fn attention_layer(
&self,
input: &[f32],
sequence_lengths: &[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_segmented(
&query_weight,
&key_weight,
&value_weight,
input,
sequence_lengths,
)?;
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_batch(&query, &key, &value, sequence_lengths)?;
let output_weight = self.gguf.tensor_from_handle(&layer.attention_output);
let mut projected = matrix_matrix_segmented(&output_weight, &attended, sequence_lengths)?;
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],
sequence_lengths: &[usize],
layer: &LayerWeights,
) -> Result<Vec<f32>> {
let up_weight = self.gguf.tensor_from_handle(&layer.feed_forward_up);
let mut intermediate = matrix_matrix_segmented(&up_weight, input, sequence_lengths)?;
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_segmented(&down_weight, &intermediate, sequence_lengths)?;
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 embedding_batch_end(tokenized: &[Vec<u32>], start: usize) -> Result<usize> {
ensure!(
start < tokenized.len(),
"embedding batch start is out of range"
);
let mut end = start;
let mut tokens = 0_usize;
while let Some(sequence) = tokenized.get(end) {
if end != start && tokens.saturating_add(sequence.len()) > BATCH_TOKEN_BUDGET {
break;
}
tokens = tokens.saturating_add(sequence.len());
end += 1;
}
Ok(end)
}
fn sequence_token_count(sequence_lengths: &[usize]) -> Result<usize> {
ensure!(
!sequence_lengths.is_empty() && sequence_lengths.iter().all(|tokens| *tokens != 0),
"BERT sequence batch is empty"
);
sequence_lengths.iter().try_fold(0_usize, |total, tokens| {
total
.checked_add(*tokens)
.context("BERT sequence token count overflow")
})
}
fn bidirectional_attention_batch(
query: &[f32],
key: &[f32],
value: &[f32],
sequence_lengths: &[usize],
) -> Result<Vec<f32>> {
let tokens = sequence_token_count(sequence_lengths)?;
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"
);
if let [tokens] = sequence_lengths {
return bidirectional_attention(query, key, value, *tokens);
}
let mut output = Vec::with_capacity(expected);
let mut token_offset = 0_usize;
for tokens in sequence_lengths {
let start = token_offset
.checked_mul(HIDDEN)
.context("attention offset overflow")?;
token_offset = token_offset
.checked_add(*tokens)
.context("attention offset overflow")?;
let end = token_offset
.checked_mul(HIDDEN)
.context("attention offset overflow")?;
output.extend(bidirectional_attention(
&query[start..end],
&key[start..end],
&value[start..end],
*tokens,
)?);
}
Ok(output)
}
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 head_dimension =
f32::from(u16::try_from(HEAD_DIMENSION).context("attention head dimension exceeds u16")?);
let scale = head_dimension.sqrt().recip();
let mut output = vec![0.0; expected];
output
.par_chunks_exact_mut(HIDDEN)
.enumerate()
.try_for_each(|(query_token, token_output)| -> Result<()> {
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;
}
}
}
Ok(())
})?;
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();
let hidden = f32::from(u16::try_from(HIDDEN).context("layer norm width exceeds u16")?);
output.par_chunks_exact_mut(HIDDEN).for_each(|row| {
let mean = row.iter().sum::<f32>() / hidden;
let variance = row
.iter()
.map(|value| {
let centered = value - mean;
centered * centered
})
.sum::<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"
);
for value in &mut embedding {
*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(())
}