Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 4 additions & 1 deletion astrbot/core/knowledge_base/retrieval/rank_fusion.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,10 @@ async def fuse(
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
Expand Down
8 changes: 7 additions & 1 deletion astrbot/core/knowledge_base/retrieval/sparse_retriever.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@ class SparseResult:
kb_id: str
content: str
score: float
rank: int | None = None


class SparseRetriever:
Expand Down Expand Up @@ -87,7 +88,9 @@ async def retrieve(
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(
Expand All @@ -97,6 +100,7 @@ async def retrieve(
kb_id=kb_id,
content=doc["text"],
score=-float(doc["score"]),
rank=rank,
),
)

Expand Down Expand Up @@ -172,5 +176,7 @@ async def _retrieve_with_bm25(
)

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]
59 changes: 59 additions & 0 deletions tests/unit/test_rank_fusion.py
Original file line number Diff line number Diff line change
@@ -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)
65 changes: 63 additions & 2 deletions tests/unit/test_sparse_retriever.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,15 +6,20 @@
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,
"metadata": json.dumps(
{
"chunk_index": chunk_index,
"kb_doc_id": f"doc-{chunk_index}",
"kb_id": "kb-1",
"kb_id": kb_id,
},
),
}
Expand Down Expand Up @@ -59,6 +64,14 @@ async def get_documents(self, metadata_filters: dict, limit: int | None, offset)
]


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()
Expand Down Expand Up @@ -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}