Save PyTorch training checkpoints with near-zero overhead and recover from failures without recalculating lost iterations using EasyCkpt with Megatron or DeepSpeed.
Foundation model training on large GPU clusters faces frequent interruptions from hardware failures, system issues, and network errors. Traditional checkpointing takes minutes to complete for models with tens to hundreds of billions of parameters, pausing training during each save. When training is interrupted, all iterations since the last checkpoint must be recalculated.
EasyCkpt uses asynchronous hierarchical checkpointing to save model state without pausing training. It recovers from failures using data replicas stored across GPUs, eliminating recalculation.
Architecture
EasyCkpt uses three techniques to save model state without pausing training:
Asynchronous hierarchical checkpointing copies state to memory first, then writes to disk in the background
Checkpoint-computation overlap continues training while checkpoint operations run asynchronously
Network-aware asynchronous checkpointing uses available network interfaces despite partial failures
The design leverages three failure characteristics:
Failures affect specific workers. A failure typically originates from one or two machines, impacting only a few workers in the distributed training job, not the entire job.
Failures affect specific components. GPU errors do not affect CPU and memory operations. Errors occur only on specific network interfaces, so nodes can still communicate despite partial failures.
Failures affect specific model replicas. Foundation model training uses 3D parallelism (data, tensor, pipeline) or Zero Redundancy Optimizer (ZeRO), maintaining multiple data replicas across GPUs. When one GPU fails, training recovers from replicas on other machines.
Prerequisites
Python 3.6, 3.8, or 3.10
PyTorch with Megatron or DeepSpeed
AIMaster SDK (see Install AIMaster SDK)
Install AIMaster SDK
Run pip install for your Python version:
# Python 3.6
pip install -U http://odps-release.cn-hangzhou.oss.aliyun-inc.com/aimaster/pai_aimaster-1.2.1-cp36-cp36m-linux_x86_64.whl
# Python 3.8
pip install -U http://odps-release.cn-hangzhou.oss.aliyun-inc.com/aimaster/pai_aimaster-1.2.1-cp38-cp38-linux_x86_64.whl
# Python 3.10
pip install -U http://odps-release.cn-hangzhou.oss.aliyun-inc.com/aimaster/pai_aimaster-1.2.1-cp310-cp310-linux_x86_64.whlIntegrate with Megatron
Integrate EasyCkpt by modifying your training loop file (e.g., training.py) and importing a hook in your model file (e.g., pretrain_gpt.py).
Modify training loop
Replace the native Megatron load_checkpoint import with EasyCkpt's version, then add initialize_easyckpt and save_checkpoint_if_needed calls to the training loop.
from megatron.core.utils import get_model_config
from megatron import print_rank_0
from megatron import print_rank_last
# Replace the native Megatron load_checkpoint with EasyCkpt's version
# from megatron.checkpointing import load_checkpoint
from megatron.checkpointing import save_checkpoint
from megatron.model import Float16Module
from megatron.model import GPTModel
from megatron.utils import report_memory
from megatron.model.vision.knn_monitor import compute_feature_bank
# Import EasyCkpt functions
from aimaster.python.torch.easyckpt.megatron import (load_checkpoint,
initialize_easyckpt,
save_checkpoint_if_needed)
def print_datetime(string):
"""Note that this call will sync across all ranks."""
timers('interval-time', log_level=0).start(barrier=True)
print_datetime('before the start of training step')
report_memory_flag = True
# Initialize EasyCkpt before the training loop.
# save_mem_interval=1: copy to memory every iteration
# save_storage_interval=5: write to storage every 5 iterations
# max_ckpt_num=5: keep up to 5 checkpoints on disk
initialize_easyckpt(save_mem_interval=1,
save_storage_interval=5,
max_ckpt_num=5,
log_file_path='./test.log')
while iteration < args.train_iters:
if args.profile and \
iteration == args.profile_step_start and \
args.micro_batch_size * \
get_num_microbatches()
# Trigger an in-memory checkpoint if conditions are met
save_checkpoint_if_needed(iteration, model, optimizer, opt_param_scheduler)
# Logging.
loss_scale = optimizer.get_loss_scale().item()
params_norm = NoneImport EasyCkpt hook
In your model training file (e.g., pretrain_gpt.py), import the EasyCkpt Megatron hook:
from megatron.utils import average_losses_across_data_parallel_group
from megatron.arguments import core_transformer_config_from_args
# Import the EasyCkpt hook to enable async checkpointing
import aimaster.python.torch.easyckpt.megatron.hook
def model_provider(pre_process=True, post_process=True):
"""Build the model."""API reference
EasyCkpt provides three functions for Megatron integration:
initialize_easyckpt
Initializes EasyCkpt. Call this before the training loop.
initialize_easyckpt(save_mem_interval, save_storage_interval, max_ckpt_num, log_file_path=None)Parameter | Type | Default | Description |
| int |
| Iteration frequency for copying model state to memory |
| int |
| Iteration frequency for writing checkpoints to persistent storage asynchronously |
| int |
| Maximum number of checkpoints to retain on disk. Older checkpoints are deleted automatically |
| str |
| Path to a log file for checkpoint operation logs |
load_checkpoint
Extends the native Megatron load_checkpoint() function with a concat parameter for distributed optimizer support.
load_checkpoint(model, optimizer, opt_param_scheduler, load_arg='load', strict=True, concat=False)Parameter | Type | Default | Description |
| Model | Required | The Megatron model instance |
| Optimizer | Required | The optimizer instance |
| Scheduler | Required | The optimizer parameter scheduler |
| str |
| Argument name for the load directory |
| bool |
| Whether to enforce strict state dict matching |
| bool |
| Set to |
For Megatron 2305 or 2306 with distributed optimizer enabled, set concat=True in load_checkpoint() when changing instance count during loading or merging distributed optimizer parameters.
For Megatron 2304, replace the native load_checkpoint with EasyCkpt's version.
save_checkpoint_if_needed
Triggers an in-memory checkpoint. Call this inside the training loop.
save_checkpoint_if_needed(iteration, model, optimizer, opt_param_scheduler)Parameter | Type | Default | Description |
| int | Required | Current training iteration number |
| Model | Required | The Megatron model instance |
| Optimizer | Required | The optimizer instance |
| Scheduler | Required | The optimizer parameter scheduler |
All parameters are existing Megatron variables in the training loop.
Integrate with DeepSpeed
For DeepSpeed tasks using Transformers Trainer, EasyCkpt integrates through a TrainerWrapper class.
Configure checkpoint parameters
EasyCkpt reuses standard Transformers checkpoint parameters. Add them to your launch script:
--max_steps=10 \
--block_size=2048 \
--num_train_examples=100000 \
--gradient_checkpointing=false \
--save_strategy="steps" \
--save_steps="2" \
--save_total_limit="2"This example saves a checkpoint every 2 steps and retains a maximum of 2 checkpoints.
Wrap the Trainer
Import TrainerWrapper from EasyCkpt's Transformers module. Wrap your Trainer instance and enable checkpoint resumption:
import datasets
import transformers
# Import the EasyCkpt wrapper for Transformers Trainer
from aimaster.python.torch.easyckpt.transformers import TrainerWrapper
logger = logging.getLogger(__name__)
tokenizer=tokenizer,
data_collator=transformers.default_data_collator,
)
# Wrap the Trainer with EasyCkpt to enable async checkpointing
trainer = TrainerWrapper(trainer)
# Resume from the latest checkpoint automatically
trainer.train(resume_from_checkpoint=True)
if __name__ == "__main__":
main()API reference
TrainerWrapper
Wraps a Transformers Trainer instance with asynchronous checkpointing capabilities.
from aimaster.python.torch.easyckpt.transformers import TrainerWrapper
trainer = TrainerWrapper(trainer)
trainer.train(resume_from_checkpoint=True)Checkpoint parameters
These parameters follow standard Transformers TrainingArguments definitions:
Parameter | Type | Valid values | Description |
| str |
| When to save checkpoints. |
| int | Positive integer | Step interval for saving checkpoints. Only applies when |
| int | Positive integer | Maximum number of retained checkpoints. Older checkpoints are deleted automatically when this limit is reached |
When save_total_limit is enabled, outdated checkpoints are deleted automatically. Back up any data stored in those folders before deletion. For details, see the official Transformers documentation.
Data security
EasyCkpt reads and writes data only within storage directories that you specify. It may delete older checkpoints to enforce the maximum checkpoint count.
PAI enforces the following guarantees:
All operations are limited to configured save and load directories.
All save and delete operations are logged.
Do not store other data in the model's save or load directory. EasyCkpt may not function correctly, and you are responsible for any resulting data loss.