feat: 接入 BGE-M3 本地嵌入默认配置 + 新增 Qwen3-Reranker 重排
- .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
This commit is contained in:
+10
-3
@@ -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
|
||||
|
||||
|
||||
@@ -36,3 +36,7 @@ logs/
|
||||
.pytest_cache/
|
||||
.coverage
|
||||
htmlcov/
|
||||
|
||||
# Project memory & Trae IDE spec (切勿提交)
|
||||
.workbuddy/
|
||||
.trae/
|
||||
|
||||
@@ -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
|
||||
|
||||
+26
-3
@@ -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 仅上下文标注)"""
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
+3
-2
@@ -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
|
||||
|
||||
@@ -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"
|
||||
+64
-1
@@ -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<c3<c2
|
||||
retriever = _make_retriever(self._full_qdrant(), _route(), reranker=reranker)
|
||||
|
||||
resp = await retriever.search(SearchRequest(query="测试查询"))
|
||||
|
||||
# 重排后顺序为 c2, c3, c1
|
||||
assert [h.text for h in resp.hits] == ["text-c2", "text-c3", "text-c1"]
|
||||
assert resp.hits[0].score == 0.9
|
||||
|
||||
async def test_disabled_keeps_rrf_order(self):
|
||||
"""重排未注入(关闭)时保持 RRF 融合原始顺序"""
|
||||
retriever = _make_retriever(self._full_qdrant(), _route(), reranker=None)
|
||||
|
||||
resp = await retriever.search(SearchRequest(query="测试查询"))
|
||||
|
||||
assert [h.text for h in resp.hits] == ["text-c1", "text-c2", "text-c3"]
|
||||
|
||||
async def test_failure_falls_back_to_rrf(self):
|
||||
"""重排调用异常时退化为 RRF 顺序,检索不中断"""
|
||||
reranker = FakeReranker(exc=RuntimeError("ollama 不可用"))
|
||||
retriever = _make_retriever(self._full_qdrant(), _route(), reranker=reranker)
|
||||
|
||||
resp = await retriever.search(SearchRequest(query="测试查询"))
|
||||
|
||||
assert [h.text for h in resp.hits] == ["text-c1", "text-c2", "text-c3"]
|
||||
|
||||
Reference in New Issue
Block a user