(self,
dataset='ExampleDataset',
query_file='rag_dataset.json',
experiment_type = 'retrieval_infer',
generate_vlm='qwen-vl-max',
embed_model_name='BAAI/bge-m3',
embed_model_name_vl=None, # openbmb/VisRAG-Ret vidore/colqwen2-v1.0
embed_model_name_text=None, # nvidia/NV-Embed-v2 BAAI/bge-m3
workers_num = 1,
topk=10)
| 17 | |
| 18 | class MMRAG: |
| 19 | def __init__(self, |
| 20 | dataset='ExampleDataset', |
| 21 | query_file='rag_dataset.json', |
| 22 | experiment_type = 'retrieval_infer', |
| 23 | generate_vlm='qwen-vl-max', |
| 24 | embed_model_name='BAAI/bge-m3', |
| 25 | embed_model_name_vl=None, # openbmb/VisRAG-Ret vidore/colqwen2-v1.0 |
| 26 | embed_model_name_text=None, # nvidia/NV-Embed-v2 BAAI/bge-m3 |
| 27 | workers_num = 1, |
| 28 | topk=10): |
| 29 | self.experiment_type = experiment_type |
| 30 | self.workers_num = workers_num |
| 31 | self.top_k = topk |
| 32 | self.dataset = dataset |
| 33 | self.query_file = query_file |
| 34 | self.dataset_dir = os.path.join('./data', dataset) |
| 35 | self.img_dir = os.path.join(self.dataset_dir, "img") |
| 36 | self.results_dir = os.path.join(self.dataset_dir, "results") |
| 37 | os.makedirs(self.results_dir, exist_ok=True) |
| 38 | |
| 39 | self.vlm = LLM(model_name=generate_vlm) |
| 40 | self.evaluator = Evaluator() |
| 41 | |
| 42 | # load search_engine |
| 43 | if embed_model_name_vl is not None and embed_model_name_text is not None: |
| 44 | self.search_engine = HybridSearchEngine(self.dataset, |
| 45 | embed_model_name_vl=embed_model_name_vl, |
| 46 | embed_model_name_text=embed_model_name_text, |
| 47 | topk=topk) |
| 48 | else: |
| 49 | self.search_engine = SearchEngine(self.dataset, embed_model_name=embed_model_name) |
| 50 | |
| 51 | # retrieval only |
| 52 | if experiment_type == 'retrieval_infer': |
| 53 | self.eval_func = self.retrieval_infer |
| 54 | self.output_file_name = f'base_retrieval_{embed_model_name}.jsonl' |
| 55 | # hybrid retrieval |
| 56 | elif experiment_type == 'dynamic_hybird_retrieval_infer': |
| 57 | self.eval_func = self.retrieval_infer |
| 58 | self.search_engine.gmm = True |
| 59 | self.output_file_name = f'dynamic_hybird_retrieval_{embed_model_name_vl}_{embed_model_name_text}.jsonl' |
| 60 | # vidorag |
| 61 | elif experiment_type == 'vidorag': |
| 62 | self.agents = ViDoRAG_Agents(self.vlm) |
| 63 | self.eval_func = self.vidorag |
| 64 | self.search_engine.gmm = True |
| 65 | self.search_engine.gmm_candidate_length = True |
| 66 | self.output_file_name = f'vidorag_{generate_vlm}.jsonl' |
| 67 | |
| 68 | self.output_file_path = os.path.join(self.results_dir, self.output_file_name.replace("/","-")) |
| 69 | |
| 70 | def retrieval_infer(self,sample): |
| 71 | query = sample['query'] |
nothing calls this directly
no test coverage detected