Resume from saved checkpoints :param checkpoint_path: Checkpoint path to be resumed
(model, cfg, optimizer=None, lr_scheduler=None, logger=None)
| 47 | |
| 48 | |
| 49 | def load_ckpt(model, cfg, optimizer=None, lr_scheduler=None, logger=None): |
| 50 | """ |
| 51 | Resume from saved checkpoints |
| 52 | :param checkpoint_path: Checkpoint path to be resumed |
| 53 | """ |
| 54 | if logger is None: |
| 55 | logger = get_logger() |
| 56 | checkpoints = cfg["Global"].get("checkpoints") |
| 57 | pretrained_model = cfg["Global"].get("pretrained_model") |
| 58 | |
| 59 | status = {} |
| 60 | if checkpoints and os.path.exists(checkpoints): |
| 61 | checkpoint = torch.load(checkpoints, map_location=torch.device("cpu")) |
| 62 | model.load_state_dict(checkpoint["state_dict"], strict=True) |
| 63 | if optimizer is not None: |
| 64 | optimizer.load_state_dict(checkpoint["optimizer"]) |
| 65 | if lr_scheduler is not None: |
| 66 | lr_scheduler.load_state_dict(checkpoint["scheduler"]) |
| 67 | logger.info(f"resume from checkpoint {checkpoints} (epoch {checkpoint['epoch']})") |
| 68 | |
| 69 | status["global_step"] = checkpoint["global_step"] |
| 70 | status["epoch"] = checkpoint["epoch"] + 1 |
| 71 | status["metrics"] = checkpoint["metrics"] |
| 72 | elif pretrained_model and os.path.exists(pretrained_model): |
| 73 | load_pretrained_params(model, pretrained_model, logger) |
| 74 | logger.info(f"finetune from checkpoint {pretrained_model}") |
| 75 | else: |
| 76 | logger.info("train from scratch") |
| 77 | return status |
| 78 | |
| 79 | |
| 80 | def load_pretrained_params(model, pretrained_model, logger): |
no test coverage detected