llm_kernel/embedding/
qwen3.rs1use crate::embedding::types::{EmbeddingProvider, EmbeddingResult};
18use crate::error::{KernelError, Result};
19
20pub struct Qwen3Provider {
26 inner: fastembed::Qwen3TextEmbedding,
27 model_id: String,
28 dim: usize,
29}
30
31pub const QWEN3_EMBEDDING_0_6B: &str = "Qwen/Qwen3-Embedding-0.6B";
33
34pub const QWEN3_EMBEDDING_8B: &str = "Qwen/Qwen3-Embedding-8B";
36
37pub const QWEN3_VL_EMBEDDING_2B: &str = "Qwen/Qwen3-VL-Embedding-2B";
39
40const DEFAULT_MAX_LENGTH: usize = 512;
42
43impl Qwen3Provider {
44 pub fn new(model_id: &str) -> Result<Self> {
48 Self::with_options(
49 model_id,
50 candle_core::Device::Cpu,
51 candle_core::DType::F32,
52 DEFAULT_MAX_LENGTH,
53 )
54 }
55
56 #[cfg(all(feature = "embedding-metal", target_os = "macos"))]
63 pub fn new_metal(model_id: &str) -> Result<Self> {
64 Self::with_options(
65 model_id,
66 candle_core::Device::new_metal(0)
67 .map_err(|e| KernelError::Embedding(format!("metal device init: {e}")))?,
68 candle_core::DType::F16,
69 DEFAULT_MAX_LENGTH,
70 )
71 }
72
73 pub fn with_options(
75 model_id: &str,
76 device: candle_core::Device,
77 dtype: candle_core::DType,
78 max_length: usize,
79 ) -> Result<Self> {
80 let te = fastembed::Qwen3TextEmbedding::from_hf(model_id, &device, dtype, max_length)
81 .map_err(KernelError::embedding)?;
82 let dim = te.config().hidden_size;
83 Ok(Self {
84 inner: te,
85 model_id: model_id.to_string(),
86 dim,
87 })
88 }
89
90 pub fn model_id(&self) -> &str {
92 &self.model_id
93 }
94}
95
96impl EmbeddingProvider for Qwen3Provider {
97 fn dim(&self) -> usize {
98 self.dim
99 }
100
101 fn name(&self) -> &str {
102 &self.model_id
103 }
104
105 fn embed(&self, text: &str) -> Result<EmbeddingResult> {
106 let embeddings = self.inner.embed(&[text]).map_err(KernelError::embedding)?;
107 let vector = embeddings
108 .into_iter()
109 .next()
110 .ok_or_else(|| KernelError::Embedding("empty embedding output".into()))?;
111
112 let preview = if text.len() > 64 {
113 format!("{}…", &text[..64])
114 } else {
115 text.to_string()
116 };
117 Ok(EmbeddingResult {
118 vector,
119 text_preview: preview,
120 })
121 }
122
123 fn embed_batch(&self, texts: &[&str]) -> Result<Vec<EmbeddingResult>> {
124 if texts.is_empty() {
125 return Ok(vec![]);
126 }
127 let embeddings = self.inner.embed(texts).map_err(KernelError::embedding)?;
128 Ok(embeddings
129 .into_iter()
130 .zip(texts.iter())
131 .map(|(vector, &text)| {
132 let preview = if text.len() > 64 {
133 format!("{}…", &text[..64])
134 } else {
135 text.to_string()
136 };
137 EmbeddingResult {
138 vector,
139 text_preview: preview,
140 }
141 })
142 .collect())
143 }
144}
145
146#[cfg(test)]
147mod tests {
148 use super::*;
149
150 #[test]
151 fn model_id_constants() {
152 assert_eq!(QWEN3_EMBEDDING_0_6B, "Qwen/Qwen3-Embedding-0.6B");
153 assert_eq!(QWEN3_EMBEDDING_8B, "Qwen/Qwen3-Embedding-8B");
154 assert_eq!(QWEN3_VL_EMBEDDING_2B, "Qwen/Qwen3-VL-Embedding-2B");
155 }
156
157 #[cfg(all(feature = "embedding-metal", target_os = "macos"))]
160 #[test]
161 fn metal_device_initialises() {
162 assert!(
163 candle_core::Device::new_metal(0).is_ok(),
164 "Metal device failed to init — new_metal() would error"
165 );
166 }
167
168 #[test]
169 #[ignore = "requires model download"]
170 fn embed_with_qwen3_0_6b() {
171 let provider = Qwen3Provider::new(QWEN3_EMBEDDING_0_6B).unwrap();
172 let result = provider.embed("hello world").unwrap();
173 assert!(!result.vector.is_empty());
175 assert_eq!(result.vector.len(), provider.dim());
176 }
177
178 #[cfg(all(feature = "embedding-metal", target_os = "macos"))]
183 #[test]
184 #[ignore = "requires model download + Metal (macOS)"]
185 fn embed_with_qwen3_metal() {
186 let provider = Qwen3Provider::new_metal(QWEN3_EMBEDDING_0_6B).unwrap();
187 let result = provider.embed("hello world").unwrap();
188 assert!(!result.vector.is_empty());
189 assert_eq!(result.vector.len(), provider.dim());
190 eprintln!(
191 "metal embed ok: dim={} preview={:?}",
192 result.vector.len(),
193 &result.vector[..3.min(result.vector.len())]
194 );
195 }
196
197 #[test]
198 #[ignore = "requires model download"]
199 fn embed_batch_with_qwen3() {
200 let provider = Qwen3Provider::new(QWEN3_EMBEDDING_0_6B).unwrap();
201 let results = provider
202 .embed_batch(&["hello", "world", "foo bar"])
203 .unwrap();
204 assert_eq!(results.len(), 3);
205 for r in &results {
206 assert_eq!(r.vector.len(), provider.dim());
207 }
208 }
209}