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

Function add_checkpoint_callback_policy

codegeex/mindspore/finetune.py:72–99  ·  view source on GitHub ↗

r""" Add checkpoint policy to callback.

(args_param, callback, rank_id)

Source from the content-addressed store, hash-verified

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

Callers 1

run_trainFunction · 0.70

Calls 1

Tested by

no test coverage detected