feat: 新增 dashscope rerank/embeddings 模型支持并重构 reranker 基类 #341

- 在 models.py 中添加 dashscope 的 text-embedding-v4 和 gte-rerank-v2/qwen3-rerank 模型配置
- 重构 reranker 模块,提取 BaseReranker 抽象基类并实现 OpenAIReranker 和 DashscopeReranker
- 在 embed.py 中优化错误提示信息并统一测试连接方法
- 更新 roadmap.md 文档记录新增功能
This commit is contained in:
Wenjie Zhang 2025-11-17 23:04:43 +08:00
parent 4d5c1c0a83
commit 2047170540
4 changed files with 110 additions and 48 deletions

View File

@ -27,7 +27,7 @@
- 新增基于知识库文件生成思维导图功能([#335](https://github.com/xerrors/Yuxi-Know/pull/335#issuecomment-3530976425)
- 新增基于知识库文件生成示例问题功能([#335](https://github.com/xerrors/Yuxi-Know/pull/335#issuecomment-3530976425)
- 新增知识库支持文件夹/压缩包上传的功能([#335](https://github.com/xerrors/Yuxi-Know/pull/335#issuecomment-3530976425)
- 新增自定义模型支持
- 新增自定义模型支持、新增 dashscope rerank/embeddings 模型的支持
### 修复
- 修复重排序模型实际未生效的问题

View File

@ -129,6 +129,29 @@ DEFAULT_CHAT_MODEL_PROVIDERS: dict[str, ChatModelProvider] = {
"anthropic/claude-sonnet-4",
],
),
# "moonshot": ChatModelProvider(
# name="月之暗面",
# url="https://platform.moonshot.cn/docs/overview",
# base_url="https://api.moonshot.cn/v1",
# default="kimi-latest",
# env="MOONSHOT_API_KEY",
# models=[
# "kimi-latest",
# "kimi-k2-thinking",
# "kimi-k2-0905-preview",
# ],
# ), # 目前适配有问题 Error code: 400 - {'error': {'message': 'Invalid request: function name is invalid, must start with a letter and can contain letters, numbers, underscores, and dashes', 'type': 'invalid_request_error'}} # noqa: E501
"modelscope": ChatModelProvider(
name="ModelScope",
url="https://www.modelscope.cn/docs/model-service/API-Inference/intro",
base_url="https://api-inference.modelscope.cn/v1/",
default="deepseek-ai/DeepSeek-V3.2-Exp",
env="MODELSCOPE_ACCESS_TOKEN",
models=[
"Qwen/Qwen3-32B",
"deepseek-ai/DeepSeek-V3.2-Exp"
],
),
}
@ -173,6 +196,12 @@ DEFAULT_EMBED_MODELS: dict[str, EmbedModelInfo] = {
base_url="http://localhost:11434/api/embed",
api_key="no_api_key",
),
"dashscope/text-embedding-v4": EmbedModelInfo(
name="text-embedding-v4",
dimension=1024,
base_url="https://dashscope.aliyuncs.com/compatible-mode/v1/embeddings",
api_key="DASHSCOPE_API_KEY",
),
}
@ -191,6 +220,16 @@ DEFAULT_RERANKERS: dict[str, RerankerInfo] = {
base_url="https://api.siliconflow.cn/v1/rerank",
api_key="SILICONFLOW_API_KEY",
),
"dashscope/gte-rerank-v2": RerankerInfo(
name="gte-rerank-v2",
base_url="https://dashscope.aliyuncs.com/api/v1/services/rerank/text-rerank/text-rerank",
api_key="DASHSCOPE_API_KEY",
),
"dashscope/qwen3-rerank": RerankerInfo(
name="qwen3-rerank",
base_url="https://dashscope.aliyuncs.com/api/v1/services/rerank/text-rerank/text-rerank",
api_key="DASHSCOPE_API_KEY",
),
"vllm/BAAI/bge-reranker-v2-m3": RerankerInfo(
name="BAAI/bge-reranker-v2-m3",
base_url="http://localhost:8000/v1/rerank",

View File

@ -89,6 +89,23 @@ class BaseEmbeddingModel(ABC):
return data
async def test_connection(self) -> tuple[bool, str]:
"""
测试embedding模型的连接性
Returns:
tuple: (success: bool, message: str)
"""
try:
# 使用简单的测试文本
test_text = ["Hello world"]
await self.aencode(test_text)
return True, "连接正常"
except Exception as e:
error_msg = str(e)
error_msg += f", maybe you can check the `{self.base_url}` end with /embeddings as examples."
logger.error(error_msg)
return False, error_msg
class OllamaEmbedding(BaseEmbeddingModel):
"""
@ -129,24 +146,7 @@ class OllamaEmbedding(BaseEmbeddingModel):
raise ValueError(f"Ollama Embedding failed: Invalid response format {result}")
return result["embeddings"]
except (httpx.RequestError, json.JSONDecodeError) as e:
logger.error(f"Ollama Embedding async request failed: {e}, {payload}")
raise ValueError(f"Ollama Embedding async request failed: {e}")
async def test_connection(self) -> tuple[bool, str]:
"""
测试embedding模型的连接性
Returns:
tuple: (success: bool, message: str)
"""
try:
# 使用简单的测试文本
test_text = ["Hello world"]
await self.aencode(test_text)
return True, "连接正常"
except Exception as e:
error_msg = str(e)
return False, error_msg
raise ValueError(f"Ollama Embedding async request failed: {e}, {payload}, {self.base_url=}")
class OtherEmbedding(BaseEmbeddingModel):
@ -181,24 +181,8 @@ class OtherEmbedding(BaseEmbeddingModel):
raise ValueError(f"Other Embedding failed: Invalid response format {result}")
return [item["embedding"] for item in result["data"]]
except (httpx.RequestError, json.JSONDecodeError) as e:
logger.error(f"Other Embedding async request failed: {e}, {payload}")
raise ValueError(f"Other Embedding async request failed: {e}")
raise ValueError(f"Other Embedding async request failed: {e}, {payload}, {self.base_url=}")
async def test_connection(self) -> tuple[bool, str]:
"""
测试embedding模型的连接性
Returns:
tuple: (success: bool, message: str)
"""
try:
# 使用简单的测试文本
test_text = ["Hello world"]
await self.aencode(test_text)
return True, "连接正常"
except Exception as e:
error_msg = str(e)
return False, error_msg
async def test_embedding_model_status(model_id: str) -> dict:

View File

@ -1,5 +1,6 @@
import asyncio
import os
from abc import ABC, abstractmethod
from collections.abc import Iterable, Sequence
from typing import Any
@ -14,7 +15,7 @@ def sigmoid(x):
return 1 / (1 + np.exp(-x))
class OnlineReranker:
class BaseReranker(ABC):
def __init__(self, model_name, api_key, base_url, **kwargs):
self.url = get_docker_safe_url(base_url)
self.model = model_name
@ -22,11 +23,20 @@ class OnlineReranker:
self.headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
self.session: aiohttp.ClientSession | None = None
self.timeout = aiohttp.ClientTimeout(total=30)
self.parameters: dict[str, Any] = dict(kwargs.get("parameters", {}))
async def _ensure_session(self) -> None:
if self.session is None or self.session.closed:
self.session = aiohttp.ClientSession(headers=self.headers, timeout=self.timeout)
@abstractmethod
def _build_payload(self, query: str, documents: list[str], max_length: int) -> dict[str, Any]:
raise NotImplementedError
@abstractmethod
def _extract_results(self, result: dict[str, Any]) -> list[dict[str, Any]]:
raise NotImplementedError
async def acompute_score(
self,
sentence_pairs: Sequence[Sequence[str]],
@ -65,16 +75,12 @@ class OnlineReranker:
return all_scores
async def _batch_rerank(self, query: str, documents: Iterable[str], max_length: int) -> list[float]:
payload = {
"model": self.model,
"query": query,
"documents": list(documents),
"max_chunks_per_doc": max_length,
}
if not payload["documents"]:
docs = list(documents)
if not docs:
return []
payload = self._build_payload(query, docs, max_length)
await self._ensure_session()
assert self.session is not None
@ -90,11 +96,10 @@ class OnlineReranker:
logger.error(f"Reranking request failed: {exc}")
raise exc
processed = sorted(result.get("results", []), key=lambda item: item.get("index", 0))
processed = sorted(self._extract_results(result), key=lambda item: item.get("index", 0))
return [float(entry.get("relevance_score", 0.0)) for entry in processed]
def compute_score(self, sentence_pairs, batch_size=256, max_length=512, normalize=False):
"""Synchronous helper retained for backwards compatibility."""
try:
_ = asyncio.get_running_loop()
except RuntimeError:
@ -119,6 +124,35 @@ class OnlineReranker:
loop.run_until_complete(self.aclose())
class OpenAIReranker(BaseReranker):
def _build_payload(self, query: str, documents: list[str], max_length: int) -> dict[str, Any]:
return {
"model": self.model,
"query": query,
"documents": documents,
"max_chunks_per_doc": max_length,
}
def _extract_results(self, result: dict[str, Any]) -> list[dict[str, Any]]:
return list(result.get("results", []))
class DashscopeReranker(BaseReranker):
def _build_payload(self, query: str, documents: list[str], max_length: int) -> dict[str, Any]:
params = {"top_n": len(documents), "return_documents": False}
instruct = self.parameters.get("instruct")
if instruct:
params["instruct"] = instruct
return {
"model": self.model,
"input": {"query": query, "documents": documents},
"parameters": params,
}
def _extract_results(self, result: dict[str, Any]) -> list[dict[str, Any]]:
return list(result.get("output", {}).get("results", []))
def get_reranker(model_id, **kwargs):
support_rerankers = config.reranker_names.keys()
assert model_id in support_rerankers, f"Unsupported Reranker: {model_id}, only support {support_rerankers}"
@ -127,4 +161,9 @@ def get_reranker(model_id, **kwargs):
base_url = model_info.base_url
api_key = os.getenv(model_info.api_key) or model_info.api_key
assert api_key, f"{model_info.name} api_key is required"
return OnlineReranker(model_name=model_info.name, api_key=api_key, base_url=base_url, **kwargs)
provider = model_id.split("/", maxsplit=1)[0] if "/" in model_id else ""
if provider in {"siliconflow", "vllm"}:
return OpenAIReranker(model_name=model_info.name, api_key=api_key, base_url=base_url, **kwargs)
if provider == "dashscope":
return DashscopeReranker(model_name=model_info.name, api_key=api_key, base_url=base_url, **kwargs)
return OpenAIReranker(model_name=model_info.name, api_key=api_key, base_url=base_url, **kwargs)