fix(retrieval): budget skip, standard RRF, fail-fast collection load
Browse filesThree hybrid-retrieval fixes:
- _apply_token_budget broke on the first over-budget result, so one
oversized rank-1 document (retrieve_doc under a small per-request
budget) returned zero context even when smaller lower-ranked chunks
fit. Skip the oversized result and keep filling the budget.
- reciprocal_rank_fusion accumulated one contribution per occurrence of
a dedupe key within a single ranked list, so a section split into
several chunks in one retriever's top-k masqueraded as
cross-retriever consensus. Standard RRF: each key counts once per
list at its best rank; representative selection is unchanged.
- LocalChromaRetriever used get_or_create_collection, silently creating
an empty collection from a broken or mismatched bundle and degrading
dense search to zero hits with no signal. Use get_collection and
raise a clear error naming the collection and db path.
Also extracts the 0.10 rerank floor as RERANK_SCORE_FLOOR for reuse by
the GraphRAG backend.
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
- app/chroma_rag.py +31 -4
- tests/test_chroma_rag.py +129 -4
|
@@ -32,6 +32,10 @@ DEFAULT_FUSION_TOP_K = 30
|
|
| 32 |
DEFAULT_RERANK_TOP_K = 5
|
| 33 |
DEFAULT_RRF_K = 60
|
| 34 |
DEFAULT_CONTEXT_TOKEN_BUDGET = 100_000
|
|
|
|
|
|
|
|
|
|
|
|
|
| 35 |
DEFAULT_EMBED_MODEL = "embed-v4.0"
|
| 36 |
DEFAULT_RERANK_MODEL = "rerank-v4.0-fast"
|
| 37 |
DEFAULT_ENCODING = "cl100k_base"
|
|
@@ -1120,9 +1124,16 @@ def reciprocal_rank_fusion(
|
|
| 1120 |
representatives: dict[str, SearchResult] = {}
|
| 1121 |
|
| 1122 |
for ranked_results in ranked_lists:
|
|
|
|
| 1123 |
for rank, result in enumerate(ranked_results, start=1):
|
| 1124 |
key = result_dedupe_key(result)
|
| 1125 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1126 |
|
| 1127 |
current = representatives.get(key)
|
| 1128 |
if current is None:
|
|
@@ -1199,7 +1210,19 @@ class LocalChromaRetriever:
|
|
| 1199 |
)
|
| 1200 |
|
| 1201 |
client = chromadb.PersistentClient(path=db_path)
|
| 1202 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1203 |
with open(document_dict_path, "rb") as handle:
|
| 1204 |
self._document_dict: dict[str, dict[str, Any]] = pickle.load(handle)
|
| 1205 |
|
|
@@ -1377,7 +1400,7 @@ class LocalChromaRetriever:
|
|
| 1377 |
filtered: list[SearchResult] = []
|
| 1378 |
total_tokens = 0
|
| 1379 |
for result in results:
|
| 1380 |
-
if result.score <
|
| 1381 |
continue
|
| 1382 |
|
| 1383 |
# disallowed_special=() so literal "<|endoftext|>" in chunks doesn't crash.
|
|
@@ -1385,7 +1408,11 @@ class LocalChromaRetriever:
|
|
| 1385 |
self._encoding.encode(result.content, disallowed_special=())
|
| 1386 |
)
|
| 1387 |
if total_tokens + result_tokens > budget:
|
| 1388 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1389 |
|
| 1390 |
total_tokens += result_tokens
|
| 1391 |
filtered.append(result)
|
|
|
|
| 32 |
DEFAULT_RERANK_TOP_K = 5
|
| 33 |
DEFAULT_RRF_K = 60
|
| 34 |
DEFAULT_CONTEXT_TOKEN_BUDGET = 100_000
|
| 35 |
+
# Reranked results below this relevance score are dropped before the token
|
| 36 |
+
# budget is filled. Shared with the GraphRAG backend so both eval arms apply
|
| 37 |
+
# the same low-relevance floor.
|
| 38 |
+
RERANK_SCORE_FLOOR = 0.10
|
| 39 |
DEFAULT_EMBED_MODEL = "embed-v4.0"
|
| 40 |
DEFAULT_RERANK_MODEL = "rerank-v4.0-fast"
|
| 41 |
DEFAULT_ENCODING = "cl100k_base"
|
|
|
|
| 1124 |
representatives: dict[str, SearchResult] = {}
|
| 1125 |
|
| 1126 |
for ranked_results in ranked_lists:
|
| 1127 |
+
scored_keys: set[str] = set()
|
| 1128 |
for rank, result in enumerate(ranked_results, start=1):
|
| 1129 |
key = result_dedupe_key(result)
|
| 1130 |
+
# Standard RRF: a key contributes once per ranked list, at its best
|
| 1131 |
+
# rank. Without this, a section split into several chunks that all
|
| 1132 |
+
# land in one retriever's top-k collects one contribution per chunk
|
| 1133 |
+
# and masquerades as cross-retriever consensus.
|
| 1134 |
+
if key not in scored_keys:
|
| 1135 |
+
scored_keys.add(key)
|
| 1136 |
+
fused_scores[key] = fused_scores.get(key, 0.0) + 1.0 / (rrf_k + rank)
|
| 1137 |
|
| 1138 |
current = representatives.get(key)
|
| 1139 |
if current is None:
|
|
|
|
| 1210 |
)
|
| 1211 |
|
| 1212 |
client = chromadb.PersistentClient(path=db_path)
|
| 1213 |
+
try:
|
| 1214 |
+
# get_collection (not get_or_create_collection): a broken or
|
| 1215 |
+
# mismatched bundle must fail loudly at startup instead of silently
|
| 1216 |
+
# creating an empty collection that degrades dense search to zero
|
| 1217 |
+
# hits forever.
|
| 1218 |
+
self._collection = client.get_collection(name=collection_name)
|
| 1219 |
+
except Exception as exc:
|
| 1220 |
+
raise RuntimeError(
|
| 1221 |
+
f"Chroma collection '{collection_name}' not found in vector db "
|
| 1222 |
+
f"at '{db_path}'. The bundle is likely missing, incomplete, or "
|
| 1223 |
+
"mismatched: delete the directory and restart to re-download "
|
| 1224 |
+
"it, or rebuild it with create_vector_stores."
|
| 1225 |
+
) from exc
|
| 1226 |
with open(document_dict_path, "rb") as handle:
|
| 1227 |
self._document_dict: dict[str, dict[str, Any]] = pickle.load(handle)
|
| 1228 |
|
|
|
|
| 1400 |
filtered: list[SearchResult] = []
|
| 1401 |
total_tokens = 0
|
| 1402 |
for result in results:
|
| 1403 |
+
if result.score < RERANK_SCORE_FLOOR:
|
| 1404 |
continue
|
| 1405 |
|
| 1406 |
# disallowed_special=() so literal "<|endoftext|>" in chunks doesn't crash.
|
|
|
|
| 1408 |
self._encoding.encode(result.content, disallowed_special=())
|
| 1409 |
)
|
| 1410 |
if total_tokens + result_tokens > budget:
|
| 1411 |
+
# An oversized result (e.g. a retrieve_doc full document under a
|
| 1412 |
+
# small per-request budget) must not cut off the whole list:
|
| 1413 |
+
# skip it and keep filling the budget with lower-ranked results
|
| 1414 |
+
# that fit, preserving rank order among the kept ones.
|
| 1415 |
+
continue
|
| 1416 |
|
| 1417 |
total_tokens += result_tokens
|
| 1418 |
filtered.append(result)
|
|
@@ -7,12 +7,15 @@ import unittest
|
|
| 7 |
from pathlib import Path
|
| 8 |
from unittest.mock import patch
|
| 9 |
|
|
|
|
|
|
|
| 10 |
from data.scraping_scripts.add_context_to_nodes import process
|
| 11 |
from data.scraping_scripts.create_vector_stores import write_retrieval_artifacts
|
| 12 |
from llama_index.core import Document
|
| 13 |
from app.chroma_rag import (
|
| 14 |
BM25Index,
|
| 15 |
ChunkRecord,
|
|
|
|
| 16 |
build_chunk_records,
|
| 17 |
heading_aware_markdown_chunks,
|
| 18 |
reciprocal_rank_fusion,
|
|
@@ -253,7 +256,66 @@ After the example.
|
|
| 253 |
self.assertEqual(fused[0].chunk_id, "overlap")
|
| 254 |
self.assertEqual(fused[0].retrieval_method, "hybrid")
|
| 255 |
|
| 256 |
-
def
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 257 |
return SearchResult(
|
| 258 |
chunk_id=chunk_id,
|
| 259 |
doc_id=chunk_id,
|
|
@@ -263,11 +325,74 @@ After the example.
|
|
| 263 |
retrieve_doc=False,
|
| 264 |
tokens=10,
|
| 265 |
score=score,
|
| 266 |
-
content=
|
| 267 |
-
chunk_content=
|
| 268 |
heading_path="section",
|
| 269 |
-
retrieval_method=
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 270 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 271 |
|
| 272 |
|
| 273 |
if __name__ == "__main__":
|
|
|
|
| 7 |
from pathlib import Path
|
| 8 |
from unittest.mock import patch
|
| 9 |
|
| 10 |
+
import chromadb
|
| 11 |
+
import tiktoken
|
| 12 |
from data.scraping_scripts.add_context_to_nodes import process
|
| 13 |
from data.scraping_scripts.create_vector_stores import write_retrieval_artifacts
|
| 14 |
from llama_index.core import Document
|
| 15 |
from app.chroma_rag import (
|
| 16 |
BM25Index,
|
| 17 |
ChunkRecord,
|
| 18 |
+
LocalChromaRetriever,
|
| 19 |
build_chunk_records,
|
| 20 |
heading_aware_markdown_chunks,
|
| 21 |
reciprocal_rank_fusion,
|
|
|
|
| 256 |
self.assertEqual(fused[0].chunk_id, "overlap")
|
| 257 |
self.assertEqual(fused[0].retrieval_method, "hybrid")
|
| 258 |
|
| 259 |
+
def test_rrf_counts_each_key_once_per_ranked_list(self) -> None:
|
| 260 |
+
# A section split into several chunks can land at multiple ranks of ONE
|
| 261 |
+
# retriever's list (same dedupe key). Standard RRF scores a key once per
|
| 262 |
+
# list, at its best rank; per-occurrence accumulation would let one
|
| 263 |
+
# retriever's duplicates masquerade as cross-retriever consensus.
|
| 264 |
+
dup_top = self._result("dup:0", 0.9, "dense", doc_id="dup-doc")
|
| 265 |
+
dup_mid = self._result("dup:1", 0.8, "dense", doc_id="dup-doc")
|
| 266 |
+
dup_low = self._result("dup:2", 0.7, "dense", doc_id="dup-doc")
|
| 267 |
+
consensus_dense = self._result("uni:0", 0.6, "dense", doc_id="consensus-doc")
|
| 268 |
+
consensus_bm25 = self._result("uni:0", 5.0, "bm25", doc_id="consensus-doc")
|
| 269 |
+
|
| 270 |
+
fused = reciprocal_rank_fusion(
|
| 271 |
+
[[dup_top, dup_mid, dup_low, consensus_dense], [consensus_bm25]],
|
| 272 |
+
top_k=5,
|
| 273 |
+
)
|
| 274 |
+
|
| 275 |
+
by_doc = {result.doc_id: result for result in fused}
|
| 276 |
+
# One contribution at the best rank (1), nothing from ranks 2-3.
|
| 277 |
+
self.assertAlmostEqual(by_doc["dup-doc"].score, 1.0 / 61)
|
| 278 |
+
# Rank 4 in dense + rank 1 in bm25.
|
| 279 |
+
self.assertAlmostEqual(by_doc["consensus-doc"].score, 1.0 / 64 + 1.0 / 61)
|
| 280 |
+
# Genuine cross-retriever consensus outranks single-list duplication.
|
| 281 |
+
self.assertEqual(fused[0].doc_id, "consensus-doc")
|
| 282 |
+
# Representative selection still works: best-scoring dense duplicate.
|
| 283 |
+
self.assertEqual(by_doc["dup-doc"].chunk_id, "dup:0")
|
| 284 |
+
|
| 285 |
+
def _result(
|
| 286 |
+
self,
|
| 287 |
+
chunk_id: str,
|
| 288 |
+
score: float,
|
| 289 |
+
method: str,
|
| 290 |
+
*,
|
| 291 |
+
doc_id: str | None = None,
|
| 292 |
+
content: str | None = None,
|
| 293 |
+
retrieve_doc: bool = False,
|
| 294 |
+
) -> SearchResult:
|
| 295 |
+
return SearchResult(
|
| 296 |
+
chunk_id=chunk_id,
|
| 297 |
+
doc_id=doc_id if doc_id is not None else chunk_id,
|
| 298 |
+
title=chunk_id,
|
| 299 |
+
url="",
|
| 300 |
+
source="test",
|
| 301 |
+
retrieve_doc=retrieve_doc,
|
| 302 |
+
tokens=10,
|
| 303 |
+
score=score,
|
| 304 |
+
content=content if content is not None else chunk_id,
|
| 305 |
+
chunk_content=content if content is not None else chunk_id,
|
| 306 |
+
heading_path="section",
|
| 307 |
+
retrieval_method=method,
|
| 308 |
+
)
|
| 309 |
+
|
| 310 |
+
|
| 311 |
+
class TokenBudgetTestCase(unittest.TestCase):
|
| 312 |
+
def _retriever(self) -> LocalChromaRetriever:
|
| 313 |
+
retriever = LocalChromaRetriever.__new__(LocalChromaRetriever)
|
| 314 |
+
retriever._encoding = tiktoken.get_encoding("cl100k_base")
|
| 315 |
+
retriever._token_budget = 100_000
|
| 316 |
+
return retriever
|
| 317 |
+
|
| 318 |
+
def _result(self, chunk_id: str, score: float, content: str) -> SearchResult:
|
| 319 |
return SearchResult(
|
| 320 |
chunk_id=chunk_id,
|
| 321 |
doc_id=chunk_id,
|
|
|
|
| 325 |
retrieve_doc=False,
|
| 326 |
tokens=10,
|
| 327 |
score=score,
|
| 328 |
+
content=content,
|
| 329 |
+
chunk_content=content,
|
| 330 |
heading_path="section",
|
| 331 |
+
retrieval_method="dense",
|
| 332 |
+
)
|
| 333 |
+
|
| 334 |
+
def test_budget_skips_oversized_result_and_fills_with_smaller(self) -> None:
|
| 335 |
+
# A rank-1 retrieve_doc result whose full document exceeds a small
|
| 336 |
+
# per-request budget must not empty the whole result list: it is
|
| 337 |
+
# skipped and the budget is filled with lower-ranked results that fit,
|
| 338 |
+
# in rank order.
|
| 339 |
+
retriever = self._retriever()
|
| 340 |
+
oversized = self._result("big", 0.9, "word " * 500)
|
| 341 |
+
small_one = self._result("small-1", 0.8, "word " * 15)
|
| 342 |
+
small_two = self._result("small-2", 0.7, "word " * 15)
|
| 343 |
+
|
| 344 |
+
kept = retriever._apply_token_budget(
|
| 345 |
+
[oversized, small_one, small_two], token_budget=50
|
| 346 |
)
|
| 347 |
+
self.assertEqual([result.chunk_id for result in kept], ["small-1", "small-2"])
|
| 348 |
+
|
| 349 |
+
# With the default (large) budget everything still fits.
|
| 350 |
+
kept_all = retriever._apply_token_budget([oversized, small_one, small_two])
|
| 351 |
+
self.assertEqual(len(kept_all), 3)
|
| 352 |
+
|
| 353 |
+
|
| 354 |
+
class CollectionOpenTestCase(unittest.TestCase):
|
| 355 |
+
def _write_document_dict(self, directory: str) -> str:
|
| 356 |
+
path = Path(directory) / "document_dict_test.pkl"
|
| 357 |
+
with open(path, "wb") as handle:
|
| 358 |
+
pickle.dump({}, handle)
|
| 359 |
+
return str(path)
|
| 360 |
+
|
| 361 |
+
def test_init_fails_loudly_when_collection_missing(self) -> None:
|
| 362 |
+
# A broken/mismatched bundle must raise at startup instead of silently
|
| 363 |
+
# creating an empty collection that returns zero dense hits forever.
|
| 364 |
+
with tempfile.TemporaryDirectory() as temp_dir:
|
| 365 |
+
document_dict_path = self._write_document_dict(temp_dir)
|
| 366 |
+
|
| 367 |
+
with self.assertRaises(RuntimeError) as ctx:
|
| 368 |
+
LocalChromaRetriever(
|
| 369 |
+
db_path=temp_dir,
|
| 370 |
+
collection_name="missing-collection",
|
| 371 |
+
document_dict_path=document_dict_path,
|
| 372 |
+
cohere_api_key="fake",
|
| 373 |
+
)
|
| 374 |
+
|
| 375 |
+
message = str(ctx.exception)
|
| 376 |
+
self.assertIn("missing-collection", message)
|
| 377 |
+
self.assertIn(temp_dir, message)
|
| 378 |
+
|
| 379 |
+
def test_init_opens_collection_created_beforehand(self) -> None:
|
| 380 |
+
# Mirrors production: create_vector_stores creates the collection; the
|
| 381 |
+
# retriever only opens it.
|
| 382 |
+
with tempfile.TemporaryDirectory() as temp_dir:
|
| 383 |
+
chromadb.PersistentClient(path=temp_dir).create_collection(
|
| 384 |
+
name="test-collection"
|
| 385 |
+
)
|
| 386 |
+
document_dict_path = self._write_document_dict(temp_dir)
|
| 387 |
+
|
| 388 |
+
retriever = LocalChromaRetriever(
|
| 389 |
+
db_path=temp_dir,
|
| 390 |
+
collection_name="test-collection",
|
| 391 |
+
document_dict_path=document_dict_path,
|
| 392 |
+
cohere_api_key="fake",
|
| 393 |
+
)
|
| 394 |
+
|
| 395 |
+
self.assertEqual(retriever._collection.name, "test-collection")
|
| 396 |
|
| 397 |
|
| 398 |
if __name__ == "__main__":
|