Spaces:
Sleeping
Sleeping
perf: skip CPU cross-encoder reranking by default
Browse files- app/agents/nodes.py +18 -4
- app/config.py +5 -1
- app/main.py +3 -2
- app/services/vector_store.py +6 -0
app/agents/nodes.py
CHANGED
|
@@ -213,15 +213,24 @@ async def retrieve_node(state: AgentState) -> dict[str, Any]:
|
|
| 213 |
|
| 214 |
# ========== Node 4: rerank ==========
|
| 215 |
async def rerank_node(state: AgentState) -> dict[str, Any]:
|
| 216 |
-
"""
|
| 217 |
started = time.time()
|
| 218 |
hits = state.get("retrieved") or []
|
| 219 |
query = state.get("query_rewritten") or _last_user_query(state)
|
| 220 |
if not hits:
|
| 221 |
return {"reranked": [], "citations": [], "relevance_score": 0.0, "relevance_verdict": "irrelevant"}
|
| 222 |
|
| 223 |
-
|
| 224 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 225 |
|
| 226 |
# 构造引用 (前 5 个, 按 rerank 分数)
|
| 227 |
citations: list[dict[str, Any]] = []
|
|
@@ -238,7 +247,12 @@ async def rerank_node(state: AgentState) -> dict[str, Any]:
|
|
| 238 |
})
|
| 239 |
|
| 240 |
top_score = reranked[0].rerank_score if reranked else 0.0
|
| 241 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 242 |
verdict = "relevant"
|
| 243 |
elif top_score < 0.3:
|
| 244 |
verdict = "irrelevant"
|
|
|
|
| 213 |
|
| 214 |
# ========== Node 4: rerank ==========
|
| 215 |
async def rerank_node(state: AgentState) -> dict[str, Any]:
|
| 216 |
+
"""可选 CrossEncoder 精排;CPU 默认沿用混合检索排序并产出引用。"""
|
| 217 |
started = time.time()
|
| 218 |
hits = state.get("retrieved") or []
|
| 219 |
query = state.get("query_rewritten") or _last_user_query(state)
|
| 220 |
if not hits:
|
| 221 |
return {"reranked": [], "citations": [], "relevance_score": 0.0, "relevance_verdict": "irrelevant"}
|
| 222 |
|
| 223 |
+
if settings.enable_reranker:
|
| 224 |
+
reranker = get_reranker_service()
|
| 225 |
+
reranked = await reranker.rerank(query, hits, top_n=settings.rerank_top_n)
|
| 226 |
+
else:
|
| 227 |
+
# BGE-M3 dense + sparse 已在 hybrid_query 中完成 RRF 排序。免费 CPU Space
|
| 228 |
+
# 上再跑 CrossEncoder 实测会额外阻塞约 73 秒;直接取 top-N,并用保留下来的
|
| 229 |
+
# dense cosine 作为可解释的相关性分数。GPU 部署可通过环境变量恢复精排。
|
| 230 |
+
reranked = list(hits[: settings.rerank_top_n])
|
| 231 |
+
for rank, hit in enumerate(reranked):
|
| 232 |
+
hit.original_rank = rank
|
| 233 |
+
hit.rerank_score = max(0.0, min(1.0, float(hit.dense_score)))
|
| 234 |
|
| 235 |
# 构造引用 (前 5 个, 按 rerank 分数)
|
| 236 |
citations: list[dict[str, Any]] = []
|
|
|
|
| 247 |
})
|
| 248 |
|
| 249 |
top_score = reranked[0].rerank_score if reranked else 0.0
|
| 250 |
+
relevance_threshold = (
|
| 251 |
+
settings.crag_relevance_threshold
|
| 252 |
+
if settings.enable_reranker
|
| 253 |
+
else settings.hybrid_relevance_threshold
|
| 254 |
+
)
|
| 255 |
+
if top_score >= relevance_threshold:
|
| 256 |
verdict = "relevant"
|
| 257 |
elif top_score < 0.3:
|
| 258 |
verdict = "irrelevant"
|
app/config.py
CHANGED
|
@@ -61,12 +61,15 @@ class Settings(BaseSettings):
|
|
| 61 |
reranker_model: str = "BAAI/bge-reranker-v2-m3"
|
| 62 |
embedding_device: Literal["cpu", "cuda", "mps"] = "cpu"
|
| 63 |
use_fp16: bool = True
|
|
|
|
|
|
|
|
|
|
| 64 |
|
| 65 |
# ========== 向量库 / ChromaDB ==========
|
| 66 |
chroma_persist_dir: str = "./data/chroma"
|
| 67 |
chroma_collection: str = "docs"
|
| 68 |
# ⚡ A 改良版: 默认关闭 ColBERT. 三路融合 (dense+sparse+colbert) 每次查询
|
| 69 |
-
# 多扫 30 个 .npy 文件 + matmul, 0.3-1.5s 开销.
|
| 70 |
# 想恢复三路融合, 设环境变量 ENABLE_COLBERT=true.
|
| 71 |
enable_colbert: bool = False
|
| 72 |
|
|
@@ -97,6 +100,7 @@ class Settings(BaseSettings):
|
|
| 97 |
rerank_top_n: int = 5
|
| 98 |
crag_max_iterations: int = 2
|
| 99 |
crag_relevance_threshold: float = 0.7
|
|
|
|
| 100 |
|
| 101 |
# ========== LLM 缓存 ==========
|
| 102 |
llm_cache_enabled: bool = True
|
|
|
|
| 61 |
reranker_model: str = "BAAI/bge-reranker-v2-m3"
|
| 62 |
embedding_device: Literal["cpu", "cuda", "mps"] = "cpu"
|
| 63 |
use_fp16: bool = True
|
| 64 |
+
# HF 免费 CPU 上对 20 个候选做 CrossEncoder 重排实测约 73 秒。
|
| 65 |
+
# 默认使用 BGE-M3 dense+sparse 混合排序;GPU Space 可设 ENABLE_RERANKER=true。
|
| 66 |
+
enable_reranker: bool = False
|
| 67 |
|
| 68 |
# ========== 向量库 / ChromaDB ==========
|
| 69 |
chroma_persist_dir: str = "./data/chroma"
|
| 70 |
chroma_collection: str = "docs"
|
| 71 |
# ⚡ A 改良版: 默认关闭 ColBERT. 三路融合 (dense+sparse+colbert) 每次查询
|
| 72 |
+
# 多扫 30 个 .npy 文件 + matmul, 0.3-1.5s 开销. 默认保留 dense+sparse 双路融合.
|
| 73 |
# 想恢复三路融合, 设环境变量 ENABLE_COLBERT=true.
|
| 74 |
enable_colbert: bool = False
|
| 75 |
|
|
|
|
| 100 |
rerank_top_n: int = 5
|
| 101 |
crag_max_iterations: int = 2
|
| 102 |
crag_relevance_threshold: float = 0.7
|
| 103 |
+
hybrid_relevance_threshold: float = 0.45
|
| 104 |
|
| 105 |
# ========== LLM 缓存 ==========
|
| 106 |
llm_cache_enabled: bool = True
|
app/main.py
CHANGED
|
@@ -58,7 +58,7 @@ async def lifespan(app: FastAPI):
|
|
| 58 |
except Exception as e: # noqa: BLE001
|
| 59 |
logger.warning("ChromaDB init failed (will retry on first request): %s", e)
|
| 60 |
|
| 61 |
-
# BGE-M3 & Reranker 预热.
|
| 62 |
# 必须在主线程 *同步* 执行, 不能 run_in_executor 异步跑 — 因为 FlagEmbedding
|
| 63 |
# 在 torch 2.2 + transformers 4.57 组合下, 子线程首次 .to(device) 会撞 meta tensor
|
| 64 |
# 错 "Cannot copy out of meta tensor; no data!"; 必须等它 meta→cpu 转移完再 await
|
|
@@ -66,7 +66,8 @@ async def lifespan(app: FastAPI):
|
|
| 66 |
if settings.embedding_model and not os.environ.get("TESTING"):
|
| 67 |
try:
|
| 68 |
warm_embedder()
|
| 69 |
-
|
|
|
|
| 70 |
except Exception as e: # noqa: BLE001
|
| 71 |
logger.warning("Embedder/reranker warm-up failed: %s", e)
|
| 72 |
|
|
|
|
| 58 |
except Exception as e: # noqa: BLE001
|
| 59 |
logger.warning("ChromaDB init failed (will retry on first request): %s", e)
|
| 60 |
|
| 61 |
+
# BGE-M3 & 可选 Reranker 预热.
|
| 62 |
# 必须在主线程 *同步* 执行, 不能 run_in_executor 异步跑 — 因为 FlagEmbedding
|
| 63 |
# 在 torch 2.2 + transformers 4.57 组合下, 子线程首次 .to(device) 会撞 meta tensor
|
| 64 |
# 错 "Cannot copy out of meta tensor; no data!"; 必须等它 meta→cpu 转移完再 await
|
|
|
|
| 66 |
if settings.embedding_model and not os.environ.get("TESTING"):
|
| 67 |
try:
|
| 68 |
warm_embedder()
|
| 69 |
+
if settings.enable_reranker:
|
| 70 |
+
warm_reranker()
|
| 71 |
except Exception as e: # noqa: BLE001
|
| 72 |
logger.warning("Embedder/reranker warm-up failed: %s", e)
|
| 73 |
|
app/services/vector_store.py
CHANGED
|
@@ -371,6 +371,9 @@ def hybrid_query(
|
|
| 371 |
|
| 372 |
# RRF 融合
|
| 373 |
fused = rrf_fuse(dense, sparse, colbert_ranked, k=60)[:k]
|
|
|
|
|
|
|
|
|
|
| 374 |
|
| 375 |
# 构造 RetrievalHit
|
| 376 |
hits: list[RetrievalHit] = []
|
|
@@ -385,6 +388,9 @@ def hybrid_query(
|
|
| 385 |
heading=meta.get("heading"),
|
| 386 |
context_prefix=meta.get("context_prefix"),
|
| 387 |
meta=meta,
|
|
|
|
|
|
|
|
|
|
| 388 |
))
|
| 389 |
|
| 390 |
logger.debug("hybrid_query returned %d hits in %dms", len(hits), int((time.time() - started) * 1000))
|
|
|
|
| 371 |
|
| 372 |
# RRF 融合
|
| 373 |
fused = rrf_fuse(dense, sparse, colbert_ranked, k=60)[:k]
|
| 374 |
+
dense_scores = {cid: score for cid, score, _payload in dense}
|
| 375 |
+
sparse_scores_by_id = {cid: score for cid, score, _payload in sparse}
|
| 376 |
+
colbert_scores = {cid: score for cid, score, _payload in colbert_ranked}
|
| 377 |
|
| 378 |
# 构造 RetrievalHit
|
| 379 |
hits: list[RetrievalHit] = []
|
|
|
|
| 388 |
heading=meta.get("heading"),
|
| 389 |
context_prefix=meta.get("context_prefix"),
|
| 390 |
meta=meta,
|
| 391 |
+
dense_score=float(dense_scores.get(cid, 0.0)),
|
| 392 |
+
sparse_score=float(sparse_scores_by_id.get(cid, 0.0)),
|
| 393 |
+
colbert_score=float(colbert_scores.get(cid, 0.0)),
|
| 394 |
))
|
| 395 |
|
| 396 |
logger.debug("hybrid_query returned %d hits in %dms", len(hits), int((time.time() - started) * 1000))
|