ReflectionEraser/ReflectionEraserApp
0
1from .net_utils import *
2from .util import *
3
4
5def load_checkpoint(model, ckpt_path):
6 checkpoint = torch.load(ckpt_path)
7 if 'model' in checkpoint:
8 checkpoint = checkpoint['model']
9 if 'state_dict' in checkpoint:
10 checkpoint = checkpoint['state_dict']
11 ckpt = {}
12 for k, v in checkpoint.items():
13 if k.startswith('module.'):
14 ckpt[k[7:]] = v
15 else:
16 ckpt[k] = v
17 model.load_state_dict(ckpt,strict=False) #strict=False by nami
18 