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
169 changes: 144 additions & 25 deletions src/retrievers/csv_chroma.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,10 @@
import asyncio
from collections.abc import Coroutine
from pathlib import Path
from typing import Annotated, Any, TypedDict
from typing import Any, TypedDict

import chromadb.config
from langchain.chains.query_constructor.schema import AttributeInfo
from langchain.retrievers import EnsembleRetriever, MultiQueryRetriever
from langchain.retrievers.self_query.base import SelfQueryRetriever
from langchain_chroma.vectorstores import Chroma
from langchain_community.document_loaders.csv_loader import CSVLoader
Expand All @@ -17,10 +16,12 @@
from langchain_core.documents import Document
from langchain_core.embeddings import Embeddings
from langchain_core.language_models.chat_models import BaseChatModel
from langchain_core.output_parsers import BaseOutputParser
from langchain_core.prompts.prompt import PromptTemplate
from langchain_core.retrievers import BaseRetriever
from langchain_core.runnables import Runnable
from nltk.tokenize import word_tokenize
from pydantic import AfterValidator, Field
from pydantic.json_schema import SkipJsonSchema
from pydantic import ConfigDict

chroma_settings = chromadb.config.Settings(anonymized_telemetry=False)

Expand Down Expand Up @@ -56,11 +57,6 @@
)


ExcludedField = SkipJsonSchema[
Annotated[Any, Field(default=None, exclude=True), AfterValidator(lambda x: None)]
]


RESULTS_PER_RETRIEVER = 10
# The vector store is asked for more than we intend to keep, because one
# Reactome entity can occupy several rows -- a reaction appears once per
Expand Down Expand Up @@ -111,6 +107,85 @@ def dedupe_by_entity(docs: list[Document], limit: int) -> list[Document]:
return kept


RRF_K = 60
"""Reciprocal Rank Fusion constant, from Cormack et al. (SIGIR 2009).

Copied from the value LangChain's EnsembleRetriever uses, so that vendoring the
maths here changes nothing. It damps the difference between adjacent ranks: at
k=60 a first-place hit scores 1/61 and an eleventh-place hit 1/71.
"""


def reciprocal_rank_fusion(
doc_lists: list[list[Document]], weights: list[float] | None = None
) -> list[Document]:
"""Fuse ranked lists by Reciprocal Rank Fusion.

Vendored rather than borrowed. Reaching LangChain's implementation required
constructing `EnsembleRetriever(retrievers=[])` -- an ensemble with no
retrievers -- purely to call one method, on every fusion. Ranking is the
product's core retrieval quality, so a library upgrade should not be able to
reorder results without anyone noticing.

Behaviour is identical to `EnsembleRetriever.weighted_reciprocal_rank`, which
the tests in tests/retrievers/ pin:

- a document's score is the sum over lists of `weight / (rank + RRF_K)`,
with rank counted from 1
- documents are identified by `page_content`, so the same text found by two
retrievers accumulates both scores
- ties keep the order of first appearance across the concatenated lists,
because the sort is stable
"""
if weights is None:
weights = [1 / len(doc_lists)] * len(doc_lists)
if len(doc_lists) != len(weights):
raise ValueError(
f"Got {len(doc_lists)} document lists and {len(weights)} weights; "
"they must correspond one to one."
)

scores: dict[str, float] = {}
for docs, weight in zip(doc_lists, weights, strict=True):
for rank, doc in enumerate(docs, start=1):
scores[doc.page_content] = scores.get(doc.page_content, 0.0) + weight / (
rank + RRF_K
)

seen: set[str] = set()
unique: list[Document] = []
for docs in doc_lists:
for doc in docs:
if doc.page_content not in seen:
seen.add(doc.page_content)
unique.append(doc)
return sorted(unique, key=lambda d: scores[d.page_content], reverse=True)


def unique_documents(documents: list[Document]) -> list[Document]:
"""Drop repeats, keeping first appearance.

Matches MultiQueryRetriever's final `unique_union`, which compares whole
Document objects rather than just their text.
"""
seen: list[Document] = []
for doc in documents:
if doc not in seen:
seen.append(doc)
return seen


class LineListOutputParser(BaseOutputParser[list[str]]):
"""Split an LLM response into non-empty lines.

Same behaviour as LangChain's parser of the same name, which the query
expansion prompt was written against.
"""

def parse(self, text: str) -> list[str]:
return [line for line in text.strip().split("\n") if line]


def list_chroma_subdirectories(directory: Path) -> list[str]:
return [
chroma_file.parent.name for chroma_file in directory.glob("*/chroma.sqlite3")
Expand Down Expand Up @@ -140,9 +215,29 @@ class RetrieverDict(TypedDict):
vector: SelfQueryRetriever


class HybridRetriever(MultiQueryRetriever):
retriever: ExcludedField = None
_retrievers: dict[str, RetrieverDict]
class HybridRetriever(BaseRetriever):
"""BM25 and vector search over each Chroma collection, fused by RRF.

A plain BaseRetriever. It previously subclassed MultiQueryRetriever and
reached into LangChain internals in five places -- overriding a required
field to None through SkipJsonSchema, building a throwaway
MultiQueryRetriever to steal its llm_chain, assigning a private attribute
outside pydantic, instantiating an empty EnsembleRetriever to borrow one
method, and overriding two internal methods. None of those are public API,
which is what blocked the LangChain upgrade.

Query expansion is now done here rather than inherited: one LLM call turns
the question into alternates, and `include_original` appends the original
LAST, matching the order MultiQueryRetriever used. Order matters, because
RRF resolves ties by first appearance.
"""

query_expander: Runnable[dict[str, str], list[str]]
include_original: bool = False
collection_retrievers: dict[str, RetrieverDict]

# BM25Retriever and SelfQueryRetriever are not pydantic models.
model_config = ConfigDict(arbitrary_types_allowed=True)

@classmethod
def from_subdirectory(
Expand Down Expand Up @@ -189,29 +284,53 @@ def from_subdirectory(
"bm25": bm25_retriever,
"vector": selfq_retriever,
}
llm_chain = MultiQueryRetriever.from_llm(
bm25_retriever, llm, multi_query_prompt, None, include_original
).llm_chain
hybrid_retriever = cls(
llm_chain=llm_chain,
# The expansion chain, built directly. This used to be extracted from a
# throwaway MultiQueryRetriever constructed only to reach its .llm_chain.
return cls(
query_expander=multi_query_prompt | llm | LineListOutputParser(),
include_original=include_original,
_retrievers={},
collection_retrievers=_retrievers,
)

def _get_relevant_documents(
self, query: str, *, run_manager: CallbackManagerForRetrieverRun
) -> list[Document]:
"""Expand the question, retrieve for every variant, fuse, de-duplicate.

Reproduces what MultiQueryRetriever._get_relevant_documents did, so this
stage changes structure without changing results. Note the original query
is appended AFTER the generated ones: RRF breaks ties by first
appearance, so reordering here would silently change the ranking.
"""
queries = self.query_expander.invoke(
{"question": query}, config={"callbacks": run_manager.get_child()}
)
if self.include_original:
queries.append(query)
return unique_documents(self.retrieve_documents(queries, run_manager))

async def _aget_relevant_documents(
self, query: str, *, run_manager: AsyncCallbackManagerForRetrieverRun
) -> list[Document]:
"""Async twin of the above; must agree with it document for document."""
queries = await self.query_expander.ainvoke(
{"question": query}, config={"callbacks": run_manager.get_child()}
)
hybrid_retriever._retrievers = _retrievers
return hybrid_retriever
if self.include_original:
queries.append(query)
return unique_documents(await self.aretrieve_documents(queries, run_manager))

def weighted_reciprocal_rank(
self, doc_lists: list[list[Document]]
) -> list[Document]:
return EnsembleRetriever(
retrievers=[], weights=[1 / len(doc_lists)] * len(doc_lists)
).weighted_reciprocal_rank(doc_lists)
"""Kept as a method so the existing tests address it unchanged."""
return reciprocal_rank_fusion(doc_lists)

def retrieve_documents(
self, queries: list[str], run_manager: CallbackManagerForRetrieverRun
) -> list[Document]:
subdirectory_docs: list[Document] = []
for subdirectory, retrievers in self._retrievers.items():
for subdirectory, retrievers in self.collection_retrievers.items():
bm25_retriever = retrievers["bm25"]
vector_retriever = retrievers["vector"]
doc_lists: list[list[Document]] = []
Expand Down Expand Up @@ -250,7 +369,7 @@ async def aretrieve_documents(
run_manager: AsyncCallbackManagerForRetrieverRun,
) -> list[Document]:
subdirectory_results: dict[str, list[Coroutine[Any, Any, list[Document]]]] = {}
for subdirectory, retrievers in self._retrievers.items():
for subdirectory, retrievers in self.collection_retrievers.items():
bm25_retriever = retrievers["bm25"]
vector_retriever = retrievers["vector"]
subdirectory_results[subdirectory] = []
Expand Down
72 changes: 72 additions & 0 deletions tests/retrievers/test_rrf_equivalence.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,72 @@
"""The vendored RRF must behave exactly like the LangChain method it replaces.

`reciprocal_rank_fusion` was copied out of `EnsembleRetriever` so that ranking --
the core of retrieval quality -- is owned here and cannot be reordered by a
library upgrade without anyone noticing. That is only safe if the copy is
faithful, so this compares the two implementations directly rather than trusting
that the maths was transcribed correctly.

This test is the reason Stage 1 of the retriever rewrite can claim "no behaviour
change". An end-to-end diff cannot show it: SelfQueryRetriever makes an LLM call,
so the same code run twice already differs on roughly a quarter of questions.
"""

import random

import pytest
from langchain_core.documents import Document

from retrievers.csv_chroma import RRF_K, reciprocal_rank_fusion

pytest.importorskip("langchain", reason="retrieval stack not installed")

from langchain.retrievers import EnsembleRetriever # noqa: E402

pytestmark = pytest.mark.requires_retrieval_stack


def _langchain_rrf(doc_lists: list[list[Document]]) -> list[str]:
n = len(doc_lists)
ensemble = EnsembleRetriever(retrievers=[], weights=[1 / n] * n)
return [d.page_content for d in ensemble.weighted_reciprocal_rank(doc_lists)]


def _ours(doc_lists: list[list[Document]]) -> list[str]:
return [d.page_content for d in reciprocal_rank_fusion(doc_lists)]


def test_the_constant_matches_the_source_it_was_copied_from() -> None:
assert EnsembleRetriever(retrievers=[], weights=[1.0]).c == RRF_K


@pytest.mark.parametrize("seed", range(25))
def test_equivalent_on_random_inputs(seed: int) -> None:
"""Random shapes: varying list counts, overlaps, orderings and empties."""
# S311: generating test data, not keys.
rng = random.Random(seed) # noqa: S311
for _ in range(80):
n_lists = rng.randint(1, 5)
pool = [f"doc-{i}" for i in range(rng.randint(1, 12))]
lists = [
[
Document(page_content=c)
for c in rng.sample(pool, rng.randint(0, len(pool)))
]
for _ in range(n_lists)
]
assert _ours(lists) == _langchain_rrf(lists), f"diverged on {lists}"


def test_equivalent_when_every_list_is_empty() -> None:
assert _ours([[], []]) == _langchain_rrf([[], []]) == []


def test_equivalent_when_one_retriever_returns_nothing() -> None:
"""A retriever finding nothing must not drop the other's results."""
found = [Document(page_content="a"), Document(page_content="b")]
assert _ours([found, []]) == _langchain_rrf([found, []]) == ["a", "b"]


def test_mismatched_weights_are_rejected() -> None:
with pytest.raises(ValueError, match="one to one"):
reciprocal_rank_fusion([[Document(page_content="a")]], weights=[0.5, 0.5])
Loading