llm_kernel/embedding/
nomic_moe.rs1use crate::embedding::types::{EmbeddingProvider, EmbeddingResult};
18use crate::error::{KernelError, Result};
19
20pub struct NomicMoeProvider {
26 inner: fastembed::NomicV2MoeTextEmbedding,
27 model_id: String,
28 dim: usize,
29}
30
31pub const NOMIC_EMBED_TEXT_V2_MOE: &str = "nomic-ai/nomic-embed-text-v2-moe";
33
34const DEFAULT_MAX_LENGTH: usize = 512;
36
37impl NomicMoeProvider {
38 pub fn new() -> Result<Self> {
42 Self::with_options(
43 NOMIC_EMBED_TEXT_V2_MOE,
44 candle_core::Device::Cpu,
45 candle_core::DType::F32,
46 DEFAULT_MAX_LENGTH,
47 )
48 }
49
50 #[cfg(all(feature = "embedding-metal", target_os = "macos"))]
57 pub fn new_metal() -> Result<Self> {
58 Self::with_options(
59 NOMIC_EMBED_TEXT_V2_MOE,
60 candle_core::Device::new_metal(0)
61 .map_err(|e| KernelError::Embedding(format!("metal device init: {e}")))?,
62 candle_core::DType::F16,
63 DEFAULT_MAX_LENGTH,
64 )
65 }
66
67 pub fn with_options(
69 model_id: &str,
70 device: candle_core::Device,
71 dtype: candle_core::DType,
72 max_length: usize,
73 ) -> Result<Self> {
74 let te = fastembed::NomicV2MoeTextEmbedding::from_hf(model_id, &device, dtype, max_length)
75 .map_err(KernelError::embedding)?;
76 let dim = te.config().hidden_size;
77 Ok(Self {
78 inner: te,
79 model_id: model_id.to_string(),
80 dim,
81 })
82 }
83
84 pub fn model_id(&self) -> &str {
86 &self.model_id
87 }
88}
89
90impl EmbeddingProvider for NomicMoeProvider {
91 fn dim(&self) -> usize {
92 self.dim
93 }
94
95 fn name(&self) -> &str {
96 &self.model_id
97 }
98
99 fn embed(&self, text: &str) -> Result<EmbeddingResult> {
100 let embeddings = self.inner.embed(&[text]).map_err(KernelError::embedding)?;
101 let vector = embeddings
102 .into_iter()
103 .next()
104 .ok_or_else(|| KernelError::Embedding("empty embedding output".into()))?;
105
106 let preview = if text.len() > 64 {
107 format!("{}…", &text[..64])
108 } else {
109 text.to_string()
110 };
111 Ok(EmbeddingResult {
112 vector,
113 text_preview: preview,
114 })
115 }
116
117 fn embed_batch(&self, texts: &[&str]) -> Result<Vec<EmbeddingResult>> {
118 if texts.is_empty() {
119 return Ok(vec![]);
120 }
121 let embeddings = self.inner.embed(texts).map_err(KernelError::embedding)?;
122 Ok(embeddings
123 .into_iter()
124 .zip(texts.iter())
125 .map(|(vector, &text)| {
126 let preview = if text.len() > 64 {
127 format!("{}…", &text[..64])
128 } else {
129 text.to_string()
130 };
131 EmbeddingResult {
132 vector,
133 text_preview: preview,
134 }
135 })
136 .collect())
137 }
138}
139
140#[cfg(test)]
141mod tests {
142 use super::*;
143
144 #[test]
145 fn model_id_constant() {
146 assert_eq!(NOMIC_EMBED_TEXT_V2_MOE, "nomic-ai/nomic-embed-text-v2-moe");
147 }
148
149 #[test]
150 #[ignore = "requires model download"]
151 fn embed_with_nomic_moe() {
152 let provider = NomicMoeProvider::new().unwrap();
153 let result = provider.embed("hello world").unwrap();
154 assert_eq!(result.vector.len(), 768);
156 assert_eq!(result.vector.len(), provider.dim());
157 }
158
159 #[test]
160 #[ignore = "requires model download"]
161 fn embed_batch_with_nomic_moe() {
162 let provider = NomicMoeProvider::new().unwrap();
163 let results = provider
164 .embed_batch(&["hello", "world", "foo bar"])
165 .unwrap();
166 assert_eq!(results.len(), 3);
167 for r in &results {
168 assert_eq!(r.vector.len(), 768);
169 }
170 }
171}