Single training step.
(data_iterator, model, optimizer, lr_scheduler, args, timers, forward_step_func, mems=None,
single_step=False)
| 322 | |
| 323 | |
| 324 | def train_step(data_iterator, model, optimizer, lr_scheduler, args, timers, forward_step_func, mems=None, |
| 325 | single_step=False): |
| 326 | """Single training step.""" |
| 327 | lm_loss_total, count = 0.0, 0 |
| 328 | mems = [] if mems is None else mems |
| 329 | if not args.deepspeed: |
| 330 | optimizer.zero_grad() |
| 331 | while True: |
| 332 | skipped_iter, complete = 0, False |
| 333 | # Forward model for one step. |
| 334 | timers('forward').start() |
| 335 | lm_loss, mems, _ = forward_step_func(data_iterator, model, args, timers, mems) |
| 336 | timers('forward').stop() |
| 337 | # print_rank_0("Forward step") |
| 338 | if not args.deepspeed: |
| 339 | lm_loss /= args.gradient_accumulation_steps |
| 340 | |
| 341 | reduced_loss = lm_loss.detach().clone().view(1) |
| 342 | torch.distributed.all_reduce(reduced_loss.data, group=mpu.get_data_parallel_group()) |
| 343 | reduced_loss.data = reduced_loss.data / (args.world_size / args.model_parallel_size) |
| 344 | |
| 345 | if not DynamicLossScaler._has_inf_or_nan(reduced_loss): |
| 346 | lm_loss_total += reduced_loss |
| 347 | count += 1 |
| 348 | |
| 349 | # Calculate gradients, reduce across processes, and clip. |
| 350 | timers('backward').start() |
| 351 | backward_step(optimizer, model, lm_loss, args, timers) |
| 352 | timers('backward').stop() |
| 353 | # print_rank_0("Backward step") |
| 354 | # Update parameters. |
| 355 | timers('optimizer').start() |
| 356 | if args.deepspeed: |
| 357 | if model.is_gradient_accumulation_boundary(): |
| 358 | model.step() |
| 359 | complete = True |
| 360 | if not (args.fp16 and optimizer.overflow): |
| 361 | lr_scheduler.step() |
| 362 | else: |
| 363 | skipped_iter = 1 |
| 364 | else: |
| 365 | model.step() |
| 366 | else: |
| 367 | if count == args.gradient_accumulation_steps: |
| 368 | optimizer.step() |
| 369 | complete = True |
| 370 | # Update learning rate. |
| 371 | if not (args.fp16 and optimizer.overflow): |
| 372 | lr_scheduler.step() |
| 373 | else: |
| 374 | skipped_iter = 1 |
| 375 | # print_rank_0("Optimizer step") |
| 376 | timers('optimizer').stop() |
| 377 | if complete: |
| 378 | break |
| 379 | else: |
| 380 | print_rank_0("Found NaN loss, skip backward") |
| 381 | del lm_loss, reduced_loss |
no test coverage detected