use std::path::Path;
use std::sync::Mutex;
use tokenizers::Tokenizer;
use frankensearch_core::error::{SearchError, SearchResult};
use frankensearch_core::generation::{EmbeddingIdentityBundleV1, QuantizationFormat};
use frankensearch_core::traits::{ModelCategory, SyncEmbed};
use frankensearch_embed::model_manifest::ModelArtifactManifestV1;
use crate::native::{
DEFAULT_MAX_LENGTH, Model, SAFETENSORS_FALLBACK, TOKENIZER_JSON, build_model, parse_weights,
};
const DEFAULT_MODEL_NAME: &str = "all-minilm-l6-v2";
const DEFAULT_EMBEDDER_ID: &str = "minilm-384-native";
const MULTILINGUAL_MODEL_NAME: &str = "paraphrase-multilingual-minilm-l12-v2";
const MULTILINGUAL_EMBEDDER_ID: &str = "paraphrase-multilingual-minilm-l12-v2-384-native";
const DIM: usize = 384;
const IDENTITY_DIMENSION: u32 = 384;
const IDENTITY_SEQUENCE_POLICY: &str = "max-length=512;longest-first;no-padding";
const IDENTITY_POOLING: &str = "mean-all-returned-tokens-including-specials-no-padding-v1";
const IDENTITY_OUTPUT_NORMALIZATION: &str = "l2-f32-if-norm-gt-zero-else-unchanged-v1";
const MAX_BATCH_TOKENS: usize = 2048;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum NativeEmbeddingModel {
AllMiniLmL6V2,
ParaphraseMultilingualMiniLmL12V2,
}
impl NativeEmbeddingModel {
const fn model_name(self) -> &'static str {
match self {
Self::AllMiniLmL6V2 => DEFAULT_MODEL_NAME,
Self::ParaphraseMultilingualMiniLmL12V2 => MULTILINGUAL_MODEL_NAME,
}
}
const fn embedder_id(self) -> &'static str {
match self {
Self::AllMiniLmL6V2 => DEFAULT_EMBEDDER_ID,
Self::ParaphraseMultilingualMiniLmL12V2 => MULTILINGUAL_EMBEDDER_ID,
}
}
const fn encoder_layers(self) -> usize {
match self {
Self::AllMiniLmL6V2 => 6,
Self::ParaphraseMultilingualMiniLmL12V2 => 12,
}
}
fn manifest(self) -> SearchResult<ModelArtifactManifestV1> {
match self {
Self::AllMiniLmL6V2 => ModelArtifactManifestV1::minilm_native_frankentorch(),
Self::ParaphraseMultilingualMiniLmL12V2 => {
ModelArtifactManifestV1::multilingual_minilm_native_frankentorch()
}
}
}
}
pub struct NativeEmbedder {
inner: Mutex<Model>,
tokenizer: Tokenizer,
max_length: usize,
name: String,
id: String,
identity: EmbeddingIdentityBundleV1,
}
impl std::fmt::Debug for NativeEmbedder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("NativeEmbedder")
.field("name", &self.name)
.field("max_length", &self.max_length)
.finish_non_exhaustive()
}
}
impl NativeEmbedder {
pub fn load(model_dir: impl AsRef<Path>) -> SearchResult<Self> {
Self::load_model(model_dir, NativeEmbeddingModel::AllMiniLmL6V2)
}
pub fn load_multilingual(model_dir: impl AsRef<Path>) -> SearchResult<Self> {
Self::load_model(
model_dir,
NativeEmbeddingModel::ParaphraseMultilingualMiniLmL12V2,
)
}
pub fn load_model(
model_dir: impl AsRef<Path>,
profile: NativeEmbeddingModel,
) -> SearchResult<Self> {
let dir = model_dir.as_ref();
let model_name = profile.model_name();
let verified = profile.manifest()?.verify_dir(dir)?;
let identity = verified.identity_bundle(QuantizationFormat::F32, "in-memory-f32-v1")?;
if identity.space.dimension != IDENTITY_DIMENSION {
return Err(SearchError::ModelLoadFailed {
path: dir.to_path_buf(),
source: format!(
"registered dimension {} disagrees with native backend dimension {DIM}",
identity.space.dimension
)
.into(),
});
}
for (field, actual, expected) in [
(
"sequence policy",
identity.space.sequence_policy.as_str(),
IDENTITY_SEQUENCE_POLICY,
),
("pooling", identity.space.pooling.as_str(), IDENTITY_POOLING),
(
"output normalization",
identity.space.output_normalization.as_str(),
IDENTITY_OUTPUT_NORMALIZATION,
),
] {
if actual != expected {
return Err(SearchError::ModelLoadFailed {
path: dir.to_path_buf(),
source: format!(
"registered {field} disagrees with the native backend contract"
)
.into(),
});
}
}
let tok_path = dir.join(TOKENIZER_JSON);
if !tok_path.is_file() {
return Err(SearchError::ModelNotFound {
name: format!(
"{model_name} (missing {TOKENIZER_JSON} in {})",
dir.display()
),
});
}
let mut tokenizer =
Tokenizer::from_file(&tok_path).map_err(|e| SearchError::ModelLoadFailed {
path: tok_path.clone(),
source: format!("tokenizer load failed: {e}").into(),
})?;
tokenizer
.with_truncation(Some(tokenizers::TruncationParams {
max_length: DEFAULT_MAX_LENGTH,
..Default::default()
}))
.map_err(|e| SearchError::ModelLoadFailed {
path: tok_path.clone(),
source: format!("failed to enable truncation: {e}").into(),
})?;
tokenizer.with_padding(None);
let weights_path = dir.join(SAFETENSORS_FALLBACK);
if !weights_path.is_file() {
return Err(SearchError::ModelNotFound {
name: format!(
"{model_name} (missing verified {SAFETENSORS_FALLBACK} in {})",
dir.display()
),
});
}
let shared = parse_weights(&weights_path)?;
let model = build_model(shared)?;
if model.encoder_layers() != profile.encoder_layers() {
return Err(SearchError::ModelLoadFailed {
path: weights_path,
source: format!(
"registered model requires {} encoder layers, weights contain {}",
profile.encoder_layers(),
model.encoder_layers()
)
.into(),
});
}
tracing::info!(
model = model_name,
dimension = DIM,
encoder_layers = model.encoder_layers(),
max_length = DEFAULT_MAX_LENGTH,
manifest = %verified.frozen().fingerprint,
identity = %identity.fingerprint(),
"native frankentorch MiniLM embedder loaded (int8 linear, mean-pool + L2)"
);
Ok(Self {
inner: Mutex::new(model),
tokenizer,
max_length: DEFAULT_MAX_LENGTH,
name: model_name.to_owned(),
id: profile.embedder_id().to_owned(),
identity,
})
}
fn tokenize(&self, text: &str) -> SearchResult<Vec<i64>> {
let encoding =
self.tokenizer
.encode(text, true)
.map_err(|e| SearchError::EmbeddingFailed {
model: self.name.clone(),
source: format!("tokenize failed: {e}").into(),
})?;
Ok(crate::ids_to_truncated_i64(
encoding.get_ids(),
self.max_length,
))
}
fn lock_model(&self) -> SearchResult<std::sync::MutexGuard<'_, Model>> {
self.inner.lock().map_err(|e| SearchError::EmbeddingFailed {
model: self.name.clone(),
source: format!("embedder mutex poisoned: {e}").into(),
})
}
}
impl SyncEmbed for NativeEmbedder {
fn embed_sync(&self, text: &str) -> SearchResult<Vec<f32>> {
let ids = self.tokenize(text)?;
let mut model = self.lock_model()?;
let mut out = model.embed_forward(&[ids])?;
drop(model);
let vector = out.pop().ok_or_else(|| SearchError::EmbeddingFailed {
model: self.name.clone(),
source: "native backend returned no embedding".into(),
})?;
if vector.len() != DIM {
return Err(SearchError::EmbeddingFailed {
model: self.name.clone(),
source: format!(
"native backend returned dimension {}, expected {DIM}",
vector.len()
)
.into(),
});
}
Ok(vector)
}
fn embed_batch_sync(&self, texts: &[&str]) -> SearchResult<Vec<Vec<f32>>> {
if texts.is_empty() {
return Ok(Vec::new());
}
let token_batches: Vec<Vec<i64>> = texts
.iter()
.map(|t| self.tokenize(t))
.collect::<SearchResult<_>>()?;
let mut model = self.lock_model()?;
let mut out = Vec::with_capacity(texts.len());
let mut start = 0usize;
while start < token_batches.len() {
let mut end = start;
let mut tok = 0usize;
while end < token_batches.len() {
let len = token_batches[end].len().max(1);
if end > start && tok + len > MAX_BATCH_TOKENS {
break;
}
tok += len;
end += 1;
}
out.extend(model.embed_forward(&token_batches[start..end])?);
start = end;
}
drop(model);
if out.len() != texts.len() || out.iter().any(|vector| vector.len() != DIM) {
return Err(SearchError::EmbeddingFailed {
model: self.name.clone(),
source:
"native backend returned a batch shape inconsistent with its attested identity"
.into(),
});
}
Ok(out)
}
fn dimension(&self) -> usize {
DIM
}
fn identity(&self) -> SearchResult<&EmbeddingIdentityBundleV1> {
Ok(&self.identity)
}
fn id(&self) -> &str {
&self.id
}
fn model_name(&self) -> &str {
&self.name
}
fn is_semantic(&self) -> bool {
true
}
fn category(&self) -> ModelCategory {
ModelCategory::TransformerEmbedder
}
}
#[cfg(test)]
mod tests {
use super::*;
const fn assert_sync_embed<T: SyncEmbed>() {}
const _: () = assert_sync_embed::<NativeEmbedder>();
#[test]
fn registered_identity_matches_native_backend_contract() {
let identity = ModelArtifactManifestV1::minilm_native_frankentorch()
.expect("registered native MiniLM manifest")
.declared_identity_bundle(QuantizationFormat::F32, "in-memory-f32-v1")
.expect("derive native MiniLM identity");
assert_eq!(identity.space.dimension, IDENTITY_DIMENSION);
assert_eq!(identity.space.sequence_policy, IDENTITY_SEQUENCE_POLICY);
assert_eq!(identity.space.pooling, IDENTITY_POOLING);
assert_eq!(
identity.space.output_normalization,
IDENTITY_OUTPUT_NORMALIZATION
);
}
#[test]
fn multilingual_identity_is_distinct_from_same_dimension_minilm() {
let baseline = ModelArtifactManifestV1::minilm_native_frankentorch()
.expect("registered native MiniLM manifest")
.declared_identity_bundle(QuantizationFormat::F32, "in-memory-f32-v1")
.expect("derive native MiniLM identity");
let multilingual = ModelArtifactManifestV1::multilingual_minilm_native_frankentorch()
.expect("registered multilingual MiniLM manifest")
.declared_identity_bundle(QuantizationFormat::F32, "in-memory-f32-v1")
.expect("derive multilingual MiniLM identity");
assert_eq!(baseline.space.dimension, multilingual.space.dimension);
assert_ne!(
baseline.space.fingerprint(),
multilingual.space.fingerprint()
);
assert!(
baseline.verify_exact_producer_with(&multilingual).is_err(),
"same dimensionality must not admit vectors from a different model space"
);
}
#[test]
#[ignore = "requires a local all-MiniLM-L6-v2 model dir via MINILM_FIXTURE_DIR"]
fn embeds_unit_vector_from_fixture() {
let dir = std::env::var("MINILM_FIXTURE_DIR")
.expect("set MINILM_FIXTURE_DIR to an all-MiniLM-L6-v2 model directory");
let embedder = NativeEmbedder::load(&dir).expect("load native MiniLM embedder");
assert_eq!(embedder.dimension(), DIM);
let v = embedder.embed_sync("hello world").expect("embed");
assert_eq!(v.len(), DIM, "embedding dimensionality");
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!(
(norm - 1.0).abs() < 1e-3,
"expected L2-normalized unit vector, got norm {norm}"
);
let batch = embedder
.embed_batch_sync(&["hello world", "a second sentence"])
.expect("batch embed");
assert_eq!(batch.len(), 2);
assert_eq!(batch[0].len(), DIM);
let cos: f32 = v.iter().zip(&batch[0]).map(|(a, b)| a * b).sum();
assert!(
cos > 0.999,
"single vs batch embedding mismatch (cos {cos})"
);
}
#[test]
#[ignore = "requires a verified all-MiniLM-L6-v2 model dir via MINILM_FIXTURE_DIR"]
fn conformance_certificate_matches_fixture() {
let dir = std::env::var("MINILM_FIXTURE_DIR")
.expect("set MINILM_FIXTURE_DIR to an all-MiniLM-L6-v2 model directory");
let manifest = ModelArtifactManifestV1::minilm_native_frankentorch()
.expect("registered native MiniLM manifest");
let expected_identity = manifest
.declared_identity_bundle(QuantizationFormat::F32, "in-memory-f32-v1")
.expect("derive registered native MiniLM identity");
let embedder = NativeEmbedder::load(&dir).expect("load native MiniLM embedder");
assert_eq!(embedder.identity().unwrap(), &expected_identity);
let texts = &frankensearch_embed::model_manifest::MODEL_CONFORMANCE_TEXTS_V1;
let vectors = embedder
.embed_batch_sync(texts)
.expect("embed bounded conformance corpus");
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,
"native MiniLM output bits drifted from the registered producer certificate"
);
}
fn cosine(left: &[f32], right: &[f32]) -> f32 {
left.iter().zip(right).map(|(a, b)| a * b).sum()
}
#[test]
#[ignore = "requires a verified multilingual MiniLM model dir via MULTILINGUAL_MINILM_FIXTURE_DIR"]
fn multilingual_fixture_proves_cross_language_retrieval_and_determinism() {
let dir = std::env::var("MULTILINGUAL_MINILM_FIXTURE_DIR")
.expect("set MULTILINGUAL_MINILM_FIXTURE_DIR to paraphrase-multilingual-MiniLM-L12-v2");
let load_started = std::time::Instant::now();
let embedder = NativeEmbedder::load_multilingual(&dir)
.expect("load verified multilingual MiniLM embedder");
let load_elapsed = load_started.elapsed();
assert_eq!(embedder.dimension(), DIM);
assert_eq!(embedder.id(), MULTILINGUAL_EMBEDDER_ID);
assert_eq!(
embedder.lock_model().expect("lock model").encoder_layers(),
12
);
let chinese_ids = embedder
.tokenize("如何修复数据库事务死锁?")
.expect("tokenize Chinese query");
assert!(
chinese_ids.len() > 4,
"native multilingual tokenizer collapsed Chinese input"
);
let texts = [
"如何在 Rust 中处理任务取消和结构化并发?",
"In Rust, structured concurrency keeps child tasks scoped and propagates cancellation safely.",
"A sourdough starter needs flour, water, and a warm kitchen.",
"How should a database transaction deadlock be resolved?",
"数据库事务发生死锁时,应回滚其中一个事务,并按固定顺序重试锁操作。",
"这份食谱介绍如何烤制苹果派和准备奶油馅料。",
"修复 Rust async cancellation bug in worker_queue.rs",
"worker_queue.rs 必须在 async 任务取消时归还 reservation,避免消息丢失。",
"The watercolor landscape uses blue pigment and cold-press paper.",
];
let first_started = std::time::Instant::now();
let first = embedder
.embed_batch_sync(&texts)
.expect("embed multilingual retrieval fixture");
let first_elapsed = first_started.elapsed();
let repeat_started = std::time::Instant::now();
let second = embedder
.embed_batch_sync(&texts)
.expect("repeat multilingual retrieval fixture");
let repeat_elapsed = repeat_started.elapsed();
assert_eq!(
first, second,
"native multilingual output must be bit-exact"
);
for (query, relevant, distractor, label) in [
(0, 1, 2, "Chinese query to English discussion"),
(3, 4, 5, "English query to Chinese discussion"),
(6, 7, 8, "mixed Chinese/code query"),
] {
let relevant_score = cosine(&first[query], &first[relevant]);
let distractor_score = cosine(&first[query], &first[distractor]);
assert!(
relevant_score > distractor_score + 0.05,
"{label} failed: relevant={relevant_score}, distractor={distractor_score}"
);
}
let manifest = ModelArtifactManifestV1::multilingual_minilm_native_frankentorch()
.expect("registered multilingual MiniLM manifest");
let conformance_texts = &frankensearch_embed::model_manifest::MODEL_CONFORMANCE_TEXTS_V1;
let conformance_started = std::time::Instant::now();
let vectors = embedder
.embed_batch_sync(conformance_texts)
.expect("embed bounded conformance corpus");
let conformance_elapsed = conformance_started.elapsed();
let observed = frankensearch_core::generation::GoldenVectorCertificateV1::from_exact_f32(
conformance_texts,
&vectors,
)
.expect("compute exact multilingual conformance certificate");
assert_eq!(observed, manifest.execution.golden_vectors);
eprintln!(
"multilingual_native_metrics load_ms={} first_9_ms={} repeat_9_ms={} conformance_4_ms={}",
load_elapsed.as_millis(),
first_elapsed.as_millis(),
repeat_elapsed.as_millis(),
conformance_elapsed.as_millis()
);
}
}