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