Single training step.
(forward_step_func, data_iterator,
model, optimizer, opt_param_scheduler, config)
| 305 | return model |
| 306 | |
| 307 | def train_step(forward_step_func, data_iterator, |
| 308 | model, optimizer, opt_param_scheduler, config): |
| 309 | """Single training step.""" |
| 310 | args = get_args() |
| 311 | timers = get_timers() |
| 312 | |
| 313 | # Set grad to zero. |
| 314 | for partition in model: |
| 315 | try: |
| 316 | partition.zero_grad_buffer() |
| 317 | except: |
| 318 | partition.zero_grad_buffer(zero_buffer=(not args.use_distributed_optimizer)) |
| 319 | optimizer.zero_grad() |
| 320 | |
| 321 | # Forward pass. |
| 322 | forward_backward_func = get_forward_backward_func() |
| 323 | losses_reduced = forward_backward_func( |
| 324 | forward_step_func=forward_step_func, |
| 325 | data_iterator=data_iterator, |
| 326 | model=model, |
| 327 | num_microbatches=get_num_microbatches(), |
| 328 | seq_length=args.seq_length, |
| 329 | micro_batch_size=args.micro_batch_size, |
| 330 | decoder_seq_length=args.decoder_seq_length, |
| 331 | forward_only=False) |
| 332 | |
| 333 | # Empty unused memory. |
| 334 | if args.empty_unused_memory_level >= 1: |
| 335 | torch.cuda.empty_cache() |
| 336 | |
| 337 | # Vision gradients. |
| 338 | if args.vision_pretraining and args.vision_pretraining_type == "dino": |
| 339 | unwrapped_model = unwrap_model(model[0]) |
| 340 | unwrapped_model.cancel_gradients_last_layer(args.curr_iteration) |
| 341 | |
| 342 | # Update parameters. |
| 343 | timers('optimizer', log_level=1).start(barrier=args.barrier_with_L1_time) |
| 344 | update_successful, grad_norm, num_zeros_in_grad = optimizer.step(args, timers) |
| 345 | timers('optimizer').stop() |
| 346 | |
| 347 | try: |
| 348 | if update_successful: |
| 349 | optimizer.gather_model_params(args, timers) |
| 350 | except: |
| 351 | pass |
| 352 | |
| 353 | # Vision momentum. |
| 354 | if args.vision_pretraining and args.vision_pretraining_type == "dino": |
| 355 | unwrapped_model = unwrap_model(model[0]) |
| 356 | unwrapped_model.update_momentum(args.curr_iteration) |
| 357 | |
| 358 | # Update learning rate. |
| 359 | if update_successful: |
| 360 | increment = get_num_microbatches() * \ |
| 361 | args.micro_batch_size * \ |
| 362 | args.data_parallel_size |
| 363 | opt_param_scheduler.step(increment=increment) |
| 364 | skipped_iter = 0 |