From ed4811a01fc10caf42d47caea3d50b6e3f0ef394 Mon Sep 17 00:00:00 2001 From: jhqxxx <18280426169@163.com> Date: Mon, 6 Apr 2026 17:48:15 +0800 Subject: [PATCH] add qwen3-reranker --- src/cli/mod.rs | 5 ++ src/exec/mod.rs | 1 + src/exec/qwen3_reranker.rs | 72 ++++++++++++++++++++++++++++ src/models/common/mod.rs | 1 + src/models/common/model_mapping.rs | 9 ++++ src/models/common/modules.rs | 10 ++++ src/models/common/reranker.rs | 25 ++++++++++ src/models/mod.rs | 25 +++++++++- src/models/qwen3_embedding/mod.rs | 47 +++++++++++-------- src/models/qwen3_reranker/mod.rs | 31 ++++++++++++ src/params/rerank.rs | 16 +++---- src/server/mod.rs | 6 +++ src/server/reranker.rs | 75 ++++++++++++++++++++++++++++++ tests/test_all_minilm_l6_v2.rs | 4 +- tests/test_qwen3_rerank.rs | 26 +++++++++++ 15 files changed, 320 insertions(+), 33 deletions(-) create mode 100644 src/exec/qwen3_reranker.rs create mode 100644 src/models/common/reranker.rs create mode 100644 src/models/qwen3_reranker/mod.rs create mode 100644 src/server/reranker.rs create mode 100644 tests/test_qwen3_rerank.rs diff --git a/src/cli/mod.rs b/src/cli/mod.rs index 3cc26f6..7d963f1 100644 --- a/src/cli/mod.rs +++ b/src/cli/mod.rs @@ -278,6 +278,11 @@ pub(crate) fn run_run(args: RunArgs) -> anyhow::Result<()> { | WhichModel::Qwen3Embedding8B => { qwen3_embedding::Qwen3EmbeddingExec::run(&input, output.as_deref(), &weight_path)?; } + WhichModel::Qwen3Reranker0_6B + | WhichModel::Qwen3Reranker4B + | WhichModel::Qwen3Reranker8B => { + qwen3_reranker::Qwen3RerankerExec::run(&input, output.as_deref(), &weight_path)?; + } WhichModel::Qwen3VL2B | WhichModel::Qwen3VL4B | WhichModel::Qwen3VL8B diff --git a/src/exec/mod.rs b/src/exec/mod.rs index 9681442..40947cc 100644 --- a/src/exec/mod.rs +++ b/src/exec/mod.rs @@ -18,6 +18,7 @@ pub mod qwen3; pub mod qwen3_5; pub mod qwen3_asr; pub mod qwen3_embedding; +pub mod qwen3_reranker; pub mod qwen3vl; pub mod rmbg2_0; pub mod voxcpm; diff --git a/src/exec/qwen3_reranker.rs b/src/exec/qwen3_reranker.rs new file mode 100644 index 0000000..66d83ee --- /dev/null +++ b/src/exec/qwen3_reranker.rs @@ -0,0 +1,72 @@ +use crate::{ + exec::ExecModel, + models::{common::reranker::TextRerank, qwen3_reranker::Qwen3Reranker}, +}; +use anyhow::{Result, anyhow}; +use std::time::Instant; + +pub struct Qwen3RerankerExec; + +impl ExecModel for Qwen3RerankerExec { + fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()> { + if input.len() < 2 { + return Err(anyhow!( + "reranker run requires two inputs: " + )); + } + let query = input[0].clone(); + let docs_source = input[1].clone(); + let documents = parse_documents_source(&docs_source)?; + if documents.is_empty() { + return Err(anyhow!("documents list is empty")); + } + + let i_start = Instant::now(); + let mut model = Qwen3Reranker::init(weight_path, None, None)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in load model is: {:?}", i_duration); + + let i_start = Instant::now(); + let scores = model.rerank(&query, &documents)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in generate is: {:?}", i_duration); + + println!("Result: {:?}", scores); + + if let Some(out) = output { + std::fs::write(out, format!("{:?}", scores))?; + println!("Output saved to: {}", out); + } + + Ok(()) + } +} + +fn parse_documents_source(source: &str) -> Result> { + if source.starts_with("file://") { + let path = source.trim_start_matches("file://"); + return read_documents_file(path); + } + if std::path::Path::new(source).exists() { + return read_documents_file(source); + } + + let docs = source + .split("|||") + .map(str::trim) + .filter(|x| !x.is_empty()) + .map(|x| x.to_string()) + .collect::>(); + Ok(docs) +} + +fn read_documents_file(path: &str) -> Result> { + let content = std::fs::read_to_string(path)?; + let docs = content + .lines() + .map(str::trim) + .filter(|line| !line.is_empty()) + .map(|line| line.to_string()) + .collect::>(); + Ok(docs) +} diff --git a/src/models/common/mod.rs b/src/models/common/mod.rs index 33aa1ec..5286926 100644 --- a/src/models/common/mod.rs +++ b/src/models/common/mod.rs @@ -5,6 +5,7 @@ pub mod generate; pub mod gguf; pub mod model_mapping; pub mod modules; +pub mod reranker; /// 多模态模型的特征数据 /// 每个模型数据不一样 diff --git a/src/models/common/model_mapping.rs b/src/models/common/model_mapping.rs index ec747ee..dd99fc2 100644 --- a/src/models/common/model_mapping.rs +++ b/src/models/common/model_mapping.rs @@ -44,6 +44,12 @@ pub enum WhichModel { Qwen3Embedding4B, #[value(name = "Qwen/Qwen3-Embedding-8B")] Qwen3Embedding8B, + #[value(name = "Qwen/Qwen3-Reranker-0.6B")] + Qwen3Reranker0_6B, + #[value(name = "Qwen/Qwen3-Reranker-4B")] + Qwen3Reranker4B, + #[value(name = "Qwen/Qwen3-Reranker-8B")] + Qwen3Reranker8B, #[value(name = "Qwen/Qwen3-VL-2B-Instruct")] Qwen3VL2B, #[value(name = "Qwen/Qwen3-VL-4B-Instruct")] @@ -165,6 +171,9 @@ impl WhichModel { | WhichModel::Qwen3Embedding4B | WhichModel::Qwen3Embedding8B | WhichModel::AllMiniLML6V2 => "embedding", + WhichModel::Qwen3Reranker0_6B + | WhichModel::Qwen3Reranker4B + | WhichModel::Qwen3Reranker8B => "reranker", } } } diff --git a/src/models/common/modules.rs b/src/models/common/modules.rs index cad29e3..ab364f3 100644 --- a/src/models/common/modules.rs +++ b/src/models/common/modules.rs @@ -1453,6 +1453,16 @@ pub fn cosine_similarity(query_vector: &Tensor, matrix: &Tensor) -> Result Result { + // query_vector: (n, dim) + // matrix: (m, dim) + // return (n, m) + let similarity = query_vector + .matmul(&matrix.transpose(D::Minus1, D::Minus2)?)? + .squeeze(D::Minus1)?; + Ok(similarity) +} + pub fn quick_gelu(xs: &Tensor) -> Result { let x = xs.affine(1.702, 0.0)?; let x = sigmoid(&x)?; diff --git a/src/models/common/reranker.rs b/src/models/common/reranker.rs new file mode 100644 index 0000000..c962578 --- /dev/null +++ b/src/models/common/reranker.rs @@ -0,0 +1,25 @@ +use anyhow::Result; +use candle_core::Tensor; + +use crate::models::common::modules::{cosine_similarity, cosine_similarity_no_l2}; +pub trait TextRerank { + fn rerank(&mut self, query: &str, documents: &[String]) -> Result>; +} + +pub enum RerankerSimilarity { + Cosine, +} + +impl RerankerSimilarity { + pub fn similar(&self, query_vector: &Tensor, matrix: &Tensor, need_l2: bool) -> Result { + match self { + RerankerSimilarity::Cosine => { + if need_l2 { + cosine_similarity(query_vector, matrix) + } else { + cosine_similarity_no_l2(query_vector, matrix) + } + } + } + } +} diff --git a/src/models/mod.rs b/src/models/mod.rs index b7b18e9..74a675c 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -19,6 +19,7 @@ pub mod qwen3; pub mod qwen3_5; pub mod qwen3_asr; pub mod qwen3_embedding; +pub mod qwen3_reranker; pub mod qwen3vl; pub mod rmbg2_0; pub mod voxcpm; @@ -27,8 +28,9 @@ pub mod w2v_bert_2_0; use crate::{ models::{ all_minilm_l6_v2::AllMiniLML6V2Embedding, - common::{embedding::TextEmbedding, model_mapping::WhichModel}, + common::{embedding::TextEmbedding, model_mapping::WhichModel, reranker::TextRerank}, qwen3_embedding::Qwen3Embedding, + qwen3_reranker::Qwen3Reranker, }, params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, }; @@ -73,6 +75,7 @@ pub enum ModelInstance<'a> { Qwen3_5(Qwen3_5GenerateModel<'a>), Qwen3ASR(Qwen3AsrGenerateModel<'a>), Qwen3Embedding(Qwen3Embedding), + Qwen3Reranker(Qwen3Reranker), Qwen3VL(Box>), DeepSeekOCR(DeepseekOCRGenerateModel), HunyuanOCR(HunyuanOCRGenerateModel<'a>), @@ -100,6 +103,9 @@ impl<'a> GenerateModel for ModelInstance<'a> { Err(anyhow!("embedding model does not support chat completions")) } ModelInstance::Qwen3ASR(model) => model.generate(mes), + ModelInstance::Qwen3Reranker(_) => { + Err(anyhow!("reranker model does not support chat completions")) + } ModelInstance::Qwen3VL(model) => model.generate(mes), ModelInstance::DeepSeekOCR(model) => model.generate(mes), ModelInstance::HunyuanOCR(model) => model.generate(mes), @@ -133,11 +139,14 @@ impl<'a> GenerateModel for ModelInstance<'a> { ModelInstance::Qwen2_5VL(model) => model.generate_stream(mes), ModelInstance::Qwen3(model) => model.generate_stream(mes), ModelInstance::Qwen3_5(model) => model.generate_stream(mes), + ModelInstance::Qwen3ASR(model) => model.generate_stream(mes), ModelInstance::Qwen3Embedding(_) => Err(anyhow!( "embedding model does not support streaming chat completions" )), + ModelInstance::Qwen3Reranker(_) => { + Err(anyhow!("reranker model does not support chat completions")) + } ModelInstance::Qwen3VL(model) => model.generate_stream(mes), - ModelInstance::Qwen3ASR(model) => model.generate_stream(mes), ModelInstance::DeepSeekOCR(model) => model.generate_stream(mes), ModelInstance::HunyuanOCR(model) => model.generate_stream(mes), ModelInstance::PaddleOCRVL(model) => model.generate_stream(mes), @@ -158,6 +167,12 @@ impl<'a> ModelInstance<'a> { _ => Err(anyhow!("current model does not support embeddings")), } } + pub fn rerank(&mut self, query: &str, documents: &[String]) -> Result> { + match self { + ModelInstance::Qwen3Reranker(model) => model.rerank(query, documents), + _ => Err(anyhow!("current model does not support rerank")), + } + } } #[allow(unused)] @@ -229,6 +244,12 @@ pub fn load_model<'a>( let model = Qwen3Embedding::init(path, device, dtype)?; ModelInstance::Qwen3Embedding(model) } + WhichModel::Qwen3Reranker0_6B + | WhichModel::Qwen3Reranker4B + | WhichModel::Qwen3Reranker8B => { + let model = Qwen3Reranker::init(path, device, dtype)?; + ModelInstance::Qwen3Reranker(model) + } WhichModel::Qwen3VL2B | WhichModel::Qwen3VL4B | WhichModel::Qwen3VL8B diff --git a/src/models/qwen3_embedding/mod.rs b/src/models/qwen3_embedding/mod.rs index 04bc9e6..ab9130d 100644 --- a/src/models/qwen3_embedding/mod.rs +++ b/src/models/qwen3_embedding/mod.rs @@ -7,7 +7,7 @@ use crate::{ utils::{find_type_files, get_device, get_dtype}, }; use anyhow::{Result, anyhow}; -use candle_core::{DType, Device}; +use candle_core::{DType, Device, Tensor}; use candle_nn::VarBuilder; pub struct Qwen3Embedding { @@ -35,32 +35,39 @@ impl Qwen3Embedding { }) } - fn embed_one(&mut self, text: &str) -> Result> { - let input_ids = self.tokenizer.text_encode(text.to_string(), &self.device)?; - let hidden = self - .model - .forward_hidden(Some(&input_ids), None, 0)? - .squeeze(0)? - .to_dtype(DType::F32)?; - let norm = self - .normalize - .normalize(&hidden, hidden.rank() - 1)? - .squeeze(0)?; - let norm = norm.to_vec1::()?; - Ok(norm) - } -} - -impl TextEmbedding for Qwen3Embedding { - fn embed_texts(&mut self, input: &[String]) -> Result>> { + pub fn embed_multi(&mut self, input: &[String]) -> Result { if input.is_empty() { return Err(anyhow!("embedding input cannot be empty")); } let mut out = Vec::with_capacity(input.len()); for text in input { out.push(self.embed_one(text)?); - self.model.clear_kv_cache(); } + let out = Tensor::stack(&out, 0)?; Ok(out) } + + pub fn embed_one(&mut self, text: &str) -> Result { + let input_ids = self.tokenizer.text_encode(text.to_string(), &self.device)?; + let hidden = self + .model + .forward_hidden(Some(&input_ids), None, 0)? + .squeeze(0)? + .to_dtype(DType::F32)?; + + self.model.clear_kv_cache(); + let norm = self + .normalize + .normalize(&hidden, hidden.rank() - 1)? + .squeeze(0)?; + Ok(norm) + } +} + +impl TextEmbedding for Qwen3Embedding { + fn embed_texts(&mut self, input: &[String]) -> Result>> { + let embeds = self.embed_multi(input)?; + let embeds = embeds.to_vec2::()?; + Ok(embeds) + } } diff --git a/src/models/qwen3_reranker/mod.rs b/src/models/qwen3_reranker/mod.rs new file mode 100644 index 0000000..255d9dc --- /dev/null +++ b/src/models/qwen3_reranker/mod.rs @@ -0,0 +1,31 @@ +use crate::models::{ + common::reranker::{RerankerSimilarity, TextRerank}, + qwen3_embedding::Qwen3Embedding, +}; +use anyhow::Result; +use candle_core::{DType, Device}; + +pub struct Qwen3Reranker { + embedding: Qwen3Embedding, + similar: RerankerSimilarity, +} + +impl Qwen3Reranker { + pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result { + let embedding = Qwen3Embedding::init(path, device, dtype)?; + Ok(Self { + embedding, + similar: RerankerSimilarity::Cosine, + }) + } +} + +impl TextRerank for Qwen3Reranker { + fn rerank(&mut self, query: &str, documents: &[String]) -> Result> { + let query = self.embedding.embed_one(query)?.unsqueeze(0)?; + let documents_matrix = self.embedding.embed_multi(documents)?; + let score = self.similar.similar(&query, &documents_matrix, false)?; + let score = score.squeeze(0)?.to_vec1::()?; + Ok(score) + } +} diff --git a/src/params/rerank.rs b/src/params/rerank.rs index 303ab70..f7ab75f 100644 --- a/src/params/rerank.rs +++ b/src/params/rerank.rs @@ -9,15 +9,15 @@ pub(crate) struct RerankRequest { } #[derive(Debug, Serialize)] -struct RerankResult { - index: usize, - relevance_score: f32, - document: String, +pub(crate) struct RerankResult { + pub index: usize, + pub relevance_score: f32, + pub document: String, } #[derive(Debug, Serialize)] -struct RerankResponse { - object: String, - model: String, - results: Vec, +pub(crate) struct RerankResponse { + pub object: String, + pub model: String, + pub results: Vec, } diff --git a/src/server/mod.rs b/src/server/mod.rs index 3e26d6b..c01a93a 100644 --- a/src/server/mod.rs +++ b/src/server/mod.rs @@ -12,6 +12,7 @@ pub(crate) mod api; pub(crate) mod asr; pub(crate) mod embedding; pub(crate) mod process; +pub(crate) mod reranker; pub(crate) async fn start_http_server( address: String, @@ -64,6 +65,11 @@ pub(crate) async fn start_http_server( // /embeddings and /v1/embeddings (OpenAI-compatible embeddings endpoint) builder = builder.mount("/", routes![embedding::embeddings]); builder = builder.mount("/v1", routes![embedding::embeddings]); + + // /rerank and /v1/rerank (OpenAI-compatible embeddings endpoint) + builder = builder.mount("/", routes![reranker::rerank]); + builder = builder.mount("/v1", routes![reranker::rerank]); + // Health check and model info endpoints builder = builder.mount("/", routes![api::health, api::models]); // Shutdown endpoint diff --git a/src/server/reranker.rs b/src/server/reranker.rs new file mode 100644 index 0000000..00de56b --- /dev/null +++ b/src/server/reranker.rs @@ -0,0 +1,75 @@ +use rocket::{http::Status, post, serde::json::Json}; +use serde_json::Value; + +use crate::{ + params::rerank::{RerankRequest, RerankResponse, RerankResult}, + server::api::MODEL, +}; + +fn validate_rerank_input(query: &str, documents: &[String]) -> anyhow::Result<()> { + if query.trim().is_empty() { + return Err(anyhow::anyhow!("rerank query cannot be empty")); + } + if documents.is_empty() { + return Err(anyhow::anyhow!("rerank documents cannot be empty")); + } + if documents.iter().any(|doc| doc.trim().is_empty()) { + return Err(anyhow::anyhow!( + "rerank documents cannot contain empty strings" + )); + } + Ok(()) +} + +#[post("/rerank", data = "")] +pub(crate) async fn rerank(req: Json) -> (Status, Json) { + let req = req.into_inner(); + if let Err(e) = validate_rerank_input(&req.query, &req.documents) { + return ( + Status::BadRequest, + Json(serde_json::json!({ "error": e.to_string() })), + ); + } + + let model_ref = match MODEL.get().cloned() { + Some(v) => v, + None => { + return ( + Status::ServiceUnavailable, + Json(serde_json::json!({ "error": "model not init" })), + ); + } + }; + + let mut guard = model_ref.write().await; + let scores = match guard.instance.rerank(&req.query, &req.documents) { + Ok(v) => v, + Err(e) => { + return ( + Status::BadRequest, + Json(serde_json::json!({ "error": e.to_string() })), + ); + } + }; + + let mut results = scores + .into_iter() + .enumerate() + .map(|(index, relevance_score)| RerankResult { + index, + relevance_score, + document: req.documents[index].clone(), + }) + .collect::>(); + results.sort_by(|a, b| b.relevance_score.total_cmp(&a.relevance_score)); + if let Some(top_n) = req.top_n { + results.truncate(top_n.min(results.len())); + } + + let response = RerankResponse { + object: "list".to_string(), + model: guard.which_model.as_string(), + results, + }; + (Status::Ok, Json(serde_json::to_value(response).unwrap())) +} diff --git a/tests/test_all_minilm_l6_v2.rs b/tests/test_all_minilm_l6_v2.rs index 5e37d18..08a9c91 100644 --- a/tests/test_all_minilm_l6_v2.rs +++ b/tests/test_all_minilm_l6_v2.rs @@ -1,6 +1,4 @@ -use aha::models::{ - all_minilm_l6_v2::AllMiniLML6V2Embedding, common::embedding::TextEmbedding, -}; +use aha::models::{all_minilm_l6_v2::AllMiniLML6V2Embedding, common::embedding::TextEmbedding}; use anyhow::Result; use std::time::Instant; diff --git a/tests/test_qwen3_rerank.rs b/tests/test_qwen3_rerank.rs new file mode 100644 index 0000000..832595d --- /dev/null +++ b/tests/test_qwen3_rerank.rs @@ -0,0 +1,26 @@ +use aha::models::{common::reranker::TextRerank, qwen3_reranker::Qwen3Reranker}; +use anyhow::Result; +use std::time::Instant; + +#[test] +fn qwen3_rerank() -> Result<()> { + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_qwen3_rerank qwen3_rerank -r -- --nocapture + + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/Qwen/Qwen3-Reranker-0.6B/", save_dir); + + let i_start = Instant::now(); + let mut model = Qwen3Reranker::init(&model_path, None, None)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in load model is: {:?}", i_duration); + let docs = vec![ + "Rust async requests are commonly built with reqwest and tokio.".to_string(), + "Paris is the capital of France.".to_string(), + ]; + let input_texts = "How to make async HTTP calls in Rust?"; + let score = model.rerank(input_texts, &docs)?; + println!("result: {:?}", score); + + Ok(()) +}