| 161 | json.dump(results, f,indent=2,ensure_ascii=False) |
| 162 | |
| 163 | def arg_parse(): |
| 164 | import argparse |
| 165 | parser = argparse.ArgumentParser() |
| 166 | parser.add_argument("--dataset", type=str, default='ExampleDataset', help="The name of dataset") |
| 167 | parser.add_argument("--query_file", type=str, default='rag_dataset.json', help="The name of anno_file") |
| 168 | parser.add_argument("--experiment_type", type=str, default='retrieval_infer', help="The type of experiment") |
| 169 | parser.add_argument("--embed_model_name", type=str, default='BAAI/bge-m3', help="The name of embedding model") |
| 170 | parser.add_argument("--workers_num", type=int, default=1, help="The number of workers") |
| 171 | parser.add_argument("--topk", type=int, default=10, help="The number of topk") |
| 172 | parser.add_argument("--embed_model_name_vl", type=str, default=None, help="The name of embedding model for vl") |
| 173 | parser.add_argument("--embed_model_name_text", type=str, default=None, help="The name of embedding model for text") |
| 174 | parser.add_argument("--generate_vlm", type=str, default='qwen-vl-max', help="The name of VLM model") |
| 175 | args = parser.parse_args() |
| 176 | return args |
| 177 | |
| 178 | if __name__ == "__main__": |
| 179 | args = arg_parse() |