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

1# SPDX-FileCopyrightText: 2022-present deepset GmbH <info@deepset.ai> 

2# 

3# SPDX-License-Identifier: Apache-2.0 

4 

5import asyncio 

6from concurrent.futures import ThreadPoolExecutor 

7from typing import Any 

8 

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 

15 

16 

17@component 

18class MultiQueryEmbeddingRetriever: 

19 """ 

20 A component that retrieves documents using multiple queries in parallel with an embedding-based retriever. 

21 

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. 

25 

26 ### Usage example 

27 

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 

37 

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 ] 

46 

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) 

53 

54 # Run the multi-query retriever 

55 in_memory_retriever = InMemoryEmbeddingRetriever(document_store=doc_store, top_k=1) 

56 query_embedder = OpenAITextEmbedder() 

57 

58 multi_query_retriever = MultiQueryEmbeddingRetriever( 

59 retriever=in_memory_retriever, 

60 query_embedder=query_embedder, 

61 max_workers=3 

62 ) 

63 

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 

76 

77 def __init__(self, *, retriever: EmbeddingRetriever, query_embedder: TextEmbedder, max_workers: int = 3) -> None: 

78 """ 

79 Initialize MultiQueryEmbeddingRetriever. 

80 

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 

88 

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() 

96 

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() 

106 

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() 

114 

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() 

124 

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. 

129 

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 {} 

138 

139 self.warm_up() 

140 

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) 

147 

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} 

152 

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. 

159 

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. 

162 

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 {} 

170 

171 await self.warm_up_async() 

172 

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} 

179 

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. 

183 

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 

195 

196 async def _run_one_async(self, query: str, retriever_kwargs: dict[str, Any]) -> list[Document] | None: 

197 """ 

198 Process a single query asynchronously. 

199 

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) 

206 

207 query_embedding = embedding_result["embedding"] 

208 

209 result = await _execute_component_async(self.retriever, query_embedding=query_embedding, **retriever_kwargs) 

210 

211 if result and "documents" in result: 

212 return result["documents"] 

213 return None 

214 

215 def to_dict(self) -> dict[str, Any]: 

216 """ 

217 Serializes the component to a dictionary. 

218 

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 ) 

228 

229 @classmethod 

230 def from_dict(cls, data: dict[str, Any]) -> "MultiQueryEmbeddingRetriever": 

231 """ 

232 Deserializes the component from a dictionary. 

233 

234 :param data: The dictionary to deserialize from. 

235 :returns: 

236 The deserialized component. 

237 """ 

238 return default_from_dict(cls, data)