diff --git a/src/retrievers/csv_chroma.py b/src/retrievers/csv_chroma.py index 30333d2d..e07ca537 100644 --- a/src/retrievers/csv_chroma.py +++ b/src/retrievers/csv_chroma.py @@ -4,8 +4,6 @@ from typing import Any, TypedDict import chromadb.config -from langchain.chains.query_constructor.schema import AttributeInfo -from langchain.retrievers.self_query.base import SelfQueryRetriever from langchain_chroma.vectorstores import Chroma from langchain_community.document_loaders.csv_loader import CSVLoader from langchain_community.retrievers import BM25Retriever @@ -196,23 +194,18 @@ def create_bm25_chroma_ensemble_retriever( llm: BaseChatModel, embedding: Embeddings, embeddings_directory: Path, - *, - descriptions_info: dict[str, str], - field_info: dict[str, list[AttributeInfo]], ) -> "HybridRetriever": return HybridRetriever.from_subdirectory( llm, embedding, embeddings_directory, - descriptions_info=descriptions_info, - field_info=field_info, include_original=True, ) class RetrieverDict(TypedDict): bm25: BM25Retriever - vector: SelfQueryRetriever + vector: BaseRetriever class HybridRetriever(BaseRetriever): @@ -236,7 +229,7 @@ class HybridRetriever(BaseRetriever): include_original: bool = False collection_retrievers: dict[str, RetrieverDict] - # BM25Retriever and SelfQueryRetriever are not pydantic models. + # BM25Retriever and the Chroma retriever are not pydantic models. model_config = ConfigDict(arbitrary_types_allowed=True) @classmethod @@ -246,8 +239,6 @@ def from_subdirectory( embedding: Embeddings, embeddings_directory: Path, *, - descriptions_info: dict[str, str], - field_info: dict[str, list[AttributeInfo]], include_original: bool = False, ) -> "HybridRetriever": _retrievers: dict[str, RetrieverDict] = {} @@ -265,24 +256,22 @@ def from_subdirectory( ) bm25_retriever.k = RESULTS_PER_RETRIEVER - # set up vectorstore SelfQuery retriever + # Plain semantic search: similarity against the stored vectors, with + # no LLM in the loop. This replaced SelfQueryRetriever, which spent + # one LLM call per collection per query variant -- 20 per message -- + # translating the question into a Chroma metadata filter. vectordb = Chroma( persist_directory=str(embeddings_directory / subdirectory), embedding_function=embedding, client_settings=chroma_settings, ) - - selfq_retriever = SelfQueryRetriever.from_llm( - llm=llm, - vectorstore=vectordb, - document_contents=descriptions_info[subdirectory], - metadata_field_info=field_info[subdirectory], - search_kwargs={"k": RESULTS_PER_RETRIEVER * VECTOR_OVERFETCH}, + vector_retriever = vectordb.as_retriever( + search_kwargs={"k": RESULTS_PER_RETRIEVER * VECTOR_OVERFETCH} ) _retrievers[subdirectory] = { "bm25": bm25_retriever, - "vector": selfq_retriever, + "vector": vector_retriever, } # The expansion chain, built directly. This used to be extracted from a # throwaway MultiQueryRetriever constructed only to reach its .llm_chain. diff --git a/src/retrievers/plantreactome/rag.py b/src/retrievers/plantreactome/rag.py index 22a9537d..17eafe1e 100644 --- a/src/retrievers/plantreactome/rag.py +++ b/src/retrievers/plantreactome/rag.py @@ -5,10 +5,6 @@ from langchain_core.runnables import Runnable from retrievers.csv_chroma import create_bm25_chroma_ensemble_retriever -from retrievers.plantreactome.metadata_info import ( - plantreactome_descriptions_info, - plantreactome_field_info, -) from retrievers.plantreactome.prompt import plantreactome_qa_prompt from retrievers.rag_chain import create_rag_chain from util.embedding_environment import EmbeddingEnvironment @@ -27,8 +23,6 @@ def create_plantreactome_rag( llm, embedding, embeddings_directory, - descriptions_info=plantreactome_descriptions_info, - field_info=plantreactome_field_info, ) if streaming: diff --git a/src/retrievers/reactome/rag.py b/src/retrievers/reactome/rag.py index 0611e270..5ce5a3b3 100644 --- a/src/retrievers/reactome/rag.py +++ b/src/retrievers/reactome/rag.py @@ -6,10 +6,6 @@ from retrievers.csv_chroma import create_bm25_chroma_ensemble_retriever from retrievers.rag_chain import create_rag_chain -from retrievers.reactome.metadata_info import ( - reactome_descriptions_info, - reactome_field_info, -) from retrievers.reactome.prompt import reactome_qa_prompt from util.embedding_environment import EmbeddingEnvironment @@ -27,8 +23,6 @@ def create_reactome_rag( llm, embedding, embeddings_directory, - descriptions_info=reactome_descriptions_info, - field_info=reactome_field_info, ) if streaming: diff --git a/src/retrievers/uniprot/rag.py b/src/retrievers/uniprot/rag.py index 676d3453..dc2d9aa5 100644 --- a/src/retrievers/uniprot/rag.py +++ b/src/retrievers/uniprot/rag.py @@ -6,10 +6,6 @@ from retrievers.csv_chroma import create_bm25_chroma_ensemble_retriever from retrievers.rag_chain import create_rag_chain -from retrievers.uniprot.metadata_info import ( - uniprot_descriptions_info, - uniprot_field_info, -) from retrievers.uniprot.prompt import uniprot_qa_prompt from util.embedding_environment import EmbeddingEnvironment @@ -27,8 +23,6 @@ def create_uniprot_rag( llm, embedding, embeddings_directory, - descriptions_info=uniprot_descriptions_info, - field_info=uniprot_field_info, ) if streaming: diff --git a/tests/retrievers/test_hybrid_retriever.py b/tests/retrievers/test_hybrid_retriever.py index f24b20dc..990b2287 100644 --- a/tests/retrievers/test_hybrid_retriever.py +++ b/tests/retrievers/test_hybrid_retriever.py @@ -11,7 +11,7 @@ """ from pathlib import Path -from typing import Any, cast +from typing import Any, cast, get_type_hints import pytest @@ -19,10 +19,13 @@ pytest.importorskip("chromadb", reason="retrieval stack not installed") from langchain_core.documents import Document # noqa: E402 +from langchain_core.retrievers import BaseRetriever # noqa: E402 +from retrievers import csv_chroma # noqa: E402 from retrievers.csv_chroma import ( # noqa: E402 MAX_DOCUMENTS_PER_COLLECTION, HybridRetriever, + RetrieverDict, dedupe_by_entity, list_chroma_subdirectories, ) @@ -200,3 +203,32 @@ def test_fused_results_are_capped_per_collection() -> None: assert ( MAX_DOCUMENTS_PER_COLLECTION < 50 ), "the cap, applied by the caller, is what bounds the prompt" + + +def test_the_vector_side_makes_no_llm_call() -> None: + """D1: SelfQueryRetriever is gone, so retrieval costs one LLM call, not 21. + + SelfQueryRetriever translated each question into a Chroma metadata filter with + an LLM call, once per collection per query variant -- 4 x 5 = 20 per message, + plus the expansion. The vector side is now plain similarity search. + + Asserted structurally rather than by counting calls, because a call counter + would pass just as well against a cached or mocked LLM. + """ + source = Path(csv_chroma.__file__).read_text() + # Checks the import, not the word: the name still appears in a comment + # explaining what was replaced, and that comment is worth keeping. + assert ( + "from langchain.retrievers.self_query" not in source + ), "the vector side must not reintroduce an LLM-backed retriever" + assert "as_retriever(" in source, "plain similarity search is expected" + + +def test_retriever_dict_accepts_any_base_retriever() -> None: + """The vector slot is typed to the contract, not to one implementation. + + It was `SelfQueryRetriever`, which meant swapping the implementation was a + type change as well as a behaviour change. + """ + hints = get_type_hints(RetrieverDict) + assert hints["vector"] is BaseRetriever