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

Class SearchEngine

search_engine.py:61–216  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

59 return valid_recall_result
60
61class SearchEngine:
62 def __init__(self,dataset, node_dir_prefix=None,embed_model_name='BAAI/bge-m3'):# nvidia/NV-Embed-v2 "vidore/colqwen2-v0.1"
63 Settings.llm = None
64 self.gmm=False
65 self.gmm_candidate_length = False
66 self.return_raw = False
67 self.input_gmm = 20
68 self.max_output_gmm = 10
69 self.min_output_gmm = 5
70 self.dataset = dataset
71 self.dataset_dir = os.path.join('./data', dataset)
72 if node_dir_prefix is None:
73 if 'bge' in embed_model_name:
74 node_dir_prefix = 'bge_ingestion'
75 elif 'NV-Embed' in embed_model_name:
76 node_dir_prefix = 'nv_ingestion'
77 elif 'colqwen' in embed_model_name:
78 node_dir_prefix = 'colqwen_ingestion'
79 elif 'openbmb' in embed_model_name:
80 node_dir_prefix = 'visrag_ingestion'
81 elif 'colpali' in embed_model_name:
82 node_dir_prefix = 'colpali_ingestion'
83 else:
84 raise ValueError('Please specify the node_dir_prefix')
85
86 if node_dir_prefix in ['colqwen_ingestion','visrag_ingestion','colpali_ingestion']:
87 self.vl_ret = True
88 else:
89 self.vl_ret = False
90
91 self.node_dir = os.path.join(self.dataset_dir, node_dir_prefix)
92 self.rag_dataset_path = os.path.join(self.dataset_dir, 'rag_dataset.json')
93 self.workers = 1
94 self.embed_model_name = embed_model_name
95 if 'vidore' in embed_model_name or 'openbmb' in embed_model_name:
96 if self.vl_ret:
97 self.vector_embed_model = VL_Embedding(model=embed_model_name, mode='image')
98 else:
99 self.vector_embed_model = VL_Embedding(model=embed_model_name, mode='text')
100 else:
101 self.vector_embed_model = HuggingFaceEmbedding(model_name=self.embed_model_name, embed_batch_size=10, max_length=512, trust_remote_code=True, device='cuda')
102 self.recall_num = 100
103 self.query_engine = self.load_query_engine()
104 self.output_dir = os.path.join(self.dataset_dir, 'search_output')
105 # os.makedirs(self.output_dir, exist_ok=True)
106
107 def online_search(self,query,node_list,topk=9):
108 nodes = [TextNode.from_dict(node['node']) for node in node_list]
109 vector_index = VectorStoreIndex(nodes, embed_model= self.vector_embed_model, show_progress=True, use_async=False, insert_batch_size=2048)
110 vector_retriever = vector_index.as_retriever(similarity_top_k=topk)
111 node_postprocessors = self.load_node_postprocessors()
112 query_engine = RetrieverQueryEngine(
113 retriever=vector_retriever,
114 node_postprocessors=node_postprocessors
115 )
116 query_bundle = QueryBundle(query_str=query)
117 recall_results = query_engine.retrieve(query_bundle)
118 return nodes2dict(recall_results)

Callers 2

__init__Method · 0.90
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected