|
1 | 1 | import asyncio |
2 | 2 | from collections.abc import Coroutine |
3 | 3 | from pathlib import Path |
4 | | -from typing import Annotated, Any, TypedDict |
| 4 | +from typing import Any, TypedDict |
5 | 5 |
|
6 | 6 | import chromadb.config |
7 | 7 | from langchain.chains.query_constructor.schema import AttributeInfo |
8 | | -from langchain.retrievers import EnsembleRetriever, MultiQueryRetriever |
9 | 8 | from langchain.retrievers.self_query.base import SelfQueryRetriever |
10 | 9 | from langchain_chroma.vectorstores import Chroma |
11 | 10 | from langchain_community.document_loaders.csv_loader import CSVLoader |
|
17 | 16 | from langchain_core.documents import Document |
18 | 17 | from langchain_core.embeddings import Embeddings |
19 | 18 | from langchain_core.language_models.chat_models import BaseChatModel |
| 19 | +from langchain_core.output_parsers import BaseOutputParser |
20 | 20 | from langchain_core.prompts.prompt import PromptTemplate |
| 21 | +from langchain_core.retrievers import BaseRetriever |
| 22 | +from langchain_core.runnables import Runnable |
21 | 23 | from nltk.tokenize import word_tokenize |
22 | | -from pydantic import AfterValidator, Field |
23 | | -from pydantic.json_schema import SkipJsonSchema |
| 24 | +from pydantic import ConfigDict |
24 | 25 |
|
25 | 26 | chroma_settings = chromadb.config.Settings(anonymized_telemetry=False) |
26 | 27 |
|
|
56 | 57 | ) |
57 | 58 |
|
58 | 59 |
|
59 | | -ExcludedField = SkipJsonSchema[ |
60 | | - Annotated[Any, Field(default=None, exclude=True), AfterValidator(lambda x: None)] |
61 | | -] |
62 | | - |
63 | | - |
64 | 60 | RESULTS_PER_RETRIEVER = 10 |
65 | 61 | # The vector store is asked for more than we intend to keep, because one |
66 | 62 | # Reactome entity can occupy several rows -- a reaction appears once per |
@@ -111,6 +107,85 @@ def dedupe_by_entity(docs: list[Document], limit: int) -> list[Document]: |
111 | 107 | return kept |
112 | 108 |
|
113 | 109 |
|
| 110 | +RRF_K = 60 |
| 111 | +"""Reciprocal Rank Fusion constant, from Cormack et al. (SIGIR 2009). |
| 112 | +
|
| 113 | +Copied from the value LangChain's EnsembleRetriever uses, so that vendoring the |
| 114 | +maths here changes nothing. It damps the difference between adjacent ranks: at |
| 115 | +k=60 a first-place hit scores 1/61 and an eleventh-place hit 1/71. |
| 116 | +""" |
| 117 | + |
| 118 | + |
| 119 | +def reciprocal_rank_fusion( |
| 120 | + doc_lists: list[list[Document]], weights: list[float] | None = None |
| 121 | +) -> list[Document]: |
| 122 | + """Fuse ranked lists by Reciprocal Rank Fusion. |
| 123 | +
|
| 124 | + Vendored rather than borrowed. Reaching LangChain's implementation required |
| 125 | + constructing `EnsembleRetriever(retrievers=[])` -- an ensemble with no |
| 126 | + retrievers -- purely to call one method, on every fusion. Ranking is the |
| 127 | + product's core retrieval quality, so a library upgrade should not be able to |
| 128 | + reorder results without anyone noticing. |
| 129 | +
|
| 130 | + Behaviour is identical to `EnsembleRetriever.weighted_reciprocal_rank`, which |
| 131 | + the tests in tests/retrievers/ pin: |
| 132 | +
|
| 133 | + - a document's score is the sum over lists of `weight / (rank + RRF_K)`, |
| 134 | + with rank counted from 1 |
| 135 | + - documents are identified by `page_content`, so the same text found by two |
| 136 | + retrievers accumulates both scores |
| 137 | + - ties keep the order of first appearance across the concatenated lists, |
| 138 | + because the sort is stable |
| 139 | + """ |
| 140 | + if weights is None: |
| 141 | + weights = [1 / len(doc_lists)] * len(doc_lists) |
| 142 | + if len(doc_lists) != len(weights): |
| 143 | + raise ValueError( |
| 144 | + f"Got {len(doc_lists)} document lists and {len(weights)} weights; " |
| 145 | + "they must correspond one to one." |
| 146 | + ) |
| 147 | + |
| 148 | + scores: dict[str, float] = {} |
| 149 | + for docs, weight in zip(doc_lists, weights, strict=True): |
| 150 | + for rank, doc in enumerate(docs, start=1): |
| 151 | + scores[doc.page_content] = scores.get(doc.page_content, 0.0) + weight / ( |
| 152 | + rank + RRF_K |
| 153 | + ) |
| 154 | + |
| 155 | + seen: set[str] = set() |
| 156 | + unique: list[Document] = [] |
| 157 | + for docs in doc_lists: |
| 158 | + for doc in docs: |
| 159 | + if doc.page_content not in seen: |
| 160 | + seen.add(doc.page_content) |
| 161 | + unique.append(doc) |
| 162 | + return sorted(unique, key=lambda d: scores[d.page_content], reverse=True) |
| 163 | + |
| 164 | + |
| 165 | +def unique_documents(documents: list[Document]) -> list[Document]: |
| 166 | + """Drop repeats, keeping first appearance. |
| 167 | +
|
| 168 | + Matches MultiQueryRetriever's final `unique_union`, which compares whole |
| 169 | + Document objects rather than just their text. |
| 170 | + """ |
| 171 | + seen: list[Document] = [] |
| 172 | + for doc in documents: |
| 173 | + if doc not in seen: |
| 174 | + seen.append(doc) |
| 175 | + return seen |
| 176 | + |
| 177 | + |
| 178 | +class LineListOutputParser(BaseOutputParser[list[str]]): |
| 179 | + """Split an LLM response into non-empty lines. |
| 180 | +
|
| 181 | + Same behaviour as LangChain's parser of the same name, which the query |
| 182 | + expansion prompt was written against. |
| 183 | + """ |
| 184 | + |
| 185 | + def parse(self, text: str) -> list[str]: |
| 186 | + return [line for line in text.strip().split("\n") if line] |
| 187 | + |
| 188 | + |
114 | 189 | def list_chroma_subdirectories(directory: Path) -> list[str]: |
115 | 190 | return [ |
116 | 191 | chroma_file.parent.name for chroma_file in directory.glob("*/chroma.sqlite3") |
@@ -140,9 +215,29 @@ class RetrieverDict(TypedDict): |
140 | 215 | vector: SelfQueryRetriever |
141 | 216 |
|
142 | 217 |
|
143 | | -class HybridRetriever(MultiQueryRetriever): |
144 | | - retriever: ExcludedField = None |
145 | | - _retrievers: dict[str, RetrieverDict] |
| 218 | +class HybridRetriever(BaseRetriever): |
| 219 | + """BM25 and vector search over each Chroma collection, fused by RRF. |
| 220 | +
|
| 221 | + A plain BaseRetriever. It previously subclassed MultiQueryRetriever and |
| 222 | + reached into LangChain internals in five places -- overriding a required |
| 223 | + field to None through SkipJsonSchema, building a throwaway |
| 224 | + MultiQueryRetriever to steal its llm_chain, assigning a private attribute |
| 225 | + outside pydantic, instantiating an empty EnsembleRetriever to borrow one |
| 226 | + method, and overriding two internal methods. None of those are public API, |
| 227 | + which is what blocked the LangChain upgrade. |
| 228 | +
|
| 229 | + Query expansion is now done here rather than inherited: one LLM call turns |
| 230 | + the question into alternates, and `include_original` appends the original |
| 231 | + LAST, matching the order MultiQueryRetriever used. Order matters, because |
| 232 | + RRF resolves ties by first appearance. |
| 233 | + """ |
| 234 | + |
| 235 | + query_expander: Runnable[dict[str, str], list[str]] |
| 236 | + include_original: bool = False |
| 237 | + collection_retrievers: dict[str, RetrieverDict] |
| 238 | + |
| 239 | + # BM25Retriever and SelfQueryRetriever are not pydantic models. |
| 240 | + model_config = ConfigDict(arbitrary_types_allowed=True) |
146 | 241 |
|
147 | 242 | @classmethod |
148 | 243 | def from_subdirectory( |
@@ -189,29 +284,53 @@ def from_subdirectory( |
189 | 284 | "bm25": bm25_retriever, |
190 | 285 | "vector": selfq_retriever, |
191 | 286 | } |
192 | | - llm_chain = MultiQueryRetriever.from_llm( |
193 | | - bm25_retriever, llm, multi_query_prompt, None, include_original |
194 | | - ).llm_chain |
195 | | - hybrid_retriever = cls( |
196 | | - llm_chain=llm_chain, |
| 287 | + # The expansion chain, built directly. This used to be extracted from a |
| 288 | + # throwaway MultiQueryRetriever constructed only to reach its .llm_chain. |
| 289 | + return cls( |
| 290 | + query_expander=multi_query_prompt | llm | LineListOutputParser(), |
197 | 291 | include_original=include_original, |
198 | | - _retrievers={}, |
| 292 | + collection_retrievers=_retrievers, |
| 293 | + ) |
| 294 | + |
| 295 | + def _get_relevant_documents( |
| 296 | + self, query: str, *, run_manager: CallbackManagerForRetrieverRun |
| 297 | + ) -> list[Document]: |
| 298 | + """Expand the question, retrieve for every variant, fuse, de-duplicate. |
| 299 | +
|
| 300 | + Reproduces what MultiQueryRetriever._get_relevant_documents did, so this |
| 301 | + stage changes structure without changing results. Note the original query |
| 302 | + is appended AFTER the generated ones: RRF breaks ties by first |
| 303 | + appearance, so reordering here would silently change the ranking. |
| 304 | + """ |
| 305 | + queries = self.query_expander.invoke( |
| 306 | + {"question": query}, config={"callbacks": run_manager.get_child()} |
| 307 | + ) |
| 308 | + if self.include_original: |
| 309 | + queries.append(query) |
| 310 | + return unique_documents(self.retrieve_documents(queries, run_manager)) |
| 311 | + |
| 312 | + async def _aget_relevant_documents( |
| 313 | + self, query: str, *, run_manager: AsyncCallbackManagerForRetrieverRun |
| 314 | + ) -> list[Document]: |
| 315 | + """Async twin of the above; must agree with it document for document.""" |
| 316 | + queries = await self.query_expander.ainvoke( |
| 317 | + {"question": query}, config={"callbacks": run_manager.get_child()} |
199 | 318 | ) |
200 | | - hybrid_retriever._retrievers = _retrievers |
201 | | - return hybrid_retriever |
| 319 | + if self.include_original: |
| 320 | + queries.append(query) |
| 321 | + return unique_documents(await self.aretrieve_documents(queries, run_manager)) |
202 | 322 |
|
203 | 323 | def weighted_reciprocal_rank( |
204 | 324 | self, doc_lists: list[list[Document]] |
205 | 325 | ) -> list[Document]: |
206 | | - return EnsembleRetriever( |
207 | | - retrievers=[], weights=[1 / len(doc_lists)] * len(doc_lists) |
208 | | - ).weighted_reciprocal_rank(doc_lists) |
| 326 | + """Kept as a method so the existing tests address it unchanged.""" |
| 327 | + return reciprocal_rank_fusion(doc_lists) |
209 | 328 |
|
210 | 329 | def retrieve_documents( |
211 | 330 | self, queries: list[str], run_manager: CallbackManagerForRetrieverRun |
212 | 331 | ) -> list[Document]: |
213 | 332 | subdirectory_docs: list[Document] = [] |
214 | | - for subdirectory, retrievers in self._retrievers.items(): |
| 333 | + for subdirectory, retrievers in self.collection_retrievers.items(): |
215 | 334 | bm25_retriever = retrievers["bm25"] |
216 | 335 | vector_retriever = retrievers["vector"] |
217 | 336 | doc_lists: list[list[Document]] = [] |
@@ -250,7 +369,7 @@ async def aretrieve_documents( |
250 | 369 | run_manager: AsyncCallbackManagerForRetrieverRun, |
251 | 370 | ) -> list[Document]: |
252 | 371 | subdirectory_results: dict[str, list[Coroutine[Any, Any, list[Document]]]] = {} |
253 | | - for subdirectory, retrievers in self._retrievers.items(): |
| 372 | + for subdirectory, retrievers in self.collection_retrievers.items(): |
254 | 373 | bm25_retriever = retrievers["bm25"] |
255 | 374 | vector_retriever = retrievers["vector"] |
256 | 375 | subdirectory_results[subdirectory] = [] |
|
0 commit comments