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