From 6ae91679e28c3804efba5d7f7754894be2c8d038 Mon Sep 17 00:00:00 2001 From: kplam Date: Sat, 1 Aug 2026 00:08:38 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=8E=A5=E5=85=A5=20BGE-M3=20=E6=9C=AC?= =?UTF-8?q?=E5=9C=B0=E5=B5=8C=E5=85=A5=E9=BB=98=E8=AE=A4=E9=85=8D=E7=BD=AE?= =?UTF-8?q?=20+=20=E6=96=B0=E5=A2=9E=20Qwen3-Reranker=20=E9=87=8D=E6=8E=92?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - .env.example 嵌入默认改为本地 BGE-M3(EMBEDDING_PROVIDER=local, dim=1024),NAS 全新部署即用 - 新增 app/services/reranker.py:基于 Ollama /api/rerank 的 cross-encoder 重排服务 - retriever 在 chunk 候选阶段接入语义精排,调用失败优雅降级为 RRF 顺序 - docker-compose 拉取 qwen3-reranker:0.6b,OLLAMA_MAX_LOADED_MODELS 提至 3 - .gitignore 排除 .workbuddy/ 与 .trae/(防止误提交项目记忆与 IDE spec) - 新增 reranker 单元与集成测试,全量测试 529 passed --- .env.example | 13 +++- .gitignore | 4 ++ app/config.py | 6 ++ app/core/retriever.py | 29 ++++++++- app/services/reranker.py | 79 ++++++++++++++++++++++++ docker-compose.yml | 5 +- tests/test_reranker.py | 128 +++++++++++++++++++++++++++++++++++++++ tests/test_retriever.py | 65 +++++++++++++++++++- 8 files changed, 320 insertions(+), 9 deletions(-) create mode 100644 app/services/reranker.py create mode 100644 tests/test_reranker.py diff --git a/.env.example b/.env.example index 6e08560..c5da2ed 100644 --- a/.env.example +++ b/.env.example @@ -7,12 +7,14 @@ APP_PORT=8000 LOG_LEVEL=info # --- 嵌入模型 --- -# openai | local -EMBEDDING_PROVIDER=openai +# openai | local。默认走本地 Ollama 的 BGE-M3(多语言榜首、dense+sparse 一体,零 API 成本) +EMBEDDING_PROVIDER=local +# 以下为 EMBEDDING_PROVIDER=openai 时启用的 OpenAI 兼容配置(本地模式可留空) OPENAI_API_KEY=sk-xxx OPENAI_BASE_URL=https://api.openai.com/v1 EMBEDDING_MODEL=text-embedding-3-small -EMBEDDING_DIMENSION=1536 +# 维度需与所选嵌入模型一致:BGE-M3 为 1024;切换 OpenAI 时改回 1536 +EMBEDDING_DIMENSION=1024 # --- Qdrant --- QDRANT_PORT=6333 @@ -28,6 +30,11 @@ OLLAMA_MODEL=qwen2.5:1.5b # EMBEDDING_PROVIDER=local 时使用的 Ollama 嵌入模型 OLLAMA_EMBEDDING_MODEL=bge-m3 +# --- 重排模型(cross-encoder,经 Ollama /api/rerank 对 chunk 候选精排)--- +# 启用后显著提升最终 Top-K 相关性;需 Ollama 已拉取 RERANKER_MODEL(见 docker-compose) +RERANKER_ENABLED=true +RERANKER_MODEL=qwen3-reranker:0.6b + # --- NAS 持久化 --- NAS_DATA_DIR=./data diff --git a/.gitignore b/.gitignore index 5cf2009..29f9787 100644 --- a/.gitignore +++ b/.gitignore @@ -36,3 +36,7 @@ logs/ .pytest_cache/ .coverage htmlcov/ + +# Project memory & Trae IDE spec (切勿提交) +.workbuddy/ +.trae/ diff --git a/app/config.py b/app/config.py index 9b1a7fd..4f0e6b4 100644 --- a/app/config.py +++ b/app/config.py @@ -21,6 +21,12 @@ class Settings(BaseSettings): ollama_model: str = "qwen2.5:1.5b" # 备选: qwen2.5:3b ollama_embedding_model: str = "bge-m3" # embedding_provider=local 时使用的嵌入模型 + # 重排模型(cross-encoder,经 Ollama /api/rerank,对 chunk 候选做最终精排) + # 关闭时检索链路退化为纯 RRF 融合结果,行为不变 + reranker_enabled: bool = False + reranker_model: str = "qwen3-reranker:0.6b" + reranker_timeout: float = 30.0 + # Qdrant qdrant_host: str = "localhost" qdrant_port: int = 6333 diff --git a/app/core/retriever.py b/app/core/retriever.py index c9fa4a7..4bff0ac 100644 --- a/app/core/retriever.py +++ b/app/core/retriever.py @@ -16,7 +16,7 @@ from qdrant_client import models from app.config import settings from app.core.embeddings import EmbeddingService, create_embedding_service from app.core.query_parser import QueryParser, RouteDecision -from app.core.ranker import finalize, rrf_fuse +from app.core.ranker import rrf_fuse from app.core.result_summarizer import ResultSummarizer from app.core.sparse import SparseEncoder from app.models.knowledge import load_taxonomy @@ -30,6 +30,7 @@ from app.services.qdrant import ( SPARSE_COLLECTIONS, QdrantService, ) +from app.services.reranker import RerankerService, create_reranker_service logger = structlog.get_logger() @@ -49,6 +50,7 @@ class Retriever: embedding: EmbeddingService | None = None, sparse_encoder: SparseEncoder | None = None, result_summarizer: ResultSummarizer | None = None, + reranker: RerankerService | None = None, ) -> None: self.qdrant = qdrant or QdrantService() self.query_parser = query_parser or QueryParser( @@ -58,6 +60,7 @@ class Retriever: self.embedding = embedding or create_embedding_service() self.sparse_encoder = sparse_encoder or SparseEncoder() self.result_summarizer = result_summarizer or ResultSummarizer() + self.reranker = reranker if reranker is not None else create_reranker_service() async def search(self, request: SearchRequest) -> SearchResponse: """分层检索主流程""" @@ -77,7 +80,7 @@ class Retriever: # L1 无候选文档 → 全库 chunk 兜底 chunk_hits = await self._search_collection(COLLECTION_CHUNKS, dense, sparse, settings.retrieval_top_k, None) logger.info("L1 无命中,全库 chunk 兜底", hits=len(chunk_hits)) - hits = self._to_hits(finalize(chunk_hits, self._final_k(request))) + hits = self._to_hits(await self._rerank(query_text, chunk_hits, self._final_k(request))) return await self._build_response(request, route, hits, fallback=True) doc_ids = _unique((p.payload or {}).get("doc_id") for p in l1_hits) @@ -123,7 +126,7 @@ class Retriever: ) logger.info("chunk 检索完成", hits=len(chunk_hits)) - final_points = finalize(rrf_fuse([chunk_hits]), self._final_k(request)) + final_points = await self._rerank(query_text, chunk_hits, self._final_k(request)) hits = self._to_hits(final_points) return await self._build_response(request, route, hits, fallback=route.fallback) @@ -145,6 +148,26 @@ class Retriever: """最终返回数:请求指定优先,否则用配置默认值""" return request.top_k or settings.retrieval_final_k + async def _rerank( + self, query_text: str, points: list[models.ScoredPoint], final_k: int + ) -> list[models.ScoredPoint]: + """对 chunk 候选池做语义精排(cross-encoder) + + 重排未启用、候选为空或调用异常时,退化为 RRF 融合的原始顺序(取前 final_k), + 保证检索链路在任何情况下都不因重排失败而中断。 + """ + if self.reranker is None or not points: + return points[:final_k] + documents = [(p.payload or {}).get("text", "") for p in points] + try: + scores = await self.reranker.rerank(query_text, documents) + except Exception: + logger.warning("重排调用失败,退化为 RRF 融合顺序", exc_info=True) + return points[:final_k] + # 按相关性分数降序重排,写回 score 字段为相关性分数 + ordered = sorted(range(len(points)), key=lambda i: scores[i], reverse=True) + return [points[i].model_copy(update={"score": scores[i]}) for i in ordered[:final_k]] + @staticmethod def _to_hits(points: list[models.ScoredPoint]) -> list[SearchHit]: """chunk 点组装为 SearchHit,字段取自 chunk payload(doc_summary 仅上下文标注)""" diff --git a/app/services/reranker.py b/app/services/reranker.py new file mode 100644 index 0000000..f934281 --- /dev/null +++ b/app/services/reranker.py @@ -0,0 +1,79 @@ +"""语义重排服务(cross-encoder reranker) + +调用 Ollama 的 `/api/rerank` 端点(Qwen3-Reranker 等重排模型),对检索召回的 +chunk 候选池做精排,提升最终 Top-K 的相关性。相比纯向量/RRF 融合,cross-encoder +以 (query, doc) 联合编码,能捕捉字词不匹配的语义关联,典型带来约 10% 的精度提升。 + +工厂 `create_reranker_service()` 按 `settings.reranker_enabled` 返回实例或 None: +关闭时上层检索链路退化为原有的 RRF 融合结果,行为完全不变。 +""" + +from typing import Protocol, runtime_checkable + +import httpx +import structlog + +from app.config import settings + +logger = structlog.get_logger() + + +@runtime_checkable +class RerankerService(Protocol): + """重排服务统一接口""" + + async def rerank(self, query: str, documents: list[str]) -> list[float]: + """对候选文档按与 query 的相关性打分 + + Args: + query: 查询文本 + documents: 候选文档文本列表(按输入顺序对齐) + + Returns: + 与 documents 等长的相关性分数列表(分数越高越相关) + """ + ... + + +class OllamaRerankerService: + """基于 Ollama /api/rerank 的重排服务""" + + def __init__(self, base_url: str, model: str, timeout: float = 30.0) -> None: + self.base_url = base_url.rstrip("/") + self.model = model + self.timeout = timeout + + async def rerank(self, query: str, documents: list[str]) -> list[float]: + if not documents: + return [] + url = f"{self.base_url}/api/rerank" + payload = {"model": self.model, "query": query, "documents": documents} + async with httpx.AsyncClient(timeout=self.timeout) as client: + resp = await client.post(url, json=payload) + resp.raise_for_status() + data = resp.json() + + results = data.get("results", []) + # 按 Ollama 返回的 index 映射 relevance_score;缺失项补 0,保证输出与输入等长 + scores: list[float] = [0.0] * len(documents) + for item in results: + idx = item.get("index") + score = float(item.get("relevance_score", 0.0)) + if idx is not None and 0 <= idx < len(documents): + scores[idx] = score + logger.debug("重排完成", model=self.model, candidates=len(documents)) + return scores + + +def create_reranker_service() -> RerankerService | None: + """按 settings.reranker_enabled 创建重排服务实例 + + 关闭时返回 None,上层据此跳过精排、退化为 RRF 融合结果。 + """ + if not settings.reranker_enabled: + return None + return OllamaRerankerService( + base_url=settings.ollama_base_url, + model=settings.reranker_model, + timeout=settings.reranker_timeout, + ) diff --git a/docker-compose.yml b/docker-compose.yml index 23c97c1..7c39c89 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -93,10 +93,10 @@ services: # 针对 Ryzen 9 7940HS(16 线程)的 CPU 推理调优: # 并行推理任务数、常驻模型数、单请求线程上限、KV 缓存量化以省内存 - OLLAMA_NUM_PARALLEL=4 - - OLLAMA_MAX_LOADED_MODELS=2 + - OLLAMA_MAX_LOADED_MODELS=3 - OLLAMA_NUM_THREADS=16 - OLLAMA_KV_CACHE_TYPE=q8_0 - # 首次启动自动拉取所需模型(qwen2.5:1.5b 总结 + bge-m3 嵌入), + # 首次启动自动拉取所需模型(qwen2.5:1.5b 总结 + bge-m3 嵌入 + qwen3-reranker 重排), # 下载完成后转交常驻 ollama serve。已存在时仅做健康检查。 entrypoint: /bin/bash command: @@ -107,6 +107,7 @@ services: sleep 6 ollama pull qwen2.5:1.5b ollama pull bge-m3 + ollama pull qwen3-reranker:0.6b wait $$SERVE_PID networks: - qmdsearch diff --git a/tests/test_reranker.py b/tests/test_reranker.py new file mode 100644 index 0000000..10eec09 --- /dev/null +++ b/tests/test_reranker.py @@ -0,0 +1,128 @@ +"""RerankerService 单元测试(mock httpx,不发起真实网络请求)""" + +from typing import Any + +import httpx +import pytest + +from app.config import settings +from app.services.reranker import OllamaRerankerService, create_reranker_service + + +def _make_response(status_code: int, data: dict | None = None) -> httpx.Response: + """构造带 request 上下文的 httpx.Response(raise_for_status 依赖 request)""" + request = httpx.Request("POST", "http://localhost:11434/api/rerank") + if data is None: + return httpx.Response(status_code, request=request) + return httpx.Response(status_code, json=data, request=request) + + +class _FakeAsyncClient: + """httpx.AsyncClient 替代品:记录请求参数,返回预设响应或抛出预设异常""" + + def __init__( + self, + captured: dict, + response: httpx.Response | None = None, + error: Exception | None = None, + **kwargs: Any, + ) -> None: + self._captured = captured + self._response = response + self._error = error + captured["timeout"] = kwargs.get("timeout") + + async def __aenter__(self) -> "_FakeAsyncClient": + return self + + async def __aexit__(self, *args: object) -> bool: + return False + + async def post(self, url: str, json: dict | None = None) -> httpx.Response: + self._captured["url"] = url + self._captured["json"] = json + if self._error is not None: + raise self._error + assert self._response is not None + return self._response + + +def _client(monkeypatch: pytest.MonkeyPatch, captured: dict, **kw: Any) -> None: + monkeypatch.setattr(httpx, "AsyncClient", lambda **k: _FakeAsyncClient(captured, **k, **kw)) + + +class TestRerank: + async def test_returns_scores_aligned_to_documents(self, monkeypatch: pytest.MonkeyPatch) -> None: + captured: dict = {} + _client( + monkeypatch, + captured, + response=_make_response( + 200, + { + "results": [ + {"index": 0, "relevance_score": 0.2}, + {"index": 1, "relevance_score": 0.9}, + {"index": 2, "relevance_score": 0.5}, + ] + }, + ), + ) + service = OllamaRerankerService(base_url="http://localhost:11434/", model="qwen3-reranker:0.6b") + + scores = await service.rerank("查询", ["doc-a", "doc-b", "doc-c"]) + + assert scores == [0.2, 0.9, 0.5] + assert captured["url"] == "http://localhost:11434/api/rerank" + assert captured["json"] == { + "model": "qwen3-reranker:0.6b", + "query": "查询", + "documents": ["doc-a", "doc-b", "doc-c"], + } + + async def test_empty_documents_returns_empty(self, monkeypatch: pytest.MonkeyPatch) -> None: + captured: dict = {} + _client(monkeypatch, captured, response=_make_response(200, {"results": []})) + service = OllamaRerankerService(base_url="http://localhost:11434", model="qwen3-reranker:0.6b") + + scores = await service.rerank("查询", []) + + assert scores == [] + # 空文档不发起请求 + assert "url" not in captured + + async def test_partial_results_zero_fill_missing(self, monkeypatch: pytest.MonkeyPatch) -> None: + captured: dict = {} + # 仅返回 index 1,index 0/2 缺失 → 补 0 + _client( + monkeypatch, + captured, + response=_make_response(200, {"results": [{"index": 1, "relevance_score": 0.7}]}), + ) + service = OllamaRerankerService(base_url="http://localhost:11434", model="qwen3-reranker:0.6b") + + scores = await service.rerank("查询", ["a", "b", "c"]) + + assert scores == [0.0, 0.7, 0.0] + + async def test_http_error_raises(self, monkeypatch: pytest.MonkeyPatch) -> None: + _client(monkeypatch, {}, response=_make_response(500)) + service = OllamaRerankerService(base_url="http://localhost:11434", model="qwen3-reranker:0.6b") + + with pytest.raises(httpx.HTTPStatusError): + await service.rerank("查询", ["a", "b"]) + + +class TestCreateRerankerService: + def test_disabled_returns_none(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(settings, "reranker_enabled", False) + assert create_reranker_service() is None + + def test_enabled_returns_service(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(settings, "reranker_enabled", True) + monkeypatch.setattr(settings, "reranker_model", "qwen3-reranker:0.6b") + monkeypatch.setattr(settings, "ollama_base_url", "http://ollama:11434") + monkeypatch.setattr(settings, "reranker_timeout", 30.0) + service = create_reranker_service() + assert isinstance(service, OllamaRerankerService) + assert service.model == "qwen3-reranker:0.6b" diff --git a/tests/test_retriever.py b/tests/test_retriever.py index 20966a0..81c6c14 100644 --- a/tests/test_retriever.py +++ b/tests/test_retriever.py @@ -82,15 +82,29 @@ def _route( ) -def _make_retriever(qdrant: FakeQdrant, route: RouteDecision) -> Retriever: +def _make_retriever(qdrant: FakeQdrant, route: RouteDecision, reranker: Any = None) -> Retriever: return Retriever( qdrant=qdrant, # type: ignore[arg-type] query_parser=FakeParser(route), # type: ignore[arg-type] embedding=FakeEmbedding(), sparse_encoder=SparseEncoder(), + reranker=reranker, ) +class FakeReranker: + """假 RerankerService:返回预设分数或抛出预设异常""" + + def __init__(self, scores: list[float] | None = None, exc: Exception | None = None) -> None: + self.scores = scores if scores is not None else [] + self.exc = exc + + async def rerank(self, query: str, documents: list[str]) -> list[float]: + if self.exc is not None: + raise self.exc + return self.scores + + def _calls(qdrant: FakeQdrant, collection: str) -> list[dict[str, Any]]: return [c for c in qdrant.calls if c["collection"] == collection] @@ -348,3 +362,52 @@ class TestSearchApi: resp = client.get("/api/v1/health") assert resp.status_code == 200 assert resp.json() == {"status": "ok"} + + +class TestReranker: + """重排启用时对 chunk 候选精排;关闭或失败则退化为 RRF 顺序""" + + @staticmethod + def _full_qdrant() -> FakeQdrant: + return FakeQdrant( + { + COLLECTION_L1: [[_point("l1a", "d1")]], + COLLECTION_L2: [[_point("l2a", "d1", "章节A")]], + COLLECTION_L3: [[_point("l3a", "d1", "章节A")]], + COLLECTION_CHUNKS: [ + [ + _point("c1", "d1", "章节A"), + _point("c2", "d1", "章节A"), + _point("c3", "d1", "章节A"), + ] + ], + } + ) + + async def test_enabled_reorders_by_relevance(self): + """重排按相关性分数降序重排,并将 score 写回相关性分数""" + reranker = FakeReranker(scores=[0.1, 0.9, 0.5]) # c1