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:
parent
4d5c1c0a83
commit
2047170540
@ -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 模型的支持
|
||||
|
||||
### 修复
|
||||
- 修复重排序模型实际未生效的问题
|
||||
|
||||
@ -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",
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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)
|
||||
|
||||
Loading…
Reference in New Issue
Block a user