MCPcopy Create free account
hub / github.com/SwinTransformer/Transformer-SSL / load_checkpoint

Function load_checkpoint

utils.py:37–61  ·  view source on GitHub ↗
(config, model, optimizer, lr_scheduler, logger)

Source from the content-addressed store, hash-verified

35
36
37def load_checkpoint(config, model, optimizer, lr_scheduler, logger):
38 logger.info(f"==============> Resuming form {config.MODEL.RESUME}....................")
39 if config.MODEL.RESUME.startswith('https'):
40 checkpoint = torch.hub.load_state_dict_from_url(
41 config.MODEL.RESUME, map_location='cpu', check_hash=True)
42 else:
43 checkpoint = torch.load(config.MODEL.RESUME, map_location='cpu')
44 msg = model.load_state_dict(checkpoint['model'], strict=False)
45 logger.info(msg)
46 max_accuracy = 0.0
47 if not config.EVAL_MODE and 'optimizer' in checkpoint and 'lr_scheduler' in checkpoint and 'epoch' in checkpoint:
48 optimizer.load_state_dict(checkpoint['optimizer'])
49 lr_scheduler.load_state_dict(checkpoint['lr_scheduler'])
50 config.defrost()
51 config.TRAIN.START_EPOCH = checkpoint['epoch'] + 1
52 config.freeze()
53 if 'amp' in checkpoint and config.AMP_OPT_LEVEL != "O0" and checkpoint['config'].AMP_OPT_LEVEL != "O0":
54 amp.load_state_dict(checkpoint['amp'])
55 logger.info(f"=> loaded successfully '{config.MODEL.RESUME}' (epoch {checkpoint['epoch']})")
56 if 'max_accuracy' in checkpoint:
57 max_accuracy = checkpoint['max_accuracy']
58
59 del checkpoint
60 torch.cuda.empty_cache()
61 return max_accuracy
62
63
64def save_checkpoint(config, epoch, model, max_accuracy, optimizer, lr_scheduler, logger):

Callers 3

mainFunction · 0.90
mainFunction · 0.90
mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected