@@ -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 ,
0 commit comments