diff --git a/docs/latest/changelog/roadmap.md b/docs/latest/changelog/roadmap.md index 35c65ab6..b770fc2c 100644 --- a/docs/latest/changelog/roadmap.md +++ b/docs/latest/changelog/roadmap.md @@ -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 模型的支持 ### 修复 - 修复重排序模型实际未生效的问题 diff --git a/src/config/static/models.py b/src/config/static/models.py index 4874c25e..f6b9a069 100644 --- a/src/config/static/models.py +++ b/src/config/static/models.py @@ -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", diff --git a/src/models/embed.py b/src/models/embed.py index 528e1ccf..7d231c9f 100644 --- a/src/models/embed.py +++ b/src/models/embed.py @@ -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: diff --git a/src/models/rerank.py b/src/models/rerank.py index 3b42124b..3be61dc3 100644 --- a/src/models/rerank.py +++ b/src/models/rerank.py @@ -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)