(args)
| 219 | |
| 220 | |
| 221 | def update_train_iters(args): |
| 222 | |
| 223 | # For iteration-based training, we don't need to do anything |
| 224 | if args.train_iters: |
| 225 | return |
| 226 | |
| 227 | # Constant batch size with sample-based training. |
| 228 | if args.rampup_batch_size is None: |
| 229 | args.train_iters = args.train_samples // args.global_batch_size |
| 230 | |
| 231 | else: |
| 232 | # Sample based training with rampup batch size. |
| 233 | iterations = 0 |
| 234 | consumed_samples = 0 |
| 235 | # Rampup phase. |
| 236 | while consumed_samples <= int(args.rampup_batch_size[2]): |
| 237 | update_num_microbatches(consumed_samples, consistency_check=False) |
| 238 | consumed_samples += get_current_global_batch_size() |
| 239 | iterations += 1 |
| 240 | # Reset |
| 241 | update_num_microbatches(0, consistency_check=False) |
| 242 | # Constant phase |
| 243 | # Note that we throw away any partial last batch. |
| 244 | iterations += (args.train_samples - consumed_samples) // args.global_batch_size |
| 245 | args.train_iters = iterations |
| 246 | |
| 247 | print_rank_0("setting training iterations to {}".format(args.train_iters)) |
| 248 | |
| 249 | |
| 250 | def get_model(model_provider_func): |
no test coverage detected