Load a model checkpoint.
(model, args, load_path=None, prefix='', specific_iteration=None)
| 261 | |
| 262 | |
| 263 | def load_checkpoint(model, args, load_path=None, prefix='', specific_iteration=None): |
| 264 | """Load a model checkpoint.""" |
| 265 | if load_path is None: |
| 266 | load_path = args.load |
| 267 | |
| 268 | if load_path.endswith('.pt'): |
| 269 | checkpoint_name = load_path |
| 270 | iter_str = checkpoint_name.split('/')[-2] |
| 271 | if iter_str.isdigit(): |
| 272 | iteration = int(iter_str) |
| 273 | else: |
| 274 | iteration = int(1e10) |
| 275 | else: |
| 276 | # If model-only mode, set necessary args. |
| 277 | if not hasattr(args, 'mode'): |
| 278 | from copy import deepcopy |
| 279 | args = deepcopy(args) |
| 280 | args.mode = 'inference' |
| 281 | |
| 282 | iteration, release, success = get_checkpoint_iteration(load_path) |
| 283 | if specific_iteration is not None: |
| 284 | assert type(specific_iteration) == int and specific_iteration > 0 |
| 285 | print_rank0('Overriding checkpoint iteration to {}'.format(specific_iteration)) |
| 286 | iteration = specific_iteration |
| 287 | |
| 288 | if not success: |
| 289 | return 0 |
| 290 | |
| 291 | checkpoint_name = get_checkpoint_name(load_path, iteration, release, use_ema=args.use_ema) |
| 292 | if mpu.get_data_parallel_rank() == 0: |
| 293 | print_all('global rank {} is loading checkpoint {}'.format( |
| 294 | torch.distributed.get_rank(), checkpoint_name)) |
| 295 | |
| 296 | # load state_dict into CPU |
| 297 | # sd = torch.load(checkpoint_name, map_location='cpu') |
| 298 | sd = torch.load(checkpoint_name, map_location='cpu', weights_only=False) |
| 299 | |
| 300 | # if given `prefix`, load a speficic prefix in the checkpoint, e.g. encoder |
| 301 | new_sd = {'module':{}} |
| 302 | for k in sd: |
| 303 | if k != 'module': |
| 304 | new_sd[k] = sd[k] |
| 305 | for k in sd['module']: |
| 306 | if k.startswith(prefix): |
| 307 | new_sd['module'][k[len(prefix):]] = sd['module'][k] |
| 308 | sd = new_sd |
| 309 | |
| 310 | if hasattr(model, 'module'): |
| 311 | module = model.module |
| 312 | else: # inference without deepspeed |
| 313 | module = model |
| 314 | |
| 315 | # * Remove pretrained pos embedding (the pos embs are not learned, but sin-cos) |
| 316 | REMOVE_POS = True |
| 317 | if REMOVE_POS: |
| 318 | pos_emb_str = 'model.diffusion_model.mixins.pos_embed.pos_embedding' |
| 319 | if pos_emb_str in sd['module']: |
| 320 | del sd['module'][pos_emb_str] |
no test coverage detected