MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / add_checkpoint_callback_policy

Function add_checkpoint_callback_policy

codegeex/mindspore/train.py:74–101  ·  view source on GitHub ↗

r""" Add checkpoint policy to callback.

(args_param, callback, rank_id)

Source from the content-addressed store, hash-verified

72
73
74def add_checkpoint_callback_policy(args_param, callback, rank_id):
75 r"""
76 Add checkpoint policy to callback.
77 """
78 if args_param.save_checkpoint:
79 # checkpoint store epoch_num and step_num info
80 ckpt_append_info = [{"epoch_num": args_param.has_trained_epoches, "step_num": args_param.has_trained_steps}]
81 ckpt_config = CheckpointConfig(
82 save_checkpoint_steps=args_param.save_checkpoint_steps,
83 keep_checkpoint_max=args_param.keep_checkpoint_max,
84 integrated_save=False,
85 append_info=ckpt_append_info,
86 )
87
88 # save checkpoint into rank directory
89 ckpoint_cb = ModelCheckpoint(prefix=args_param.ckpt_name_prefix + str(rank_id),
90 directory=os.path.join(args_param.save_checkpoint_path, f"rank_{rank_id}"),
91 config=ckpt_config)
92
93 callback.append(ckpoint_cb)
94
95 saveckpt_cb = SaveCheckpointCallback(cache_dir=args_param.save_checkpoint_path,
96 bucket=args_param.save_checkpoint_obs_path,
97 local_rank=rank_id,
98 has_trained_epoch=args_param.has_trained_epoches,
99 has_trained_step=args_param.has_trained_steps,
100 syn_times=args_param.save_checkpoint_steps)
101 callback.append(saveckpt_cb)
102
103
104def set_parallel_context(args_opt):

Callers 1

run_trainFunction · 0.70

Calls 1

Tested by

no test coverage detected