appQQQ commited on
Commit
09da783
·
verified ·
1 Parent(s): 356cfcc

chore: upload app/services/reranker.py

Browse files
Files changed (1) hide show
  1. app/services/reranker.py +107 -0
app/services/reranker.py ADDED
@@ -0,0 +1,107 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """BGE-reranker-v2-m3 精排服务.
2
+
3
+ 输入: query + list[RetrievalHit]
4
+ 输出: 同长度 list, 按相关性分数重排, 返回 top_n
5
+ """
6
+ from __future__ import annotations
7
+
8
+ import asyncio
9
+ import logging
10
+ import threading
11
+ from functools import lru_cache
12
+
13
+ from app.config import settings
14
+ from app.services.vector_store import RetrievalHit
15
+
16
+ logger = logging.getLogger(__name__)
17
+
18
+
19
+ @lru_cache(maxsize=1)
20
+ def get_reranker():
21
+ """懒加载 FlagReranker. CPU 上 fp16 强制 fp32."""
22
+ from FlagEmbedding import FlagReranker
23
+
24
+ use_fp16 = settings.use_fp16 and settings.embedding_device in ("cuda", "mps")
25
+ logger.info(
26
+ "Loading reranker: model=%s device=%s fp16=%s",
27
+ settings.reranker_model, settings.embedding_device, use_fp16,
28
+ )
29
+ model = FlagReranker(
30
+ settings.reranker_model,
31
+ use_fp16=use_fp16,
32
+ device=settings.embedding_device,
33
+ cache_dir=str(settings.hf_cache_dir),
34
+ )
35
+ # 预热一次 compute_score, 强制 meta→cpu device 转移 (避免后续 "Cannot copy out of meta tensor")
36
+ try:
37
+ with __import__("torch").no_grad():
38
+ model.compute_score([["warmup query", "warmup passage"]], normalize=True)
39
+ logger.info("Reranker meta→cpu device transfer done")
40
+ except Exception as e: # noqa: BLE001
41
+ logger.warning("Reranker warmup failed: %s", e)
42
+ return model
43
+
44
+
45
+ def warm_up() -> None:
46
+ try:
47
+ get_reranker()
48
+ logger.info("Reranker warmed up")
49
+ except Exception as e: # noqa: BLE001
50
+ logger.warning("Reranker warm-up failed: %s", e)
51
+
52
+
53
+ class RerankerService:
54
+ """封装 Reranker, 异步 + 限流 + 截断长文本."""
55
+
56
+ _sem = None # 全局信号量, 避免并发打爆 CPU
57
+
58
+ def __init__(self, max_concurrency: int = 2) -> None:
59
+ if RerankerService._sem is None:
60
+ RerankerService._sem = asyncio.Semaphore(max_concurrency)
61
+
62
+ async def rerank(
63
+ self,
64
+ query: str,
65
+ hits: list[RetrievalHit],
66
+ top_n: int | None = None,
67
+ max_length: int = 512,
68
+ ) -> list[RetrievalHit]:
69
+ """对 hits 按 query 相关性重排, 返回 top_n. 原顺序保留在 original_rank."""
70
+ if not hits:
71
+ return []
72
+ top_n = top_n or settings.rerank_top_n
73
+
74
+ # 截断过长的 text (BGE-reranker max 512 token)
75
+ pairs = [[query, (h.text or "")[: max_length * 2]] for h in hits]
76
+
77
+ loop = asyncio.get_running_loop()
78
+ scores = await loop.run_in_executor(None, self._rerank_sync, pairs)
79
+
80
+ # 记 original_rank + 新分数
81
+ for h, s in zip(hits, scores):
82
+ h.original_rank = hits.index(h)
83
+ h.rerank_score = float(s)
84
+
85
+ # 按 rerank_score 降序
86
+ reranked = sorted(hits, key=lambda h: h.rerank_score, reverse=True)
87
+ return reranked[:top_n]
88
+
89
+ def _rerank_sync(self, pairs: list[list[str]]) -> list[float]:
90
+ model = get_reranker()
91
+ # FlagReranker.compute_score 直接吃 list[list[str]]
92
+ scores = model.compute_score(pairs, normalize=True)
93
+ # 单条输入时返回 float, 统一为 list
94
+ if isinstance(scores, (int, float)):
95
+ scores = [float(scores)]
96
+ return [float(s) for s in scores]
97
+
98
+
99
+ # 进程内单例
100
+ _reranker: RerankerService | None = None
101
+
102
+
103
+ def get_reranker_service() -> RerankerService:
104
+ global _reranker
105
+ if _reranker is None:
106
+ _reranker = RerankerService()
107
+ return _reranker