Single training step.
(forward_step_func, data_iterator,
model, optimizer, lr_scheduler)
| 359 | |
| 360 | |
| 361 | def train_step(forward_step_func, data_iterator, |
| 362 | model, optimizer, lr_scheduler): |
| 363 | """Single training step.""" |
| 364 | args = get_args() |
| 365 | timers = get_timers() |
| 366 | |
| 367 | #see_memory_usage(f'before forward {model.global_steps}', force=True) |
| 368 | # Forward model for one step. |
| 369 | timers('forward').start() |
| 370 | loss, loss_reduced = forward_step_func(data_iterator, model) |
| 371 | timers('forward').stop() |
| 372 | |
| 373 | #see_memory_usage(f'before backward {model.global_steps}', force=True) |
| 374 | # Calculate gradients, reduce across processes, and clip. |
| 375 | timers('backward').start() |
| 376 | backward_step(optimizer, model, loss) |
| 377 | timers('backward').stop() |
| 378 | |
| 379 | #print('[Rank=',torch.distributed.get_rank(),'] backward complete!') |
| 380 | #see_memory_usage(f'before optimizer {model.global_steps}', force=True) |
| 381 | # Update parameters. |
| 382 | skipped_iter = 0 |
| 383 | timers('optimizer').start() |
| 384 | if args.deepspeed: |
| 385 | model.step() |
| 386 | else: |
| 387 | optimizer.step() |
| 388 | # Update learning rate. |
| 389 | if not (args.fp16 and optimizer.overflow): |
| 390 | lr_scheduler.step() |
| 391 | else: |
| 392 | skipped_iter = 1 |
| 393 | timers('optimizer').stop() |
| 394 | |
| 395 | return loss_reduced, skipped_iter |
| 396 | |
| 397 | |
| 398 | def training_log(loss_dict, total_loss_dict, learning_rate, iteration, |
no test coverage detected