use std::path::PathBuf;
use std::sync::Mutex;
use fastembed::{Bgem3Embedding, Bgem3InitOptions, Bgem3Model, SparseEmbedding};
use crate::embedding::sparse::SparseVector;
use crate::error::{KernelError, Result};
pub const BGEM3_DENSE_DIM: usize = 1024;
pub const BGEM3_VOCAB_SIZE: usize = 250_002;
const DEFAULT_BATCH_CAP: usize = 32;
#[derive(Debug, Clone)]
pub struct JointEmbedding {
pub dense: Vec<f32>,
pub sparse: SparseVector,
}
fn sparse_from_fastembed(sp: SparseEmbedding, top_k: Option<usize>) -> Result<SparseVector> {
let indices = sp
.indices
.into_iter()
.map(|i| {
u32::try_from(i)
.map_err(|_| KernelError::Embedding(format!("sparse index {i} exceeds u32")))
})
.collect::<Result<Vec<u32>>>()?;
let sv = SparseVector::new(indices, sp.values).ok_or_else(|| {
KernelError::Embedding("BGE-M3 sparse indices/values length mismatch".into())
})?;
Ok(match top_k {
Some(k) => sv.prune_top_k(k),
None => sv,
})
}
pub struct Bgem3Provider {
inner: Mutex<Bgem3Embedding>,
batch_cap: usize,
sparse_top_k: Option<usize>,
}
impl Bgem3Provider {
pub fn new(cache_dir: Option<PathBuf>) -> Result<Self> {
Self::build(Bgem3InitOptions::new(Bgem3Model::BGEM3Q), cache_dir)
}
pub fn with_max_length(max_length: usize, cache_dir: Option<PathBuf>) -> Result<Self> {
Self::build(
Bgem3InitOptions::new(Bgem3Model::BGEM3Q).with_max_length(max_length),
cache_dir,
)
}
fn build(mut options: Bgem3InitOptions, cache_dir: Option<PathBuf>) -> Result<Self> {
if let Some(dir) = cache_dir {
options = options.with_cache_dir(dir);
}
let model = Bgem3Embedding::try_new(options).map_err(KernelError::embedding)?;
Ok(Self {
inner: Mutex::new(model),
batch_cap: DEFAULT_BATCH_CAP,
sparse_top_k: None,
})
}
pub fn with_batch_cap(mut self, cap: usize) -> Self {
self.batch_cap = cap.max(1);
self
}
pub fn with_sparse_top_k(mut self, k: usize) -> Self {
self.sparse_top_k = Some(k);
self
}
pub fn dense_dim(&self) -> usize {
BGEM3_DENSE_DIM
}
pub fn vocab_size(&self) -> usize {
BGEM3_VOCAB_SIZE
}
pub fn embed(&self, texts: &[&str]) -> Result<Vec<JointEmbedding>> {
let mut out = Vec::with_capacity(texts.len());
for run in texts.chunks(self.batch_cap) {
let batch = {
let mut model = self
.inner
.lock()
.map_err(|e| KernelError::Embedding(format!("lock: {e}")))?;
model
.embed(run, Some(self.batch_cap))
.map_err(KernelError::embedding)?
};
if batch.dense.len() != batch.sparse.len() {
return Err(KernelError::Embedding(format!(
"BGE-M3 returned {} dense and {} sparse vectors",
batch.dense.len(),
batch.sparse.len()
)));
}
for (dense, sp) in batch.dense.into_iter().zip(batch.sparse) {
out.push(JointEmbedding {
dense,
sparse: sparse_from_fastembed(sp, self.sparse_top_k)?,
});
}
}
Ok(out)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn fe_sparse(indices: Vec<usize>, values: Vec<f32>) -> SparseEmbedding {
SparseEmbedding { indices, values }
}
#[test]
fn converts_fastembed_sparse() {
let sv = sparse_from_fastembed(fe_sparse(vec![7, 1], vec![0.25, 0.75]), None).unwrap();
assert_eq!(sv.indices(), &[1, 7]);
assert_eq!(sv.values(), &[0.75, 0.25]);
}
#[test]
fn prunes_to_top_k_when_requested() {
let sv =
sparse_from_fastembed(fe_sparse(vec![1, 2, 3], vec![0.1, 0.9, 0.5]), Some(2)).unwrap();
assert_eq!(sv.nnz(), 2);
assert_eq!(sv.indices(), &[2, 3]);
}
#[test]
fn rejects_length_mismatch() {
let err = sparse_from_fastembed(fe_sparse(vec![1, 2], vec![0.5]), None);
assert!(err.is_err());
}
#[test]
fn vocab_and_dense_dims_are_bgem3_shapes() {
assert_eq!(BGEM3_DENSE_DIM, 1024);
assert_eq!(BGEM3_VOCAB_SIZE, 250_002);
}
}