Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions megatron/core/dist_checkpointing/strategies/torch.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.

""" Strategies using PyTorch distributed.checkpoint as an underlying format. """
import inspect
import io
import os
import pickle
Expand Down Expand Up @@ -600,6 +601,7 @@ def __init__(
thread_count: int = 1,
cached_metadata: bool = False,
separation_hint: Optional[str] = None,
cpu_shm_mode: bool = False,
):
"""Adds parameters specific to PyT Distributed format
Args:
Expand All @@ -614,6 +616,10 @@ def __init__(
gathering local metadata every checkpointing invocation
separation_hint(str, optional): If provided, all tensors whose keys have this
prefix will be saved to a separate file.
cpu_shm_mode (bool, optional): Copy GPU tensors to CPU shared-memory in the
training process before handing off to the async worker. Avoids CUDA IPC /
NVLink fabric handles in the worker subprocess. Only applies with nvrx async
strategy.
"""
self.backend = backend
self.version = version
Expand All @@ -640,6 +646,7 @@ def __init__(
self.cached_global_metadata: Optional[Metadata] = None

self.separation_hint = separation_hint
self.cpu_shm_mode = cpu_shm_mode

self.validated_loaded_metadata_reuse = False

Expand Down Expand Up @@ -698,6 +705,18 @@ def async_save(
self._metadata_cache.set_cached_global_metadata(self.cached_global_metadata)
# Define additional arguments
async_writer_kwargs["use_cached_data_structure"] = self.use_cached_ckpt_structure
if self.cpu_shm_mode:
if (
"use_cpu_shm_for_gpu_tensors"
in inspect.signature(async_writer.__init__).parameters
):
async_writer_kwargs["use_cpu_shm_for_gpu_tensors"] = True
else:
raise AssertionError(
"Installed nvidia-resiliency-ext does not support "
"use_cpu_shm_for_gpu_tensors. Update nvidia-resiliency-ext "
"to enable cpu_shm_mode."
)
state_dict_saver_kwargs["enable_cache"] = self.use_cached_ckpt_structure
state_dict_saver_kwargs["metadata_cache"] = self._metadata_cache
else:
Expand Down
28 changes: 24 additions & 4 deletions megatron/training/async_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,9 @@
This module provides a singleton instance of AsyncCallsQueue which manages
the async checkpoint save calls.
"""
import inspect
import logging
import time

from abc import ABC

from megatron.core.dist_checkpointing.strategies.async_utils import AsyncRequest
Expand All @@ -17,10 +17,14 @@
try:
from nvidia_resiliency_ext.checkpointing.async_ckpt.core import AsyncRequest as NVRxAsyncRequest
from nvidia_resiliency_ext.checkpointing.async_ckpt.filesystem_async import _results_queue
from nvidia_resiliency_ext.checkpointing.async_ckpt.state_dict_saver import save_state_dict_async_finalize
from nvidia_resiliency_ext.checkpointing.async_ckpt.state_dict_saver import (
save_state_dict_async_finalize,
)
except (ImportError, ModuleNotFoundError):
from megatron.core.dist_checkpointing.strategies.filesystem_async import _results_queue
from megatron.core.dist_checkpointing.strategies.state_dict_saver import save_state_dict_async_finalize
from megatron.core.dist_checkpointing.strategies.state_dict_saver import (
save_state_dict_async_finalize,
)

NVRxAsyncRequest = ABC

Expand Down Expand Up @@ -56,12 +60,28 @@ def init_persistent_async_worker(rank: int, mp_mode: str = 'spawn'):
time_start = time.time()
if rank == 0:
print(f"init_persistent_async_worker: {rank}, Starting Async Caller", flush=True)
_async_calls_queue = AsyncCallsQueue(persistent=True)
_async_calls_queue = AsyncCallsQueue(
persistent=True,
**(
{"cpu_shm_mode": args.async_ckpt_use_cpu_shm}
if async_strategy == "nvrx"
and "cpu_shm_mode" in inspect.signature(AsyncCallsQueue.__init__).parameters
else {}
),
)
# initialize the persistent caller with QoS priorities from args
kwargs = {}
if async_strategy == "mcore":
# Note: nvidia-resiliency-ext uses is_daemon instead of mp_mode (always spawns)
kwargs["mp_mode"] = mp_mode
elif async_strategy == "nvrx":
if "cpu_shm_mode" in inspect.signature(AsyncCallsQueue.warmup_persistent_caller).parameters:
kwargs["cpu_shm_mode"] = args.async_ckpt_use_cpu_shm
elif args.async_ckpt_use_cpu_shm:
raise AssertionError(
"Installed nvidia-resiliency-ext does not support cpu_shm_mode. "
"Update nvidia-resiliency-ext to use --async-ckpt-use-cpu-shm."
)
AsyncCallsQueue.warmup_persistent_caller(
rank,
cpu_priority=args.async_ckpt_cpu_priority,
Expand Down
57 changes: 44 additions & 13 deletions megatron/training/checkpointing.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
"""Input/output checkpointing."""

import contextlib
import inspect
import multiprocessing
import os
import random
Expand All @@ -16,41 +17,45 @@
from logging import getLogger
from pathlib import Path
from time import time
from typing import Any, Dict, List, Optional, Union

import numpy as np
import torch
from typing import Optional, Union, List, Dict, Any
from torch.distributed.checkpoint import FileSystemReader, default_planner

from megatron.core import dist_checkpointing, mpu, tensor_parallel
from megatron.core.dist_checkpointing.mapping import ShardedObject
from megatron.core.dist_checkpointing.strategies.torch import TorchDistLoadShardedStrategy, TorchDistSaveShardedStrategy
from megatron.core.dist_checkpointing.strategies.async_utils import _disable_gc
from megatron.core.dist_checkpointing.strategies.fully_parallel import (
FullyParallelLoadStrategyWrapper,
FullyParallelSaveStrategyWrapper,
)
from megatron.core.dist_checkpointing.strategies.torch import (
TorchDistLoadShardedStrategy,
TorchDistSaveShardedStrategy,
)
from megatron.core.msc_utils import MultiStorageClientFeature, open_file
from megatron.core.num_microbatches_calculator import update_num_microbatches
from megatron.core.utils import get_pg_rank, get_pg_size
from megatron.core.optimizer import DistributedOptimizer
from megatron.core.rerun_state_machine import get_rerun_state_machine
from megatron.core.utils import get_torch_version, is_torch_min_version
from megatron.core.utils import get_pg_rank, get_pg_size, get_torch_version, is_torch_min_version

from ..core.dist_checkpointing.utils import _clean_metadata_for_serialization
from . import ft_integration, wandb_utils
from .async_utils import get_save_and_finalize_callbacks, is_empty_async_queue, schedule_async_save
from megatron.core.dist_checkpointing.strategies.async_utils import _disable_gc
from .global_vars import get_args
from .one_logger_utils import on_save_checkpoint_start, on_save_checkpoint_success
from .utils import append_to_progress_log, is_last_rank, print_rank_0, unwrap_model

try:
from megatron.core.distributed.fsdp.src.megatron_fsdp.uneven_dtensor import preprocess_state_dict_for_uneven_dtensor
from megatron.core.distributed.fsdp.src.megatron_fsdp.uneven_dtensor import (
preprocess_state_dict_for_uneven_dtensor,
)
from megatron.core.transformer.fsdp_dtensor_checkpoint import (
print_diff_in_state_dicts,
handle_experts_in_state_dict,
handle_fp8_extra_state_case,
handle_swiglu_in_state_dict,
handle_experts_in_state_dict,
print_diff_in_state_dicts,
)
HAVE_MEGATRON_FSDP = True
except ImportError:
Expand All @@ -60,15 +65,20 @@
# [ModelOpt]: Import
try:
from modelopt.torch.opt.plugins import save_modelopt_state, save_sharded_modelopt_state

from megatron.post_training.utils import print_distributed_quant_summary
has_nvidia_modelopt = True
except Exception:
has_nvidia_modelopt = False


try:
from nvidia_resiliency_ext.checkpointing.async_ckpt.filesystem_async import FileSystemWriterAsync
from nvidia_resiliency_ext.checkpointing.async_ckpt.state_dict_saver import save_state_dict_async_plan
from nvidia_resiliency_ext.checkpointing.async_ckpt.filesystem_async import (
FileSystemWriterAsync,
)
from nvidia_resiliency_ext.checkpointing.async_ckpt.state_dict_saver import (
save_state_dict_async_plan,
)

HAVE_NVRX = True
except (ImportError, ModuleNotFoundError):
Expand Down Expand Up @@ -631,7 +641,9 @@ def save_checkpoint(iteration, model, optimizer, opt_param_scheduler, num_floati
validate_sharding_integrity = not args.ckpt_assume_constant_structure
else:
validate_sharding_integrity = True
save_strategy = TorchDistSaveShardedStrategy()
save_strategy = TorchDistSaveShardedStrategy(
cpu_shm_mode=getattr(args, 'async_ckpt_use_cpu_shm', False)
)
if args.ckpt_assume_constant_structure and args.ckpt_format == 'torch_dist':
save_strategy.use_cached_ckpt_structure = args.ckpt_assume_constant_structure
if args.async_save:
Expand Down Expand Up @@ -681,8 +693,25 @@ def save_checkpoint(iteration, model, optimizer, opt_param_scheduler, num_floati
if args.async_save:
planner = torch.distributed.checkpoint.DefaultSavePlanner()
coordinator_rank = 0
_cpu_shm = getattr(args, 'async_ckpt_use_cpu_shm', False)
_writer_kwargs = {}
if _cpu_shm:
if (
"use_cpu_shm_for_gpu_tensors"
in inspect.signature(FileSystemWriterAsync.__init__).parameters
):
_writer_kwargs["use_cpu_shm_for_gpu_tensors"] = True
else:
raise AssertionError(
"Installed nvidia-resiliency-ext does not support "
"use_cpu_shm_for_gpu_tensors. Update nvidia-resiliency-ext "
"to use --async-ckpt-use-cpu-shm."
)
fs_storage_writer = FileSystemWriterAsync(
checkpoint_name, thread_count=args.dist_ckpt_workers, use_msc=args.enable_msc
checkpoint_name,
thread_count=args.dist_ckpt_workers,
use_msc=args.enable_msc,
**_writer_kwargs,
)

save_state_dict_ret = save_state_dict_async_plan(
Expand All @@ -707,7 +736,9 @@ def save_checkpoint(iteration, model, optimizer, opt_param_scheduler, num_floati
logger.debug(f"rank: {rank}, takes {end_ckpt - start_ckpt} to prepare state dict for ckpt ")
if ckpt_type == CheckpointType.LOCAL:
try:
from megatron.core.dist_checkpointing.tensor_aware_state_dict import MCoreTensorAwareStateDict
from megatron.core.dist_checkpointing.tensor_aware_state_dict import (
MCoreTensorAwareStateDict,
)
except ModuleNotFoundError:
raise RuntimeError("The 'nvidia_resiliency_ext' module is required for local "
"checkpointing but was not found. Please ensure it is installed.")
Expand Down
6 changes: 6 additions & 0 deletions megatron/training/config/training_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -477,6 +477,12 @@ class CheckpointConfig:
async_ckpt_io_priority: Optional[int] = 3
"""I/O scheduling class (0-3, 3=idle) for the async checkpoint writer process."""

async_ckpt_use_cpu_shm: bool = False
"""Copy GPU tensors to CPU shared-memory in the training process before handing off to
the async checkpoint worker. Avoids CUDA IPC / NVLink fabric handles in the worker
subprocess. Useful on MNNVL systems where fabric resources are exhausted.
Only applies with the nvrx async strategy."""

fully_parallel_load: bool = field(default=False, metadata={"argparse_meta": {"arg_names": ["--ckpt-fully-parallel-load"], "dest": "ckpt_fully_parallel_load"}})
"""Apply full load parallelization across DP for distributed checkpoints."""

Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -205,7 +205,7 @@ flash_mla = [
transformer-engine = { git = "https://github.com/NVIDIA/TransformerEngine.git", rev = "f031cf87bd054c7558b887df7bed93975456667f" }
nemo-run = { git = "https://github.com/NVIDIA-NeMo/Run.git", rev = "17ae86b64d7f75653351664f5d8c9e466faede00" }
emerging_optimizers = { git = "https://github.com/NVIDIA-NeMo/Emerging-Optimizers.git", rev = "v0.2.0" }
nvidia-resiliency-ext = { git = "https://github.com/NVIDIA/nvidia-resiliency-ext.git", rev = "15a851565a4ce846c04431ecb0cf09903ab4837e" }
nvidia-resiliency-ext = { git = "https://github.com/NVIDIA/nvidia-resiliency-ext.git", rev = "b2bb3d728a18795807d9f76c535e005a609a1b01" }

[tool.isort]
profile = "black" # black-compatible
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -50,11 +50,13 @@ MODEL_ARGS:
--ckpt-format: torch_dist
--dist-ckpt-optim-fully-reshardable: true
--dist-ckpt-strictness: log_all # backward compatibility for TE changes
--ckpt-assume-constant-structure: true
--data-cache-path: ${DATA_CACHE_PATH}
--bf16: true
--log-memory-to-tensorboard: true
--async-save: true
--async-strategy: mcore
--use-persistent-ckpt-worker: true
--async-ckpt-use-cpu-shm: true
TEST_TYPE: ckpt-resume
LAUNCHER: ft_launcher
Original file line number Diff line number Diff line change
Expand Up @@ -66,5 +66,6 @@ MODEL_ARGS:
--log-memory-to-tensorboard: true
--async-save: true
--use-persistent-ckpt-worker: true
--async-ckpt-use-cpu-shm: true
TEST_TYPE: ckpt-resume
LAUNCHER: ft_launcher
Loading
Loading