MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / load_checkpoint

Function load_checkpoint

SwissArmyTransformer/sat/training/model_io.py:263–374  ·  view source on GitHub ↗

Load a model checkpoint.

(model, args, load_path=None, prefix='', specific_iteration=None)

Source from the content-addressed store, hash-verified

261
262
263def 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]

Callers 8

mainFunction · 0.90
from_pretrained_baseMethod · 0.90
from_pretrained_baseMethod · 0.90
sampling_mainFunction · 0.90
__init__Method · 0.85
from_pretrainedMethod · 0.85
from_pretrainedMethod · 0.85
training_mainFunction · 0.85

Calls 8

print_rank0Function · 0.90
print_allFunction · 0.90
get_checkpoint_iterationFunction · 0.85
get_checkpoint_nameFunction · 0.85
appendMethod · 0.80
loadMethod · 0.45
load_state_dictMethod · 0.45
reinitMethod · 0.45

Tested by

no test coverage detected