From 635c01a799b9f9c8b3f23f3d2de78ceb555b56df Mon Sep 17 00:00:00 2001 From: Louis Maddox Date: Sat, 27 Dec 2025 15:26:56 +0000 Subject: [PATCH 1/3] feat: late interaction and MUVERA --- Cargo.toml | 10 + README.md | 5 + examples/muvera_demo.rs | 74 ++++++ src/late_interaction/impl.rs | 464 ++++++++++++++++++++++++++++++++++ src/late_interaction/init.rs | 139 +++++++++++ src/late_interaction/mod.rs | 9 + src/lib.rs | 13 + src/models/text_embedding.rs | 4 +- src/postprocess/mod.rs | 7 + src/postprocess/muvera.rs | 466 +++++++++++++++++++++++++++++++++++ tests/late-interaction.rs | 362 +++++++++++++++++++++++++++ 11 files changed, 1552 insertions(+), 1 deletion(-) create mode 100644 examples/muvera_demo.rs create mode 100644 src/late_interaction/impl.rs create mode 100644 src/late_interaction/init.rs create mode 100644 src/late_interaction/mod.rs create mode 100644 src/postprocess/mod.rs create mode 100644 src/postprocess/muvera.rs create mode 100644 tests/late-interaction.rs diff --git a/Cargo.toml b/Cargo.toml index 35f815c..ed093b4 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -31,6 +31,9 @@ ort = { version = "=2.0.0-rc.10", default-features = false, features = [ ] } serde_json = { version = "1" } tokenizers = { version = "0.22.0", default-features = false, features = ["onig"] } +rand = { version = "0.8", optional = true } +rand_chacha = { version = "0.3", optional = true } +rand_distr = { version = "0.4", optional = true } [features] default = ["ort-download-binaries", "hf-hub-native-tls", "image-models"] @@ -44,9 +47,16 @@ ort-load-dynamic = ["ort/load-dynamic"] image-models = ["image"] +# MUVERA post-processing support +muvera = ["dep:rand", "dep:rand_chacha", "dep:rand_distr"] + # This feature does not change any code, but is used to limit tests if # the user does not have `optimum-cli` or even python installed. optimum-cli = [] # For compatibility recommend using hf-hub-native-tls online = ["hf-hub-native-tls"] + +[[examples]] +name = "muvera_demo" +features = ["hf-hub", "muvera"] diff --git a/README.md b/README.md index 4655b12..b6ac1a6 100644 --- a/README.md +++ b/README.md @@ -71,6 +71,11 @@ Quantized versions are also available for several models above (append `Q` to th - [**jinaai/jina-reranker-v1-turbo-en**](https://huggingface.co/jinaai/jina-reranker-v1-turbo-en) - [**jinaai/jina-reranker-v2-base-multiligual**](https://huggingface.co/jinaai/jina-reranker-v2-base-multilingual) +### Late Interaction Embedding + +- [**colbert-ir/colbertv2.0**](https://huggingface.co/colbert-ir/colbertv2.0) - Default, 128-dim embeddings +- [**answerdotai/answerai-colbert-small-v1**](https://huggingface.co/answerdotai/answerai-colbert-small-v1) - 96-dim embeddings + ## ✊ Support To support the library, please donate to our primary upstream dependency, [`ort`](https://github.com/pykeio/ort?tab=readme-ov-file#-sponsor-ort) - The Rust wrapper for the ONNX runtime. diff --git a/examples/muvera_demo.rs b/examples/muvera_demo.rs new file mode 100644 index 0000000..70ff16c --- /dev/null +++ b/examples/muvera_demo.rs @@ -0,0 +1,74 @@ +// examples/muvera_demo.rs +use anyhow::Result; +use fastembed::{ + LateInteractionInitOptions, LateInteractionModel, LateInteractionTextEmbedding, Muvera, +}; + +/// Compute ColBERT MaxSim score between query and document embeddings +fn maxsim(query_emb: &[Vec], doc_emb: &[Vec]) -> f32 { + // For each query token, find max similarity with any document token + query_emb + .iter() + .map(|q| { + doc_emb + .iter() + .map(|d| q.iter().zip(d.iter()).map(|(a, b)| a * b).sum::()) + .fold(f32::NEG_INFINITY, f32::max) + }) + .sum() +} + +fn main() -> Result<()> { + // 1. Initialize the ColBERT model + println!("Loading ColBERT model..."); + let mut model = LateInteractionTextEmbedding::try_new(LateInteractionInitOptions::new( + LateInteractionModel::ColBERTV2, + ))?; + + // 2. Create MUVERA postprocessor from the model + let muvera = Muvera::from_late_interaction_model( + &model, + Some(5), // k_sim: 2^5 = 32 clusters + Some(16), // dim_proj: project to 16 dimensions per cluster + Some(20), // r_reps: 20 repetitions for robustness + Some(42), // random_seed + )?; + + // 3. Define documents and queries + let documents = vec![ + "Machine learning is a subset of artificial intelligence.", + "Python is a popular programming language.", + ]; + let queries = vec!["What is machine learning?"]; + + // 4. Get multi-vector embeddings + let doc_embeddings = model.embed(&documents, None)?; + let query_embeddings = model.query_embed(&queries, None)?; + + println!("Document 0 shape: ({}, {})", doc_embeddings[0].len(), doc_embeddings[0][0].len()); + println!("Query 0 shape: ({}, {})", query_embeddings[0].len(), query_embeddings[0][0].len()); + println!("FDE size: {}", muvera.embedding_size()); + + // 5. Convert to Fixed Dimensional Encodings (FDEs) + let doc_fdes: Vec> = doc_embeddings + .iter() + .map(|emb| muvera.process_document(emb)) + .collect(); + let query_fde = muvera.process_query(&query_embeddings[0]); + + println!("Doc FDE shape: ({},)", doc_fdes[0].len()); + + // 6. Compute MUVERA similarities (for candidate retrieval) + for (i, doc_fde) in doc_fdes.iter().enumerate() { + let similarity: f32 = query_fde.iter().zip(doc_fde.iter()).map(|(a, b)| a * b).sum(); + println!("Query-Doc{} similarity (MUVERA): {:.4}", i, similarity); + } + + // 7. Compute ColBERT MaxSim scores (for reranking) + for (i, doc_emb) in doc_embeddings.iter().enumerate() { + let score = maxsim(&query_embeddings[0], doc_emb); + println!("Query-Doc{} MaxSim score: {:.4}", i, score); + } + + Ok(()) +} \ No newline at end of file diff --git a/src/late_interaction/impl.rs b/src/late_interaction/impl.rs new file mode 100644 index 0000000..b079be7 --- /dev/null +++ b/src/late_interaction/impl.rs @@ -0,0 +1,464 @@ +use std::{collections::HashSet, fmt::Display, str::FromStr, thread::available_parallelism}; + +use anyhow::{Context, Result}; +#[cfg(feature = "hf-hub")] +use hf_hub::api::sync::ApiRepo; +use ndarray::Array; +use ort::{ + session::{builder::GraphOptimizationLevel, Session}, + value::Value, +}; +#[cfg(feature = "hf-hub")] +use std::path::PathBuf; +use tokenizers::Tokenizer; + +use crate::common::load_tokenizer; +#[cfg(feature = "hf-hub")] +use crate::common::load_tokenizer_hf_hub; + +use super::{ + LateInteractionEmbedding, LateInteractionInitOptions, LateInteractionInitOptionsUserDefined, + LateInteractionModel, LateInteractionModelInfo, LateInteractionTextEmbedding, + UserDefinedLateInteractionModel, DEFAULT_BATCH_SIZE, +}; + +/// All punctuation characters for skip list +fn get_punctuation_chars() -> Vec { + "!\"#$%&'()*+,-./:;<=>?@[\\]^_`{|}~".chars().collect() +} + +fn models_list() -> Vec { + vec![ + LateInteractionModelInfo { + model: LateInteractionModel::ColBERTV2, + dim: 128, + description: String::from( + "Text embeddings, Unimodal (text), English, 512 input tokens truncation, 2023 year", + ), + model_code: String::from("colbert-ir/colbertv2.0"), + model_file: String::from("model.onnx"), + additional_files: Vec::new(), + query_marker_token_id: 1, + document_marker_token_id: 2, + mask_token: String::from("[MASK]"), + min_query_length: 31, + }, + LateInteractionModelInfo { + model: LateInteractionModel::AnswerAIColBERTSmallV1, + dim: 96, + description: String::from( + "Text embeddings, Unimodal (text), English, 512 input tokens truncation, 2024 year", + ), + model_code: String::from("answerdotai/answerai-colbert-small-v1"), + model_file: String::from("vespa_colbert.onnx"), + additional_files: Vec::new(), + query_marker_token_id: 1, + document_marker_token_id: 2, + mask_token: String::from("[MASK]"), + min_query_length: 31, + }, + ] +} + +impl Display for LateInteractionModel { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let model_info = models_list() + .into_iter() + .find(|m| m.model == *self) + .ok_or(std::fmt::Error)?; + write!(f, "{}", model_info.model_code) + } +} + +impl FromStr for LateInteractionModel { + type Err = String; + + fn from_str(s: &str) -> Result { + models_list() + .into_iter() + .find(|m| m.model_code.eq_ignore_ascii_case(s)) + .map(|m| m.model) + .ok_or_else(|| format!("Unknown late interaction model: {s}")) + } +} + +impl TryFrom for LateInteractionModel { + type Error = String; + + fn try_from(value: String) -> Result { + value.parse() + } +} + +impl LateInteractionTextEmbedding { + /// Try to create a new LateInteractionTextEmbedding instance + #[cfg(feature = "hf-hub")] + pub fn try_new(options: LateInteractionInitOptions) -> Result { + let LateInteractionInitOptions { + max_length, + model_name, + execution_providers, + cache_dir, + show_download_progress, + } = options; + + let threads = available_parallelism()?.get(); + let model_info = Self::get_model_info(&model_name); + + let model_repo = Self::retrieve_model( + model_name.clone(), + cache_dir.clone(), + show_download_progress, + )?; + + let model_file_reference = model_repo + .get(&model_info.model_file) + .context(format!("Failed to retrieve {}", model_info.model_file))?; + + for file in &model_info.additional_files { + model_repo + .get(file) + .context(format!("Failed to retrieve {}", file))?; + } + + let session = Session::builder()? + .with_execution_providers(execution_providers)? + .with_optimization_level(GraphOptimizationLevel::Level3)? + .with_intra_threads(threads)? + .commit_from_file(model_file_reference)?; + + // Load tokenizer for documents (truncates at max_length - 1 to leave room for marker) + let tokenizer = load_tokenizer_hf_hub(model_repo, max_length - 1)?; + + // Load query tokenizer with MASK padding + let mut query_tokenizer = tokenizer.clone(); + + // Get mask token id + let mask_token_id = query_tokenizer + .token_to_id(&model_info.mask_token) + .ok_or_else(|| anyhow::anyhow!("Could not find mask token in vocabulary"))?; + + // Get pad token id + let pad_token_id = tokenizer.get_padding().map(|p| p.pad_id).unwrap_or(0); + + // Enable padding for query tokenizer with MASK token + // with_padding takes &mut self and returns &mut Self, mutating in place + query_tokenizer.with_padding(Some(tokenizers::PaddingParams { + strategy: tokenizers::PaddingStrategy::Fixed(model_info.min_query_length), + pad_token: model_info.mask_token.clone(), + pad_id: mask_token_id, + ..Default::default() + })); + + // Build skip list from punctuation + let skip_list = Self::build_skip_list(&tokenizer); + + Ok(Self::new( + tokenizer, + query_tokenizer, + session, + model_info.query_marker_token_id, + model_info.document_marker_token_id, + mask_token_id, + pad_token_id, + skip_list, + model_info.min_query_length, + model_info.dim, + )) + } + + /// Create from user-defined model + pub fn try_new_from_user_defined( + model: UserDefinedLateInteractionModel, + options: LateInteractionInitOptionsUserDefined, + ) -> Result { + let LateInteractionInitOptionsUserDefined { + execution_providers, + max_length, + } = options; + + let threads = available_parallelism()?.get(); + + let session = Session::builder()? + .with_execution_providers(execution_providers)? + .with_optimization_level(GraphOptimizationLevel::Level3)? + .with_intra_threads(threads)? + .commit_from_memory(&model.onnx_file)?; + + let tokenizer = load_tokenizer(model.tokenizer_files.clone(), max_length - 1)?; + let mut query_tokenizer = load_tokenizer(model.tokenizer_files, max_length - 1)?; + + let mask_token_id = query_tokenizer + .token_to_id(&model.mask_token) + .ok_or_else(|| anyhow::anyhow!("Could not find mask token in vocabulary"))?; + + let pad_token_id = tokenizer.get_padding().map(|p| p.pad_id).unwrap_or(0); + + query_tokenizer.with_padding(Some(tokenizers::PaddingParams { + strategy: tokenizers::PaddingStrategy::Fixed(model.min_query_length), + pad_token: model.mask_token.clone(), + pad_id: mask_token_id, + ..Default::default() + })); + + let skip_list = Self::build_skip_list(&tokenizer); + + Ok(Self::new( + tokenizer, + query_tokenizer, + session, + model.query_marker_token_id, + model.document_marker_token_id, + mask_token_id, + pad_token_id, + skip_list, + model.min_query_length, + model.dim, + )) + } + + #[allow(clippy::too_many_arguments)] + fn new( + tokenizer: Tokenizer, + query_tokenizer: Tokenizer, + session: Session, + query_marker_token_id: u32, + document_marker_token_id: u32, + mask_token_id: u32, + pad_token_id: u32, + skip_list: HashSet, + min_query_length: usize, + dim: usize, + ) -> Self { + let need_token_type_ids = session + .inputs + .iter() + .any(|input| input.name == "token_type_ids"); + + Self { + tokenizer, + query_tokenizer, + session, + need_token_type_ids, + query_marker_token_id, + document_marker_token_id, + mask_token_id, + pad_token_id, + skip_list, + min_query_length, + dim, + } + } + + fn build_skip_list(tokenizer: &Tokenizer) -> HashSet { + let mut skip_list = HashSet::new(); + for c in get_punctuation_chars() { + if let Some(id) = tokenizer.token_to_id(&c.to_string()) { + skip_list.insert(id); + } + } + skip_list + } + + #[cfg(feature = "hf-hub")] + fn retrieve_model( + model: LateInteractionModel, + cache_dir: PathBuf, + show_download_progress: bool, + ) -> Result { + use crate::common::pull_from_hf; + pull_from_hf(model.to_string(), cache_dir, show_download_progress) + } + + pub fn list_supported_models() -> Vec { + models_list() + } + + pub fn get_model_info(model: &LateInteractionModel) -> LateInteractionModelInfo { + Self::list_supported_models() + .into_iter() + .find(|m| &m.model == model) + .expect("Model not found in supported models list") + } + + /// Get the embedding dimension + pub fn dim(&self) -> usize { + self.dim + } + + /// Embed documents (passages) + pub fn embed + Send + Sync>( + &mut self, + documents: impl AsRef<[S]>, + batch_size: Option, + ) -> Result> { + self.embed_internal(documents, batch_size, false) + } + + /// Embed queries + pub fn query_embed + Send + Sync>( + &mut self, + queries: impl AsRef<[S]>, + batch_size: Option, + ) -> Result> { + self.embed_internal(queries, batch_size, true) + } + + fn embed_internal + Send + Sync>( + &mut self, + texts: impl AsRef<[S]>, + batch_size: Option, + is_query: bool, + ) -> Result> { + let texts = texts.as_ref(); + let batch_size = batch_size.unwrap_or(DEFAULT_BATCH_SIZE); + + let mut all_embeddings = Vec::with_capacity(texts.len()); + + for batch in texts.chunks(batch_size) { + let tokenizer = if is_query { + &self.query_tokenizer + } else { + &self.tokenizer + }; + + let inputs: Vec<&str> = batch.iter().map(|t| t.as_ref()).collect(); + let encodings = tokenizer + .encode_batch(inputs, true) + .map_err(|e| anyhow::anyhow!("Failed to encode batch: {}", e))?; + + let encoding_length = encodings + .first() + .ok_or_else(|| anyhow::anyhow!("Empty encodings"))? + .len(); + + let batch_size_actual = batch.len(); + // +1 for the marker token we'll insert + let final_length = encoding_length + 1; + + let mut ids_array = Vec::with_capacity(batch_size_actual * final_length); + let mut mask_array = Vec::with_capacity(batch_size_actual * final_length); + let mut type_ids_array = Vec::with_capacity(batch_size_actual * final_length); + + let marker_token = if is_query { + self.query_marker_token_id + } else { + self.document_marker_token_id + }; + + for encoding in &encodings { + let ids = encoding.get_ids(); + let mask = encoding.get_attention_mask(); + let type_ids = encoding.get_type_ids(); + + // Insert marker token after first token (position 1) + // [CLS, marker, token1, token2, ...] + ids_array.push(ids[0] as i64); + ids_array.push(marker_token as i64); + ids_array.extend(ids[1..].iter().map(|&x| x as i64)); + + mask_array.push(mask[0] as i64); + mask_array.push(1i64); // Marker always attended + mask_array.extend(mask[1..].iter().map(|&x| x as i64)); + + type_ids_array.push(type_ids[0] as i64); + type_ids_array.push(0i64); + type_ids_array.extend(type_ids[1..].iter().map(|&x| x as i64)); + } + + let input_ids_array = + Array::from_shape_vec((batch_size_actual, final_length), ids_array)?; + let attention_mask_array = + Array::from_shape_vec((batch_size_actual, final_length), mask_array)?; + let token_type_ids_array = + Array::from_shape_vec((batch_size_actual, final_length), type_ids_array)?; + + let mut session_inputs = ort::inputs![ + "input_ids" => Value::from_array(input_ids_array.clone())?, + "attention_mask" => Value::from_array(attention_mask_array.clone())?, + ]; + + if self.need_token_type_ids { + session_inputs.push(( + "token_type_ids".into(), + Value::from_array(token_type_ids_array)?.into(), + )); + } + + let outputs = self.session.run(session_inputs)?; + + // Get the first output + let (_, output_value) = outputs + .iter() + .next() + .ok_or_else(|| anyhow::anyhow!("No output from model"))?; + + // Extract as ndarray ArrayView + let (shape, data) = output_value.try_extract_tensor::()?; + // Shape is (batch_size, seq_len, dim) + let shape_vec: Vec = shape.iter().map(|&d| d as usize).collect(); + let seq_len = shape_vec[1]; + let dim = shape_vec[2]; + + // Process each item in batch + for batch_idx in 0..batch_size_actual { + let mut embeddings: Vec> = Vec::new(); + + if is_query { + // For queries: return ALL token embeddings (including MASK padding) + for seq_idx in 0..seq_len { + let start = batch_idx * seq_len * dim + seq_idx * dim; + let end = start + dim; + let token_embedding: Vec = data[start..end].to_vec(); + + // L2 normalize + let norm: f32 = token_embedding.iter().map(|x| x * x).sum::().sqrt(); + let norm = norm.max(1e-12); + let normalized: Vec = + token_embedding.iter().map(|x| x / norm).collect(); + + embeddings.push(normalized); + } + } else { + // For documents: mask out punctuation and pad tokens, filter by attention mask + let mut attention_mask_vec: Vec = attention_mask_array + .row(batch_idx) + .iter() + .copied() + .collect(); + + let input_ids_row = input_ids_array.row(batch_idx); + for (j, &token_id) in input_ids_row.iter().enumerate() { + let token_id_u32 = token_id as u32; + if self.skip_list.contains(&token_id_u32) + || token_id_u32 == self.pad_token_id + { + attention_mask_vec[j] = 0; + } + } + + for (seq_idx, &mask_val) in attention_mask_vec.iter().enumerate().take(seq_len) { + if mask_val == 1 { + let start = batch_idx * seq_len * dim + seq_idx * dim; + let end = start + dim; + let token_embedding: Vec = data[start..end].to_vec(); + + // L2 normalize + let norm: f32 = + token_embedding.iter().map(|x| x * x).sum::().sqrt(); + let norm = norm.max(1e-12); + let normalized: Vec = + token_embedding.iter().map(|x| x / norm).collect(); + + embeddings.push(normalized); + } + } + } + + all_embeddings.push(embeddings); + } + } + + Ok(all_embeddings) + } +} diff --git a/src/late_interaction/init.rs b/src/late_interaction/init.rs new file mode 100644 index 0000000..1438d15 --- /dev/null +++ b/src/late_interaction/init.rs @@ -0,0 +1,139 @@ +use crate::common::TokenizerFiles; +use crate::init::{HasMaxLength, InitOptionsWithLength}; +use ort::{execution_providers::ExecutionProviderDispatch, session::Session}; +use std::collections::HashSet; +use tokenizers::Tokenizer; + +use super::DEFAULT_MAX_LENGTH; + +/// Supported late interaction models +#[derive(Debug, Default, Clone, PartialEq, Eq, Hash)] +pub enum LateInteractionModel { + /// colbert-ir/colbertv2.0 + #[default] + ColBERTV2, + /// answerdotai/answerai-colbert-small-v1 + AnswerAIColBERTSmallV1, +} + +impl HasMaxLength for LateInteractionModel { + const MAX_LENGTH: usize = DEFAULT_MAX_LENGTH; +} + +/// Options for initializing late interaction models +pub type LateInteractionInitOptions = InitOptionsWithLength; + +/// Options for user-defined late interaction models +#[derive(Debug, Clone)] +#[non_exhaustive] +pub struct LateInteractionInitOptionsUserDefined { + pub execution_providers: Vec, + pub max_length: usize, +} + +impl Default for LateInteractionInitOptionsUserDefined { + fn default() -> Self { + Self { + execution_providers: Default::default(), + max_length: DEFAULT_MAX_LENGTH, + } + } +} + +impl LateInteractionInitOptionsUserDefined { + pub fn new() -> Self { + Self::default() + } + + pub fn with_execution_providers( + mut self, + execution_providers: Vec, + ) -> Self { + self.execution_providers = execution_providers; + self + } + + pub fn with_max_length(mut self, max_length: usize) -> Self { + self.max_length = max_length; + self + } +} + +impl From for LateInteractionInitOptionsUserDefined { + fn from(options: LateInteractionInitOptions) -> Self { + Self { + execution_providers: options.execution_providers, + max_length: options.max_length, + } + } +} + +/// User-defined late interaction model +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct UserDefinedLateInteractionModel { + pub onnx_file: Vec, + pub tokenizer_files: TokenizerFiles, + pub query_marker_token_id: u32, + pub document_marker_token_id: u32, + pub mask_token: String, + pub min_query_length: usize, + pub dim: usize, +} + +impl UserDefinedLateInteractionModel { + pub fn new( + onnx_file: Vec, + tokenizer_files: TokenizerFiles, + query_marker_token_id: u32, + document_marker_token_id: u32, + mask_token: String, + min_query_length: usize, + dim: usize, + ) -> Self { + Self { + onnx_file, + tokenizer_files, + query_marker_token_id, + document_marker_token_id, + mask_token, + min_query_length, + dim, + } + } +} + +/// Data struct for late interaction model info +#[derive(Debug, Clone)] +#[non_exhaustive] +pub struct LateInteractionModelInfo { + pub model: LateInteractionModel, + pub dim: usize, + pub description: String, + pub model_code: String, + pub model_file: String, + pub additional_files: Vec, + pub query_marker_token_id: u32, + pub document_marker_token_id: u32, + pub mask_token: String, + pub min_query_length: usize, +} + +/// Late interaction embedding output - variable length per document +pub type LateInteractionEmbedding = Vec>; + +/// Rust representation of a late interaction embedding model +pub struct LateInteractionTextEmbedding { + pub(crate) tokenizer: Tokenizer, + pub(crate) query_tokenizer: Tokenizer, + pub(crate) session: Session, + pub(crate) need_token_type_ids: bool, + pub(crate) query_marker_token_id: u32, + pub(crate) document_marker_token_id: u32, + #[expect(dead_code, reason = "Used to initialise tokeniser")] + pub(crate) mask_token_id: u32, + pub(crate) pad_token_id: u32, + pub(crate) skip_list: HashSet, + #[expect(dead_code, reason = "Used to initialise tokeniser")] + pub(crate) min_query_length: usize, + pub(crate) dim: usize, +} diff --git a/src/late_interaction/mod.rs b/src/late_interaction/mod.rs new file mode 100644 index 0000000..88a2846 --- /dev/null +++ b/src/late_interaction/mod.rs @@ -0,0 +1,9 @@ +//! Late interaction text embedding models (ColBERT-style). + +const DEFAULT_BATCH_SIZE: usize = 256; +const DEFAULT_MAX_LENGTH: usize = 512; + +mod init; +pub use init::*; + +mod r#impl; diff --git a/src/lib.rs b/src/lib.rs index d31949d..4dc7ce3 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -58,9 +58,11 @@ mod common; #[cfg(feature = "image-models")] mod image_embedding; mod init; +mod late_interaction; mod models; pub mod output; mod pooling; +mod postprocess; mod reranking; mod sparse_text_embedding; mod text_embedding; @@ -105,3 +107,14 @@ pub use crate::reranking::{ OnnxSource, RerankInitOptions, RerankInitOptionsUserDefined, RerankResult, TextRerank, UserDefinedRerankingModel, }; + +// For Late Interaction +pub use crate::late_interaction::{ + LateInteractionEmbedding, LateInteractionInitOptions, LateInteractionInitOptionsUserDefined, + LateInteractionModel, LateInteractionModelInfo, LateInteractionTextEmbedding, + UserDefinedLateInteractionModel, +}; + +// For Post-processing (MUVERA) +#[cfg(feature = "muvera")] +pub use crate::postprocess::{Muvera, SimHashProjection}; diff --git a/src/models/text_embedding.rs b/src/models/text_embedding.rs index fc4dd9e..5a28a39 100644 --- a/src/models/text_embedding.rs +++ b/src/models/text_embedding.rs @@ -457,7 +457,9 @@ fn init_models_map() -> HashMap> { ModelInfo { model: EmbeddingModel::SnowflakeArcticEmbedMLongQ, dim: 768, - description: String::from("Quantized Snowflake Arctic embed model, medium with 2048 context"), + description: String::from( + "Quantized Snowflake Arctic embed model, medium with 2048 context", + ), model_code: String::from("snowflake/snowflake-arctic-embed-m-long"), model_file: String::from("onnx/model_quantized.onnx"), additional_files: Vec::new(), diff --git a/src/postprocess/mod.rs b/src/postprocess/mod.rs new file mode 100644 index 0000000..82a767d --- /dev/null +++ b/src/postprocess/mod.rs @@ -0,0 +1,7 @@ +//! Post-processing utilities for embeddings. + +#[cfg(feature = "muvera")] +mod muvera; + +#[cfg(feature = "muvera")] +pub use muvera::*; diff --git a/src/postprocess/muvera.rs b/src/postprocess/muvera.rs new file mode 100644 index 0000000..8fa5712 --- /dev/null +++ b/src/postprocess/muvera.rs @@ -0,0 +1,466 @@ +//! MUVERA (Multi-Vector Retrieval Architecture) implementation. +//! +//! Converts variable-length multi-vector embeddings into fixed-dimensional encodings +//! using SimHash clustering and random projections. + +use anyhow::Result; +use rand::Rng; +use rand_distr::{Distribution, StandardNormal}; + +use crate::LateInteractionTextEmbedding; + +/// Maximum Hamming distance value (64 bits + 1) +const MAX_HAMMING_DISTANCE: u32 = 65; + +/// Precomputed popcount lookup table for bytes +const POPCOUNT_LUT: [u8; 256] = { + let mut table = [0u8; 256]; + let mut i = 0usize; + while i < 256 { + table[i] = (i as u8).count_ones() as u8; + i += 1; + } + table +}; + +/// Compute Hamming distance between two u64 values +#[inline] +fn hamming_distance(a: u64, b: u64) -> u32 { + let xor = a ^ b; + let bytes = xor.to_ne_bytes(); + bytes.iter().map(|&b| POPCOUNT_LUT[b as usize] as u32).sum() +} + +/// Compute full Hamming distance matrix for cluster IDs 0..n +fn hamming_distance_matrix(n: usize) -> Vec> { + let mut matrix = vec![vec![0u32; n]; n]; + for (i, row) in matrix.iter_mut().enumerate() { + for (j, cell) in row.iter_mut().enumerate() { + *cell = hamming_distance(i as u64, j as u64); + } + } + matrix +} + +/// SimHash projection component for MUVERA clustering. +/// +/// Uses random hyperplanes to partition the vector space into 2^k_sim clusters. +#[derive(Debug, Clone)] +pub struct SimHashProjection { + /// Random hyperplane normal vectors of shape (dim, k_sim) + simhash_vectors: Vec>, + k_sim: usize, + dim: usize, +} + +impl SimHashProjection { + /// Create a new SimHash projection with random hyperplanes. + /// + /// # Arguments + /// * `k_sim` - Number of SimHash functions (creates 2^k_sim clusters) + /// * `dim` - Dimensionality of input vectors + /// * `rng` - Random number generator for reproducibility + pub fn new(k_sim: usize, dim: usize, rng: &mut impl Rng) -> Self { + // Generate k_sim random hyperplanes from standard normal distribution + // Shape: (dim, k_sim) - each column is a hyperplane normal vector + let mut simhash_vectors = vec![vec![0.0f32; k_sim]; dim]; + for row in simhash_vectors.iter_mut().take(dim) { + for cell in row.iter_mut().take(k_sim) { + *cell = StandardNormal.sample(rng); + } + } + + Self { + simhash_vectors, + k_sim, + dim, + } + } + + /// Compute cluster IDs for a batch of vectors using SimHash. + /// + /// # Arguments + /// * `vectors` - Input vectors of shape (n_vectors, dim) + /// + /// # Returns + /// Vector of cluster IDs in range [0, 2^k_sim - 1] + pub fn get_cluster_ids(&self, vectors: &[Vec]) -> Vec { + vectors + .iter() + .map(|vec| { + assert_eq!(vec.len(), self.dim, "Vector dimension mismatch"); + + // Compute dot product with each hyperplane: vec @ simhash_vectors + // Result shape: (k_sim,) + let mut dot_products = vec![0.0f32; self.k_sim]; + for (d, &value) in vec.iter().enumerate().take(self.dim) { + for (k, dp) in dot_products.iter_mut().enumerate().take(self.k_sim) { + *dp += value * self.simhash_vectors[d][k]; + } + } + + // Convert signs to cluster ID: (dot_product > 0) @ (1 << arange(k_sim)) + let mut cluster_id = 0u64; + for (k, &dp) in dot_products.iter().enumerate().take(self.k_sim) { + if dp > 0.0 { + cluster_id |= 1u64 << k; + } + } + + cluster_id + }) + .collect() + } +} + +/// MUVERA (Multi-Vector Retrieval Architecture) algorithm implementation. +/// +/// Creates Fixed Dimensional Encodings (FDEs) from variable-length sequences +/// of vectors using SimHash clustering and random projections. +#[derive(Debug, Clone)] +pub struct Muvera { + /// Number of SimHash functions per projection + k_sim: usize, + /// Input vector dimensionality + dim: usize, + /// Output dimensionality after random projection + dim_proj: usize, + /// Number of random projection repetitions + r_reps: usize, + /// SimHash projections for each repetition + simhash_projections: Vec, + /// Random projection matrices: (r_reps, dim, dim_proj) with values in {-1, +1} + dim_reduction_projections: Vec>>, + /// Precomputed Hamming distance matrix for cluster centers + hamming_matrix: Vec>, + /// Number of partitions (2^k_sim) + num_partitions: usize, +} + +impl Muvera { + /// Create a new MUVERA instance. + /// + /// # Arguments + /// * `dim` - Dimensionality of individual input vectors + /// * `k_sim` - Number of SimHash functions (creates 2^k_sim clusters). Default: 5 + /// * `dim_proj` - Dimensionality after random projection (must be <= dim). Default: 16 + /// * `r_reps` - Number of random projection repetitions. Default: 20 + /// * `random_seed` - Seed for random number generator. Default: 42 + /// + /// # Errors + /// Returns error if dim_proj > dim + pub fn new( + dim: usize, + k_sim: Option, + dim_proj: Option, + r_reps: Option, + random_seed: Option, + ) -> Result { + let k_sim = k_sim.unwrap_or(5); + let dim_proj = dim_proj.unwrap_or(16); + let r_reps = r_reps.unwrap_or(20); + let random_seed = random_seed.unwrap_or(42); + + if dim_proj > dim { + return Err(anyhow::anyhow!( + "Cannot project to higher dimensionality (dim_proj={} > dim={})", + dim_proj, + dim + )); + } + + use rand::SeedableRng; + use rand_chacha::ChaCha8Rng; + + let mut rng = ChaCha8Rng::seed_from_u64(random_seed); + + // Create r_reps independent SimHash projections + let simhash_projections: Vec = (0..r_reps) + .map(|_| SimHashProjection::new(k_sim, dim, &mut rng)) + .collect(); + + // Create random projection matrices with entries from {-1, +1} (Rademacher) + let dim_reduction_projections: Vec>> = (0..r_reps) + .map(|_| { + (0..dim) + .map(|_| { + (0..dim_proj) + .map(|_| if rng.gen::() { 1.0f32 } else { -1.0f32 }) + .collect() + }) + .collect() + }) + .collect(); + + let num_partitions = 1 << k_sim; + let hamming_matrix = hamming_distance_matrix(num_partitions); + + Ok(Self { + k_sim, + dim, + dim_proj, + r_reps, + simhash_projections, + dim_reduction_projections, + hamming_matrix, + num_partitions, + }) + } + + /// Create a Muvera instance from a late interaction embedding model. + /// + /// This extracts the embedding dimension from the model automatically. + pub fn from_late_interaction_model( + model: &LateInteractionTextEmbedding, + k_sim: Option, + dim_proj: Option, + r_reps: Option, + random_seed: Option, + ) -> Result { + Self::new(model.dim(), k_sim, dim_proj, r_reps, random_seed) + } + + /// Get the output embedding size. + pub fn embedding_size(&self) -> usize { + self.r_reps * self.num_partitions * self.dim_proj + } + + /// Get the number of SimHash functions (k_sim parameter) + pub fn k_sim(&self) -> usize { + self.k_sim + } + + /// Process a document's vectors into a Fixed Dimensional Encoding (FDE). + /// + /// Uses document-specific settings: normalizes cluster centers by vector count + /// and fills empty clusters using Hamming distance-based selection. + /// + /// # Arguments + /// * `vectors` - Document vectors of shape (n_tokens, dim) + /// + /// # Returns + /// Fixed dimensional encoding of length (r_reps * 2^k_sim * dim_proj) + pub fn process_document(&self, vectors: &[Vec]) -> Vec { + self.process(vectors, true, true) + } + + /// Process a query's vectors into a Fixed Dimensional Encoding (FDE). + /// + /// Uses query-specific settings: no normalization by count and no empty + /// cluster filling to preserve query vector magnitudes. + /// + /// # Arguments + /// * `vectors` - Query vectors of shape (n_tokens, dim) + /// + /// # Returns + /// Fixed dimensional encoding of length (r_reps * 2^k_sim * dim_proj) + pub fn process_query(&self, vectors: &[Vec]) -> Vec { + self.process(vectors, false, false) + } + + /// Core processing method. + /// + /// # Arguments + /// * `vectors` - Input vectors of shape (n_vectors, dim) + /// * `fill_empty_clusters` - Whether to fill empty clusters using nearest vectors + /// * `normalize_by_count` - Whether to normalize cluster centers by count + pub fn process( + &self, + vectors: &[Vec], + fill_empty_clusters: bool, + normalize_by_count: bool, + ) -> Vec { + assert!( + vectors.iter().all(|v| v.len() == self.dim), + "All vectors must have dimension {}", + self.dim + ); + + let mut output = Vec::with_capacity(self.embedding_size()); + + for proj_idx in 0..self.r_reps { + let simhash = &self.simhash_projections[proj_idx]; + + // Initialize cluster centers and track which vectors belong to each cluster + let mut cluster_centers = vec![vec![0.0f32; self.dim]; self.num_partitions]; + let mut cluster_vector_indices: Vec> = vec![Vec::new(); self.num_partitions]; + + // Assign vectors to clusters and accumulate centers (sum) + let cluster_ids = simhash.get_cluster_ids(vectors); + for (vec_idx, &cluster_id) in cluster_ids.iter().enumerate() { + let cluster_idx = cluster_id as usize; + for d in 0..self.dim { + cluster_centers[cluster_idx][d] += vectors[vec_idx][d]; + } + cluster_vector_indices[cluster_idx].push(vec_idx); + } + + // Compute cluster counts and empty mask + let cluster_counts: Vec = + cluster_vector_indices.iter().map(|v| v.len()).collect(); + let empty_mask: Vec = cluster_counts.iter().map(|&c| c == 0).collect(); + + // Normalize by count if requested (only non-empty clusters) + if normalize_by_count { + for (cluster_idx, &count) in cluster_counts.iter().enumerate() { + if count > 0 { + for val in cluster_centers[cluster_idx].iter_mut() { + *val /= count as f32; + } + } + } + } + + // Fill empty clusters using Hamming distance + if fill_empty_clusters { + // For each cluster (row i), find the nearest NON-EMPTY cluster (column j) + // by masking empty columns with MAX_HAMMING_DISTANCE + // This matches Python: masked_hamming = np.where(empty_mask[None, :], MAX, hamming) + // Then argmin along axis=1 + let mut nearest_non_empty: Vec = vec![0; self.num_partitions]; + for (i, nearest) in nearest_non_empty.iter_mut().enumerate() { + let mut min_dist = MAX_HAMMING_DISTANCE; + let mut best_j = 0usize; + for (j, &is_empty) in empty_mask.iter().enumerate() { + let dist = if is_empty { + MAX_HAMMING_DISTANCE + } else { + self.hamming_matrix[i][j] + }; + if dist < min_dist { + min_dist = dist; + best_j = j; + } + } + *nearest = best_j; + } + + // Now fill empty clusters: for each empty cluster i, + // use first vector from nearest_non_empty[i] + for cluster_idx in 0..self.num_partitions { + if empty_mask[cluster_idx] { + let source_cluster = nearest_non_empty[cluster_idx]; + if !cluster_vector_indices[source_cluster].is_empty() { + let fill_vec_idx = cluster_vector_indices[source_cluster][0]; + cluster_centers[cluster_idx] = vectors[fill_vec_idx].clone(); + } + } + } + } + + // Apply random projection for dimensionality reduction + let dim_reduction = &self.dim_reduction_projections[proj_idx]; + let scale = 1.0 / (self.dim_proj as f32).sqrt(); + + for center in cluster_centers.iter().take(self.num_partitions) { + if self.dim_proj < self.dim { + // Project: (1/sqrt(dim_proj)) * (cluster_center @ projection_matrix) + let mut projected = vec![0.0f32; self.dim_proj]; + for (d, row) in dim_reduction.iter().enumerate().take(self.dim) { + for (p, slot) in projected.iter_mut().enumerate() { + *slot += center[d] * row[p]; + } + } + output.extend(projected.into_iter().map(|v| v * scale)); + } else { + // No projection needed (dim_proj == dim) + output.extend_from_slice(center); + } + } + } + + output + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_hamming_distance() { + assert_eq!(hamming_distance(0, 0), 0); + assert_eq!(hamming_distance(0, 1), 1); + assert_eq!(hamming_distance(0b1111, 0b0000), 4); + assert_eq!(hamming_distance(0b1010, 0b0101), 4); + } + + #[test] + fn test_muvera_output_size() { + let muvera = Muvera::new(128, Some(5), Some(16), Some(20), Some(42)).unwrap(); + // r_reps * 2^k_sim * dim_proj = 20 * 32 * 16 = 10240 + assert_eq!(muvera.embedding_size(), 20 * 32 * 16); + } + + #[test] + fn test_muvera_process() { + let muvera = Muvera::new(128, Some(4), Some(8), Some(10), Some(42)).unwrap(); + + // Create some test vectors + let vectors: Vec> = (0..10) + .map(|i| (0..128).map(|j| ((i * 128 + j) as f32).sin()).collect()) + .collect(); + + let doc_encoding = muvera.process_document(&vectors); + let query_encoding = muvera.process_query(&vectors); + + assert_eq!(doc_encoding.len(), muvera.embedding_size()); + assert_eq!(query_encoding.len(), muvera.embedding_size()); + } + + #[test] + fn test_dim_proj_validation() { + let result = Muvera::new(128, Some(5), Some(256), Some(20), Some(42)); + assert!(result.is_err()); + } + + #[test] + fn test_muvera_deterministic() { + let muvera1 = Muvera::new(128, Some(5), Some(16), Some(20), Some(42)).unwrap(); + let muvera2 = Muvera::new(128, Some(5), Some(16), Some(20), Some(42)).unwrap(); + + let vectors: Vec> = (0..10) + .map(|i| (0..128).map(|j| ((i * 128 + j) as f32).sin()).collect()) + .collect(); + + let fde1 = muvera1.process_document(&vectors); + let fde2 = muvera2.process_document(&vectors); + + assert_eq!(fde1, fde2, "MUVERA should be deterministic with same seed"); + } + + #[test] + fn test_muvera_different_seeds() { + let muvera1 = Muvera::new(128, Some(5), Some(16), Some(20), Some(42)).unwrap(); + let muvera2 = Muvera::new(128, Some(5), Some(16), Some(20), Some(123)).unwrap(); + + let vectors: Vec> = (0..10) + .map(|i| (0..128).map(|j| ((i * 128 + j) as f32).sin()).collect()) + .collect(); + + let fde1 = muvera1.process_document(&vectors); + let fde2 = muvera2.process_document(&vectors); + + assert_ne!( + fde1, fde2, + "Different seeds should produce different results" + ); + } + + #[test] + fn test_empty_cluster_filling() { + // Test with very few vectors to ensure some clusters are empty + let muvera = Muvera::new(8, Some(3), Some(4), Some(2), Some(42)).unwrap(); + + // Only 2 vectors, but 2^3 = 8 clusters, so most will be empty + let vectors: Vec> = vec![vec![1.0; 8], vec![-1.0; 8]]; + + let doc_encoding = muvera.process_document(&vectors); + + // Should not panic and should produce valid output + assert_eq!(doc_encoding.len(), muvera.embedding_size()); + + // Verify no NaN or Inf values + assert!(doc_encoding.iter().all(|&v| v.is_finite())); + } +} diff --git a/tests/late-interaction.rs b/tests/late-interaction.rs new file mode 100644 index 0000000..04f4b6c --- /dev/null +++ b/tests/late-interaction.rs @@ -0,0 +1,362 @@ +//! Tests for late interaction (ColBERT-style) embeddings and MUVERA post-processing. + +#![cfg(feature = "hf-hub")] + +#[cfg(feature = "muvera")] +use fastembed::Muvera; +use fastembed::{LateInteractionInitOptions, LateInteractionModel, LateInteractionTextEmbedding}; + +// Canonical values for "Hello World" with colbert-ir/colbertv2.0 +// First 5 columns of first 5 tokens +const CANONICAL_DOC_VALUES_COLBERT: [[f32; 5]; 5] = [ + [0.0759, 0.0841, -0.0299, 0.0374, 0.0254], + [0.0005, -0.0163, -0.0127, 0.2165, 0.1517], + [-0.0257, -0.0575, 0.0135, 0.2202, 0.1896], + [0.0846, 0.0122, 0.0032, -0.0109, -0.1041], + [0.0477, 0.1078, -0.0314, 0.016, 0.0156], +]; + +const CANONICAL_QUERY_VALUES_COLBERT: [[f32; 5]; 5] = [ + [0.0824, 0.0872, -0.0324, 0.0418, 0.024], + [-0.0007, -0.0154, -0.0113, 0.2277, 0.1528], + [-0.0251, -0.0565, 0.0136, 0.2236, 0.1838], + [0.0848, 0.0056, 0.0041, -0.0036, -0.1032], + [0.0574, 0.1072, -0.0332, 0.0233, 0.0209], +]; + +const CANONICAL_DOC_VALUES_ANSWERAI: [[f32; 5]; 5] = [ + [-0.07281, 0.04632, -0.04711, 0.00762, -0.07374], + [-0.04464, 0.04426, -0.074, 0.01801, -0.05233], + [0.09936, -0.05123, -0.04925, -0.05276, -0.08944], + [0.01644, 0.0203, -0.03789, 0.03165, -0.06501], + [-0.07281, 0.04633, -0.04711, 0.00762, -0.07374], +]; + +fn assert_close(actual: &[f32], expected: &[f32], atol: f32) { + assert_eq!(actual.len(), expected.len(), "Length mismatch"); + for (i, (a, e)) in actual.iter().zip(expected.iter()).enumerate() { + assert!( + (a - e).abs() < atol, + "Mismatch at index {}: actual={}, expected={}, diff={}", + i, + a, + e, + (a - e).abs() + ); + } +} + +// Late Interaction Embedding Tests + +#[test] +fn test_colbert_document_embedding_canonical_values() { + let mut model = LateInteractionTextEmbedding::try_new(LateInteractionInitOptions::new( + LateInteractionModel::ColBERTV2, + )) + .unwrap(); + + let embeddings = model.embed(&["Hello World"], None).unwrap(); + + assert_eq!(embeddings.len(), 1); + assert!(embeddings[0].len() >= 5, "Expected at least 5 tokens"); + assert_eq!(embeddings[0][0].len(), 128, "ColBERT dim should be 128"); + + for (token_idx, expected_row) in CANONICAL_DOC_VALUES_COLBERT.iter().enumerate() { + let actual: Vec = embeddings[0][token_idx][..5].to_vec(); + assert_close(&actual, expected_row, 2e-3); + } +} + +#[test] +fn test_colbert_query_embedding_canonical_values() { + let mut model = LateInteractionTextEmbedding::try_new(LateInteractionInitOptions::new( + LateInteractionModel::ColBERTV2, + )) + .unwrap(); + + let embeddings = model.query_embed(&["Hello World"], None).unwrap(); + + assert_eq!(embeddings.len(), 1); + assert_eq!( + embeddings[0].len(), + 32, + "Query should be padded to 32 tokens" + ); + assert_eq!(embeddings[0][0].len(), 128); + + for (token_idx, expected_row) in CANONICAL_QUERY_VALUES_COLBERT.iter().enumerate() { + let actual: Vec = embeddings[0][token_idx][..5].to_vec(); + assert_close(&actual, expected_row, 2e-3); + } +} + +#[test] +fn test_answerai_document_embedding_canonical_values() { + let mut model = LateInteractionTextEmbedding::try_new(LateInteractionInitOptions::new( + LateInteractionModel::AnswerAIColBERTSmallV1, + )) + .unwrap(); + + let embeddings = model.embed(&["Hello World"], None).unwrap(); + + assert_eq!(embeddings.len(), 1); + assert_eq!(embeddings[0][0].len(), 96, "AnswerAI dim should be 96"); + + for (token_idx, expected_row) in CANONICAL_DOC_VALUES_ANSWERAI.iter().enumerate() { + let actual: Vec = embeddings[0][token_idx][..5].to_vec(); + assert_close(&actual, expected_row, 2e-3); + } +} + +#[test] +fn test_embedding_dimension() { + let colbert = LateInteractionTextEmbedding::try_new(LateInteractionInitOptions::new( + LateInteractionModel::ColBERTV2, + )) + .unwrap(); + assert_eq!(colbert.dim(), 128); + + let answerai = LateInteractionTextEmbedding::try_new(LateInteractionInitOptions::new( + LateInteractionModel::AnswerAIColBERTSmallV1, + )) + .unwrap(); + assert_eq!(answerai.dim(), 96); +} + +#[test] +fn test_batch_size_consistency() { + let mut model = LateInteractionTextEmbedding::try_new(LateInteractionInitOptions::new( + LateInteractionModel::AnswerAIColBERTSmallV1, + )) + .unwrap(); + + let documents = vec![ + "short document", + "A bit longer document, which should not affect the size", + ]; + + let result_batch_1 = model.embed(&documents, Some(1)).unwrap(); + let result_batch_2 = model.embed(&documents, Some(2)).unwrap(); + + assert_eq!( + result_batch_1[0].len(), + result_batch_2[0].len(), + "Batch size should not affect token count" + ); + + for (t1, t2) in result_batch_1[0].iter().zip(result_batch_2[0].iter()) { + assert_close(t1, t2, 1e-5); + } +} + +// MUVERA Post-Processing Tests + +#[cfg(feature = "muvera")] +mod muvera_tests { + use super::*; + + // Canonical MUVERA output for deterministic sin-based test vectors + const MUVERA_EXPECTED_FIRST_10: [f32; 10] = [ + 2.0179653, + 1.6323578, + -1.5774617, + -0.26919794, + 3.2250175, + -2.0104198, + -2.2146697, + -1.0453973, + -1.2936, + 2.5332289, + ]; + + #[test] + fn test_muvera_canonical_values() { + let muvera = Muvera::new(128, Some(5), Some(16), Some(20), Some(42)).unwrap(); + + let vectors: Vec> = (0..10) + .map(|i| (0..128).map(|j| ((i * 128 + j) as f32).sin()).collect()) + .collect(); + + let fde = muvera.process_document(&vectors); + + assert_eq!(fde.len(), 10240); + assert_close(&fde[..10], &MUVERA_EXPECTED_FIRST_10, 1e-5); + } + + #[test] + fn test_muvera_deterministic_same_seed() { + let muvera1 = Muvera::new(128, Some(5), Some(16), Some(20), Some(42)).unwrap(); + let muvera2 = Muvera::new(128, Some(5), Some(16), Some(20), Some(42)).unwrap(); + + let vectors: Vec> = (0..10) + .map(|i| (0..128).map(|j| ((i * 128 + j) as f32).sin()).collect()) + .collect(); + + let fde1 = muvera1.process_document(&vectors); + let fde2 = muvera2.process_document(&vectors); + + assert_eq!(fde1, fde2); + } + + #[test] + fn test_muvera_different_seeds_differ() { + let muvera1 = Muvera::new(128, Some(5), Some(16), Some(20), Some(42)).unwrap(); + let muvera2 = Muvera::new(128, Some(5), Some(16), Some(20), Some(123)).unwrap(); + + let vectors: Vec> = (0..10) + .map(|i| (0..128).map(|j| ((i * 128 + j) as f32).sin()).collect()) + .collect(); + + let fde1 = muvera1.process_document(&vectors); + let fde2 = muvera2.process_document(&vectors); + + assert_ne!(fde1, fde2); + } + + #[test] + fn test_muvera_document_vs_query_processing_differs() { + let muvera = Muvera::new(128, Some(5), Some(16), Some(20), Some(42)).unwrap(); + + let vectors: Vec> = (0..10) + .map(|i| (0..128).map(|j| ((i * 128 + j) as f32).sin()).collect()) + .collect(); + + let doc_fde = muvera.process_document(&vectors); + let query_fde = muvera.process_query(&vectors); + + // Document processing: normalize_by_count=true, fill_empty_clusters=true + // Query processing: normalize_by_count=false, fill_empty_clusters=false + assert_ne!(doc_fde, query_fde); + } + + #[test] + fn test_muvera_empty_cluster_handling() { + // With only 2 vectors and 2^3=8 clusters, most clusters will be empty + let muvera = Muvera::new(8, Some(3), Some(4), Some(2), Some(42)).unwrap(); + let vectors: Vec> = vec![vec![1.0; 8], vec![-1.0; 8]]; + + let fde = muvera.process_document(&vectors); + + assert_eq!(fde.len(), muvera.embedding_size()); + assert!(fde.iter().all(|&v| v.is_finite())); + } + + #[test] + fn test_muvera_rejects_invalid_dim_proj() { + // dim_proj > dim should fail + let result = Muvera::new(128, Some(5), Some(256), Some(20), Some(42)); + assert!(result.is_err()); + } + + #[test] + fn test_muvera_embedding_size_calculation() { + let muvera = Muvera::new(128, Some(5), Some(16), Some(20), Some(42)).unwrap(); + // r_reps * 2^k_sim * dim_proj = 20 * 32 * 16 = 10240 + assert_eq!(muvera.embedding_size(), 10240); + } + + #[test] + fn test_muvera_with_colbert_model() { + let mut model = LateInteractionTextEmbedding::try_new(LateInteractionInitOptions::new( + LateInteractionModel::ColBERTV2, + )) + .unwrap(); + + let muvera = + Muvera::from_late_interaction_model(&model, Some(5), Some(16), Some(20), Some(42)) + .unwrap(); + + let doc_embeddings = model + .embed(&["This is a test document about neural networks."], None) + .unwrap(); + let query_embeddings = model + .query_embed(&["What are neural networks?"], None) + .unwrap(); + + let doc_fde = muvera.process_document(&doc_embeddings[0]); + let query_fde = muvera.process_query(&query_embeddings[0]); + + assert_eq!(doc_fde.len(), 10240); + assert_eq!(query_fde.len(), 10240); + + let similarity: f32 = doc_fde + .iter() + .zip(query_fde.iter()) + .map(|(a, b)| a * b) + .sum(); + assert!( + similarity > 0.0, + "Related query-doc should have positive similarity" + ); + } + + #[test] + fn test_muvera_preserves_retrieval_ranking() { + // Verifies MUVERA ranking matches ColBERT MaxSim ranking + let mut model = LateInteractionTextEmbedding::try_new(LateInteractionInitOptions::new( + LateInteractionModel::ColBERTV2, + )) + .unwrap(); + + let muvera = + Muvera::from_late_interaction_model(&model, Some(5), Some(16), Some(20), Some(42)) + .unwrap(); + + let documents = vec![ + "Machine learning is a subset of artificial intelligence.", + "Python is a popular programming language.", + ]; + let query = "What is machine learning?"; + + let doc_embeddings = model.embed(&documents, None).unwrap(); + let query_embeddings = model.query_embed(&[query], None).unwrap(); + + // MUVERA scores + let query_fde = muvera.process_query(&query_embeddings[0]); + let muvera_scores: Vec = doc_embeddings + .iter() + .map(|d| { + let doc_fde = muvera.process_document(d); + query_fde + .iter() + .zip(doc_fde.iter()) + .map(|(a, b)| a * b) + .sum() + }) + .collect(); + + // MaxSim scores (ground truth) + let maxsim_scores: Vec = doc_embeddings + .iter() + .map(|doc_emb| { + query_embeddings[0] + .iter() + .map(|q| { + doc_emb + .iter() + .map(|d| q.iter().zip(d.iter()).map(|(a, b)| a * b).sum::()) + .fold(f32::NEG_INFINITY, f32::max) + }) + .sum() + }) + .collect(); + + // Both should rank doc0 (ML) higher than doc1 (Python) + assert!( + muvera_scores[0] > muvera_scores[1], + "MUVERA: doc0={} should beat doc1={}", + muvera_scores[0], + muvera_scores[1] + ); + assert!( + maxsim_scores[0] > maxsim_scores[1], + "MaxSim: doc0={} should beat doc1={}", + maxsim_scores[0], + maxsim_scores[1] + ); + + // MaxSim values should match Python exactly + assert_close(&[maxsim_scores[0]], &[29.5733], 0.01); + assert_close(&[maxsim_scores[1]], &[9.9226], 0.01); + } +} From a6ae70bbac10cb318ba6f78a33625bbda6002900 Mon Sep 17 00:00:00 2001 From: Louis Maddox Date: Sat, 27 Dec 2025 22:41:53 +0000 Subject: [PATCH 2/3] test: BEIR scidocs benchmark --- .gitignore | 4 +- Cargo.toml | 4 + download_beir_scidocs_dataset.py | 28 +++ examples/beir_scidocs_benchmark.rs | 278 +++++++++++++++++++++++++++++ 4 files changed, 313 insertions(+), 1 deletion(-) create mode 100644 download_beir_scidocs_dataset.py create mode 100644 examples/beir_scidocs_benchmark.rs diff --git a/.gitignore b/.gitignore index f3d681d..db4a233 100644 --- a/.gitignore +++ b/.gitignore @@ -77,4 +77,6 @@ main.rs Cargo.lock ## Nix -/.direnv \ No newline at end of file +/.direnv + +datasets/ diff --git a/Cargo.toml b/Cargo.toml index ed093b4..a7d6810 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -60,3 +60,7 @@ online = ["hf-hub-native-tls"] [[examples]] name = "muvera_demo" features = ["hf-hub", "muvera"] + +[[examples]] +name = "beir_scidocs_benchmark" +features = ["hf-hub", "muvera"] diff --git a/download_beir_scidocs_dataset.py b/download_beir_scidocs_dataset.py new file mode 100644 index 0000000..c010ed2 --- /dev/null +++ b/download_beir_scidocs_dataset.py @@ -0,0 +1,28 @@ +# /// script +# requires-python = ">=3.10" +# dependencies = ["beir"] +# /// + +from beir import util +from pathlib import Path + +dataset = "scidocs" +out_dir = Path("datasets") + +# Download + unzip +zip_path = out_dir / f"{dataset}.zip" +data_path = util.download_and_unzip( + str(zip_path), + str(out_dir) +) + +# Validate structure exists +root = out_dir / dataset +assert (root / "corpus.jsonl").exists() +assert (root / "queries.jsonl").exists() +assert (root / "qrels" / "test.tsv").exists() + +print("Done.") +print("Corpus:", root / "corpus.jsonl") +print("Queries:", root / "queries.jsonl") +print("Qrels:", root / "qrels" / "test.tsv") diff --git a/examples/beir_scidocs_benchmark.rs b/examples/beir_scidocs_benchmark.rs new file mode 100644 index 0000000..9c861c7 --- /dev/null +++ b/examples/beir_scidocs_benchmark.rs @@ -0,0 +1,278 @@ +use anyhow::Result; +use fastembed::{ + LateInteractionInitOptions, LateInteractionModel, LateInteractionTextEmbedding, Muvera, +}; +use std::collections::HashMap; +use std::fs::File; +use std::io::{BufRead, BufReader}; +use std::path::Path; +use ort::execution_providers::{CUDAExecutionProvider, CPUExecutionProvider}; + +/// Compute ColBERT MaxSim score between query and document embeddings +fn maxsim(query_emb: &[Vec], doc_emb: &[Vec]) -> f32 { + query_emb + .iter() + .map(|q| { + doc_emb + .iter() + .map(|d| q.iter().zip(d.iter()).map(|(a, b)| a * b).sum::()) + .fold(f32::NEG_INFINITY, f32::max) + }) + .sum() +} + +/// Compute dot product between two vectors +fn dot(a: &[f32], b: &[f32]) -> f32 { + a.iter().zip(b).map(|(x, y)| x * y).sum() +} + +/// Compute recall@k +fn recall_at_k(relevant: &[String], retrieved: &[String], k: usize) -> f32 { + let top_k = &retrieved[..k.min(retrieved.len())]; + let hits = top_k.iter().filter(|id| relevant.contains(id)).count(); + hits as f32 / relevant.len().max(1) as f32 +} + +/// Load corpus from BEIR jsonl format +fn load_corpus(path: &Path) -> Result<(Vec, Vec)> { + let file = File::open(path)?; + let reader = BufReader::new(file); + + let mut ids = Vec::new(); + let mut texts = Vec::new(); + + for line in reader.lines() { + let line = line?; + let json: serde_json::Value = serde_json::from_str(&line)?; + let id = json["_id"].as_str().unwrap_or("").to_string(); + let title = json["title"].as_str().unwrap_or(""); + let text = json["text"].as_str().unwrap_or(""); + let combined = if title.is_empty() { + text.to_string() + } else { + format!("{} {}", title, text) + }; + ids.push(id); + texts.push(combined); + } + + Ok((ids, texts)) +} + +/// Load queries from BEIR jsonl format +fn load_queries(path: &Path) -> Result<(Vec, Vec)> { + let file = File::open(path)?; + let reader = BufReader::new(file); + + let mut ids = Vec::new(); + let mut texts = Vec::new(); + + for line in reader.lines() { + let line = line?; + let json: serde_json::Value = serde_json::from_str(&line)?; + let id = json["_id"].as_str().unwrap_or("").to_string(); + let text = json["text"].as_str().unwrap_or("").to_string(); + ids.push(id); + texts.push(text); + } + + Ok((ids, texts)) +} + +/// Load qrels (relevance judgments) from TSV format +fn load_qrels(path: &Path) -> Result>> { + let file = File::open(path)?; + let reader = BufReader::new(file); + + let mut qrels: HashMap> = HashMap::new(); + + for (i, line) in reader.lines().enumerate() { + let line = line?; + if i == 0 && line.starts_with("query-id") { + continue; // Skip header + } + let parts: Vec<&str> = line.split('\t').collect(); + if parts.len() >= 3 { + let query_id = parts[0].to_string(); + let doc_id = parts[1].to_string(); + let relevance: i32 = parts[2].parse().unwrap_or(0); + if relevance > 0 { + qrels.entry(query_id).or_default().push(doc_id); + } + } + } + + Ok(qrels) +} + +fn main() -> Result<()> { + // Configuration + let dataset_path = Path::new("datasets/scidocs"); // Adjust path as needed + let batch_size = 8; + let top_n_candidates = 100; // Number of candidates to retrieve with MUVERA + + // MUVERA parameter configurations: (r_reps, k_sim, dim_proj) + let muvera_configs = vec![ + (20, 3, 8), // 1280-dim + (20, 4, 8), // 2560-dim + (20, 5, 8), // 5120-dim + (20, 5, 16), // 10240-dim + (30, 5, 16), // 15360-dim + (40, 5, 16), // 20480-dim + ]; + + // 1. Load dataset + println!("Loading SciDocs dataset..."); + let (corpus_ids, corpus_texts) = load_corpus(&dataset_path.join("corpus.jsonl"))?; + let (query_ids, query_texts) = load_queries(&dataset_path.join("queries.jsonl"))?; + let qrels = load_qrels(&dataset_path.join("qrels/test.tsv"))?; + + println!("Corpus size: {}", corpus_ids.len()); + println!("Query count: {}", query_ids.len()); + println!("Queries with relevance judgments: {}", qrels.len()); + + // 2. Initialize ColBERT model + println!("\nLoading ColBERT model..."); + + let mut model = LateInteractionTextEmbedding::try_new( + LateInteractionInitOptions::new(LateInteractionModel::ColBERTV2) + .with_execution_providers(vec![ + CUDAExecutionProvider::default().build().into(), + CPUExecutionProvider::default().build().into(), + ]) + )?; + + // 3. Embed corpus + println!("Embedding corpus..."); + let corpus_embeddings = model.embed(&corpus_texts, Some(batch_size))?; + + // Create id -> index mapping + let corpus_id_to_idx: HashMap<&str, usize> = corpus_ids + .iter() + .enumerate() + .map(|(i, id)| (id.as_str(), i)) + .collect(); + + // 4. Embed queries + println!("Embedding queries..."); + let query_embeddings = model.query_embed(&query_texts, Some(batch_size))?; + + // 5. Evaluate brute-force ColBERT (baseline) + println!("\n--- Evaluating brute-force ColBERT ---"); + let mut recalls_4 = Vec::new(); + let mut recalls_5 = Vec::new(); + let mut recalls_10 = Vec::new(); + + for (q_idx, query_id) in query_ids.iter().enumerate() { + if let Some(relevant_docs) = qrels.get(query_id) { + // Score all documents with MaxSim + let mut scores: Vec<(usize, f32)> = corpus_embeddings + .iter() + .enumerate() + .map(|(d_idx, doc_emb)| (d_idx, maxsim(&query_embeddings[q_idx], doc_emb))) + .collect(); + + // Sort by score descending + scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); + + // Get top-k document IDs + let retrieved: Vec = scores + .iter() + .take(10) + .map(|(idx, _)| corpus_ids[*idx].clone()) + .collect(); + + recalls_4.push(recall_at_k(relevant_docs, &retrieved, 4)); + recalls_5.push(recall_at_k(relevant_docs, &retrieved, 5)); + recalls_10.push(recall_at_k(relevant_docs, &retrieved, 10)); + } + } + + let avg_r4 = recalls_4.iter().sum::() / recalls_4.len() as f32; + let avg_r5 = recalls_5.iter().sum::() / recalls_5.len() as f32; + let avg_r10 = recalls_10.iter().sum::() / recalls_10.len() as f32; + + println!( + "Vector: colbert, Recall@4: {:.4}, Recall@5: {:.4}, Recall@10: {:.4}", + avg_r4, avg_r5, avg_r10 + ); + + // 6. Evaluate MUVERA configurations + for (r_reps, k_sim, dim_proj) in &muvera_configs { + let muvera = Muvera::from_late_interaction_model( + &model, + Some(*k_sim), + Some(*dim_proj), + Some(*r_reps), + Some(42), + )?; + + let embedding_size = muvera.embedding_size(); + println!("\n--- Evaluating MUVERA-{} (r={}, k={}, d={}) ---", + embedding_size, r_reps, k_sim, dim_proj); + + // Convert corpus to FDEs + let corpus_fdes: Vec> = corpus_embeddings + .iter() + .map(|emb| muvera.process_document(emb)) + .collect(); + + // Convert queries to FDEs + let query_fdes: Vec> = query_embeddings + .iter() + .map(|emb| muvera.process_query(emb)) + .collect(); + + let mut recalls_4 = Vec::new(); + let mut recalls_5 = Vec::new(); + let mut recalls_10 = Vec::new(); + + for (q_idx, query_id) in query_ids.iter().enumerate() { + if let Some(relevant_docs) = qrels.get(query_id) { + // Stage 1: ANN candidate retrieval using MUVERA dot product + let mut fde_scores: Vec<(usize, f32)> = corpus_fdes + .iter() + .enumerate() + .map(|(d_idx, doc_fde)| (d_idx, dot(&query_fdes[q_idx], doc_fde))) + .collect(); + + fde_scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); + + let candidates: Vec = fde_scores + .iter() + .take(top_n_candidates) + .map(|(idx, _)| *idx) + .collect(); + + // Stage 2: Rerank candidates with ColBERT MaxSim + let mut reranked: Vec<(usize, f32)> = candidates + .iter() + .map(|&d_idx| (d_idx, maxsim(&query_embeddings[q_idx], &corpus_embeddings[d_idx]))) + .collect(); + + reranked.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); + + let retrieved: Vec = reranked + .iter() + .take(10) + .map(|(idx, _)| corpus_ids[*idx].clone()) + .collect(); + + recalls_4.push(recall_at_k(relevant_docs, &retrieved, 4)); + recalls_5.push(recall_at_k(relevant_docs, &retrieved, 5)); + recalls_10.push(recall_at_k(relevant_docs, &retrieved, 10)); + } + } + + let avg_r4 = recalls_4.iter().sum::() / recalls_4.len() as f32; + let avg_r5 = recalls_5.iter().sum::() / recalls_5.len() as f32; + let avg_r10 = recalls_10.iter().sum::() / recalls_10.len() as f32; + + println!( + "Vector: muvera-{}, Recall@4: {:.4}, Recall@5: {:.4}, Recall@10: {:.4}", + embedding_size, avg_r4, avg_r5, avg_r10 + ); + } + + Ok(()) +} From aed1cfb9b762326222b5e3654324bf87a8406fc2 Mon Sep 17 00:00:00 2001 From: Louis Maddox Date: Sun, 28 Dec 2025 00:01:42 +0000 Subject: [PATCH 3/3] chore: get BEIR eval working --- examples/beir_scidocs_benchmark.rs | 168 +++++++++++++++++------------ 1 file changed, 99 insertions(+), 69 deletions(-) diff --git a/examples/beir_scidocs_benchmark.rs b/examples/beir_scidocs_benchmark.rs index 9c861c7..3a14507 100644 --- a/examples/beir_scidocs_benchmark.rs +++ b/examples/beir_scidocs_benchmark.rs @@ -144,58 +144,65 @@ fn main() -> Result<()> { // 3. Embed corpus println!("Embedding corpus..."); - let corpus_embeddings = model.embed(&corpus_texts, Some(batch_size))?; - - // Create id -> index mapping - let corpus_id_to_idx: HashMap<&str, usize> = corpus_ids - .iter() - .enumerate() - .map(|(i, id)| (id.as_str(), i)) - .collect(); + let mut corpus_embeddings = Vec::with_capacity(corpus_texts.len()); + for (i, text) in corpus_texts.iter().enumerate() { + if i % 1000 == 0 { + println!(" Embedded {}/{} documents", i, corpus_texts.len()); + } + let emb = model.embed(&[text], None)?; + corpus_embeddings.push(emb.into_iter().next().unwrap()); + } // 4. Embed queries println!("Embedding queries..."); - let query_embeddings = model.query_embed(&query_texts, Some(batch_size))?; - - // 5. Evaluate brute-force ColBERT (baseline) - println!("\n--- Evaluating brute-force ColBERT ---"); - let mut recalls_4 = Vec::new(); - let mut recalls_5 = Vec::new(); - let mut recalls_10 = Vec::new(); - - for (q_idx, query_id) in query_ids.iter().enumerate() { - if let Some(relevant_docs) = qrels.get(query_id) { - // Score all documents with MaxSim - let mut scores: Vec<(usize, f32)> = corpus_embeddings - .iter() - .enumerate() - .map(|(d_idx, doc_emb)| (d_idx, maxsim(&query_embeddings[q_idx], doc_emb))) - .collect(); - - // Sort by score descending - scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); - - // Get top-k document IDs - let retrieved: Vec = scores - .iter() - .take(10) - .map(|(idx, _)| corpus_ids[*idx].clone()) - .collect(); - - recalls_4.push(recall_at_k(relevant_docs, &retrieved, 4)); - recalls_5.push(recall_at_k(relevant_docs, &retrieved, 5)); - recalls_10.push(recall_at_k(relevant_docs, &retrieved, 10)); + let mut query_embeddings = Vec::with_capacity(query_texts.len()); + for (i, text) in query_texts.iter().enumerate() { + if i % 100 == 0 { + println!(" Embedded {}/{} queries", i, query_texts.len()); } + let emb = model.query_embed(&[text], None)?; + query_embeddings.push(emb.into_iter().next().unwrap()); } - - let avg_r4 = recalls_4.iter().sum::() / recalls_4.len() as f32; - let avg_r5 = recalls_5.iter().sum::() / recalls_5.len() as f32; - let avg_r10 = recalls_10.iter().sum::() / recalls_10.len() as f32; - - println!( - "Vector: colbert, Recall@4: {:.4}, Recall@5: {:.4}, Recall@10: {:.4}", - avg_r4, avg_r5, avg_r10 - ); + + // // 5. Evaluate brute-force ColBERT (baseline) + // println!("\n--- Evaluating brute-force ColBERT ---"); + // let mut recalls_4 = Vec::new(); + // let mut recalls_5 = Vec::new(); + // let mut recalls_10 = Vec::new(); + // + // for (q_idx, query_id) in query_ids.iter().enumerate() { + // if let Some(relevant_docs) = qrels.get(query_id) { + // // Score all documents with MaxSim + // let mut scores: Vec<(usize, f32)> = corpus_embeddings + // .iter() + // .enumerate() + // .map(|(d_idx, doc_emb)| (d_idx, maxsim(&query_embeddings[q_idx], doc_emb))) + // .collect(); + // + // // Sort by score descending + // scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); + // + // // Get top-k document IDs + // let retrieved: Vec = scores + // .iter() + // .take(10) + // .map(|(idx, _)| corpus_ids[*idx].clone()) + // .collect(); + // + // recalls_4.push(recall_at_k(relevant_docs, &retrieved, 4)); + // recalls_5.push(recall_at_k(relevant_docs, &retrieved, 5)); + // recalls_10.push(recall_at_k(relevant_docs, &retrieved, 10)); + // } + // } + // + // let avg_r4 = recalls_4.iter().sum::() / recalls_4.len() as f32; + // let avg_r5 = recalls_5.iter().sum::() / recalls_5.len() as f32; + // let avg_r10 = recalls_10.iter().sum::() / recalls_10.len() as f32; + // + // println!( + // "Vector: colbert, Recall@4: {:.4}, Recall@5: {:.4}, Recall@10: {:.4}", + // avg_r4, avg_r5, avg_r10 + // ); // 6. Evaluate MUVERA configurations for (r_reps, k_sim, dim_proj) in &muvera_configs { @@ -226,44 +233,67 @@ fn main() -> Result<()> { let mut recalls_4 = Vec::new(); let mut recalls_5 = Vec::new(); let mut recalls_10 = Vec::new(); - + for (q_idx, query_id) in query_ids.iter().enumerate() { if let Some(relevant_docs) = qrels.get(query_id) { - // Stage 1: ANN candidate retrieval using MUVERA dot product - let mut fde_scores: Vec<(usize, f32)> = corpus_fdes + // Just MUVERA dot product - NO reranking + let mut scores: Vec<(usize, f32)> = corpus_fdes .iter() .enumerate() .map(|(d_idx, doc_fde)| (d_idx, dot(&query_fdes[q_idx], doc_fde))) .collect(); - - fde_scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); - - let candidates: Vec = fde_scores - .iter() - .take(top_n_candidates) - .map(|(idx, _)| *idx) - .collect(); - - // Stage 2: Rerank candidates with ColBERT MaxSim - let mut reranked: Vec<(usize, f32)> = candidates - .iter() - .map(|&d_idx| (d_idx, maxsim(&query_embeddings[q_idx], &corpus_embeddings[d_idx]))) - .collect(); - - reranked.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); - - let retrieved: Vec = reranked + + scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); + + let retrieved: Vec = scores .iter() .take(10) .map(|(idx, _)| corpus_ids[*idx].clone()) .collect(); - + recalls_4.push(recall_at_k(relevant_docs, &retrieved, 4)); recalls_5.push(recall_at_k(relevant_docs, &retrieved, 5)); recalls_10.push(recall_at_k(relevant_docs, &retrieved, 10)); } } + // for (q_idx, query_id) in query_ids.iter().enumerate() { + // if let Some(relevant_docs) = qrels.get(query_id) { + // // Stage 1: ANN candidate retrieval using MUVERA dot product + // let mut fde_scores: Vec<(usize, f32)> = corpus_fdes + // .iter() + // .enumerate() + // .map(|(d_idx, doc_fde)| (d_idx, dot(&query_fdes[q_idx], doc_fde))) + // .collect(); + // + // fde_scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); + // + // let candidates: Vec = fde_scores + // .iter() + // .take(top_n_candidates) + // .map(|(idx, _)| *idx) + // .collect(); + // + // // Stage 2: Rerank candidates with ColBERT MaxSim + // let mut reranked: Vec<(usize, f32)> = candidates + // .iter() + // .map(|&d_idx| (d_idx, maxsim(&query_embeddings[q_idx], &corpus_embeddings[d_idx]))) + // .collect(); + // + // reranked.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); + // + // let retrieved: Vec = reranked + // .iter() + // .take(10) + // .map(|(idx, _)| corpus_ids[*idx].clone()) + // .collect(); + // + // recalls_4.push(recall_at_k(relevant_docs, &retrieved, 4)); + // recalls_5.push(recall_at_k(relevant_docs, &retrieved, 5)); + // recalls_10.push(recall_at_k(relevant_docs, &retrieved, 10)); + // } + // } + let avg_r4 = recalls_4.iter().sum::() / recalls_4.len() as f32; let avg_r5 = recalls_5.iter().sum::() / recalls_5.len() as f32; let avg_r10 = recalls_10.iter().sum::() / recalls_10.len() as f32;