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
23 changes: 20 additions & 3 deletions bin/retrieval_baseline
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,6 @@ import nltk
from dotenv import load_dotenv
from langchain_classic.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
from langchain_core.documents import Document
from langchain_core.retrievers import BaseRetriever
Expand All @@ -46,8 +45,10 @@ from nltk.tokenize import word_tokenize
sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src"))

from agent.models import get_embedding, get_llm
from data_generation.metadata_csv_loader import MetaDataCSVLoader
from retrievers.csv_chroma import (
VECTOR_OVERFETCH,
_csv_column_names,
chroma_settings,
dedupe_by_entity,
list_chroma_subdirectories,
Expand Down Expand Up @@ -83,16 +84,32 @@ def build_retrievers(
embedding = get_embedding("openai", "text-embedding-3-large")

csv_path = embeddings_dir / "csv_files" / f"{collection}.csv"
# MetaDataCSVLoader, not CSVLoader: the plain loader gives documents metadata
# of only {row, source}, so `doc_id` below fell back to row numbers for every
# BM25 document -- defeating its own docstring, since row numbers do not
# survive a bundle rebuild and a rebuild happens every release. Measured
# before this change: 1000 of 1000 BM25 ids were row-based while all 1000
# vector ids were stable.
#
# The application had the same defect and was fixed the same way; this is the
# second caller.
bm25 = BM25Retriever.from_documents(
CSVLoader(file_path=str(csv_path)).load(),
MetaDataCSVLoader(
file_path=str(csv_path), metadata_columns=_csv_column_names(csv_path)
).load(),
preprocess_func=lambda text: word_tokenize(text.casefold(), language="english"),
)
bm25.k = k

vectordb = Chroma(
persist_directory=str(embeddings_dir / collection),
embedding_function=embedding,
client_settings=chroma_settings,
# Called, not passed: chroma_settings is a factory, because
# langchain_chroma mutates the Settings it is given and a shared object
# leaks one store's persist directory into the next. It used to be a
# constant, and this call site was not updated when it changed -- which
# broke this tool for every use, silently, since nothing exercises it.
client_settings=chroma_settings(),
)

# Both vector-side retrievers are asked for more than k and collapsed to k
Expand Down
Loading
Loading