Skip to main content

llm_kernel/embedding/
qwen3.rs

1//! Qwen3 embedding provider via fastembed-rs candle backend.
2//!
3//! Uses the candle-nn pure Rust inference engine (no ONNX Runtime).
4//! Models are downloaded from HuggingFace on first use.
5//!
6//! Supported repos: `Qwen/Qwen3-Embedding-0.6B`, `Qwen/Qwen3-Embedding-8B`,
7//! `Qwen/Qwen3-VL-Embedding-2B` (text-only mode).
8//!
9//! ```ignore
10//! use llm_kernel::embedding::Qwen3Provider;
11//! use llm_kernel::embedding::EmbeddingProvider;
12//!
13//! let provider = Qwen3Provider::new("Qwen/Qwen3-Embedding-0.6B")?;
14//! let result = provider.embed("hello world")?;
15//! ```
16
17use crate::embedding::types::{EmbeddingProvider, EmbeddingResult};
18use crate::error::{KernelError, Result};
19
20/// Qwen3 embedding provider backed by candle-nn.
21///
22/// Unlike [`FastembedProvider`](super::FastembedProvider) (ONNX), this uses
23/// candle for pure Rust GPU/CPU inference. The `embed()` method takes `&self`,
24/// so no `Mutex` is needed.
25pub struct Qwen3Provider {
26    inner: fastembed::Qwen3TextEmbedding,
27    model_id: String,
28    dim: usize,
29}
30
31/// Default HuggingFace repo for Qwen3-Embedding-0.6B.
32pub const QWEN3_EMBEDDING_0_6B: &str = "Qwen/Qwen3-Embedding-0.6B";
33
34/// Default HuggingFace repo for Qwen3-Embedding-8B.
35pub const QWEN3_EMBEDDING_8B: &str = "Qwen/Qwen3-Embedding-8B";
36
37/// Default HuggingFace repo for Qwen3-VL-Embedding-2B (text-only mode).
38pub const QWEN3_VL_EMBEDDING_2B: &str = "Qwen/Qwen3-VL-Embedding-2B";
39
40/// Default max sequence length for Qwen3 models.
41const DEFAULT_MAX_LENGTH: usize = 512;
42
43impl Qwen3Provider {
44    /// Create a new provider using CPU with F32 precision.
45    ///
46    /// Downloads the model from HuggingFace on first call (cached locally).
47    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    /// Create using the Apple Silicon GPU (Metal) with F16 precision.
57    ///
58    /// Requires the `embedding-metal` feature and macOS. Routes candle
59    /// inference to the Metal device — typically several× faster than CPU on
60    /// Apple Silicon for these encoder models. If Metal is unavailable this
61    /// returns an error (no automatic CPU fallback).
62    #[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    /// Create with custom device (GPU), dtype, and max sequence length.
74    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    /// The HuggingFace model repo ID.
91    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    // Verifies the Metal device initialises — the only thing new_metal() adds
158    // beyond with_options(). No model download. macOS + embedding-metal only.
159    #[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        // Qwen3-Embedding-0.6B has hidden_size that the config reports
174        assert!(!result.vector.is_empty());
175        assert_eq!(result.vector.len(), provider.dim());
176    }
177
178    // End-to-end Metal verification: downloads Qwen3-Embedding-0.6B and runs a
179    // real embed on Device::Metal. candle executes ops on the device's backend
180    // or errors — no silent CPU fallback — so a passing embed on a Metal device
181    // proves Metal kernels ran. macOS + embedding-metal only.
182    #[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}