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

Method __init__

search_engine.py:62–105  ·  view source on GitHub ↗
(self,dataset, node_dir_prefix=None,embed_model_name='BAAI/bge-m3')

Source from the content-addressed store, hash-verified

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]

Callers

nothing calls this directly

Calls 2

load_query_engineMethod · 0.95
VL_EmbeddingClass · 0.90

Tested by

no test coverage detected