MCPcopy Create free account
hub / github.com/Alibaba-NLP/ViDoRAG / HybridSearchEngine

Class HybridSearchEngine

search_engine.py:218–380  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

216 json.dump(results, json_file, indent=2, ensure_ascii=False)
217
218class HybridSearchEngine:
219
220 def __init__(self,
221 dataset,
222 node_dir_prefix_vl = None,
223 node_dir_prefix_text = None,
224 embed_model_name_vl = 'vidore/colqwen2-v1.0',
225 embed_model_name_text = 'BAAI/bge-m3',
226 topk=10,
227 gmm=False):
228 self.dataset = dataset
229 self.dataset_dir = os.path.join('./data', dataset)
230 self.img_dir = os.path.join(self.dataset_dir, 'img')
231 self.ppocr_dir = os.path.join(self.dataset_dir, 'bge_ingestion')
232 self.engine_vl = SearchEngine(dataset,node_dir_prefix=node_dir_prefix_vl,embed_model_name=embed_model_name_vl)
233 self.engine_text = SearchEngine(dataset,node_dir_prefix=node_dir_prefix_text,embed_model_name=embed_model_name_text)
234 self.topk = topk
235 self.gmm = gmm
236
237 def search(self,query):
238 union_result = False
239 if union_result:
240 if self.gmm:
241 self.engine_vl.gmm = True
242 self.engine_text.gmm = True
243 self.engine_vl.input_gmm = self.topk *2
244 self.engine_text.input_gmm = self.topk *2
245 self.engine_vl.max_output_gmm = self.topk
246 self.engine_text.max_output_gmm = self.topk
247 self.engine_vl.min_output_gmm = self.topk//2
248 self.engine_text.min_output_gmm = self.topk//2
249 result_vl = self.engine_vl.search(query)
250 result_text = self.engine_text.search(query)
251 result_vl['source_nodes'] = result_vl['source_nodes'][:self.topk]
252 result_text['source_nodes'] = result_text['source_nodes'][:self.topk]
253
254 result_docs = dict()
255 for node in result_vl['source_nodes']:
256 file = os.path.basename(node['node']['image_path']).split('.')[0]
257 doc = '_'.join(file.split('_')[:-1])
258 page = file.split('_')[-1]
259 if doc not in result_docs:
260 result_docs[doc] = [int(page)]
261 else:
262 if int(page) not in result_docs[doc]:
263 result_docs[doc].append(int(page))
264
265 for node in result_text['source_nodes']:
266 file = node['node']['metadata']['filename'].split('.')[0]
267 doc = '_'.join(file.split('_')[:-1])
268 page = file.split('_')[-1]
269 if doc not in result_docs:
270 result_docs[doc] = [int(page)]
271 else:
272 if int(page) not in result_docs[doc]:
273 result_docs[doc].append(int(page))
274
275 recall_result = []

Callers 2

__init__Method · 0.90
search_engine.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected