r""" Add checkpoint policy to callback.
(args_param, callback, rank_id)
| 70 | |
| 71 | |
| 72 | def 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 | |
| 102 | def set_parallel_context(args_opt): |
no test coverage detected