use std::fmt;
use std::path::{Path, PathBuf};
#[cfg(test)]
use std::sync::atomic::{AtomicBool, Ordering};
use asupersync::Cx;
use rayon::prelude::*;
use safetensors::SafeTensors;
use tokenizers::Tokenizer;
use tracing::instrument;
use crate::model_manifest::{
MODEL2VEC_OUTPUT_NORMALIZATION_V1, MODEL2VEC_POOLING_V1, MODEL2VEC_PREPROCESSING_V1,
MODEL2VEC_SEQUENCE_POLICY_V1, ModelArtifactManifestV1,
};
use crate::model_registry::{ensure_model_storage_layout, model_directory_variants};
use frankensearch_core::error::{SearchError, SearchResult};
use frankensearch_core::generation::{EmbeddingIdentityBundleV1, QuantizationFormat};
use frankensearch_core::traits::{Embedder, ModelCategory, SearchFuture};
const REQUIRED_FILES: [&str; 2] = ["tokenizer.json", "model.safetensors"];
const PARALLEL_BATCH_MIN: usize = 8;
const TENSOR_NAME_CANDIDATES: [&str; 5] =
["embeddings", "embedding", "word_embeddings", "embed", "emb"];
const DEFAULT_MODEL_NAME: &str = "potion-multilingual-128M";
const DEFAULT_HF_ID: &str = "minishlab/potion-multilingual-128M";
pub struct Model2VecEmbedder {
tokenizer: Tokenizer,
embeddings: Vec<f32>,
dimensions: usize,
vocab_size: usize,
name: String,
model_dir: PathBuf,
identity: EmbeddingIdentityBundleV1,
#[cfg(test)]
last_tokenizer_route_was_offset_free: AtomicBool,
}
impl fmt::Debug for Model2VecEmbedder {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Model2VecEmbedder")
.field("name", &self.name)
.field("dimensions", &self.dimensions)
.field("vocab_size", &self.vocab_size)
.field("model_dir", &"<redacted>")
.field("identity", &self.identity.fingerprint())
.finish_non_exhaustive()
}
}
impl Model2VecEmbedder {
#[instrument(skip_all, fields(model = DEFAULT_MODEL_NAME))]
pub fn load(model_dir: impl AsRef<Path>) -> SearchResult<Self> {
Self::load_with_name(model_dir, DEFAULT_MODEL_NAME)
}
pub fn load_with_name(model_dir: impl AsRef<Path>, name: &str) -> SearchResult<Self> {
let model_dir = model_dir.as_ref();
#[cfg(test)]
{
Self::load_explicit_test_model(model_dir, name)
}
#[cfg(not(test))]
{
let verified = ModelArtifactManifestV1::potion_128m_native()?.verify_dir(model_dir)?;
let identity = verified.identity_bundle(QuantizationFormat::F32, "in-memory-f32-v1")?;
validate_registered_execution_contract(&identity)?;
Self::load_preverified(model_dir, name, identity)
}
}
fn load_preverified(
model_dir: &Path,
name: &str,
mut identity: EmbeddingIdentityBundleV1,
) -> SearchResult<Self> {
for filename in &REQUIRED_FILES {
let path = model_dir.join(filename);
if !path.exists() {
return Err(SearchError::ModelNotFound {
name: format!("{name} (missing {filename} in {})", model_dir.display()),
});
}
}
let tokenizer_path = model_dir.join("tokenizer.json");
let tokenizer =
Tokenizer::from_file(&tokenizer_path).map_err(|e| SearchError::ModelLoadFailed {
path: tokenizer_path,
source: format!("failed to load tokenizer: {e}").into(),
})?;
let safetensors_path = model_dir.join("model.safetensors");
let safetensors_data =
std::fs::read(&safetensors_path).map_err(|e| SearchError::ModelLoadFailed {
path: safetensors_path.clone(),
source: Box::new(e),
})?;
let safetensors = SafeTensors::deserialize(&safetensors_data).map_err(|e| {
SearchError::ModelLoadFailed {
path: safetensors_path.clone(),
source: format!("failed to parse safetensors: {e}").into(),
}
})?;
let tensor_name = discover_tensor_name(&safetensors).ok_or_else(|| {
let available: Vec<_> = safetensors.names().into_iter().collect();
SearchError::ModelLoadFailed {
path: safetensors_path.clone(),
source: format!(
"no embedding tensor found. Tried: {TENSOR_NAME_CANDIDATES:?}. Available: {available:?}"
)
.into(),
}
})?;
let tensor =
safetensors
.tensor(&tensor_name)
.map_err(|e| SearchError::ModelLoadFailed {
path: safetensors_path.clone(),
source: format!("failed to get tensor '{tensor_name}': {e}").into(),
})?;
let shape = tensor.shape();
if shape.len() != 2 {
return Err(SearchError::ModelLoadFailed {
path: safetensors_path,
source: format!(
"expected 2D tensor, got {}D with shape {shape:?}",
shape.len()
)
.into(),
});
}
let vocab_size = shape[0];
let dimensions = shape[1];
let parsed_dimension =
u32::try_from(dimensions).map_err(|_| SearchError::InvalidConfig {
field: "model2vec.dimension".to_owned(),
value: dimensions.to_string(),
reason: "parsed tensor dimension exceeds the identity schema".to_owned(),
})?;
if identity.producer.backend == "explicit-test-backend" {
identity.space.dimension = parsed_dimension;
identity.storage.dimension = parsed_dimension;
identity.producer.golden_vectors.dimension = parsed_dimension;
identity.producer.space_fingerprint = identity.space.fingerprint();
}
if identity.space.dimension != parsed_dimension {
return Err(SearchError::ModelLoadFailed {
path: safetensors_path.clone(),
source: format!(
"parsed embedding dimension {parsed_dimension} disagrees with attested dimension {}",
identity.space.dimension
)
.into(),
});
}
identity.validate()?;
let embeddings = parse_f32_matrix(tensor.data(), vocab_size, dimensions).map_err(|e| {
SearchError::ModelLoadFailed {
path: safetensors_path,
source: e.into(),
}
})?;
tracing::info!(
model = DEFAULT_MODEL_NAME,
vocab_size,
dimensions,
manifest = %identity.producer.provenance_manifest_fingerprint,
identity = %identity.fingerprint(),
"Model2Vec model loaded"
);
Ok(Self {
tokenizer,
embeddings,
dimensions,
vocab_size,
name: name.to_owned(),
model_dir: model_dir.to_owned(),
identity,
#[cfg(test)]
last_tokenizer_route_was_offset_free: AtomicBool::new(false),
})
}
#[cfg(test)]
fn load_explicit_test_model(model_dir: &Path, name: &str) -> SearchResult<Self> {
Self::load_preverified(
model_dir,
name,
EmbeddingIdentityBundleV1::explicit_test_model(name, 1),
)
}
pub fn embed_sync(&self, text: &str) -> SearchResult<Vec<f32>> {
if text.is_empty() {
return Ok(vec![0.0; self.dimensions]);
}
let encoding =
self.tokenizer
.encode_fast(text, false)
.map_err(|e| SearchError::EmbeddingFailed {
model: self.name.clone(),
source: format!("tokenization failed: {e}").into(),
})?;
#[cfg(test)]
self.last_tokenizer_route_was_offset_free.store(
!encoding.get_ids().is_empty()
&& encoding.get_tokens().iter().all(String::is_empty)
&& encoding
.get_offsets()
.iter()
.all(|&offset| offset == (0, 0))
&& encoding.get_word_ids().iter().all(Option::is_none),
Ordering::Relaxed,
);
Ok(self.embed_token_ids(encoding.get_ids()))
}
#[inline]
fn embed_token_ids(&self, token_ids: &[u32]) -> Vec<f32> {
if token_ids.is_empty() {
return vec![0.0; self.dimensions];
}
let mut sum = vec![0.0_f32; self.dimensions];
let count = crate::simd::accumulate_model2vec_rows(
&mut sum,
&self.embeddings,
token_ids,
self.vocab_size,
);
if count == 0 {
return vec![0.0; self.dimensions];
}
finish_mean_pool_and_normalize(&mut sum, count);
sum
}
#[cfg(test)]
fn last_tokenizer_route_was_offset_free(&self) -> bool {
self.last_tokenizer_route_was_offset_free
.load(Ordering::Relaxed)
}
#[cfg(feature = "bench-internals")]
#[doc(hidden)]
pub fn benchmark_embed_sync_former_encode(&self, text: &str) -> SearchResult<Vec<f32>> {
if text.is_empty() {
return Ok(vec![0.0; self.dimensions]);
}
let encoding =
self.tokenizer
.encode(text, false)
.map_err(|e| SearchError::EmbeddingFailed {
model: self.name.clone(),
source: format!("tokenization failed: {e}").into(),
})?;
Ok(self.embed_token_ids(encoding.get_ids()))
}
#[cfg(feature = "bench-internals")]
#[doc(hidden)]
pub fn benchmark_embed_sync_former_finish(&self, text: &str) -> SearchResult<Vec<f32>> {
if text.is_empty() {
return Ok(vec![0.0; self.dimensions]);
}
let encoding =
self.tokenizer
.encode_fast(text, false)
.map_err(|e| SearchError::EmbeddingFailed {
model: self.name.clone(),
source: format!("tokenization failed: {e}").into(),
})?;
Ok(self.embed_token_ids_with_former_finish(encoding.get_ids()))
}
#[cfg(feature = "bench-internals")]
fn embed_token_ids_with_former_finish(&self, token_ids: &[u32]) -> Vec<f32> {
if token_ids.is_empty() {
return vec![0.0; self.dimensions];
}
let mut sum = vec![0.0_f32; self.dimensions];
let count = crate::simd::accumulate_model2vec_rows(
&mut sum,
&self.embeddings,
token_ids,
self.vocab_size,
);
if count == 0 {
return vec![0.0; self.dimensions];
}
finish_mean_pool_and_normalize_former(&mut sum, count);
sum
}
pub fn embed_batch_sync(&self, texts: &[&str]) -> SearchResult<Vec<Vec<f32>>> {
if texts.len() >= PARALLEL_BATCH_MIN {
texts.par_iter().map(|text| self.embed_sync(text)).collect()
} else {
let mut results = Vec::with_capacity(texts.len());
for text in texts {
results.push(self.embed_sync(text)?);
}
Ok(results)
}
}
#[must_use]
pub fn model_dir(&self) -> &Path {
&self.model_dir
}
#[must_use]
pub const fn vocab_size(&self) -> usize {
self.vocab_size
}
}
#[inline]
fn finish_mean_pool_and_normalize(sum: &mut [f32], count: usize) {
#[allow(clippy::cast_precision_loss)]
let inv = 1.0 / count as f32;
let mut norm_sq = 0.0_f32;
for value in sum.iter_mut() {
*value *= inv;
norm_sq += *value * *value;
}
if norm_sq.is_finite() && norm_sq > f32::EPSILON {
let inv_norm = 1.0 / norm_sq.sqrt();
for value in sum {
*value *= inv_norm;
}
} else {
sum.fill(0.0);
}
}
#[cfg(any(test, feature = "bench-internals"))]
fn finish_mean_pool_and_normalize_former(sum: &mut [f32], count: usize) {
#[allow(clippy::cast_precision_loss)]
let inv = 1.0 / count as f32;
for value in sum.iter_mut() {
*value *= inv;
}
let norm_sq: f32 = sum.iter().map(|value| value * value).sum();
if norm_sq.is_finite() && norm_sq > f32::EPSILON {
let inv_norm = 1.0 / norm_sq.sqrt();
for value in sum {
*value *= inv_norm;
}
} else {
sum.fill(0.0);
}
}
fn validate_registered_execution_contract(
identity: &EmbeddingIdentityBundleV1,
) -> SearchResult<()> {
for (field, actual, expected) in [
(
"model preprocessing",
identity.space.model_preprocessing.as_str(),
MODEL2VEC_PREPROCESSING_V1,
),
(
"sequence policy",
identity.space.sequence_policy.as_str(),
MODEL2VEC_SEQUENCE_POLICY_V1,
),
(
"pooling",
identity.space.pooling.as_str(),
MODEL2VEC_POOLING_V1,
),
(
"output normalization",
identity.space.output_normalization.as_str(),
MODEL2VEC_OUTPUT_NORMALIZATION_V1,
),
] {
if actual != expected {
return Err(SearchError::InvalidConfig {
field: "model2vec.execution_contract".to_owned(),
value: identity.space.logical_model_id.clone(),
reason: format!("registered {field} disagrees with the native Model2Vec backend"),
});
}
}
Ok(())
}
fn embed_checkpoint(cx: &Cx, phase: &'static str) -> SearchResult<()> {
cx.checkpoint().map_err(|error| SearchError::Cancelled {
phase: phase.to_owned(),
reason: cx
.cancel_reason()
.map_or_else(|| error.to_string(), |reason| reason.to_string()),
})
}
impl Embedder for Model2VecEmbedder {
fn embed<'a>(&'a self, cx: &'a Cx, text: &'a str) -> SearchFuture<'a, Vec<f32>> {
Box::pin(async move {
embed_checkpoint(cx, "model2vec.embed")?;
self.embed_sync(text)
})
}
fn embed_batch<'a>(
&'a self,
cx: &'a Cx,
texts: &'a [&'a str],
) -> SearchFuture<'a, Vec<Vec<f32>>> {
Box::pin(async move {
embed_checkpoint(cx, "model2vec.embed_batch")?;
self.embed_batch_sync(texts)
})
}
fn identity(&self) -> SearchResult<&EmbeddingIdentityBundleV1> {
Ok(&self.identity)
}
fn dimension(&self) -> usize {
self.dimensions
}
fn id(&self) -> &str {
&self.name
}
fn model_name(&self) -> &str {
&self.name
}
fn is_semantic(&self) -> bool {
true
}
fn category(&self) -> ModelCategory {
ModelCategory::StaticEmbedder
}
}
fn discover_tensor_name(safetensors: &SafeTensors<'_>) -> Option<String> {
let names = safetensors.names();
for candidate in &TENSOR_NAME_CANDIDATES {
if names.iter().any(|n| n == candidate) {
return Some((*candidate).to_owned());
}
}
if names.len() == 1 {
return Some(names[0].to_owned());
}
None
}
fn parse_f32_matrix(data: &[u8], vocab_size: usize, dimensions: usize) -> Result<Vec<f32>, String> {
let expected_elements = vocab_size
.checked_mul(dimensions)
.ok_or_else(|| format!("matrix size overflow for [{vocab_size} x {dimensions}]"))?;
let expected_bytes = expected_elements
.checked_mul(4)
.ok_or_else(|| format!("byte size overflow for [{vocab_size} x {dimensions}] f32"))?;
if data.len() != expected_bytes {
return Err(format!(
"tensor data size mismatch: expected {expected_bytes} bytes for [{vocab_size} x {dimensions}] f32, got {}",
data.len()
));
}
let mut matrix = Vec::with_capacity(expected_elements);
for &bytes in data.as_chunks::<4>().0 {
matrix.push(f32::from_le_bytes(bytes));
}
if matrix.len() != expected_elements {
return Err(format!(
"parsed element count mismatch: expected {}, got {}",
expected_elements,
matrix.len()
));
}
Ok(matrix)
}
#[must_use]
pub fn find_model_dir(model_name: &str) -> Option<PathBuf> {
find_model_dir_with_hf_id(model_name, DEFAULT_HF_ID)
}
#[must_use]
pub fn find_model_dir_with_hf_id(model_name: &str, hf_id: &str) -> Option<PathBuf> {
let mut candidates = Vec::new();
if let Ok(dir) = std::env::var("FRANKENSEARCH_MODEL_DIR") {
let base = PathBuf::from(dir);
for variant in model_directory_variants(model_name) {
candidates.push(base.join(variant));
}
candidates.push(base);
}
let model_root = ensure_model_storage_layout();
for variant in model_directory_variants(model_name) {
candidates.push(model_root.join(variant));
}
if let Some(cache_dir) = frankensearch_core::platform_dirs::cache_dir() {
let hf_dir = cache_dir
.join("huggingface/hub")
.join(format!("models--{}", hf_id.replace('/', "--")));
if let Ok(snapshots) = std::fs::read_dir(hf_dir.join("snapshots")) {
for entry in snapshots.flatten() {
candidates.push(entry.path());
}
}
}
for candidate in &candidates {
if has_required_files(candidate) {
return Some(candidate.clone());
}
}
None
}
fn has_required_files(dir: &Path) -> bool {
REQUIRED_FILES.iter().all(|f| dir.join(f).exists())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::simd::{Model2VecAccumulationRoute, last_model2vec_accumulation_route_for_test};
use std::fs;
fn create_test_model(dir: &Path, vocab_size: usize, dimensions: usize) {
let tokenizer_json = serde_json::json!({
"version": "1.0",
"truncation": null,
"padding": null,
"added_tokens": [
{
"id": 0,
"content": "[UNK]",
"single_word": false,
"lstrip": false,
"rstrip": false,
"normalized": false,
"special": true
}
],
"normalizer": {
"type": "Lowercase"
},
"pre_tokenizer": {
"type": "Whitespace"
},
"post_processor": null,
"decoder": null,
"model": {
"type": "WordLevel",
"vocab": create_test_vocab(vocab_size),
"unk_token": "[UNK]"
}
});
fs::write(
dir.join("tokenizer.json"),
serde_json::to_string_pretty(&tokenizer_json).unwrap(),
)
.unwrap();
create_test_safetensors(dir, vocab_size, dimensions);
}
fn create_tokenizer_parity_model(dir: &Path) {
let mut vocab = create_test_vocab(16)
.as_object()
.expect("test vocabulary is an object")
.clone();
vocab.insert("café".to_owned(), serde_json::Value::from(11));
let tokenizer_json = serde_json::json!({
"version": "1.0",
"truncation": {
"direction": "Right",
"max_length": 512,
"strategy": "LongestFirst",
"stride": 0
},
"padding": null,
"added_tokens": [
{
"id": 0,
"content": "[UNK]",
"single_word": false,
"lstrip": false,
"rstrip": false,
"normalized": false,
"special": true
},
{
"id": 12,
"content": "<added>",
"single_word": false,
"lstrip": false,
"rstrip": false,
"normalized": true,
"special": false
},
{
"id": 13,
"content": "[SPECIAL]",
"single_word": false,
"lstrip": false,
"rstrip": false,
"normalized": false,
"special": true
}
],
"normalizer": {
"type": "Lowercase"
},
"pre_tokenizer": {
"type": "Whitespace"
},
"post_processor": null,
"decoder": null,
"model": {
"type": "WordLevel",
"vocab": vocab,
"unk_token": "[UNK]"
}
});
fs::write(
dir.join("tokenizer.json"),
serde_json::to_string_pretty(&tokenizer_json).unwrap(),
)
.unwrap();
create_test_safetensors(dir, 16, 256);
}
fn create_test_vocab(vocab_size: usize) -> serde_json::Value {
let mut vocab = serde_json::Map::new();
vocab.insert("[UNK]".to_owned(), serde_json::Value::from(0));
let test_words = [
"hello", "world", "test", "rust", "search", "embed", "vector", "model", "fast", "query",
];
for (i, word) in test_words.iter().enumerate() {
if i + 1 < vocab_size {
vocab.insert((*word).to_owned(), serde_json::Value::from(i + 1));
}
}
serde_json::Value::Object(vocab)
}
fn create_test_safetensors(dir: &Path, vocab_size: usize, dimensions: usize) {
use std::collections::HashMap;
let mut data = Vec::with_capacity(vocab_size * dimensions * 4);
for row in 0..vocab_size {
for col in 0..dimensions {
#[allow(clippy::cast_precision_loss)]
let val = (row as f32).mul_add(0.1, (col as f32) * 0.01);
data.extend_from_slice(&val.to_le_bytes());
}
}
let mut tensors = HashMap::new();
tensors.insert(
"embeddings".to_owned(),
safetensors::tensor::TensorView::new(
safetensors::Dtype::F32,
vec![vocab_size, dimensions],
&data,
)
.unwrap(),
);
let serialized = safetensors::tensor::serialize(&tensors, None).unwrap();
fs::write(dir.join("model.safetensors"), serialized).unwrap();
}
#[test]
fn load_valid_model() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let embedder = Model2VecEmbedder::load_with_name(dir.path(), "test-model").unwrap();
assert_eq!(embedder.dimensions, 8);
assert_eq!(embedder.vocab_size, 12);
assert_eq!(embedder.name, "test-model");
}
#[test]
fn load_preverified_rejects_tensor_dimension_drift() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let mut identity = EmbeddingIdentityBundleV1::explicit_test_model("dimension-drift", 7);
identity.producer.backend = "attested-fixture-backend".to_owned();
identity.validate().unwrap();
let error = Model2VecEmbedder::load_preverified(dir.path(), "dimension-drift", identity)
.expect_err("parsed tensor width must agree with the attested identity");
assert!(matches!(error, SearchError::ModelLoadFailed { .. }));
}
#[test]
fn registered_identity_matches_native_execution_contract() {
let identity = ModelArtifactManifestV1::potion_128m_native()
.unwrap()
.declared_identity_bundle(QuantizationFormat::F32, "in-memory-f32-v1")
.unwrap();
validate_registered_execution_contract(&identity).unwrap();
let mut drifted = identity;
drifted.space.model_preprocessing.push_str("-drift");
assert!(validate_registered_execution_contract(&drifted).is_err());
}
#[test]
fn embed_batch_sync_matches_serial_across_parallel_boundary() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let embedder = Model2VecEmbedder::load_with_name(dir.path(), "test-model").unwrap();
for &batch_size in &[0_usize, 1, PARALLEL_BATCH_MIN - 1, PARALLEL_BATCH_MIN, 17] {
let docs: Vec<String> = (0..batch_size)
.map(|i| format!("hello world test rust search {i}"))
.collect();
let texts: Vec<&str> = docs.iter().map(String::as_str).collect();
let serial: Vec<Vec<f32>> = texts
.iter()
.map(|t| former_embed_sync(&embedder, t))
.collect();
let batched = embedder.embed_batch_sync(&texts).unwrap();
assert_eq!(batched.len(), serial.len(), "len at n={batch_size}");
for (index, (batched, former)) in batched.iter().zip(&serial).enumerate() {
assert_f32_bits_eq(
batched,
former,
&format!("batch order or output diverged at n={batch_size}, index={index}"),
);
}
}
}
#[test]
fn load_missing_tokenizer() {
let dir = tempfile::tempdir().unwrap();
create_test_safetensors(dir.path(), 10, 4);
let result = Model2VecEmbedder::load(dir.path());
assert!(result.is_err());
let err = result.unwrap_err();
assert!(
matches!(err, SearchError::ModelNotFound { .. }),
"expected ModelNotFound, got {err:?}"
);
}
#[test]
fn load_missing_safetensors() {
let dir = tempfile::tempdir().unwrap();
fs::write(dir.path().join("tokenizer.json"), "{}").unwrap();
let result = Model2VecEmbedder::load(dir.path());
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
SearchError::ModelNotFound { .. }
));
}
#[test]
fn load_nonexistent_directory() {
let result = Model2VecEmbedder::load("/nonexistent/path/to/model");
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
SearchError::ModelNotFound { .. }
));
}
#[test]
fn embed_produces_correct_dimension() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let embedder = Model2VecEmbedder::load(dir.path()).unwrap();
let vec = embedder.embed_sync("hello world").unwrap();
assert_eq!(vec.len(), 8);
}
#[test]
fn embed_output_is_l2_normalized() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let embedder = Model2VecEmbedder::load(dir.path()).unwrap();
let vec = embedder.embed_sync("hello world").unwrap();
let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-5, "expected unit norm, got {norm}");
}
#[test]
fn embed_empty_string_returns_zero_vector() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let embedder = Model2VecEmbedder::load(dir.path()).unwrap();
let vec = embedder.embed_sync("").unwrap();
assert_eq!(vec.len(), 8);
assert!(vec.iter().all(|&x| x == 0.0));
}
#[test]
fn embed_deterministic() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let embedder = Model2VecEmbedder::load(dir.path()).unwrap();
let a = embedder.embed_sync("hello world").unwrap();
let b = embedder.embed_sync("hello world").unwrap();
assert_eq!(a, b, "same input must produce same output");
}
#[test]
fn embed_different_inputs_different_outputs() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let embedder = Model2VecEmbedder::load(dir.path()).unwrap();
let a = embedder.embed_sync("hello").unwrap();
let b = embedder.embed_sync("world").unwrap();
assert_ne!(a, b, "different inputs should produce different embeddings");
}
#[test]
fn embed_sync_observes_offset_free_tokenizer_route() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let embedder = Model2VecEmbedder::load(dir.path()).unwrap();
let former = embedder.tokenizer.encode("hello world", false).unwrap();
assert!(
former.get_tokens().iter().any(|token| !token.is_empty()),
"the former encode route must materialize token text for this planted route witness"
);
embedder.embed_sync("hello world").unwrap();
assert!(
embedder.last_tokenizer_route_was_offset_free(),
"shipping Model2Vec must retain the encode_fast offset-free tokenizer route"
);
}
#[test]
fn encode_fast_token_ids_and_vectors_match_former_oracle() {
let dir = tempfile::tempdir().unwrap();
create_tokenizer_parity_model(dir.path());
let embedder = Model2VecEmbedder::load(dir.path()).unwrap();
let mut inputs = crate::model_manifest::MODEL_CONFORMANCE_TEXTS_V1
.iter()
.map(|text| (*text).to_owned())
.collect::<Vec<_>>();
inputs.extend([
String::new(),
"HELLO world".to_owned(),
"CAFÉ cafe\u{301}".to_owned(),
"hello 東京 world".to_owned(),
"definitely-oov-token".to_owned(),
"hello <ADDED> [SPECIAL] world".to_owned(),
]);
for tokens in [511_usize, 512, 513] {
inputs.push(
std::iter::repeat_n("hello", tokens)
.collect::<Vec<_>>()
.join(" "),
);
}
for input in &inputs {
assert_fast_ids_and_former_vector_bits(&embedder, input, input);
}
let special_ids = former_token_ids(&embedder, "hello <ADDED> [SPECIAL] world");
assert!(
special_ids.contains(&12) && special_ids.contains(&13),
"fixture must exercise both the normalized added token and the added special token"
);
for tokens in [511_usize, 512, 513] {
let input = std::iter::repeat_n("hello", tokens)
.collect::<Vec<_>>()
.join(" ");
let expected_len = tokens.min(512);
let former_ids = former_token_ids(&embedder, &input);
let fast_ids = embedder
.tokenizer
.encode_fast(input.as_str(), false)
.unwrap()
.get_ids()
.to_vec();
assert_eq!(
former_ids.len(),
expected_len,
"former tokenizer truncation at {tokens} input tokens"
);
assert_eq!(
fast_ids.len(),
expected_len,
"fast tokenizer truncation at {tokens} input tokens"
);
}
let texts = inputs.iter().map(String::as_str).collect::<Vec<_>>();
let former = texts
.iter()
.map(|text| former_embed_sync(&embedder, text))
.collect::<Vec<_>>();
let batched = embedder.embed_batch_sync(&texts).unwrap();
assert_eq!(batched.len(), former.len());
for (index, (actual, expected)) in batched.iter().zip(&former).enumerate() {
assert_f32_bits_eq(
actual,
expected,
&format!("batch order or vector bits changed at index={index}"),
);
}
}
fn former_embed_sync(embedder: &Model2VecEmbedder, text: &str) -> Vec<f32> {
if text.is_empty() {
return vec![0.0; embedder.dimensions];
}
let token_ids = former_token_ids(embedder, text);
former_embed_token_ids(embedder, &token_ids)
}
fn former_embed_token_ids(embedder: &Model2VecEmbedder, token_ids: &[u32]) -> Vec<f32> {
let mut sum = vec![0.0_f32; embedder.dimensions];
let mut count = 0_usize;
for &token_id in token_ids {
let index = token_id as usize;
if index < embedder.vocab_size {
let start = index * embedder.dimensions;
crate::simd::accumulate_f32_into(
&mut sum,
&embedder.embeddings[start..start + embedder.dimensions],
);
count += 1;
}
}
if count == 0 {
return vec![0.0; embedder.dimensions];
}
finish_mean_pool_and_normalize_former(&mut sum, count);
sum
}
fn former_token_ids(embedder: &Model2VecEmbedder, text: &str) -> Vec<u32> {
embedder
.tokenizer
.encode(text, false)
.unwrap()
.get_ids()
.to_vec()
}
fn assert_fast_ids_and_former_vector_bits(
embedder: &Model2VecEmbedder,
text: &str,
scenario: &str,
) {
let former_ids = former_token_ids(embedder, text);
let fast_ids = embedder
.tokenizer
.encode_fast(text, false)
.unwrap()
.get_ids()
.to_vec();
assert_eq!(fast_ids, former_ids, "token IDs diverged for {scenario}");
let expected = former_embed_sync(embedder, text);
let actual = embedder.embed_sync(text).unwrap();
assert_f32_bits_eq(&actual, &expected, scenario);
}
fn assert_f32_bits_eq(actual: &[f32], expected: &[f32], scenario: &str) {
assert_eq!(
actual
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>(),
expected
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>(),
"{scenario}"
);
}
fn expected_native_256_route(token_count: usize) -> Model2VecAccumulationRoute {
#[cfg(target_arch = "x86_64")]
{
if token_count < 512 {
if std::is_x86_feature_detected!("avx2") {
Model2VecAccumulationRoute::Native256ShortAvx2
} else {
Model2VecAccumulationRoute::Base
}
} else {
Model2VecAccumulationRoute::Prefetched
}
}
#[cfg(not(target_arch = "x86_64"))]
{
let _ = token_count;
Model2VecAccumulationRoute::Base
}
}
#[test]
fn native_256_embed_sync_matches_former_pool_and_finish_bits() {
const DIMENSIONS: usize = 256;
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, DIMENSIONS);
let mut embedder = Model2VecEmbedder::load(dir.path()).unwrap();
for &tokens in &[0_usize, 1, 2, 3, 4, 8, 16, 32, 64, 511, 512, 513] {
let text = (0..tokens)
.map(|position| match position % 4 {
0 => "hello",
1 => "world",
2 => "missing-token",
_ => "hello",
})
.collect::<Vec<_>>()
.join(" ");
let expected = former_embed_sync(&embedder, &text);
let actual = embedder.embed_sync(&text).unwrap();
assert_f32_bits_eq(&actual, &expected, &format!("tokens={tokens}"));
assert_eq!(
embedder
.tokenizer
.encode_fast(text.as_str(), false)
.unwrap()
.get_ids(),
former_token_ids(&embedder, &text),
"token IDs at native-256 boundary tokens={tokens}"
);
if !text.is_empty() {
let token_count = embedder
.tokenizer
.encode(text.as_str(), false)
.unwrap()
.len();
assert_eq!(
last_model2vec_accumulation_route_for_test(),
expected_native_256_route(token_count),
"shipping embed_sync route for {tokens} input words ({token_count} token IDs)"
);
}
}
for text in [
"hello caf\u{e9} world",
"hello \u{6771}\u{4eac} hello",
"HELLO hello HELLO",
] {
let expected = former_embed_sync(&embedder, text);
let actual = embedder.embed_sync(text).unwrap();
assert_f32_bits_eq(&actual, &expected, text);
}
let hello = DIMENSIONS..DIMENSIONS * 2;
embedder.embeddings[hello.clone()].fill(-0.0);
let expected = former_embed_sync(&embedder, "hello hello");
let actual = embedder.embed_sync("hello hello").unwrap();
assert_f32_bits_eq(&actual, &expected, "signed-zero row");
embedder.embeddings[hello.clone()].fill(1.0e-20);
let expected = former_embed_sync(&embedder, "hello");
let actual = embedder.embed_sync("hello").unwrap();
assert_f32_bits_eq(&actual, &expected, "below normalization guard");
embedder.embeddings[hello.clone()].fill(1.0e-4);
let expected = former_embed_sync(&embedder, "hello");
let actual = embedder.embed_sync("hello").unwrap();
assert_f32_bits_eq(&actual, &expected, "above normalization guard");
embedder.embeddings[hello].fill(f32::NAN);
let expected = former_embed_sync(&embedder, "hello world hello");
let actual = embedder.embed_sync("hello world hello").unwrap();
assert_f32_bits_eq(&actual, &expected, "non-finite pooled row");
}
#[test]
fn fused_mean_and_ordered_norm_finish_matches_former_bits() {
let arbitrary_initial_sum = (0_u32..256)
.map(|index| {
let value = f32::from_bits(0x3f80_0000 + index);
if index % 2 == 0 { value } else { -value }
})
.collect::<Vec<_>>();
let cases = vec![
("arbitrary finite sum", arbitrary_initial_sum),
(
"normalization guard boundary",
vec![f32::MIN_POSITIVE, f32::EPSILON.sqrt(), -f32::EPSILON.sqrt()],
),
("signed zero", vec![-0.0, 0.0, -0.0, 0.0]),
(
"non-finite values",
vec![f32::NAN, f32::INFINITY, f32::NEG_INFINITY, f32::MAX],
),
];
for (label, initial_sum) in cases {
for &count in &[1_usize, 2, 3, 4, 511, 512, 513] {
let mut former = initial_sum.clone();
let mut fused = initial_sum.clone();
finish_mean_pool_and_normalize_former(&mut former, count);
finish_mean_pool_and_normalize(&mut fused, count);
assert_f32_bits_eq(&fused, &former, &format!("{label}, count={count}"));
}
}
}
#[test]
fn full_embed_sync_fused_finish_matches_former_bits_across_shapes_and_values() {
for &dimensions in &[1_usize, 255, 256, 257] {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, dimensions);
let mut embedder = Model2VecEmbedder::load(dir.path()).unwrap();
let mut full_embed_sync_corpus =
vec![String::new(), "hello definitely-oov-token hello".to_owned()];
for token_count in [1_usize, 2, 3, 4, 511, 512, 513] {
full_embed_sync_corpus.push(
std::iter::repeat_n("hello", token_count)
.collect::<Vec<_>>()
.join(" "),
);
}
let oov_grouping_ids = former_token_ids(&embedder, &full_embed_sync_corpus[1]);
assert!(
oov_grouping_ids.contains(&0),
"the full embed_sync OOV grouping corpus must emit the [UNK] token at dim={dimensions}"
);
let hello = dimensions..dimensions * 2;
for (lane, value) in embedder.embeddings[hello.clone()].iter_mut().enumerate() {
let lane = u32::try_from(lane).expect("fixture dimension fits u32");
let magnitude = f32::from_bits(0x3f80_0000 + lane);
*value = if lane % 2 == 0 { magnitude } else { -magnitude };
}
for text in &full_embed_sync_corpus {
assert_fast_ids_and_former_vector_bits(
&embedder,
text,
&format!("finite full embed_sync corpus dim={dimensions}"),
);
}
#[allow(clippy::cast_precision_loss)]
let guard_center = (f32::EPSILON / dimensions as f32).sqrt();
let below_guard = f32::from_bits(guard_center.to_bits() - 1);
let above_guard = f32::from_bits(guard_center.to_bits() + 1);
for (scenario, value) in [
("subnormal", f32::from_bits(1)),
("below guard", below_guard),
("above guard", above_guard),
("signed zero", -0.0),
("NaN", f32::from_bits(0x7fc0_0001)),
("infinity", f32::INFINITY),
] {
embedder.embeddings[hello.clone()].fill(value);
for text in &full_embed_sync_corpus {
assert_fast_ids_and_former_vector_bits(
&embedder,
text,
&format!("{scenario} full embed_sync corpus dim={dimensions}"),
);
}
}
let invalid_token_ids = [
1_u32,
u32::try_from(embedder.vocab_size).expect("fixture vocabulary fits u32"),
u32::MAX,
2_u32,
];
let former = former_embed_token_ids(&embedder, &invalid_token_ids);
let fused = embedder.embed_token_ids(&invalid_token_ids);
assert_f32_bits_eq(
&fused,
&former,
&format!("mixed valid/OOV token IDs dim={dimensions}"),
);
let all_oov_token_ids = [
u32::try_from(embedder.vocab_size).expect("fixture vocabulary fits u32"),
u32::MAX,
];
let former = former_embed_token_ids(&embedder, &all_oov_token_ids);
let fused = embedder.embed_token_ids(&all_oov_token_ids);
assert_f32_bits_eq(
&fused,
&former,
&format!("all OOV token IDs dim={dimensions}"),
);
}
}
#[test]
fn embed_all_oov_returns_zero_vector() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let embedder = Model2VecEmbedder::load(dir.path()).unwrap();
let vec = embedder.embed_sync("xyzxyzxyz qqqqq").unwrap();
assert_eq!(vec.len(), 8);
}
#[test]
fn trait_is_semantic() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let embedder = Model2VecEmbedder::load(dir.path()).unwrap();
assert!(embedder.is_semantic());
}
#[test]
fn trait_category_is_static() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let embedder = Model2VecEmbedder::load(dir.path()).unwrap();
assert_eq!(embedder.category(), ModelCategory::StaticEmbedder);
}
#[test]
fn trait_dimension() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let embedder = Model2VecEmbedder::load(dir.path()).unwrap();
assert_eq!(embedder.dimension(), 8);
}
#[test]
fn trait_does_not_infer_mrl_from_model2vec_backend() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let embedder = Model2VecEmbedder::load(dir.path()).unwrap();
assert!(!embedder.supports_mrl());
}
#[test]
fn trait_id_and_name() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let embedder = Model2VecEmbedder::load_with_name(dir.path(), "my-model").unwrap();
assert_eq!(embedder.id(), "my-model");
assert_eq!(embedder.model_name(), "my-model");
}
#[test]
fn embedder_is_send_sync() {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<Model2VecEmbedder>();
}
#[test]
#[ignore = "requires a verified potion model dir via POTION_FIXTURE_DIR"]
fn conformance_certificate_matches_fixture() {
let dir = std::env::var("POTION_FIXTURE_DIR")
.expect("set POTION_FIXTURE_DIR to a potion-multilingual-128M directory");
let manifest = crate::model_manifest::ModelArtifactManifestV1::potion_128m_native()
.expect("registered potion manifest");
let verified = manifest
.verify_dir(Path::new(&dir))
.expect("verify frozen potion artifacts");
let expected_identity = verified
.identity_bundle(QuantizationFormat::F32, "in-memory-f32-v1")
.expect("derive verified potion identity");
let embedder = Model2VecEmbedder::load_preverified(
Path::new(&dir),
DEFAULT_MODEL_NAME,
expected_identity.clone(),
)
.expect("load verified potion embedder");
assert_eq!(embedder.identity().unwrap(), &expected_identity);
let added_vocabulary = embedder.tokenizer.get_added_vocabulary().get_vocab();
let pad_id = *added_vocabulary
.get("[PAD]")
.expect("verified Potion tokenizer must retain [PAD] as an added special token");
let unk_id = *added_vocabulary
.get("[UNK]")
.expect("verified Potion tokenizer must retain [UNK] as an added special token");
let added_special_text = "hello [PAD] [UNK] world";
let long_over_512_text = std::iter::repeat_n("hello", 1024)
.collect::<Vec<_>>()
.join(" ");
let mut parity_texts = crate::model_manifest::MODEL_CONFORMANCE_TEXTS_V1
.iter()
.map(|text| (*text).to_owned())
.collect::<Vec<_>>();
parity_texts.extend([
"Caf\u{e9} na\u{ef}ve \u{2014} \u{6771}\u{4eac} \u{1f980}".to_owned(),
added_special_text.to_owned(),
"\tmetaspace boundaries\tand unseen-oov-\u{10ffff}".to_owned(),
long_over_512_text.clone(),
]);
for text in &parity_texts {
assert_fast_ids_and_former_vector_bits(
&embedder,
text,
"verified Potion tokenizer parity input",
);
}
let added_special_ids = former_token_ids(&embedder, added_special_text);
assert!(
added_special_ids.contains(&pad_id) && added_special_ids.contains(&unk_id),
"verified Potion tokenizer must emit both literal added special-token IDs"
);
let former_long = embedder
.tokenizer
.encode(long_over_512_text.as_str(), false)
.expect("encode long verified Potion input");
let fast_long = embedder
.tokenizer
.encode_fast(long_over_512_text.as_str(), false)
.expect("encode_fast long verified Potion input");
assert!(
embedder.tokenizer.get_truncation().is_none(),
"registered Potion tokenizer must preserve its configured no-truncation policy"
);
assert!(
former_long.len() > 512,
"long verified Potion input must exceed the former 512-token synthetic boundary"
);
assert_eq!(
fast_long.len(),
former_long.len(),
"verified Potion long-input token count diverged"
);
assert_eq!(
fast_long.get_ids(),
former_long.get_ids(),
"verified Potion long-input token IDs diverged"
);
assert!(
former_long.get_overflowing().is_empty() && fast_long.get_overflowing().is_empty(),
"configured no-truncation policy must not emit overflow encodings"
);
let texts = &crate::model_manifest::MODEL_CONFORMANCE_TEXTS_V1;
let former_vectors = texts
.iter()
.map(|text| {
assert_fast_ids_and_former_vector_bits(
&embedder,
text,
"verified Potion conformance input",
);
former_embed_sync(&embedder, text)
})
.collect::<Vec<_>>();
let vectors = embedder
.embed_batch_sync(texts)
.expect("embed bounded conformance corpus");
assert_eq!(vectors.len(), former_vectors.len());
for (index, (actual, expected)) in vectors.iter().zip(&former_vectors).enumerate() {
assert_f32_bits_eq(
actual,
expected,
&format!("verified Potion batch order or vector bits changed at index={index}"),
);
}
let observed = frankensearch_core::generation::GoldenVectorCertificateV1::from_exact_f32(
texts, &vectors,
)
.expect("compute exact conformance certificate");
let expected = manifest.execution.golden_vectors;
assert_eq!(
observed, expected,
"Model2Vec output bits drifted from the registered producer certificate"
);
}
#[test]
fn debug_does_not_dump_embeddings() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 12, 8);
let embedder = Model2VecEmbedder::load(dir.path()).unwrap();
let debug = format!("{embedder:?}");
assert!(debug.contains("Model2VecEmbedder"));
assert!(debug.contains("dimensions: 8"));
assert!(debug.contains("vocab_size: 12"));
assert!(!debug.contains("0.1"));
}
#[test]
fn tensor_discovery_finds_standard_name() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 4, 2);
let embedder = Model2VecEmbedder::load(dir.path()).unwrap();
assert_eq!(embedder.vocab_size, 4);
}
#[test]
fn tensor_discovery_single_tensor_fallback() {
let dir = tempfile::tempdir().unwrap();
let tokenizer_json = serde_json::json!({
"version": "1.0",
"added_tokens": [],
"model": {
"type": "WordLevel",
"vocab": {"hello": 0, "world": 1},
"unk_token": "hello"
}
});
fs::write(
dir.path().join("tokenizer.json"),
serde_json::to_string(&tokenizer_json).unwrap(),
)
.unwrap();
let mut data = vec![0u8; 2 * 3 * 4]; for (i, chunk) in data.as_chunks_mut::<4>().0.iter_mut().enumerate() {
#[allow(clippy::cast_precision_loss)]
let val = i as f32;
chunk.copy_from_slice(&val.to_le_bytes());
}
let mut tensors = std::collections::HashMap::new();
tensors.insert(
"my_custom_tensor_name".to_owned(),
safetensors::tensor::TensorView::new(safetensors::Dtype::F32, vec![2, 3], &data)
.unwrap(),
);
let serialized = safetensors::tensor::serialize(&tensors, None).unwrap();
fs::write(dir.path().join("model.safetensors"), serialized).unwrap();
let embedder = Model2VecEmbedder::load(dir.path()).unwrap();
assert_eq!(embedder.vocab_size, 2);
assert_eq!(embedder.dimensions, 3);
}
#[test]
fn has_required_files_positive() {
let dir = tempfile::tempdir().unwrap();
create_test_model(dir.path(), 4, 2);
assert!(has_required_files(dir.path()));
}
#[test]
fn has_required_files_negative() {
let dir = tempfile::tempdir().unwrap();
assert!(!has_required_files(dir.path()));
}
#[test]
fn parse_f32_matrix_correct() {
let data: Vec<u8> = [1.0_f32, 2.0, 3.0, 4.0]
.iter()
.flat_map(|f| f.to_le_bytes())
.collect();
let matrix = parse_f32_matrix(&data, 2, 2).unwrap();
assert_eq!(matrix.len(), 4);
assert_eq!(&matrix[0..2], &[1.0, 2.0]);
assert_eq!(&matrix[2..4], &[3.0, 4.0]);
}
#[test]
fn parse_f32_matrix_too_short() {
let data = vec![0u8; 4]; let result = parse_f32_matrix(&data, 2, 2);
assert!(result.is_err());
}
#[test]
fn parse_f32_matrix_too_long() {
let mut data: Vec<u8> = [1.0_f32, 2.0, 3.0, 4.0]
.iter()
.flat_map(|f| f.to_le_bytes())
.collect();
data.push(0xAA);
let result = parse_f32_matrix(&data, 2, 2);
assert!(result.is_err());
}
#[test]
fn parse_f32_matrix_size_overflow() {
let result = parse_f32_matrix(&[], usize::MAX, 2);
assert!(result.is_err());
}
}