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) - 新增基于知识库文件生成示例问题功能([#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", "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", base_url="http://localhost:11434/api/embed",
api_key="no_api_key", 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", base_url="https://api.siliconflow.cn/v1/rerank",
api_key="SILICONFLOW_API_KEY", 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( "vllm/BAAI/bge-reranker-v2-m3": RerankerInfo(
name="BAAI/bge-reranker-v2-m3", name="BAAI/bge-reranker-v2-m3",
base_url="http://localhost:8000/v1/rerank", base_url="http://localhost:8000/v1/rerank",

View File

@ -89,6 +89,23 @@ class BaseEmbeddingModel(ABC):
return data 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): class OllamaEmbedding(BaseEmbeddingModel):
""" """
@ -129,24 +146,7 @@ class OllamaEmbedding(BaseEmbeddingModel):
raise ValueError(f"Ollama Embedding failed: Invalid response format {result}") raise ValueError(f"Ollama Embedding failed: Invalid response format {result}")
return result["embeddings"] return result["embeddings"]
except (httpx.RequestError, json.JSONDecodeError) as e: 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}, {payload}, {self.base_url=}")
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
class OtherEmbedding(BaseEmbeddingModel): class OtherEmbedding(BaseEmbeddingModel):
@ -181,24 +181,8 @@ class OtherEmbedding(BaseEmbeddingModel):
raise ValueError(f"Other Embedding failed: Invalid response format {result}") raise ValueError(f"Other Embedding failed: Invalid response format {result}")
return [item["embedding"] for item in result["data"]] return [item["embedding"] for item in result["data"]]
except (httpx.RequestError, json.JSONDecodeError) as e: 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}, {payload}, {self.base_url=}")
raise ValueError(f"Other 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
async def test_embedding_model_status(model_id: str) -> dict: async def test_embedding_model_status(model_id: str) -> dict:

View File

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