(config, ckpt, verbose=False)
| 117 | |
| 118 | |
| 119 | def load_model_from_config(config, ckpt, verbose=False): |
| 120 | print(f"Loading model from {ckpt}") |
| 121 | pl_sd = torch.load(ckpt, map_location="cpu") |
| 122 | if "global_step" in pl_sd: |
| 123 | print(f"Global Step: {pl_sd['global_step']}") |
| 124 | sd = pl_sd["state_dict"] |
| 125 | model = instantiate_from_config(config.model) |
| 126 | m, u = model.load_state_dict(sd, strict=False) |
| 127 | if len(m) > 0 and verbose: |
| 128 | print("missing keys:") |
| 129 | print(m) |
| 130 | if len(u) > 0 and verbose: |
| 131 | print("unexpected keys:") |
| 132 | print(u) |
| 133 | |
| 134 | model.cuda() |
| 135 | model.eval() |
| 136 | return model |
| 137 | |
| 138 | |
| 139 | def load_img(path): |
no test coverage detected