MCPcopy Create free account
hub / github.com/Topdu/OpenOCR / load_ckpt

Function load_ckpt

tools/utils/ckpt.py:49–77  ·  view source on GitHub ↗

Resume from saved checkpoints :param checkpoint_path: Checkpoint path to be resumed

(model, cfg, optimizer=None, lr_scheduler=None, logger=None)

Source from the content-addressed store, hash-verified

47
48
49def 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
80def load_pretrained_params(model, pretrained_model, logger):

Callers 4

_init_torch_modelMethod · 0.90
mainFunction · 0.90
_init_torch_modelMethod · 0.90
__init__Method · 0.90

Calls 4

get_loggerFunction · 0.90
load_pretrained_paramsFunction · 0.85
getMethod · 0.80
load_state_dictMethod · 0.80

Tested by

no test coverage detected