llm_kernel/embedding/
bgem3.rs1use std::path::PathBuf;
9use std::sync::Mutex;
10
11use fastembed::{Bgem3Embedding, Bgem3InitOptions, Bgem3Model, SparseEmbedding};
12
13use crate::embedding::sparse::SparseVector;
14use crate::error::{KernelError, Result};
15
16pub const BGEM3_DENSE_DIM: usize = 1024;
18
19pub const BGEM3_VOCAB_SIZE: usize = 250_002;
22
23const DEFAULT_BATCH_CAP: usize = 32;
26
27#[derive(Debug, Clone)]
29pub struct JointEmbedding {
30 pub dense: Vec<f32>,
32 pub sparse: SparseVector,
34}
35
36fn sparse_from_fastembed(sp: SparseEmbedding, top_k: Option<usize>) -> Result<SparseVector> {
39 let indices = sp
40 .indices
41 .into_iter()
42 .map(|i| {
43 u32::try_from(i)
44 .map_err(|_| KernelError::Embedding(format!("sparse index {i} exceeds u32")))
45 })
46 .collect::<Result<Vec<u32>>>()?;
47 let sv = SparseVector::new(indices, sp.values).ok_or_else(|| {
48 KernelError::Embedding("BGE-M3 sparse indices/values length mismatch".into())
49 })?;
50 Ok(match top_k {
51 Some(k) => sv.prune_top_k(k),
52 None => sv,
53 })
54}
55
56pub struct Bgem3Provider {
73 inner: Mutex<Bgem3Embedding>,
74 batch_cap: usize,
75 sparse_top_k: Option<usize>,
76}
77
78impl Bgem3Provider {
79 pub fn new(cache_dir: Option<PathBuf>) -> Result<Self> {
83 Self::build(Bgem3InitOptions::new(Bgem3Model::BGEM3Q), cache_dir)
84 }
85
86 pub fn with_max_length(max_length: usize, cache_dir: Option<PathBuf>) -> Result<Self> {
89 Self::build(
90 Bgem3InitOptions::new(Bgem3Model::BGEM3Q).with_max_length(max_length),
91 cache_dir,
92 )
93 }
94
95 fn build(mut options: Bgem3InitOptions, cache_dir: Option<PathBuf>) -> Result<Self> {
96 if let Some(dir) = cache_dir {
97 options = options.with_cache_dir(dir);
98 }
99 let model = Bgem3Embedding::try_new(options).map_err(KernelError::embedding)?;
100 Ok(Self {
101 inner: Mutex::new(model),
102 batch_cap: DEFAULT_BATCH_CAP,
103 sparse_top_k: None,
104 })
105 }
106
107 pub fn with_batch_cap(mut self, cap: usize) -> Self {
113 self.batch_cap = cap.max(1);
114 self
115 }
116
117 pub fn with_sparse_top_k(mut self, k: usize) -> Self {
123 self.sparse_top_k = Some(k);
124 self
125 }
126
127 pub fn dense_dim(&self) -> usize {
129 BGEM3_DENSE_DIM
130 }
131
132 pub fn vocab_size(&self) -> usize {
134 BGEM3_VOCAB_SIZE
135 }
136
137 pub fn embed(&self, texts: &[&str]) -> Result<Vec<JointEmbedding>> {
139 let mut out = Vec::with_capacity(texts.len());
140 for run in texts.chunks(self.batch_cap) {
141 let batch = {
142 let mut model = self
143 .inner
144 .lock()
145 .map_err(|e| KernelError::Embedding(format!("lock: {e}")))?;
146 model
147 .embed(run, Some(self.batch_cap))
148 .map_err(KernelError::embedding)?
149 };
150 if batch.dense.len() != batch.sparse.len() {
151 return Err(KernelError::Embedding(format!(
152 "BGE-M3 returned {} dense and {} sparse vectors",
153 batch.dense.len(),
154 batch.sparse.len()
155 )));
156 }
157 for (dense, sp) in batch.dense.into_iter().zip(batch.sparse) {
159 out.push(JointEmbedding {
160 dense,
161 sparse: sparse_from_fastembed(sp, self.sparse_top_k)?,
162 });
163 }
164 }
165 Ok(out)
166 }
167}
168
169#[cfg(test)]
170mod tests {
171 use super::*;
172
173 fn fe_sparse(indices: Vec<usize>, values: Vec<f32>) -> SparseEmbedding {
174 SparseEmbedding { indices, values }
175 }
176
177 #[test]
178 fn converts_fastembed_sparse() {
179 let sv = sparse_from_fastembed(fe_sparse(vec![7, 1], vec![0.25, 0.75]), None).unwrap();
180 assert_eq!(sv.indices(), &[1, 7]);
182 assert_eq!(sv.values(), &[0.75, 0.25]);
183 }
184
185 #[test]
186 fn prunes_to_top_k_when_requested() {
187 let sv =
188 sparse_from_fastembed(fe_sparse(vec![1, 2, 3], vec![0.1, 0.9, 0.5]), Some(2)).unwrap();
189 assert_eq!(sv.nnz(), 2);
190 assert_eq!(sv.indices(), &[2, 3]);
191 }
192
193 #[test]
194 fn rejects_length_mismatch() {
195 let err = sparse_from_fastembed(fe_sparse(vec![1, 2], vec![0.5]), None);
196 assert!(err.is_err());
197 }
198
199 #[test]
200 fn vocab_and_dense_dims_are_bgem3_shapes() {
201 assert_eq!(BGEM3_DENSE_DIM, 1024);
202 assert_eq!(BGEM3_VOCAB_SIZE, 250_002);
203 }
204}