(args)
| 108 | |
| 109 | |
| 110 | def train(args): |
| 111 | |
| 112 | accelerator = Accelerator( |
| 113 | gradient_accumulation_steps=args.accum_iter, |
| 114 | mixed_precision="bf16", |
| 115 | kwargs_handlers=[ |
| 116 | DistributedDataParallelKwargs(find_unused_parameters=True), |
| 117 | InitProcessGroupKwargs(timeout=timedelta(seconds=6000)), |
| 118 | ], |
| 119 | ) |
| 120 | device = accelerator.device |
| 121 | |
| 122 | setup_for_distributed(accelerator) |
| 123 | |
| 124 | printer.info("output_dir: " + args.output_dir) |
| 125 | if args.output_dir: |
| 126 | Path(args.output_dir).mkdir(parents=True, exist_ok=True) |
| 127 | |
| 128 | if accelerator.is_main_process: |
| 129 | dst_dir = save_current_code(outdir=args.output_dir) |
| 130 | printer.info(f"Saving current code to {dst_dir}") |
| 131 | |
| 132 | # auto resume |
| 133 | if not args.resume: |
| 134 | last_ckpt_fname = os.path.join(args.output_dir, f"checkpoint-last.pth") |
| 135 | args.resume = last_ckpt_fname if os.path.isfile(last_ckpt_fname) else None |
| 136 | |
| 137 | printer.info("job dir: {}".format(os.path.dirname(os.path.realpath(__file__)))) |
| 138 | |
| 139 | # fix the seed |
| 140 | seed = args.seed + accelerator.state.process_index |
| 141 | printer.info( |
| 142 | f"Setting seed to {seed} for process {accelerator.state.process_index}" |
| 143 | ) |
| 144 | torch.manual_seed(seed) |
| 145 | np.random.seed(seed) |
| 146 | random.seed(seed) |
| 147 | cudnn.benchmark = args.benchmark |
| 148 | |
| 149 | # training dataset and loader |
| 150 | printer.info("Building train dataset %s", args.train_dataset) |
| 151 | # dataset and loader |
| 152 | data_loader_train = build_dataset( |
| 153 | args.train_dataset, |
| 154 | args.batch_size, |
| 155 | args.num_workers, |
| 156 | accelerator=accelerator, |
| 157 | test=False, |
| 158 | fixed_length=args.fixed_length |
| 159 | ) |
| 160 | printer.info("Building test dataset %s", args.test_dataset) |
| 161 | data_loader_test = { |
| 162 | dataset.split("(")[0]: build_dataset( |
| 163 | dataset, |
| 164 | args.batch_size, |
| 165 | args.num_workers, |
| 166 | accelerator=accelerator, |
| 167 | test=True, |
no test coverage detected