Skip to content

Commit 04abc65

Browse files
committed
Fix params
1 parent 59379b0 commit 04abc65

2 files changed

Lines changed: 165 additions & 25 deletions

File tree

haystack/components/retrievers/multi_retriever.py

Lines changed: 58 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -81,7 +81,8 @@ def __init__(
8181
*,
8282
retrievers: dict[str, TextRetriever],
8383
filters: dict[str, Any] | None = None,
84-
top_k: int = 10,
84+
top_k_per_retriever: int | None = None,
85+
top_k: int | None = None,
8586
max_workers: int = 4,
8687
join_mode: Literal["concatenate", "reciprocal_rank_fusion"] = "reciprocal_rank_fusion",
8788
) -> None:
@@ -93,8 +94,12 @@ def __init__(
9394
parallel.
9495
:param filters:
9596
A dictionary of filters to apply when retrieving documents.
97+
:param top_k_per_retriever:
98+
The maximum number of documents to return per retriever. If set, this will override the `top_k`
99+
parameter for each retriever. If None, the `top_k` parameter of retrievers will be used.
96100
:param top_k:
97-
The maximum number of documents to return per retriever.
101+
The maximum number of documents to return overall. When set, this will extract the top_k documents
102+
from the combined results of all retrievers. If None, all results are returned.
98103
:param max_workers:
99104
The maximum number of threads to use for parallel retrieval.
100105
:param join_mode:
@@ -104,21 +109,27 @@ def __init__(
104109
"""
105110
self.retrievers = retrievers
106111
self.filters = filters
112+
self.top_k_per_retriever = top_k_per_retriever
107113
self.top_k = top_k
108114
self.max_workers = max_workers
109115
self.join_mode = join_mode
110116
self._is_warmed_up = False
111117

112-
def _merge_results(self, document_lists: list[list[Document]]) -> list[Document]:
118+
def _merge_results(self, document_lists: list[list[Document]], top_k: int | None = None) -> list[Document]:
113119
"""
114120
Merge per-retriever result lists according to `join_mode`.
115121
116122
In `concatenate` mode, all lists are flattened and deduplicated. In `reciprocal_rank_fusion` mode, results
117-
are deduplicated and re-scored using RRF, then returned in descending score order.
123+
are deduplicated and re-scored using RRF, then returned in descending score order. When `top_k` is set, RRF
124+
is always used so the combined results have a consistent global ranking, and only the top `top_k` documents
125+
are returned.
118126
"""
119-
if self.join_mode == "reciprocal_rank_fusion":
127+
# When top_k is set we always use reciprocal rank fusion to merge the results, regardless of join_mode,
128+
# so that truncation is applied to a consistently ranked list.
129+
if top_k is not None or self.join_mode == "reciprocal_rank_fusion":
120130
documents = _reciprocal_rank_fusion(document_lists)
121-
return sorted(documents, key=lambda d: d.score if d.score is not None else -inf, reverse=True)
131+
merged = sorted(documents, key=lambda d: d.score if d.score is not None else -inf, reverse=True)
132+
return merged[:top_k] if top_k is not None else merged
122133
return _deduplicate_documents([doc for docs in document_lists for doc in docs])
123134

124135
def _resolve_retrievers(self, active_retrievers: list[str] | None) -> dict[str, TextRetriever]:
@@ -159,6 +170,7 @@ def run(
159170
self,
160171
query: str,
161172
filters: dict[str, Any] | None = None,
173+
top_k_per_retriever: int | None = None,
162174
top_k: int | None = None,
163175
*,
164176
active_retrievers: list[str] | None = None,
@@ -170,8 +182,14 @@ def run(
170182
The query to run the retrievers on.
171183
:param filters:
172184
Filters to apply. Defaults to the value set at initialization.
185+
:param top_k_per_retriever:
186+
The maximum number of documents to return per retriever. When set, this will override the `top_k`
187+
parameter for each retriever. If None, the `top_k` parameter set for retrievers will be used.
188+
Defaults to the value set at initialization.
173189
:param top_k:
174-
Maximum documents to return per retriever. Defaults to the value set at initialization.
190+
The maximum number of documents to return overall. When set, this will extract the top_k documents
191+
from the combined results of all retrievers. If None, all results are returned. Defaults to the
192+
value set at initialization.
175193
:param active_retrievers:
176194
Names of retrievers to run. Defaults to all. Must match keys in the `retrievers` dictionary.
177195
@@ -185,17 +203,25 @@ def run(
185203
if not self._is_warmed_up:
186204
self.warm_up()
187205

206+
resolved_top_k_per_retriever = (
207+
top_k_per_retriever if top_k_per_retriever is not None else self.top_k_per_retriever
208+
)
188209
resolved_top_k = top_k if top_k is not None else self.top_k
189210
resolved_filters = filters if filters is not None else self.filters
190211

191212
retrievers_to_run = self._resolve_retrievers(active_retrievers)
192213

193214
results_by_name: dict[str, list[Document]] = {}
194215
with ThreadPoolExecutor(max_workers=self.max_workers) as executor:
195-
future_to_name = {
196-
executor.submit(retriever.run, query=query, filters=resolved_filters, top_k=resolved_top_k): name
197-
for name, retriever in retrievers_to_run.items()
198-
}
216+
future_to_name = {}
217+
for name, retriever in retrievers_to_run.items():
218+
run_kwargs: dict[str, Any] = {"query": query}
219+
if resolved_top_k_per_retriever is not None:
220+
run_kwargs["top_k"] = resolved_top_k_per_retriever
221+
if resolved_filters is not None:
222+
run_kwargs["filters"] = resolved_filters
223+
future_to_name[executor.submit(retriever.run, **run_kwargs)] = name
224+
199225
for future in as_completed(future_to_name):
200226
name = future_to_name[future]
201227
try:
@@ -204,13 +230,14 @@ def run(
204230
raise RuntimeError(f"Retriever '{name}' failed: {e}") from e
205231

206232
document_lists = [results_by_name[name] for name in retrievers_to_run]
207-
return {"documents": self._merge_results(document_lists)}
233+
return {"documents": self._merge_results(document_lists, top_k=resolved_top_k)}
208234

209235
@component.output_types(documents=list[Document])
210236
async def run_async(
211237
self,
212238
query: str,
213239
filters: dict[str, Any] | None = None,
240+
top_k_per_retriever: int | None = None,
214241
top_k: int | None = None,
215242
*,
216243
active_retrievers: list[str] | None = None,
@@ -224,8 +251,14 @@ async def run_async(
224251
The query to run the retrievers on.
225252
:param filters:
226253
Filters to apply. Defaults to the value set at initialization.
254+
:param top_k_per_retriever:
255+
The maximum number of documents to return per retriever. When set, this will override the `top_k`
256+
parameter for each retriever. If None, the `top_k` parameter set for retrievers will be used.
257+
Defaults to the value set at initialization.
227258
:param top_k:
228-
Maximum documents to return per retriever. Defaults to the value set at initialization.
259+
The maximum number of documents to return overall. When set, this will extract the top_k documents
260+
from the combined results of all retrievers. If None, all results are returned. Defaults to the
261+
value set at initialization.
229262
:param active_retrievers:
230263
Names of retrievers to run. Defaults to all. Must match keys in the `retrievers` dictionary.
231264
@@ -239,27 +272,34 @@ async def run_async(
239272
if not self._is_warmed_up:
240273
self.warm_up()
241274

275+
resolved_top_k_per_retriever = (
276+
top_k_per_retriever if top_k_per_retriever is not None else self.top_k_per_retriever
277+
)
242278
resolved_top_k = top_k if top_k is not None else self.top_k
243279
resolved_filters = filters if filters is not None else self.filters
244280

245281
retrievers_to_run = self._resolve_retrievers(active_retrievers)
246282

283+
run_kwargs: dict[str, Any] = {"query": query}
284+
if resolved_top_k_per_retriever is not None:
285+
run_kwargs["top_k"] = resolved_top_k_per_retriever
286+
if resolved_filters is not None:
287+
run_kwargs["filters"] = resolved_filters
288+
247289
loop = asyncio.get_running_loop()
248290

249291
async def _run_one(name: str, retriever: TextRetriever) -> list[Document]:
250292
try:
251293
if hasattr(retriever, "run_async") and callable(retriever.run_async):
252-
result = await retriever.run_async(query=query, filters=resolved_filters, top_k=resolved_top_k)
294+
result = await retriever.run_async(**run_kwargs)
253295
else:
254-
result = await loop.run_in_executor(
255-
None, lambda: retriever.run(query=query, filters=resolved_filters, top_k=resolved_top_k)
256-
)
296+
result = await loop.run_in_executor(None, lambda r=retriever: r.run(**run_kwargs))
257297
return result.get("documents", [])
258298
except Exception as e:
259299
raise RuntimeError(f"Retriever '{name}' failed: {e}") from e
260300

261301
document_lists = list(await asyncio.gather(*[_run_one(name, r) for name, r in retrievers_to_run.items()]))
262-
return {"documents": self._merge_results(document_lists)}
302+
return {"documents": self._merge_results(document_lists, top_k=resolved_top_k)}
263303

264304
def to_dict(self) -> dict[str, Any]:
265305
"""

test/components/retrievers/test_multi_retriever.py

Lines changed: 107 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -161,7 +161,7 @@ def test_run_deduplicates_results(self, sample_documents):
161161
ids = [doc.id for doc in result["documents"]]
162162
assert ids.count("doc1") == 1
163163

164-
def test_run_resolves_filters_and_top_k(self):
164+
def test_run_resolves_filters_and_top_k_per_retriever(self):
165165
received: dict = {}
166166

167167
@component
@@ -173,19 +173,68 @@ def run(self, query: str, filters: dict[str, Any] | None = None, top_k: int | No
173173
return {"documents": []}
174174

175175
retriever = MultiRetriever(
176-
retrievers={"capturing": CapturingRetriever()}, filters={"field": "meta.category"}, top_k=5
176+
retrievers={"capturing": CapturingRetriever()}, filters={"field": "meta.category"}, top_k_per_retriever=5
177177
)
178178

179-
# Should use init-time values when not overridden
179+
# Should use init-time values when not overridden (top_k_per_retriever is forwarded as the retriever's top_k)
180180
retriever.run(query="energy")
181181
assert received["filters"] == {"field": "meta.category"}
182182
assert received["top_k"] == 5
183183

184184
# Should prefer run-time values when provided
185-
retriever.run(query="energy", filters={"field": "meta.other"}, top_k=2)
185+
retriever.run(query="energy", filters={"field": "meta.other"}, top_k_per_retriever=2)
186186
assert received["filters"] == {"field": "meta.other"}
187187
assert received["top_k"] == 2
188188

189+
def test_run_forwards_top_k_per_retriever_not_overall_top_k(self):
190+
received: dict = {}
191+
192+
@component
193+
class CapturingRetriever:
194+
@component.output_types(documents=list[Document])
195+
def run(self, query: str, filters: dict[str, Any] | None = None, top_k: int | None = None):
196+
received["top_k"] = top_k
197+
return {"documents": []}
198+
199+
retriever = MultiRetriever(retrievers={"capturing": CapturingRetriever()})
200+
201+
# top_k_per_retriever is forwarded to each retriever as its top_k
202+
retriever.run(query="energy", top_k_per_retriever=3)
203+
assert received["top_k"] == 3
204+
205+
# the overall top_k is applied at merge-time only, not forwarded to retrievers
206+
received.clear()
207+
retriever.run(query="energy", top_k=5)
208+
assert received.get("top_k") is None
209+
210+
def test_run_top_k_truncates_merged_results(self, sample_documents):
211+
retriever = MultiRetriever(
212+
retrievers={
213+
"a": MockRetriever(documents=sample_documents[:3]),
214+
"b": MockRetriever(documents=sample_documents[2:5]),
215+
},
216+
max_workers=2,
217+
)
218+
result = retriever.run(query="energy", top_k=2)
219+
assert len(result["documents"]) == 2
220+
scores = [doc.score for doc in result["documents"]]
221+
assert all(score is not None for score in scores)
222+
assert scores == sorted(scores, reverse=True)
223+
224+
def test_run_top_k_forces_rrf_in_concatenate_mode(self, sample_documents):
225+
# In concatenate mode there is no global ranking, so setting top_k falls back to RRF to truncate consistently
226+
retriever = MultiRetriever(
227+
retrievers={
228+
"a": MockRetriever(documents=sample_documents[:3]),
229+
"b": MockRetriever(documents=sample_documents[1:4]),
230+
},
231+
join_mode="concatenate",
232+
max_workers=2,
233+
)
234+
result = retriever.run(query="energy", top_k=2)
235+
assert len(result["documents"]) == 2
236+
assert all(doc.score is not None for doc in result["documents"])
237+
189238
def test_run_with_active_retrievers(self, sample_documents):
190239
retriever = MultiRetriever(
191240
retrievers={"a": MockRetriever([sample_documents[0]]), "b": MockRetriever([sample_documents[1]])}
@@ -379,7 +428,7 @@ async def test_run_async_rrf_assigns_scores_and_sorts(self, sample_documents):
379428
assert ids.index("doc1") < ids.index("doc3")
380429

381430
@pytest.mark.asyncio
382-
async def test_run_async_resolves_filters_and_top_k(self):
431+
async def test_run_async_resolves_filters_and_top_k_per_retriever(self):
383432
received: dict = {}
384433

385434
@component
@@ -391,17 +440,68 @@ def run(self, query: str, filters: dict[str, Any] | None = None, top_k: int | No
391440
return {"documents": []}
392441

393442
retriever = MultiRetriever(
394-
retrievers={"capturing": CapturingRetriever()}, filters={"field": "meta.category"}, top_k=5
443+
retrievers={"capturing": CapturingRetriever()}, filters={"field": "meta.category"}, top_k_per_retriever=5
395444
)
396445

446+
# top_k_per_retriever is forwarded as the retriever's top_k
397447
await retriever.run_async(query="energy")
398448
assert received["filters"] == {"field": "meta.category"}
399449
assert received["top_k"] == 5
400450

401-
await retriever.run_async(query="energy", filters={"field": "meta.other"}, top_k=2)
451+
await retriever.run_async(query="energy", filters={"field": "meta.other"}, top_k_per_retriever=2)
402452
assert received["filters"] == {"field": "meta.other"}
403453
assert received["top_k"] == 2
404454

455+
@pytest.mark.asyncio
456+
async def test_run_async_forwards_top_k_per_retriever_not_overall_top_k(self):
457+
received: dict = {}
458+
459+
@component
460+
class CapturingRetriever:
461+
@component.output_types(documents=list[Document])
462+
def run(self, query: str, filters: dict[str, Any] | None = None, top_k: int | None = None):
463+
received["top_k"] = top_k
464+
return {"documents": []}
465+
466+
retriever = MultiRetriever(retrievers={"capturing": CapturingRetriever()})
467+
468+
# top_k_per_retriever is forwarded to each retriever as its top_k
469+
await retriever.run_async(query="energy", top_k_per_retriever=3)
470+
assert received["top_k"] == 3
471+
472+
# the overall top_k is applied at merge-time only, not forwarded to retrievers
473+
received.clear()
474+
await retriever.run_async(query="energy", top_k=5)
475+
assert received.get("top_k") is None
476+
477+
@pytest.mark.asyncio
478+
async def test_run_async_top_k_truncates_merged_results(self, sample_documents):
479+
retriever = MultiRetriever(
480+
retrievers={
481+
"a": MockRetriever(documents=sample_documents[:3]),
482+
"b": MockRetriever(documents=sample_documents[2:5]),
483+
}
484+
)
485+
result = await retriever.run_async(query="energy", top_k=2)
486+
assert len(result["documents"]) == 2
487+
scores = [doc.score for doc in result["documents"]]
488+
assert all(score is not None for score in scores)
489+
assert scores == sorted(scores, reverse=True)
490+
491+
@pytest.mark.asyncio
492+
async def test_run_async_top_k_forces_rrf_in_concatenate_mode(self, sample_documents):
493+
# In concatenate mode there is no global ranking, so setting top_k falls back to RRF to truncate consistently
494+
retriever = MultiRetriever(
495+
retrievers={
496+
"a": MockRetriever(documents=sample_documents[:3]),
497+
"b": MockRetriever(documents=sample_documents[1:4]),
498+
},
499+
join_mode="concatenate",
500+
)
501+
result = await retriever.run_async(query="energy", top_k=2)
502+
assert len(result["documents"]) == 2
503+
assert all(doc.score is not None for doc in result["documents"])
504+
405505
@pytest.mark.asyncio
406506
async def test_run_async_with_active_retrievers(self, sample_documents):
407507
retriever = MultiRetriever(

0 commit comments

Comments
 (0)