add qwen3-reranker

This commit is contained in:
jhqxxx
2026-04-06 17:48:15 +08:00
parent 7142e712f2
commit ed4811a01f
15 changed files with 320 additions and 33 deletions
+5
View File
@@ -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
+1
View File
@@ -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;
+72
View File
@@ -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: <query> <documents-source>"
));
}
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<Vec<String>> {
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::<Vec<_>>();
Ok(docs)
}
fn read_documents_file(path: &str) -> Result<Vec<String>> {
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::<Vec<_>>();
Ok(docs)
}
+1
View File
@@ -5,6 +5,7 @@ pub mod generate;
pub mod gguf;
pub mod model_mapping;
pub mod modules;
pub mod reranker;
/// 多模态模型的特征数据
/// 每个模型数据不一样
+9
View File
@@ -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",
}
}
}
+10
View File
@@ -1453,6 +1453,16 @@ pub fn cosine_similarity(query_vector: &Tensor, matrix: &Tensor) -> Result<Tenso
Ok(similarity)
}
pub fn cosine_similarity_no_l2(query_vector: &Tensor, matrix: &Tensor) -> Result<Tensor> {
// 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<Tensor> {
let x = xs.affine(1.702, 0.0)?;
let x = sigmoid(&x)?;
+25
View File
@@ -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<Vec<f32>>;
}
pub enum RerankerSimilarity {
Cosine,
}
impl RerankerSimilarity {
pub fn similar(&self, query_vector: &Tensor, matrix: &Tensor, need_l2: bool) -> Result<Tensor> {
match self {
RerankerSimilarity::Cosine => {
if need_l2 {
cosine_similarity(query_vector, matrix)
} else {
cosine_similarity_no_l2(query_vector, matrix)
}
}
}
}
}
+23 -2
View File
@@ -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<Qwen3VLGenerateModel<'a>>),
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<Vec<f32>> {
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
+27 -20
View File
@@ -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<Vec<f32>> {
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::<f32>()?;
Ok(norm)
}
}
impl TextEmbedding for Qwen3Embedding {
fn embed_texts(&mut self, input: &[String]) -> Result<Vec<Vec<f32>>> {
pub fn embed_multi(&mut self, input: &[String]) -> Result<Tensor> {
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<Tensor> {
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<Vec<Vec<f32>>> {
let embeds = self.embed_multi(input)?;
let embeds = embeds.to_vec2::<f32>()?;
Ok(embeds)
}
}
+31
View File
@@ -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<DType>) -> Result<Self> {
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<Vec<f32>> {
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::<f32>()?;
Ok(score)
}
}
+8 -8
View File
@@ -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<RerankResult>,
pub(crate) struct RerankResponse {
pub object: String,
pub model: String,
pub results: Vec<RerankResult>,
}
+6
View File
@@ -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
+75
View File
@@ -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 = "<req>")]
pub(crate) async fn rerank(req: Json<RerankRequest>) -> (Status, Json<Value>) {
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::<Vec<_>>();
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()))
}
+1 -3
View File
@@ -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;
+26
View File
@@ -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(())
}