MCPcopy Create free account
hub / github.com/modelscope/modelscope / save_checkpoint

Function save_checkpoint

modelscope/models/nlp/mglm/utils.py:217–283  ·  view source on GitHub ↗

Save a model checkpoint.

(iteration,
                    model,
                    optimizer,
                    lr_scheduler,
                    args,
                    tag=None,
                    barrier=True,
                    only_changed_parameters=False,
                    no_deepspeed=False,
                    no_save_optim=False)

Source from the content-addressed store, hash-verified

215
216
217def save_checkpoint(iteration,
218 model,
219 optimizer,
220 lr_scheduler,
221 args,
222 tag=None,
223 barrier=True,
224 only_changed_parameters=False,
225 no_deepspeed=False,
226 no_save_optim=False):
227 """Save a model checkpoint."""
228 if tag is None:
229 tag = str(iteration)
230 if args.deepspeed and not no_deepspeed:
231 save_ds_checkpoint(iteration, model, lr_scheduler, args, tag=tag)
232 else:
233 # Only rank zer0 of the data parallel writes to the disk.
234
235 if mpu.get_data_parallel_rank() == 0:
236 checkpoint_name = get_checkpoint_name(args.save, tag)
237 print(
238 'global rank {} is saving checkpoint at iteration {:7d} to {}'.
239 format(torch.distributed.get_rank(), iteration,
240 checkpoint_name))
241 sd = {'iteration': iteration}
242 if args.deepspeed:
243 model = model.module
244 state_dict = model.state_dict()
245 if only_changed_parameters:
246 requires_grad_dict = {}
247 for name, parameter in model.named_parameters():
248 requires_grad_dict[name] = parameter.requires_grad
249 state_dict = {
250 key: value
251 for key, value in state_dict.items()
252 if requires_grad_dict[key]
253 }
254 sd['module'] = state_dict
255
256 # Optimizer stuff.
257 if not args.no_save_optim and not no_save_optim:
258 if optimizer is not None:
259 sd['optimizer'] = optimizer.state_dict()
260 if lr_scheduler is not None:
261 sd['lr_scheduler'] = lr_scheduler.state_dict()
262
263 # rng states.
264 if not args.no_save_rng:
265 sd['random_rng_state'] = random.getstate()
266 sd['np_rng_state'] = np.random.get_state()
267 sd['torch_rng_state'] = torch.get_rng_state()
268 sd['cuda_rng_state'] = torch.cuda.get_rng_state()
269 sd['rng_tracker_states'] = mpu.get_cuda_rng_tracker(
270 ).get_states()
271
272 ensure_directory_exists(checkpoint_name)
273 torch.save(sd, checkpoint_name)
274 print(' successfully saved {}'.format(checkpoint_name))

Callers

nothing calls this directly

Calls 10

save_ds_checkpointFunction · 0.85
get_checkpoint_nameFunction · 0.85
printFunction · 0.85
ensure_directory_existsFunction · 0.85
state_dictMethod · 0.45
named_parametersMethod · 0.45
itemsMethod · 0.45
saveMethod · 0.45
writeMethod · 0.45

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…