Loading checkpoint logic for training.
(cfg, model, optimizer, loss_scaler)
| 523 | |
| 524 | |
| 525 | def load_train_checkpoint(cfg, model, optimizer, loss_scaler): |
| 526 | """ |
| 527 | Loading checkpoint logic for training. |
| 528 | """ |
| 529 | if cfg.TRAIN.AUTO_RESUME and has_checkpoint(cfg.OUTPUT_DIR): |
| 530 | last_checkpoint = get_last_checkpoint(cfg.OUTPUT_DIR) |
| 531 | logger.info("Load from last checkpoint, {}.".format(last_checkpoint)) |
| 532 | checkpoint_epoch = load_checkpoint( |
| 533 | last_checkpoint, model, loss_scaler, cfg.NUM_GPUS > 1, optimizer |
| 534 | ) |
| 535 | start_epoch = checkpoint_epoch + 1 |
| 536 | elif cfg.TRAIN.CHECKPOINT_FILE_PATH != "": |
| 537 | logger.info("Load from given checkpoint file.") |
| 538 | checkpoint_epoch = load_checkpoint( |
| 539 | cfg.TRAIN.CHECKPOINT_FILE_PATH, |
| 540 | model, |
| 541 | loss_scaler, |
| 542 | cfg.NUM_GPUS > 1, |
| 543 | optimizer, |
| 544 | inflation=cfg.TRAIN.CHECKPOINT_INFLATE, |
| 545 | convert_from_caffe2=cfg.TRAIN.CHECKPOINT_TYPE == "caffe2", |
| 546 | epoch_reset=cfg.TRAIN.CHECKPOINT_EPOCH_RESET, |
| 547 | clear_name_pattern=cfg.TRAIN.CHECKPOINT_CLEAR_NAME_PATTERN, |
| 548 | ) |
| 549 | start_epoch = checkpoint_epoch + 1 |
| 550 | else: |
| 551 | start_epoch = 0 |
| 552 | |
| 553 | return start_epoch |
nothing calls this directly
no test coverage detected