diff --git a/embeddings/src/model/local.rs b/embeddings/src/model/local.rs index 67bab744..108a14c8 100644 --- a/embeddings/src/model/local.rs +++ b/embeddings/src/model/local.rs @@ -245,6 +245,8 @@ pub struct LocalModelInfo { pub gguf_path: Option, /// Path to ONNX file if using ONNX model pub onnx_path: Option, + /// sentence-transformers `1_Pooling/config.json`, fetched for ONNX models + pub pooling_config_path: Option, } fn model_download_lock(model_id: &str) -> Arc> { @@ -399,12 +401,21 @@ pub fn build_model_info( } }; + // Only the ONNX path pools token states itself, so only it needs the + // model's published pooling mode. + let pooling_config_path = if onnx_path.is_some() { + api.get(POOLING_CONFIG_FILE).ok() + } else { + None + }; + Ok(LocalModelInfo { config_path, tokenizer_path, weights_paths, gguf_path, onnx_path, + pooling_config_path, }) } @@ -455,12 +466,39 @@ fn try_find_gguf_file(api: &hf_hub::api::sync::ApiRepo) -> Option { api.get(&gguf_files[0]).ok() } -/// Try to find an ONNX model file in the repository -/// Looks for model.onnx at root or in onnx/ subdirectory +/// Try to find an ONNX model file in the repository. +/// Looks for model.onnx at root or in onnx/ subdirectory. External-data +/// exports keep tensors in sibling files ORT resolves relative to model.onnx, +/// so those are fetched too; a failed sidecar download fails the lookup. fn try_find_onnx_file(api: &hf_hub::api::sync::ApiRepo) -> Option { - api.get("model.onnx") - .ok() - .or_else(|| api.get("onnx/model.onnx").ok()) + for model_path in ["model.onnx", "onnx/model.onnx"] { + let Ok(onnx_path) = api.get(model_path) else { + continue; + }; + // Listing needs the network; a warm cache still loads offline. + if let Ok(repo_info) = api.info() { + let siblings: Vec = repo_info + .siblings + .iter() + .map(|s| s.rfilename.clone()) + .collect(); + for sidecar in onnx_sidecar_files(model_path, &siblings) { + api.get(&sidecar).ok()?; + } + } + return Some(onnx_path); + } + None +} + +/// External-data files of `model_path`: `model.onnx_data`, `model.onnx.data`, +/// `model.onnx_data_1`, ... share the model file name as prefix. +pub(crate) fn onnx_sidecar_files(model_path: &str, siblings: &[String]) -> Vec { + siblings + .iter() + .filter(|file| file.starts_with(model_path) && file.as_str() != model_path) + .cloned() + .collect() } /// Load tokenizer with fallback for BPE format @@ -923,6 +961,37 @@ impl QuantizedEmbeddingModel { } } +const POOLING_CONFIG_FILE: &str = "1_Pooling/config.json"; +const ONNX_INPUT_IDS: &str = "input_ids"; +const ONNX_ATTENTION_MASK: &str = "attention_mask"; +const ONNX_TOKEN_TYPE_IDS: &str = "token_type_ids"; +/// Output of sentence-transformers ONNX exports: pooled and normalized in-graph. +const ONNX_POOLED_OUTPUT: &str = "sentence_embedding"; + +/// How `[batch, seq, hidden]` token states are reduced when the graph does not +/// pool them itself. Comes from sentence-transformers `1_Pooling/config.json`. +#[derive(Debug, Clone, Copy, PartialEq)] +pub(crate) enum OnnxPooling { + Cls, + Mean, +} + +pub(crate) fn parse_pooling_config(contents: &str) -> Result> { + let config: Value = + serde_json::from_str(contents).map_err(|_| LibError::ModelConfigParseFailed)?; + let flag = |key: &str| config.get(key).and_then(Value::as_bool).unwrap_or(false); + match ( + flag("pooling_mode_cls_token"), + flag("pooling_mode_mean_tokens"), + ) { + (true, false) => Ok(OnnxPooling::Cls), + (false, true) => Ok(OnnxPooling::Mean), + // max / last-token / concatenated modes would silently yield a vector + // space different from the published model's. + _ => Err(Box::new(LibError::ModelLoadFailed)), + } +} + /// ONNX embedding model using ORT (onnxruntime) for optimized inference. /// Uses a single session with concurrent inference via SessionWrapper wrapper. pub struct OnnxEmbeddingModel { @@ -930,6 +999,10 @@ pub struct OnnxEmbeddingModel { tokenizer: Tokenizer, max_input_len: usize, hidden_size: usize, + /// Graph inputs in declaration order; exports differ (XLM-R graphs have no token_type_ids). + input_names: Vec, + output_name: String, + pooling: OnnxPooling, } impl OnnxEmbeddingModel { @@ -964,26 +1037,56 @@ impl OnnxEmbeddingModel { .commit_from_file(&onnx_path) .map_err(|_| LibError::ModelWeightsLoadFailed)?; + let input_names: Vec = session + .inputs() + .iter() + .map(|input| input.name().to_string()) + .collect(); + if input_names.iter().any(|name| { + !matches!( + name.as_str(), + ONNX_INPUT_IDS | ONNX_ATTENTION_MASK | ONNX_TOKEN_TYPE_IDS + ) + }) { + return Err(Box::new(LibError::ModelLoadFailed)); + } + let output_names: Vec<&str> = session.outputs().iter().map(|o| o.name()).collect(); + let output_name = output_names + .iter() + .copied() + .find(|name| *name == ONNX_POOLED_OUTPUT) + .or(output_names.first().copied()) + .ok_or(LibError::ModelLoadFailed)? + .to_string(); + // Token-state graphs are pooled here the way the model publishes it; + // mean pooling when the repo has no pooling config. + let pooling = match &model_info.pooling_config_path { + Some(path) if output_name != ONNX_POOLED_OUTPUT => parse_pooling_config( + &std::fs::read_to_string(path).map_err(|_| LibError::ModelConfigReadFailed)?, + )?, + _ => OnnxPooling::Mean, + }; + Ok(Self { session: SessionWrapper::new(session), tokenizer, max_input_len, hidden_size, + input_names, + output_name, + pooling, }) } /// Run a single ONNX forward pass on one batch. - fn run_batch( - session: &SessionWrapper, - batch: &[Vec], - ) -> Result>, Box> { + fn run_batch(&self, batch: &[Vec]) -> Result>, Box> { + let session = &self.session; session.with_session(|sess| -> Result>, Box> { let batch_size = batch.len(); let max_len = batch.iter().map(|c| c.len()).max().unwrap_or(0); let mut flat_ids: Vec = Vec::with_capacity(batch_size * max_len); let mut flat_mask: Vec = Vec::with_capacity(batch_size * max_len); - let mut flat_type_ids: Vec = Vec::with_capacity(batch_size * max_len); for chunk in batch { let real_len = chunk.len(); @@ -993,27 +1096,25 @@ impl OnnxEmbeddingModel { flat_ids.extend(std::iter::repeat_n(0i64, max_len - real_len)); flat_mask.extend(std::iter::repeat_n(1i64, real_len)); flat_mask.extend(std::iter::repeat_n(0i64, max_len - real_len)); - flat_type_ids.extend(std::iter::repeat_n(0i64, max_len)); } - let input_ids = ort::value::Tensor::from_array((vec![batch_size, max_len], flat_ids)) - .map_err(|_| LibError::OnnxModelEvalFailed)?; - let attention_mask = - ort::value::Tensor::from_array((vec![batch_size, max_len], flat_mask.clone())) - .map_err(|_| LibError::OnnxModelEvalFailed)?; - let token_type_ids = - ort::value::Tensor::from_array((vec![batch_size, max_len], flat_type_ids)) + let mut inputs = Vec::with_capacity(self.input_names.len()); + for name in &self.input_names { + let data = match name.as_str() { + ONNX_INPUT_IDS => std::mem::take(&mut flat_ids), + ONNX_ATTENTION_MASK => flat_mask.clone(), + ONNX_TOKEN_TYPE_IDS => vec![0i64; batch_size * max_len], + _ => unreachable!("input names are validated in new()"), + }; + let tensor = ort::value::Tensor::from_array((vec![batch_size, max_len], data)) .map_err(|_| LibError::OnnxModelEvalFailed)?; - + inputs.push((name.as_str(), tensor)); + } let outputs = sess - .run(ort::inputs![ - "input_ids" => input_ids, - "attention_mask" => attention_mask, - "token_type_ids" => token_type_ids, - ]) + .run(inputs) .map_err(|_| LibError::OnnxModelEvalFailed)?; - let (shape, data) = outputs[0] + let (shape, data) = outputs[self.output_name.as_str()] .try_extract_tensor::() .map_err(|_| LibError::OnnxModelEvalFailed)?; @@ -1032,21 +1133,30 @@ impl OnnxEmbeddingModel { let seq_len = shape[1] as usize; let hidden_dim = shape[2] as usize; for i in 0..batch_size { - let mut emb = vec![0.0f32; hidden_dim]; - let mut count = 0.0f32; - for j in 0..seq_len { - let mask_val = flat_mask[i * max_len + j] as f32; - if mask_val > 0.0 { - let offset = (i * seq_len + j) * hidden_dim; - for k in 0..hidden_dim { - emb[k] += data[offset + k]; + let mut emb = match self.pooling { + OnnxPooling::Cls => { + let offset = i * seq_len * hidden_dim; + data[offset..offset + hidden_dim].to_vec() + } + OnnxPooling::Mean => { + let mut emb = vec![0.0f32; hidden_dim]; + let mut count = 0.0f32; + for j in 0..seq_len { + let mask_val = flat_mask[i * max_len + j] as f32; + if mask_val > 0.0 { + let offset = (i * seq_len + j) * hidden_dim; + for k in 0..hidden_dim { + emb[k] += data[offset + k]; + } + count += 1.0; + } + } + if count > 0.0 { + emb.iter_mut().for_each(|v| *v /= count); } - count += 1.0; + emb } - } - if count > 0.0 { - emb.iter_mut().for_each(|v| *v /= count); - } + }; normalize(&mut emb); embeddings.push(emb); } @@ -1059,17 +1169,14 @@ impl OnnxEmbeddingModel { } /// Tokenize one batch and run inference. - fn tokenize_and_infer( - session: &SessionWrapper, - tokenizer: &Tokenizer, - texts: &[&str], - max_input: usize, - ) -> Result>, Box> { + fn tokenize_and_infer(&self, texts: &[&str]) -> Result>, Box> { + let max_input = self.max_input_len; let truncated: Vec<&str> = texts .iter() .map(|t| pre_truncate_text(t, max_input)) .collect(); - let encoded = tokenizer + let encoded = self + .tokenizer .encode_batch(truncated, true) .map_err(|_| LibError::ModelTokenizerEncodeFailed)?; let chunks: Vec> = encoded @@ -1079,7 +1186,7 @@ impl OnnxEmbeddingModel { ids[..ids.len().min(max_input)].to_vec() }) .collect(); - Self::run_batch(session, &chunks) + self.run_batch(&chunks) } /// Adaptive predict: automatically chooses the best strategy based on input size. @@ -1095,9 +1202,6 @@ impl OnnxEmbeddingModel { threads: usize, ) -> Result>, Box> { let bs = batch_size(); - let max_input = self.max_input_len; - let session = &self.session; - let tokenizer = &self.tokenizer; if texts.is_empty() { return Ok(Vec::new()); @@ -1105,7 +1209,7 @@ impl OnnxEmbeddingModel { // Small input — single tokenize + infer, no threading overhead if texts.len() <= bs { - return Self::tokenize_and_infer(session, tokenizer, texts, max_input); + return self.tokenize_and_infer(texts); } // Adaptive parallelism: scale workers with input size. @@ -1130,13 +1234,9 @@ impl OnnxEmbeddingModel { s.spawn(move || -> Result>, LibError> { let mut embeddings = Vec::with_capacity(worker_texts.len()); for text in worker_texts { - let embs = Self::tokenize_and_infer( - session, - tokenizer, - std::slice::from_ref(text), - max_input, - ) - .map_err(|_| LibError::OnnxModelEvalFailed)?; + let embs = self + .tokenize_and_infer(std::slice::from_ref(text)) + .map_err(|_| LibError::OnnxModelEvalFailed)?; embeddings.extend(embs); } Ok(embeddings) diff --git a/embeddings/src/model/local_test.rs b/embeddings/src/model/local_test.rs index 80e32706..275dc55b 100644 --- a/embeddings/src/model/local_test.rs +++ b/embeddings/src/model/local_test.rs @@ -1,4 +1,7 @@ -use super::local::{build_model_info, download_max_for, reset_download_tracker, LocalModel}; +use super::local::{ + build_model_info, download_max_for, onnx_sidecar_files, parse_pooling_config, + reset_download_tracker, LocalModel, OnnxPooling, +}; #[cfg(test)] mod tests { @@ -23,6 +26,39 @@ mod tests { // Note: These tests require actual model files to run successfully // They are designed to test the structure and error handling + #[test] + fn test_onnx_sidecar_files() { + let siblings: Vec = [ + "onnx/model.onnx", + "onnx/model.onnx_data", + "onnx/model.onnx_data_1", + "onnx/model_fp16.onnx", + "onnx/Constant_7_attr__value", + "model.onnx_data", + ] + .iter() + .map(|s| s.to_string()) + .collect(); + assert_eq!( + onnx_sidecar_files("onnx/model.onnx", &siblings), + ["onnx/model.onnx_data", "onnx/model.onnx_data_1"] + ); + assert_eq!( + onnx_sidecar_files("model.onnx", &siblings), + ["model.onnx_data"] + ); + } + + #[test] + fn test_parse_pooling_config() { + let cls = r#"{"pooling_mode_cls_token": true, "pooling_mode_mean_tokens": false}"#; + assert_eq!(parse_pooling_config(cls).unwrap(), OnnxPooling::Cls); + let mean = r#"{"pooling_mode_cls_token": false, "pooling_mode_mean_tokens": true}"#; + assert_eq!(parse_pooling_config(mean).unwrap(), OnnxPooling::Mean); + assert!(parse_pooling_config(r#"{"pooling_mode_max_tokens": true}"#).is_err()); + assert!(parse_pooling_config("not json").is_err()); + } + /// MAX_INPUT_TOKENS (manticoresearch#4816): the cap lowers the model's input limit, /// never raises it, and really truncates what gets embedded. #[test] @@ -365,6 +401,7 @@ mod tests { weights_paths: vec![], gguf_path: Some(PathBuf::from("/tmp/model.gguf")), onnx_path: None, + pooling_config_path: None, }; assert!(info.gguf_path.is_some()); @@ -770,6 +807,7 @@ mod tests { weights_paths: vec![], gguf_path: None, onnx_path: Some(PathBuf::from("/tmp/model.onnx")), + pooling_config_path: None, }; assert!(info.onnx_path.is_some()); @@ -802,6 +840,40 @@ mod tests { check_embedding_properties(&embeddings[0], local_model.get_hidden_size()); } + /// BAAI/bge-m3: XLM-R graph without token_type_ids, weights in the + /// model.onnx_data sidecar, CLS-pooled `sentence_embedding` output. + #[test] + fn test_onnx_bge_m3() { + let model_id = "BAAI/bge-m3"; + let local_model = match LocalModel::new(model_id, test_cache_path(), false, None, None) { + Ok(m) => m, + Err(e) => { + println!("bge-m3 test skipped: {}", e); + return; + } + }; + + assert_eq!(local_model.get_hidden_size(), 1024); + assert_eq!(local_model.get_max_input_len(), 8192); + + let embeddings = local_model + .predict(&["булка", "хлеб", "собака"], 0) + .expect("bge-m3 should generate embeddings"); + assert_eq!(embeddings.len(), 3); + for embedding in &embeddings { + check_embedding_properties(embedding, 1024); + } + + // Leading components of transformers' XLMRobertaModel CLS token, + // L2-normalized, for "булка". Mean pooling lands ~0.87 cosine away. + let reference = [ + 0.018139, 0.058514, -0.042902, 0.017726, -0.052208, 0.010081, -0.018116, 0.046001, + ]; + for (ours, expected) in embeddings[0].iter().zip(reference) { + assert_abs_diff_eq!(ours, &expected, epsilon = 2e-3); + } + } + #[test] fn test_onnx_embedding_consistency() { let model_id = "onnx-models/all-MiniLM-L12-v2-onnx"; diff --git a/embeddings/src/utils.rs b/embeddings/src/utils.rs index e2423f88..1c400e02 100644 --- a/embeddings/src/utils.rs +++ b/embeddings/src/utils.rs @@ -28,7 +28,7 @@ pub fn get_max_input_length(contents: &str) -> Result { .get("max_position_embeddings") .and_then(Value::as_u64) { - return Ok(max_len as usize); + return Ok(max_len as usize - roberta_position_offset(&config)); } // Try n_positions (some models use this) @@ -45,6 +45,22 @@ pub fn get_max_input_length(contents: &str) -> Result { Err(std::io::Error::other("Max position embeddings not found").into()) } +/// RoBERTa-family models number positions from `pad_token_id + 1`, so the +/// usable sequence is shorter than the position table (514 -> 512, 8194 -> 8192). +fn roberta_position_offset(config: &Value) -> usize { + match config.get("model_type").and_then(Value::as_str) { + Some("roberta" | "xlm-roberta" | "camembert" | "mpnet") => { + // HF default pad_token_id for these architectures is 1 + config + .get("pad_token_id") + .and_then(Value::as_u64) + .unwrap_or(1) as usize + + 1 + } + _ => 0, + } +} + /// Get hidden size for the current model /// Supports multiple config field names for different model architectures pub fn get_hidden_size(contents: &str) -> Result { @@ -100,6 +116,19 @@ mod tests { assert_eq!(get_max_input_length(n_positions_config).unwrap(), 1024); } + #[test] + fn test_get_max_input_length_roberta_offset() { + let bge_m3 = + r#"{"model_type": "xlm-roberta", "max_position_embeddings": 8194, "pad_token_id": 1}"#; + assert_eq!(get_max_input_length(bge_m3).unwrap(), 8192); + + let roberta_default_pad = r#"{"model_type": "roberta", "max_position_embeddings": 514}"#; + assert_eq!(get_max_input_length(roberta_default_pad).unwrap(), 512); + + let bert = r#"{"model_type": "bert", "max_position_embeddings": 512, "pad_token_id": 0}"#; + assert_eq!(get_max_input_length(bert).unwrap(), 512); + } + #[test] fn test_get_hidden_size() { let config = r#"{"hidden_size": 768}"#;