| 300 | |
| 301 | |
| 302 | def get_args(): |
| 303 | parser = argparse.ArgumentParser() |
| 304 | parser.add_argument("--test_file_path", type=str, default="data/topiocqa/rewrite/DS-R1-Distill-Qwen-7B.jsonl", nargs='+') |
| 305 | parser.add_argument("--passage_embeddings_dir_path", type=str, default="embedding/ance_topiocqa/numpy<2.0") |
| 306 | parser.add_argument("--pretrained_encoder_path", type=str, default="./ckpt/ance/ad-hoc-ance-msmarco") |
| 307 | parser.add_argument("--qrel_output_dir", type=str, default="data/topiocqa/qrel") |
| 308 | parser.add_argument("--output_trec_file", type=str, default="DS-R1-Distill-Qwen-7B_fast.trec", nargs='+') |
| 309 | parser.add_argument("--trec_gold_qrel_file_path", type=str, default="data/topiocqa/dev.trec") |
| 310 | parser.add_argument("--test_type", type=str, default="rewrite") |
| 311 | # parser.add_argument("--dataset", type=str, default="topiocqa") |
| 312 | parser.add_argument("--is_train", type=bool, default=False) |
| 313 | parser.add_argument("--top_k", type=int, default=100) |
| 314 | parser.add_argument("--n_gpu", type=int, default=3) |
| 315 | parser.add_argument("--rel_threshold", type=int, default=1) |
| 316 | parser.add_argument("--seed", type=int, default=42) |
| 317 | parser.add_argument("--per_gpu_test_batch_size", type=int, default=32) |
| 318 | parser.add_argument("--use_gpu", type=bool, default=True) |
| 319 | parser.add_argument("--max_query_length", type=int, default=64) |
| 320 | parser.add_argument("--max_doc_length", type=int, default=384) |
| 321 | parser.add_argument("--max_response_length", type=int, default=64) |
| 322 | parser.add_argument("--max_concat_length", type=int, default=512) |
| 323 | args = parser.parse_args() |
| 324 | |
| 325 | if args.use_gpu: |
| 326 | device = torch.device("cuda:0") |
| 327 | else: |
| 328 | device = torch.device("cpu") |
| 329 | args.device = device |
| 330 | |
| 331 | assert len(args.test_file_path) == len(args.output_trec_file) |
| 332 | |
| 333 | logger.info("---------------------The arguments are:---------------------") |
| 334 | logger.info(args) |
| 335 | return args |
| 336 | |
| 337 | if __name__ == '__main__': |
| 338 | main() |