mirror of
https://github.com/AstrBotDevs/AstrBot.git
synced 2026-09-24 16:39:52 +08:00
fix: preserve per-knowledge-base sparse retrieval ranks and use them in rank fusion to avoid distorted ordering across independent FTS5 indexes (#9426)
This commit is contained in:
@@ -66,7 +66,10 @@ class RankFusion:
|
||||
dense_ranks = {
|
||||
r.data["doc_id"]: (idx + 1) for idx, r in enumerate(dense_results)
|
||||
} # 这里的 doc_id 实际上是 chunk_id
|
||||
sparse_ranks = {r.chunk_id: (idx + 1) for idx, r in enumerate(sparse_results)}
|
||||
sparse_ranks = {
|
||||
r.chunk_id: r.rank if r.rank is not None else idx + 1
|
||||
for idx, r in enumerate(sparse_results)
|
||||
}
|
||||
|
||||
# 2. 收集所有唯一的 ID
|
||||
# 需要统一为 chunk_id
|
||||
|
||||
@@ -30,6 +30,7 @@ class SparseResult:
|
||||
kb_id: str
|
||||
content: str
|
||||
score: float
|
||||
rank: int | None = None
|
||||
|
||||
|
||||
class SparseRetriever:
|
||||
@@ -87,7 +88,9 @@ class SparseRetriever:
|
||||
fallback_kb_ids.append(kb_id)
|
||||
continue
|
||||
|
||||
for doc in result:
|
||||
# BM25 scores from independent FTS5 indexes are not comparable.
|
||||
# Preserve each index's local rank for the later RRF stage.
|
||||
for rank, doc in enumerate(result, start=1):
|
||||
chunk_md = json.loads(doc["metadata"])
|
||||
fts_results.append(
|
||||
SparseResult(
|
||||
@@ -97,6 +100,7 @@ class SparseRetriever:
|
||||
kb_id=kb_id,
|
||||
content=doc["text"],
|
||||
score=-float(doc["score"]),
|
||||
rank=rank,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -172,5 +176,7 @@ class SparseRetriever:
|
||||
)
|
||||
|
||||
results.sort(key=lambda x: x.score, reverse=True)
|
||||
for rank, result in enumerate(results, start=1):
|
||||
result.rank = rank
|
||||
# return results[: len(results) // len(kb_ids)]
|
||||
return results[:top_k_sparse]
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
import pytest
|
||||
|
||||
from astrbot.core.db.vec_db.base import Result
|
||||
from astrbot.core.knowledge_base.retrieval.rank_fusion import RankFusion
|
||||
from astrbot.core.knowledge_base.retrieval.sparse_retriever import SparseResult
|
||||
|
||||
|
||||
def make_dense_result(chunk_id: str, similarity: float) -> Result:
|
||||
return Result(
|
||||
similarity=similarity,
|
||||
data={
|
||||
"doc_id": chunk_id,
|
||||
"text": chunk_id,
|
||||
"metadata": "{}",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def make_sparse_result(
|
||||
chunk_id: str,
|
||||
kb_id: str,
|
||||
score: float,
|
||||
rank: int,
|
||||
) -> SparseResult:
|
||||
return SparseResult(
|
||||
chunk_index=0,
|
||||
chunk_id=chunk_id,
|
||||
doc_id=f"doc-{chunk_id}",
|
||||
kb_id=kb_id,
|
||||
content=chunk_id,
|
||||
score=score,
|
||||
rank=rank,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rank_fusion_uses_source_rank_for_independent_sparse_indexes():
|
||||
dense_results = [
|
||||
make_dense_result("small-exact", 0.99),
|
||||
make_dense_result("large-1", 0.95),
|
||||
make_dense_result("large-2", 0.90),
|
||||
]
|
||||
sparse_results = [
|
||||
make_sparse_result("large-1", "kb-large", 12.0, 1),
|
||||
make_sparse_result("large-2", "kb-large", 10.0, 2),
|
||||
make_sparse_result("small-exact", "kb-small", 0.00001, 1),
|
||||
]
|
||||
|
||||
results = await RankFusion(kb_db=None).fuse(
|
||||
dense_results=dense_results,
|
||||
sparse_results=sparse_results,
|
||||
)
|
||||
|
||||
assert [result.chunk_id for result in results] == [
|
||||
"small-exact",
|
||||
"large-1",
|
||||
"large-2",
|
||||
]
|
||||
assert results[0].score == pytest.approx(2 / 61)
|
||||
@@ -6,7 +6,12 @@ import pytest
|
||||
from astrbot.core.knowledge_base.retrieval.sparse_retriever import SparseRetriever
|
||||
|
||||
|
||||
def make_doc(chunk_id: str, text: str, chunk_index: int = 0) -> dict:
|
||||
def make_doc(
|
||||
chunk_id: str,
|
||||
text: str,
|
||||
chunk_index: int = 0,
|
||||
kb_id: str = "kb-1",
|
||||
) -> dict:
|
||||
return {
|
||||
"doc_id": chunk_id,
|
||||
"text": text,
|
||||
@@ -14,7 +19,7 @@ def make_doc(chunk_id: str, text: str, chunk_index: int = 0) -> dict:
|
||||
{
|
||||
"chunk_index": chunk_index,
|
||||
"kb_doc_id": f"doc-{chunk_index}",
|
||||
"kb_id": "kb-1",
|
||||
"kb_id": kb_id,
|
||||
},
|
||||
),
|
||||
}
|
||||
@@ -59,6 +64,14 @@ class FallbackStorage:
|
||||
]
|
||||
|
||||
|
||||
class StaticFTSStorage:
|
||||
def __init__(self, documents: list[dict]):
|
||||
self.documents = documents
|
||||
|
||||
async def search_sparse(self, query_tokens: list[str], limit: int):
|
||||
return self.documents[:limit]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sparse_retriever_uses_fts5_when_available():
|
||||
storage = FTSStorage()
|
||||
@@ -91,3 +104,51 @@ async def test_sparse_retriever_falls_back_to_bm25_when_fts5_is_unavailable():
|
||||
assert [result.chunk_id for result in results] == ["chunk-1"]
|
||||
assert storage.search_sparse_calls == 1
|
||||
assert storage.get_documents_calls == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sparse_retriever_preserves_per_kb_fts_ranks():
|
||||
large_storage = StaticFTSStorage(
|
||||
[
|
||||
{
|
||||
**make_doc("large-1", "管理员账号安全说明", 0, "kb-large"),
|
||||
"score": -12.0,
|
||||
},
|
||||
{
|
||||
**make_doc("large-2", "密码策略说明", 1, "kb-large"),
|
||||
"score": -10.0,
|
||||
},
|
||||
],
|
||||
)
|
||||
small_storage = StaticFTSStorage(
|
||||
[
|
||||
{
|
||||
**make_doc(
|
||||
"small-exact",
|
||||
"如何重置管理员密码?",
|
||||
0,
|
||||
"kb-small",
|
||||
),
|
||||
"score": -0.00001,
|
||||
},
|
||||
],
|
||||
)
|
||||
retriever = SparseRetriever(kb_db=None)
|
||||
|
||||
results = await retriever.retrieve(
|
||||
query="如何重置管理员密码?",
|
||||
kb_ids=["kb-large", "kb-small"],
|
||||
kb_options={
|
||||
"kb-large": {
|
||||
"vec_db": SimpleNamespace(document_storage=large_storage),
|
||||
"top_k_sparse": 2,
|
||||
},
|
||||
"kb-small": {
|
||||
"vec_db": SimpleNamespace(document_storage=small_storage),
|
||||
"top_k_sparse": 1,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
ranks = {result.chunk_id: result.rank for result in results}
|
||||
assert ranks == {"large-1": 1, "large-2": 2, "small-exact": 1}
|
||||
|
||||
Reference in New Issue
Block a user