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
29 changes: 9 additions & 20 deletions src/retrievers/csv_chroma.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand All @@ -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
Expand All @@ -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] = {}
Expand All @@ -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.
Expand Down
6 changes: 0 additions & 6 deletions src/retrievers/plantreactome/rag.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -27,8 +23,6 @@ def create_plantreactome_rag(
llm,
embedding,
embeddings_directory,
descriptions_info=plantreactome_descriptions_info,
field_info=plantreactome_field_info,
)

if streaming:
Expand Down
6 changes: 0 additions & 6 deletions src/retrievers/reactome/rag.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -27,8 +23,6 @@ def create_reactome_rag(
llm,
embedding,
embeddings_directory,
descriptions_info=reactome_descriptions_info,
field_info=reactome_field_info,
)

if streaming:
Expand Down
6 changes: 0 additions & 6 deletions src/retrievers/uniprot/rag.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -27,8 +23,6 @@ def create_uniprot_rag(
llm,
embedding,
embeddings_directory,
descriptions_info=uniprot_descriptions_info,
field_info=uniprot_field_info,
)

if streaming:
Expand Down
34 changes: 33 additions & 1 deletion tests/retrievers/test_hybrid_retriever.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,18 +11,21 @@
"""

from pathlib import Path
from typing import Any, cast
from typing import Any, cast, get_type_hints

import pytest

pytest.importorskip("langchain", reason="retrieval stack not installed")
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,
)
Expand Down Expand Up @@ -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
Loading