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

import chromadb.config
from langchain_chroma.vectorstores import Chroma
from langchain_community.document_loaders.csv_loader import CSVLoader
from langchain_community.retrievers import BM25Retriever
from langchain_core.callbacks import (
AsyncCallbackManagerForRetrieverRun,
Expand All @@ -21,6 +21,8 @@
from nltk.tokenize import word_tokenize
from pydantic import ConfigDict

from data_generation.metadata_csv_loader import MetaDataCSVLoader


def chroma_settings() -> chromadb.config.Settings:
"""A *fresh* Settings object for every Chroma store.
Expand Down Expand Up @@ -220,6 +222,13 @@ def list_chroma_subdirectories(directory: Path) -> list[str]:
]


def _csv_column_names(csv_path: Path) -> list[str]:
"""Every column in the file, so BM25 metadata matches what was embedded."""
with csv_path.open(newline="", encoding="utf-8") as handle:
header = csv.reader(handle)
return next(header, [])


def create_bm25_chroma_ensemble_retriever(
llm: BaseChatModel,
embedding: Embeddings,
Expand Down Expand Up @@ -282,7 +291,28 @@ def from_subdirectory(
# set up BM25 retriever
csv_file_name = subdirectory + ".csv"
reactome_csvs_dir: Path = embeddings_directory / "csv_files"
loader = CSVLoader(file_path=reactome_csvs_dir / csv_file_name)
csv_path = reactome_csvs_dir / csv_file_name
# MetaDataCSVLoader, not CSVLoader: the plain loader gives every
# document metadata of only {row, source}, so the BM25 half of this
# ensemble returned documents with no `st_id` while the Chroma half
# -- built from the same CSV, by this same loader at embedding time
# -- had one.
#
# Two things followed. Those documents could not be cited, so for one
# measured question 50 of 62 retrieved documents were unattributable
# while their stable id sat in the text as "st_id: R-HSA-...". And
# `unique_documents` keys on `st_id or page_content`, so a BM25 copy
# and its Chroma twin deduplicated to different keys and both
# survived -- 11 of those 50 were exact duplicates of a citable
# document, taking context space and contributing twice to the RRF
# fusion.
#
# Every column is promoted, from the header, so this cannot drift
# from whatever the CSV holds. Content is unchanged: the loader only
# trims content when `content_columns` is given, which it is not.
loader = MetaDataCSVLoader(
file_path=str(csv_path), metadata_columns=_csv_column_names(csv_path)
)
data = loader.load()
bm25_retriever = BM25Retriever.from_documents(
data,
Expand Down
70 changes: 70 additions & 0 deletions tests/retrievers/test_bm25_metadata.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,70 @@
"""BM25 documents must carry the same metadata the embedded ones do.

The ensemble retriever loads each CSV twice over: Chroma holds documents embedded
with `MetaDataCSVLoader`, and BM25 is built by loading the same CSV at startup.
When the second used a plain `CSVLoader`, its documents had metadata of only
{row, source} -- so a document retrieved by keyword could not be cited, even
though its stable id was sitting in the page content where nothing could reach
it. Measured on one question: 50 of 62 documents the citation path walked were
unattributable.
"""

from pathlib import Path

import pytest

from data_generation.metadata_csv_loader import MetaDataCSVLoader
from retrievers.csv_chroma import _csv_column_names

HEADER = "st_id,display_name,pathway_name,url"
ROW = "R-HSA-8863013,CDK5 binds p25,Neuronal System,https://reactome.org/x"


@pytest.fixture
def csv_file(tmp_path: Path) -> Path:
path = tmp_path / "reactions.csv"
path.write_text(f"{HEADER}\n{ROW}\n")
return path


def test_every_column_is_read_from_the_header(csv_file: Path) -> None:
assert _csv_column_names(csv_file) == [
"st_id",
"display_name",
"pathway_name",
"url",
]


def test_a_missing_or_empty_file_asks_for_no_metadata(tmp_path: Path) -> None:
"""An empty list is falsy, so the loader behaves as it did before."""
empty = tmp_path / "empty.csv"
empty.write_text("")
assert _csv_column_names(empty) == []


def test_loaded_documents_carry_the_stable_id(csv_file: Path) -> None:
"""The property the citation path depends on."""
documents = MetaDataCSVLoader(
file_path=str(csv_file), metadata_columns=_csv_column_names(csv_file)
).load()

assert len(documents) == 1
assert documents[0].metadata["st_id"] == "R-HSA-8863013"
assert documents[0].metadata["display_name"] == "CDK5 binds p25"


def test_the_content_is_not_trimmed_by_promoting_columns(csv_file: Path) -> None:
"""BM25 scores page_content, so changing it would change retrieval.

`MetaDataCSVLoader` only narrows content when `content_columns` is given, and
the retriever does not give it. This pins that: every column stays in the
text, which is why retrieval measured identical before and after the change.
"""
documents = MetaDataCSVLoader(
file_path=str(csv_file), metadata_columns=_csv_column_names(csv_file)
).load()

content = documents[0].page_content
for column in ("st_id", "display_name", "pathway_name", "url"):
assert f"{column}:" in content, f"{column} vanished from the indexed text"
Loading