Train the model function.
(forward_step_func, model, optimizer, lr_scheduler,
train_data_iterator, valid_data_iterator)
| 494 | |
| 495 | |
| 496 | def train(forward_step_func, model, optimizer, lr_scheduler, |
| 497 | train_data_iterator, valid_data_iterator): |
| 498 | """Train the model function.""" |
| 499 | args = get_args() |
| 500 | timers = get_timers() |
| 501 | |
| 502 | # Turn on training mode which enables dropout. |
| 503 | model.train() |
| 504 | |
| 505 | # Tracking loss. |
| 506 | total_loss_dict = {} |
| 507 | |
| 508 | # Iterations. |
| 509 | iteration = args.iteration |
| 510 | |
| 511 | timers('interval time').start() |
| 512 | report_memory_flag = True |
| 513 | data_parallel_size = mpu.get_data_parallel_world_size() |
| 514 | global_batch_size = args.batch_size * data_parallel_size |
| 515 | while iteration < args.train_iters and \ |
| 516 | (args.train_tokens is None or args.tokens < args.train_tokens): |
| 517 | loss_dict, skipped_iter = train_step(forward_step_func, |
| 518 | train_data_iterator, |
| 519 | model, |
| 520 | optimizer, |
| 521 | lr_scheduler) |
| 522 | iteration += 1 |
| 523 | if args.curriculum_learning: |
| 524 | args.tokens += global_batch_size * args.curriculum_seqlen |
| 525 | else: |
| 526 | args.tokens += global_batch_size * args.seq_length |
| 527 | |
| 528 | # Logging. |
| 529 | loss_scale = None |
| 530 | if args.fp16: |
| 531 | loss_scale = optimizer.cur_scale if args.deepspeed else optimizer.loss_scale |
| 532 | report_memory_flag = training_log(loss_dict, total_loss_dict, |
| 533 | optimizer.param_groups[0]['lr'], |
| 534 | iteration, loss_scale, |
| 535 | report_memory_flag, skipped_iter, |
| 536 | model=model) |
| 537 | |
| 538 | # Autoresume |
| 539 | if args.adlr_autoresume and \ |
| 540 | (iteration % args.adlr_autoresume_interval == 0): |
| 541 | check_adlr_autoresume_termination(iteration, model, optimizer, |
| 542 | lr_scheduler) |
| 543 | |
| 544 | # Checkpointing |
| 545 | if args.save and args.save_interval and \ |
| 546 | iteration % args.save_interval == 0: |
| 547 | save_checkpoint(iteration, model, optimizer, lr_scheduler) |
| 548 | |
| 549 | # Evaluation |
| 550 | # XXX temporarily disabled for ZeRO-3 |
| 551 | """ |
| 552 | if args.eval_interval and iteration % args.eval_interval == 0 and \ |
| 553 | args.do_valid: |
no test coverage detected