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
49 changes: 48 additions & 1 deletion src/agent/graph.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
import asyncio
import os
import re
from collections.abc import AsyncIterator
from collections.abc import AsyncIterator, Mapping
from dataclasses import dataclass
from typing import Any, Literal, cast

Expand Down Expand Up @@ -198,6 +198,22 @@ def resolve_llm_model(llm_config: "LLMConfig | None") -> tuple[str, str, str | N
MAX_CITATIONS = 12


def _is_sub_retriever(event: Mapping[str, Any]) -> bool:
"""True for the per-collection searches inside the hybrid retriever.

Identified by the tags this repository sets itself in
`HybridRetriever.retrieve_documents` -- `{collection}-bm25-{i}` and
`{collection}-vector-{i}` -- rather than by retriever class name. The name
is not usable: the user guide's own chain-level retriever is a
`VectorStoreRetriever` too, and filtering on that would silently drop every
user-guide citation.
"""
return any(
"-bm25-" in str(tag) or "-vector-" in str(tag)
for tag in (event.get("tags") or [])
)


def _citation_for(document: "Document") -> "AnswerEvent | None":
"""One source, from whichever identifier the collection actually carries.

Expand Down Expand Up @@ -451,7 +467,38 @@ async def astream_answer(
retrieval_done = True
continue

if kind == "on_retriever_end" and _is_sub_retriever(event):
# Skip. The hybrid retriever runs one BM25 and one vector search
# per query variant per collection, and each fires its own
# on_retriever_end -- 51 events for one question. Taking
# citations from those meant taking them from whichever fired
# first, which is always the first collection's BM25.
#
# Measured: 97 of 97 cited documents came from reactions.csv,
# including for "which diseases involve variants of the PTEN
# gene", where the disease_variants documents were retrieved and
# could never be cited. The contract calls citations "the most
# relevant few"; they were one sub-retriever's few.
continue

if kind == "on_retriever_end":
# Skip the sub-retrievers. The Reactome ensemble fires 51 of
# these -- one per collection per query variant per retriever --
# and only the last is the fused result. Citations were taken
# from the first events to arrive, so every one came from
# whichever collection happened to be iterated first: measured,
# 97 of 97 cited documents were from reactions.csv, including for
# "which diseases involve variants of the PTEN gene", whose
# disease_variants documents were retrieved and never cited.
#
# Keyed on the child tags this repo creates itself
# ({collection}-bm25-{i}, {collection}-vector-{i}) rather than on
# the retriever class: the user guide's top-level retriever is a
# VectorStoreRetriever, so filtering by class would silently drop
# its citations instead.
tags = event.get("tags") or []
if any("-bm25-" in tag or "-vector-" in tag for tag in tags):
continue
retrieval_done = True
for document in event["data"].get("output") or []:
if len(seen_citations) >= MAX_CITATIONS:
Expand Down
32 changes: 26 additions & 6 deletions src/retrievers/csv_chroma.py
Original file line number Diff line number Diff line change
Expand Up @@ -445,7 +445,7 @@ def weighted_reciprocal_rank(
def retrieve_documents(
self, queries: list[str], run_manager: CallbackManagerForRetrieverRun
) -> list[Document]:
subdirectory_docs: list[Document] = []
collection_lists: list[list[Document]] = []
chosen = resolve_collections(
selected_collections.get(), self.collection_retrievers
)
Expand Down Expand Up @@ -478,12 +478,28 @@ def retrieve_documents(
# only across query variants. See issue #170.
doc_lists.append(dedupe_by_entity(bm25_docs, RESULTS_PER_RETRIEVER))
doc_lists.append(dedupe_by_entity(vector_docs, RESULTS_PER_RETRIEVER))
subdirectory_docs.extend(
collection_lists.append(
self.weighted_reciprocal_rank(doc_lists)[
: self.max_documents_per_collection
]
)
return subdirectory_docs
# Fused across collections, not concatenated. This is issue #170 one
# level up: there, bm25 and vector results were joined end to end so the
# two were never scored against each other. Here the per-collection lists
# were joined the same way, so the output was grouped by collection
# rather than ranked -- ten reactions, then ten complexes, and so on.
#
# That made the citation cap arbitrary. Measured before this change: 97
# of 97 cited documents came from reactions.csv, including for "which
# diseases involve variants of the PTEN gene", where the disease_variants
# documents were retrieved but sat at positions 41-50 and could never be
# cited. The contract tells the website to treat citations as "the most
# relevant few", which was not true of a list ordered by iteration.
#
# The count is unchanged, so this reorders rather than re-scopes.
return self.weighted_reciprocal_rank(collection_lists)[
: self.max_documents_per_collection * len(collection_lists)
]

async def aretrieve_documents(
self,
Expand Down Expand Up @@ -521,17 +537,21 @@ async def aretrieve_documents(
subdirectory_results[subdirectory].extend(
(bm25_results, vector_results)
)
subdirectory_docs: list[Document] = []
collection_lists: list[list[Document]] = []
for subdir_results in subdirectory_results.values():
# Separate lists, de-duplicated per entity, matching the synchronous
# path above. See issues #169 and #170.
doc_lists: list[list[Document]] = [
dedupe_by_entity(docs, RESULTS_PER_RETRIEVER)
for docs in await asyncio.gather(*subdir_results)
]
subdirectory_docs.extend(
collection_lists.append(
self.weighted_reciprocal_rank(doc_lists)[
: self.max_documents_per_collection
]
)
return subdirectory_docs
# Fused across collections, matching the synchronous path. The two must
# agree document for document, which is what tests/retrievers pins.
return self.weighted_reciprocal_rank(collection_lists)[
: self.max_documents_per_collection * len(collection_lists)
]
Loading