appQQQ commited on
Commit
b5d081e
·
verified ·
1 Parent(s): 8ef10ef

perf: skip CPU cross-encoder reranking by default

Browse files
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
- """BGE-reranker 精排 top-N + 产出引用."""
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
- reranker = get_reranker_service()
224
- reranked = await reranker.rerank(query, hits, top_n=settings.rerank_top_n)
 
 
 
 
 
 
 
 
 
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
- if top_score >= settings.crag_relevance_threshold:
 
 
 
 
 
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 开销. 召回率轻微降, reranker 兜底.
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
- warm_reranker()
 
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))