| 59 | return valid_recall_result |
| 60 | |
| 61 | class 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) |