use anyhow::{Context, Result};
use candle_core::{Device, Tensor};
use candle_transformers::models::qwen2::{Config as Qwen2Config, ModelForCausalLM};
use tokenizers::Tokenizer;
use crate::models::EmbeddingModel as EmbeddingModelTrait;
pub struct Qwen3EmbeddingModel {
pub model: tokio::sync::RwLock<ModelForCausalLM>,
pub tokenizer: Tokenizer,
pub config: Qwen2Config,
pub device: Device,
pub model_name: String,
}
impl std::fmt::Debug for Qwen3EmbeddingModel {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Qwen3EmbeddingModel")
.field("model_name", &self.model_name)
.field("hidden_size", &self.config.hidden_size)
.field("device", &self.device)
.finish()
}
}
impl Qwen3EmbeddingModel {
pub fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>> {
self.embed_with_instruction(texts, None)
}
pub fn embed_with_instruction(
&self,
texts: &[String],
instruction: Option<&str>,
) -> Result<Vec<Vec<f32>>> {
let mut embeddings = Vec::with_capacity(texts.len());
for text in texts {
let processed_text = if let Some(instr) = instruction {
format!("Instruct: {}\nQuery: {}", instr, text)
} else {
text.clone()
};
let encoding = self
.tokenizer
.encode(processed_text.as_str(), true)
.map_err(|e| {
anyhow::anyhow!(
"[HUGGINGFACE] [TOKENIZE] failed: Cannot tokenize input text. Error: {}",
e
)
})?;
let token_ids: Vec<u32> = encoding.get_ids().to_vec();
let attention_mask: Vec<u32> = vec![1u32; token_ids.len()];
let input_ids = Tensor::new(&token_ids[..], &self.device)
.context("[HUGGINGFACE] [TENSOR] failed: Cannot create input_ids tensor from tokenized text")?
.unsqueeze(0)?;
let attention_mask_tensor = Tensor::new(&attention_mask[..], &self.device)
.context("[HUGGINGFACE] [TENSOR] failed: Cannot create attention_mask tensor for model input")?
.unsqueeze(0)?;
let hidden_states = {
let mut model_guard = self.model.blocking_write();
model_guard
.forward(&input_ids, 0)
.context("[HUGGINGFACE] [INFERENCE] failed: Model forward pass execution failed. Check input tensor dimensions and model state.")?
};
let embedding = self.mean_pooling(&hidden_states, &attention_mask_tensor)?;
let normalized = self.normalize_embedding(embedding)?;
embeddings.push(normalized);
}
Ok(embeddings)
}
fn mean_pooling(&self, last_hidden_states: &Tensor, attention_mask: &Tensor) -> Result<Tensor> {
let attention_mask = attention_mask.unsqueeze(2)?; let expanded_mask = attention_mask.broadcast_as(last_hidden_states.shape())?;
let masked_hidden_states = (last_hidden_states * &expanded_mask)?;
let sum_hidden_states = masked_hidden_states.sum(1)?;
let sum_mask = attention_mask.sum(1)?;
let mean_pooled = sum_hidden_states.broadcast_div(&sum_mask)?;
Ok(mean_pooled)
}
fn normalize_embedding(&self, embedding: Tensor) -> Result<Vec<f32>> {
let norm = embedding.sqr()?.sum_keepdim(1)?.sqrt()?;
let normalized = embedding.broadcast_div(&norm)?;
let normalized_vec = normalized.squeeze(0)?.to_vec1::<f32>()?;
Ok(normalized_vec)
}
}
impl EmbeddingModelTrait for Qwen3EmbeddingModel {
fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>> {
self.embed(texts)
}
fn dimensions(&self) -> usize {
self.config.hidden_size
}
fn max_sequence_length(&self) -> usize {
self.config.max_position_embeddings
}
}