omarsol Claude Fable 5 commited on
Commit
0f36b6a
·
1 Parent(s): 75e4226

fix(retrieval): budget skip, standard RRF, fail-fast collection load

Browse files

Three 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>

Files changed (2) hide show
  1. app/chroma_rag.py +31 -4
  2. tests/test_chroma_rag.py +129 -4
app/chroma_rag.py CHANGED
@@ -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
- fused_scores[key] = fused_scores.get(key, 0.0) + 1.0 / (rrf_k + rank)
 
 
 
 
 
 
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
- self._collection = client.get_or_create_collection(name=collection_name)
 
 
 
 
 
 
 
 
 
 
 
 
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 < 0.10:
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
- break
 
 
 
 
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)
tests/test_chroma_rag.py CHANGED
@@ -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 _result(self, chunk_id: str, score: float, method: str) -> SearchResult:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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=chunk_id,
267
- chunk_content=chunk_id,
268
  heading_path="section",
269
- retrieval_method=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__":