diff --git a/rust/src/embeddings/embed/text.rs b/rust/src/embeddings/embed/text.rs index 8df5c03f..e68e6b15 100644 --- a/rust/src/embeddings/embed/text.rs +++ b/rust/src/embeddings/embed/text.rs @@ -22,7 +22,6 @@ use crate::embeddings::local::{ text_embedding::ONNXModel, }; - pub enum TextEmbedder { OpenAI(OpenAIEmbedder), Cohere(CohereEmbedder), @@ -89,12 +88,14 @@ impl TextEmbedder { model_id, token, None, )?))), - "ModernBertForMaskedLM" => Ok(Self::ModernBert(Box::new(ModernBertEmbedder::new( - model_id.to_string(), - revision.map(|s| s.to_string()), - token, - dtype, - )?))), + architecture if is_modernbert_architecture(architecture) => { + Ok(Self::ModernBert(Box::new(ModernBertEmbedder::new( + model_id.to_string(), + revision.map(|s| s.to_string()), + token, + dtype, + )?))) + } "Qwen3ForCausalLM" => Ok(Self::Qwen3(Box::new(Qwen3Embedder::new( model_id, revision.map(|s| s.to_string()), @@ -196,3 +197,20 @@ pub trait TextEmbed { batch_size: Option, ) -> impl Future>>; } + +fn is_modernbert_architecture(architecture: &str) -> bool { + matches!(architecture, "ModernBertForMaskedLM" | "ModernBertModel") +} + +#[cfg(test)] +mod tests { + use super::is_modernbert_architecture; + + #[test] + fn supports_modernbert_embedding_architectures() { + for architecture in ["ModernBertForMaskedLM", "ModernBertModel"] { + assert!(is_modernbert_architecture(architecture)); + } + assert!(!is_modernbert_architecture("BertModel")); + } +}