Coverage for haystack/components/retrievers/multi_query_embedding_retriever.py: 97%
78 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-21 13:53 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-21 13:53 +0000
1# SPDX-FileCopyrightText: 2022-present deepset GmbH <info@deepset.ai>
2#
3# SPDX-License-Identifier: Apache-2.0
5import asyncio
6from concurrent.futures import ThreadPoolExecutor
7from typing import Any
9from haystack import Document, component, default_from_dict, default_to_dict
10from haystack.components.embedders.types.protocol import TextEmbedder
11from haystack.components.retrievers.types import EmbeddingRetriever
12from haystack.core.serialization import component_to_dict
13from haystack.utils.async_utils import _execute_component_async, _gather_tasks_with_cancel
14from haystack.utils.misc import _deduplicate_documents
17@component
18class MultiQueryEmbeddingRetriever:
19 """
20 A component that retrieves documents using multiple queries in parallel with an embedding-based retriever.
22 This component takes a list of text queries, converts them to embeddings using a query embedder,
23 and then uses an embedding-based retriever to find relevant documents for each query in parallel.
24 The results are combined and sorted by relevance score.
26 ### Usage example
28 ```python
29 from haystack import Document
30 from haystack.document_stores.in_memory import InMemoryDocumentStore
31 from haystack.document_stores.types import DuplicatePolicy
32 from haystack.components.embedders import OpenAITextEmbedder
33 from haystack.components.embedders import OpenAIDocumentEmbedder
34 from haystack.components.retrievers import InMemoryEmbeddingRetriever
35 from haystack.components.writers import DocumentWriter
36 from haystack.components.retrievers import MultiQueryEmbeddingRetriever
38 documents = [
39 Document(content="Renewable energy is energy that is collected from renewable resources."),
40 Document(content="Solar energy is a type of green energy that is harnessed from the sun."),
41 Document(content="Wind energy is another type of green energy that is generated by wind turbines."),
42 Document(content="Geothermal energy is heat that comes from the sub-surface of the earth."),
43 Document(content="Biomass energy is produced from organic materials, such as plant and animal waste."),
44 Document(content="Fossil fuels, such as coal, oil, and natural gas, are non-renewable energy sources."),
45 ]
47 # Populate the document store
48 doc_store = InMemoryDocumentStore()
49 doc_embedder = OpenAIDocumentEmbedder()
50 doc_writer = DocumentWriter(document_store=doc_store, policy=DuplicatePolicy.SKIP)
51 documents = doc_embedder.run(documents)["documents"]
52 doc_writer.run(documents=documents)
54 # Run the multi-query retriever
55 in_memory_retriever = InMemoryEmbeddingRetriever(document_store=doc_store, top_k=1)
56 query_embedder = OpenAITextEmbedder()
58 multi_query_retriever = MultiQueryEmbeddingRetriever(
59 retriever=in_memory_retriever,
60 query_embedder=query_embedder,
61 max_workers=3
62 )
64 queries = ["Geothermal energy", "natural gas", "turbines"]
65 result = multi_query_retriever.run(queries=queries)
66 for doc in result["documents"]:
67 print(f"Content: {doc.content}, Score: {doc.score}")
68 # >> Content: Geothermal energy is heat that comes from the sub-surface of the earth., Score: 0.8509603046266574
69 # >> Content: Renewable energy is energy that is collected from renewable resources., Score: 0.42763211298893034
70 # >> Content: Solar energy is a type of green energy that is harnessed from the sun., Score: 0.40077417016494354
71 # >> Content: Fossil fuels, such as coal, oil, and natural gas, are non-renewable energy sources., Score: 0.3774863680
72 # >> Content: Wind energy is another type of green energy that is generated by wind turbines., Score: 0.30914239725622
73 # >> Content: Biomass energy is produced from organic materials, such as plant and animal waste., Score: 0.25173074243
74 ```
75 """ # noqa E501
77 def __init__(self, *, retriever: EmbeddingRetriever, query_embedder: TextEmbedder, max_workers: int = 3) -> None:
78 """
79 Initialize MultiQueryEmbeddingRetriever.
81 :param retriever: The embedding-based retriever to use for document retrieval.
82 :param query_embedder: The query embedder to convert text queries to embeddings.
83 :param max_workers: Maximum number of worker threads for parallel processing.
84 """
85 self.retriever = retriever
86 self.query_embedder = query_embedder
87 self.max_workers = max_workers
89 def warm_up(self) -> None:
90 """
91 Warm up the query embedder and the retriever.
92 """
93 for inner in (self.query_embedder, self.retriever):
94 if hasattr(inner, "warm_up"):
95 inner.warm_up()
97 async def warm_up_async(self) -> None:
98 """
99 Warm up the query embedder and the retriever on the serving event loop.
100 """
101 for inner in (self.query_embedder, self.retriever):
102 if hasattr(inner, "warm_up_async"):
103 await inner.warm_up_async()
104 elif hasattr(inner, "warm_up"):
105 inner.warm_up()
107 def close(self) -> None:
108 """
109 Release the query embedder's and the retriever's resources.
110 """
111 for inner in (self.query_embedder, self.retriever):
112 if hasattr(inner, "close"):
113 inner.close()
115 async def close_async(self) -> None:
116 """
117 Release the query embedder's and the retriever's async resources.
118 """
119 for inner in (self.query_embedder, self.retriever):
120 if hasattr(inner, "close_async"):
121 await inner.close_async()
122 elif hasattr(inner, "close"):
123 inner.close()
125 @component.output_types(documents=list[Document])
126 def run(self, queries: list[str], retriever_kwargs: dict[str, Any] | None = None) -> dict[str, list[Document]]:
127 """
128 Retrieve documents using multiple queries in parallel.
130 :param queries: List of text queries to process.
131 :param retriever_kwargs: Optional dictionary of arguments to pass to the retriever's run method.
132 :returns:
133 A dictionary containing:
134 - `documents`: List of retrieved documents sorted by relevance score.
135 """
136 docs: list[Document] = []
137 retriever_kwargs = retriever_kwargs or {}
139 self.warm_up()
141 with ThreadPoolExecutor(max_workers=self.max_workers) as executor:
142 queries_results = executor.map(lambda query: self._run_on_thread(query, retriever_kwargs), queries)
143 for result in queries_results:
144 if not result:
145 continue
146 docs.extend(result)
148 # de-duplicate and sort
149 docs = _deduplicate_documents(docs)
150 docs.sort(key=lambda x: x.score or 0.0, reverse=True)
151 return {"documents": docs}
153 @component.output_types(documents=list[Document])
154 async def run_async(
155 self, queries: list[str], retriever_kwargs: dict[str, Any] | None = None
156 ) -> dict[str, list[Document]]:
157 """
158 Retrieve documents using multiple queries concurrently.
160 Uses each component's `run_async` method if available, otherwise falls back to running `run`
161 in a thread executor. Queries are processed concurrently using asyncio.gather.
163 :param queries: List of text queries to process.
164 :param retriever_kwargs: Optional dictionary of arguments to pass to the retriever's run method.
165 :returns:
166 A dictionary containing:
167 - `documents`: List of retrieved documents sorted by relevance score.
168 """
169 retriever_kwargs = retriever_kwargs or {}
171 await self.warm_up_async()
173 tasks = [asyncio.create_task(self._run_one_async(query, retriever_kwargs)) for query in queries]
174 results = await _gather_tasks_with_cancel(tasks)
175 docs: list[Document] = [doc for result in results if result for doc in result]
176 docs = _deduplicate_documents(docs)
177 docs.sort(key=lambda x: x.score or 0.0, reverse=True)
178 return {"documents": docs}
180 def _run_on_thread(self, query: str, retriever_kwargs: dict[str, Any] | None = None) -> list[Document] | None:
181 """
182 Process a single query on a separate thread.
184 :param query: The text query to process.
185 :param retriever_kwargs: Arguments to pass to the retriever's run method.
186 :returns:
187 List of retrieved documents or None if no results.
188 """
189 embedding_result = self.query_embedder.run(text=query)
190 query_embedding = embedding_result["embedding"]
191 result = self.retriever.run(query_embedding=query_embedding, **(retriever_kwargs or {}))
192 if result and "documents" in result:
193 return result["documents"]
194 return None
196 async def _run_one_async(self, query: str, retriever_kwargs: dict[str, Any]) -> list[Document] | None:
197 """
198 Process a single query asynchronously.
200 :param query: The text query to process.
201 :param retriever_kwargs: Arguments to pass to the retriever's run method.
202 :returns:
203 List of retrieved documents or None if no results.
204 """
205 embedding_result = await _execute_component_async(self.query_embedder, text=query)
207 query_embedding = embedding_result["embedding"]
209 result = await _execute_component_async(self.retriever, query_embedding=query_embedding, **retriever_kwargs)
211 if result and "documents" in result:
212 return result["documents"]
213 return None
215 def to_dict(self) -> dict[str, Any]:
216 """
217 Serializes the component to a dictionary.
219 :returns:
220 A dictionary representing the serialized component.
221 """
222 return default_to_dict(
223 self,
224 retriever=component_to_dict(obj=self.retriever, name="retriever"),
225 query_embedder=component_to_dict(obj=self.query_embedder, name="query_embedder"),
226 max_workers=self.max_workers,
227 )
229 @classmethod
230 def from_dict(cls, data: dict[str, Any]) -> "MultiQueryEmbeddingRetriever":
231 """
232 Deserializes the component from a dictionary.
234 :param data: The dictionary to deserialize from.
235 :returns:
236 The deserialized component.
237 """
238 return default_from_dict(cls, data)