Skip to content

Commit b25f5d8

Browse files
fix: update Multiretriever params (#11823)
Co-authored-by: bogdankostic <bogdankostic@web.de>
1 parent b4b2843 commit b25f5d8

3 files changed

Lines changed: 182 additions & 26 deletions

File tree

haystack/components/retrievers/multi_retriever.py

Lines changed: 63 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -83,7 +83,8 @@ def __init__(
8383
*,
8484
retrievers: dict[str, TextRetriever],
8585
filters: dict[str, Any] | None = None,
86-
top_k: int = 10,
86+
top_k_per_retriever: int | None = None,
87+
top_k: int | None = None,
8788
max_workers: int = 4,
8889
join_mode: Literal["concatenate", "reciprocal_rank_fusion"] = "reciprocal_rank_fusion",
8990
) -> None:
@@ -95,8 +96,14 @@ def __init__(
9596
parallel.
9697
:param filters:
9798
A dictionary of filters to apply when retrieving documents.
99+
:param top_k_per_retriever:
100+
The maximum number of documents to return per retriever. If set, this will override the `top_k`
101+
parameter for each retriever. If None, the `top_k` parameter of retrievers will be used.
98102
:param top_k:
99-
The maximum number of documents to return per retriever.
103+
The maximum number of documents to return overall, extracted from the combined results of all
104+
retrievers. When set, the results are always merged using reciprocal rank fusion (regardless of
105+
`join_mode`) so that the combined list has a consistent global ranking before it is truncated to
106+
`top_k`. If None, all results are returned.
100107
:param max_workers:
101108
The maximum number of threads to use for parallel retrieval.
102109
:param join_mode:
@@ -106,21 +113,27 @@ def __init__(
106113
"""
107114
self.retrievers = retrievers
108115
self.filters = filters
116+
self.top_k_per_retriever = top_k_per_retriever
109117
self.top_k = top_k
110118
self.max_workers = max_workers
111119
self.join_mode = join_mode
112120
self._is_warmed_up = False
113121

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

126139
def _resolve_retrievers(self, active_retrievers: list[str] | None) -> dict[str, TextRetriever]:
@@ -161,6 +174,7 @@ def run(
161174
self,
162175
query: str,
163176
filters: dict[str, Any] | None = None,
177+
top_k_per_retriever: int | None = None,
164178
top_k: int | None = None,
165179
*,
166180
active_retrievers: list[str] | None = None,
@@ -172,8 +186,15 @@ def run(
172186
The query to run the retrievers on.
173187
:param filters:
174188
Filters to apply. Defaults to the value set at initialization.
189+
:param top_k_per_retriever:
190+
The maximum number of documents to return per retriever. When set, this will override the `top_k`
191+
parameter for each retriever. If None, the `top_k` parameter set for retrievers will be used.
192+
Defaults to the value set at initialization.
175193
:param top_k:
176-
Maximum documents to return per retriever. Defaults to the value set at initialization.
194+
The maximum number of documents to return overall, extracted from the combined results of all
195+
retrievers. When set, the results are always merged using reciprocal rank fusion (regardless of
196+
`join_mode`) so that the combined list has a consistent global ranking before it is truncated to
197+
`top_k`. If None, all results are returned. Defaults to the value set at initialization.
177198
:param active_retrievers:
178199
Names of retrievers to run. Defaults to all. Must match keys in the `retrievers` dictionary.
179200
@@ -187,17 +208,25 @@ def run(
187208
if not self._is_warmed_up:
188209
self.warm_up()
189210

211+
resolved_top_k_per_retriever = (
212+
top_k_per_retriever if top_k_per_retriever is not None else self.top_k_per_retriever
213+
)
190214
resolved_top_k = top_k if top_k is not None else self.top_k
191215
resolved_filters = filters if filters is not None else self.filters
192216

193217
retrievers_to_run = self._resolve_retrievers(active_retrievers)
194218

195219
results_by_name: dict[str, list[Document]] = {}
196220
with ThreadPoolExecutor(max_workers=self.max_workers) as executor:
197-
future_to_name = {
198-
executor.submit(retriever.run, query=query, filters=resolved_filters, top_k=resolved_top_k): name
199-
for name, retriever in retrievers_to_run.items()
200-
}
221+
future_to_name = {}
222+
for name, retriever in retrievers_to_run.items():
223+
run_kwargs: dict[str, Any] = {"query": query}
224+
if resolved_top_k_per_retriever is not None:
225+
run_kwargs["top_k"] = resolved_top_k_per_retriever
226+
if resolved_filters is not None:
227+
run_kwargs["filters"] = resolved_filters
228+
future_to_name[executor.submit(retriever.run, **run_kwargs)] = name
229+
201230
for future in as_completed(future_to_name):
202231
name = future_to_name[future]
203232
try:
@@ -206,13 +235,14 @@ def run(
206235
raise RuntimeError(f"Retriever '{name}' failed: {e}") from e
207236

208237
document_lists = [results_by_name[name] for name in retrievers_to_run]
209-
return {"documents": self._merge_results(document_lists)}
238+
return {"documents": self._merge_results(document_lists, top_k=resolved_top_k)}
210239

211240
@component.output_types(documents=list[Document])
212241
async def run_async(
213242
self,
214243
query: str,
215244
filters: dict[str, Any] | None = None,
245+
top_k_per_retriever: int | None = None,
216246
top_k: int | None = None,
217247
*,
218248
active_retrievers: list[str] | None = None,
@@ -226,8 +256,15 @@ async def run_async(
226256
The query to run the retrievers on.
227257
:param filters:
228258
Filters to apply. Defaults to the value set at initialization.
259+
:param top_k_per_retriever:
260+
The maximum number of documents to return per retriever. When set, this will override the `top_k`
261+
parameter for each retriever. If None, the `top_k` parameter set for retrievers will be used.
262+
Defaults to the value set at initialization.
229263
:param top_k:
230-
Maximum documents to return per retriever. Defaults to the value set at initialization.
264+
The maximum number of documents to return overall, extracted from the combined results of all
265+
retrievers. When set, the results are always merged using reciprocal rank fusion (regardless of
266+
`join_mode`) so that the combined list has a consistent global ranking before it is truncated to
267+
`top_k`. If None, all results are returned. Defaults to the value set at initialization.
231268
:param active_retrievers:
232269
Names of retrievers to run. Defaults to all. Must match keys in the `retrievers` dictionary.
233270
@@ -241,27 +278,34 @@ async def run_async(
241278
if not self._is_warmed_up:
242279
self.warm_up()
243280

281+
resolved_top_k_per_retriever = (
282+
top_k_per_retriever if top_k_per_retriever is not None else self.top_k_per_retriever
283+
)
244284
resolved_top_k = top_k if top_k is not None else self.top_k
245285
resolved_filters = filters if filters is not None else self.filters
246286

247287
retrievers_to_run = self._resolve_retrievers(active_retrievers)
248288

289+
run_kwargs: dict[str, Any] = {"query": query}
290+
if resolved_top_k_per_retriever is not None:
291+
run_kwargs["top_k"] = resolved_top_k_per_retriever
292+
if resolved_filters is not None:
293+
run_kwargs["filters"] = resolved_filters
294+
249295
loop = asyncio.get_running_loop()
250296

251297
async def _run_one(name: str, retriever: TextRetriever) -> list[Document]:
252298
try:
253299
if hasattr(retriever, "run_async") and callable(retriever.run_async):
254-
result = await retriever.run_async(query=query, filters=resolved_filters, top_k=resolved_top_k)
300+
result = await retriever.run_async(**run_kwargs)
255301
else:
256-
result = await loop.run_in_executor(
257-
None, lambda: retriever.run(query=query, filters=resolved_filters, top_k=resolved_top_k)
258-
)
302+
result = await loop.run_in_executor(None, lambda: retriever.run(**run_kwargs))
259303
return result.get("documents", [])
260304
except Exception as e:
261305
raise RuntimeError(f"Retriever '{name}' failed: {e}") from e
262306

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

266310
def to_dict(self) -> dict[str, Any]:
267311
"""
@@ -274,6 +318,7 @@ def to_dict(self) -> dict[str, Any]:
274318
self,
275319
retrievers={name: component_to_dict(obj=r, name=name) for name, r in self.retrievers.items()},
276320
filters=self.filters,
321+
top_k_per_retriever=self.top_k_per_retriever,
277322
top_k=self.top_k,
278323
max_workers=self.max_workers,
279324
join_mode=self.join_mode,
Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,6 @@
1+
---
2+
fixes:
3+
- |
4+
Update parameters in MultiRetriever to allow for more flexible retrieval of documents. Changes include:
5+
- ``top_k_per_retriever``: Allows specifying the maximum number of documents to return per retriever
6+
- ``top_k``: Allows specifying the maximum number of documents to be retrieved from combined results of all retrievers

0 commit comments

Comments
 (0)