#![cfg(feature = "hf-hub")]
use fastembed::{SparseInitOptions, SparseModel, SparseTextEmbedding};
const EPS: f32 = 1e-3;
const BATCH_SIZE: usize = 4;
#[test]
fn test_if_splade_embeddings_match_python() {
let mut model = SparseTextEmbedding::try_new(SparseInitOptions::new(
SparseModel::OpenSearchNeuralSparseDocV3Gte,
))
.expect("Failed to initialize the inference-free SPLADE model");
let documents = vec!["Hello World"];
let expected_document_indices = [
999, 1010, 1011, 1024, 1028, 1029, 1045, 1074, 1993, 2017, 2033, 2054, 2073, 2080, 2088,
];
let expected_document_values = [
0.16544909, 0.00529129, 0.0392109, 0.12337475, 0.09640586, 0.05325737, 0.09611791,
0.03159865, 0.01349991, 0.09392473, 0.01928805, 0.05238346, 0.05515401, 0.03156782,
0.98263124,
];
let expected_query_indices = [2088, 7592];
let expected_query_values = [3.42086864, 6.93775654];
let embeddings = model
.embed(documents.clone(), Some(BATCH_SIZE))
.expect("Embedding failed");
assert_eq!(embeddings.len(), documents.len());
let document = &embeddings[0];
assert_eq!(document.indices.len(), document.values.len());
assert!(document.indices.len() > expected_document_indices.len());
assert_eq!(
document.indices[..expected_document_indices.len()],
expected_document_indices
);
for (i, expected) in expected_document_values.iter().enumerate() {
assert!(
(document.values[i] - expected).abs() < EPS,
"dimension {} is {}, expected {expected}",
document.indices[i],
document.values[i],
);
}
let model: &SparseTextEmbedding = &model;
let query_embeddings = model
.query_embed(documents)
.expect("Query embedding failed");
assert_eq!(query_embeddings.len(), 1);
let query = &query_embeddings[0];
assert_eq!(query.indices, expected_query_indices);
assert_eq!(query.values.len(), expected_query_values.len());
for (i, expected) in expected_query_values.iter().enumerate() {
assert!(
(query.values[i] - expected).abs() < EPS,
"dimension {} is {}, expected {expected}",
query.indices[i],
query.values[i],
);
}
}