| 118 | |
| 119 | |
| 120 | def parse_args(): |
| 121 | parser = argparse.ArgumentParser() |
| 122 | parser.add_argument("--model", type=str, default=0) # model path |
| 123 | parser.add_argument("--data_file", type=str, default='data/math_eval/MATH_test.jsonl') # data path |
| 124 | parser.add_argument("--start", type=int, default=0) # start index |
| 125 | parser.add_argument("--end", type=int, default=MAX_INT) # end index |
| 126 | parser.add_argument("--batch_size", type=int, default=32) # batch_size |
| 127 | parser.add_argument("--tensor_parallel_size", type=int, default=1) # tensor_parallel_size |
| 128 | parser.add_argument("--run_dir", type=str) # run_dir |
| 129 | parser.add_argument("--no_wandb", action="store_true") # no_wandb |
| 130 | |
| 131 | args = parser.parse_args() |
| 132 | |
| 133 | if args.run_dir and not args.no_wandb: |
| 134 | try: |
| 135 | with open(os.path.join(args.run_dir, "wandb_run_id.txt"), "r") as f: |
| 136 | wandb_run_id = f.read().strip() |
| 137 | wandb.init( |
| 138 | id=wandb_run_id, |
| 139 | project="project-name", |
| 140 | resume="must" |
| 141 | ) |
| 142 | except FileNotFoundError: |
| 143 | print("WandB run ID file not found, starting new run") |
| 144 | wandb.init(project="project-name") |
| 145 | |
| 146 | return args |
| 147 | |
| 148 | if __name__ == "__main__": |
| 149 | args = parse_args() |