Skip to content
Open
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
79 changes: 73 additions & 6 deletions megatron/training/checkpointing.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,8 +65,9 @@
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 .global_vars import get_args
from .global_vars import get_args, get_train_state_if_initialized
from .one_logger_utils import on_save_checkpoint_start, on_save_checkpoint_success
from .state import TRAIN_STATE_FILENAME, TrainState, load_train_state, save_train_state
from .utils import append_to_progress_log, is_last_rank, print_rank_0, print_rank_last, warn_rank_0

try:
Expand Down Expand Up @@ -563,6 +564,34 @@ class CheckpointType(Enum):
FSDP_DTENSOR = auto()


def _get_checkpoint_train_state_filename(checkpoint_name, ckpt_type):
"""Return the per-iteration train-state sidecar path for a loaded checkpoint."""
checkpoint_path = maybe_msc.Path(checkpoint_name)
if ckpt_type == CheckpointType.LEGACY:
checkpoint_path = checkpoint_path.parent.parent
return str(checkpoint_path.joinpath(TRAIN_STATE_FILENAME))


def _snapshot_train_state(args, iteration, num_floating_point_operations_so_far):
"""Capture checkpoint progress without requiring full runtime initialization."""
active_train_state = get_train_state_if_initialized()
source_train_state = active_train_state if active_train_state is not None else TrainState()
source_train_state.update_from_args(args, iteration, num_floating_point_operations_so_far)
snapshot = TrainState()
snapshot.load_state_dict(source_train_state.state_dict())
return snapshot


def _get_train_state_finalize_fn(train_state, train_state_filename):
"""Build the deferred sidecar write around an immutable state snapshot."""

def finalize_train_state():
ensure_directory_exists(train_state_filename)
save_train_state(train_state, train_state_filename)

return finalize_train_state


def _build_sharded_state_dict_metadata(
args: Namespace, dp_cp_group: Optional[torch.distributed.ProcessGroup] = None
) -> dict:
Expand Down Expand Up @@ -779,6 +808,13 @@ def save_checkpoint(
expert_rank=expert_rank,
return_base_dir=return_base_dir,
)
train_state_filename = os.path.join(
get_checkpoint_name(save_dir, iteration, release=release, return_base_dir=True),
TRAIN_STATE_FILENAME,
)
# Async finalizers may run after training advances, so persist an immutable snapshot.
train_state = _snapshot_train_state(args, iteration, num_floating_point_operations_so_far)
finalize_train_state = _get_train_state_finalize_fn(train_state, train_state_filename)

# Save distributed optimizer's custom parameter state.
if (
Expand Down Expand Up @@ -1134,6 +1170,7 @@ def _rank_and_size(explicit_rank, group, mpu_rank_fn, mpu_size_fn):

def iter_finalize_fn():
prev_iteration = 0
finalize_train_state()
save_retain_interval = getattr(
args, 'save_retain_interval', None
) # For backwards compatibility of tests.
Expand Down Expand Up @@ -3006,6 +3043,12 @@ def load_checkpoint(
# Iteration and num_floating_point_operations_so_far default to 0.
return 0, 0

loaded_train_state = None
if not args.finetune and not release and ckpt_type != CheckpointType.LOCAL:
train_state_filename = _get_checkpoint_train_state_filename(checkpoint_name, ckpt_type)
if maybe_msc.os.path.isfile(train_state_filename):
loaded_train_state = load_train_state(train_state_filename)

# Set checkpoint version.
set_checkpoint_version(state_dict.get('checkpoint_version', 0))

Expand All @@ -3017,6 +3060,8 @@ def load_checkpoint(
# Set iteration.
if args.finetune or release:
iteration = 0
elif loaded_train_state is not None:
iteration = loaded_train_state.iteration
else:
try:
iteration = state_dict['iteration']
Expand All @@ -3030,7 +3075,11 @@ def load_checkpoint(
)
)
sys.exit()
num_floating_point_operations_so_far = state_dict.get('num_floating_point_operations_so_far', 0)
num_floating_point_operations_so_far = (
loaded_train_state.num_floating_point_operations_so_far
if loaded_train_state is not None
else state_dict.get('num_floating_point_operations_so_far', 0)
)

# Check arguments.
if 'args' in state_dict and not args.finetune:
Expand All @@ -3041,13 +3090,28 @@ def load_checkpoint(
# compatibility check.
skip_args = {'num_layers'} if gpt_compat_layer_maps is not None else None
check_checkpoint_args(checkpoint_args, skip_args=skip_args)
args.consumed_train_samples = getattr(checkpoint_args, 'consumed_train_samples', 0)
args.skipped_train_samples = getattr(checkpoint_args, 'skipped_train_samples', 0)
update_num_microbatches(consumed_samples=args.consumed_train_samples, verbose=True)
args.consumed_valid_samples = getattr(checkpoint_args, 'consumed_valid_samples', 0)
if loaded_train_state is None:
args.consumed_train_samples = getattr(checkpoint_args, 'consumed_train_samples', 0)
args.skipped_train_samples = getattr(checkpoint_args, 'skipped_train_samples', 0)
args.consumed_valid_samples = getattr(checkpoint_args, 'consumed_valid_samples', 0)
args.do_train = getattr(checkpoint_args, 'do_train', False)
args.do_valid = getattr(checkpoint_args, 'do_valid', False)
args.do_test = getattr(checkpoint_args, 'do_test', False)
else:
print_rank_0('could not find arguments in the checkpoint ...')

active_train_state = get_train_state_if_initialized()
if active_train_state is None:
active_train_state = TrainState()
if loaded_train_state is not None:
active_train_state.load_state_dict(loaded_train_state.state_dict())
else:
active_train_state.update_from_args(args, iteration, num_floating_point_operations_so_far)
active_train_state.apply_to_args(args)

if not args.finetune and (loaded_train_state is not None or 'args' in state_dict):
update_num_microbatches(consumed_samples=args.consumed_train_samples, verbose=True)

# --override-ckpt-iteration: rewind the data loader to this iteration, operating on `args`
# (not state_dict) so it also works on checkpoints with no saved `args` (release / HF). The
# GBS-match check applies only when adopting this checkpoint's own args (not a finetune load).
Expand All @@ -3072,6 +3136,9 @@ def load_checkpoint(
print_rank_0(f'--override-ckpt-iteration: start at iteration {iteration} '
f'(consumed_train_samples {args.consumed_train_samples})')

# Keep the canonical runtime state consistent with any load-time iteration override.
active_train_state.update_from_args(args, iteration, num_floating_point_operations_so_far)

def load_model_state_dict(module, state_dict, strict: bool):
"""Helper function to load state dict with fallback for missing extra states."""
# GTP native-FP8 weights: load_state_dict's copy_ re-quantizes into the FP8 param, which
Expand Down
5 changes: 5 additions & 0 deletions megatron/training/global_vars.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,11 @@ def get_train_state():
return _GLOBAL_TRAIN_STATE


def get_train_state_if_initialized():
"""Return the active train state, or ``None`` for legacy initialization paths."""
return _GLOBAL_TRAIN_STATE


def get_tokenizer():
"""Return tokenizer."""
_ensure_var_is_initialized(_GLOBAL_TOKENIZER, 'tokenizer')
Expand Down
67 changes: 65 additions & 2 deletions megatron/training/state.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,16 @@
# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.

from dataclasses import dataclass
from os import PathLike
from typing import Any

import torch
from torch.distributed.checkpoint.stateful import Stateful

from megatron.core.msc_utils import maybe_msc

TRAIN_STATE_FILENAME = "train_state.pt"


@dataclass
class TrainState(Stateful):
Expand All @@ -24,6 +30,28 @@ class TrainState(Stateful):
do_valid: bool = False
do_test: bool = False

def update_from_args(
self, args: Any, iteration: int, num_floating_point_operations_so_far: int
) -> None:
"""Update the active state from legacy runtime fields during migration."""
self.iteration = iteration
self.consumed_train_samples = getattr(args, "consumed_train_samples", 0)
self.skipped_train_samples = getattr(args, "skipped_train_samples", 0)
self.consumed_valid_samples = getattr(args, "consumed_valid_samples", 0)
self.num_floating_point_operations_so_far = num_floating_point_operations_so_far
self.do_train = getattr(args, "do_train", False)
self.do_valid = getattr(args, "do_valid", False)
self.do_test = getattr(args, "do_test", False)

def apply_to_args(self, args: Any) -> None:
"""Mirror loaded state to legacy runtime fields during migration."""
args.consumed_train_samples = self.consumed_train_samples
args.skipped_train_samples = self.skipped_train_samples
args.consumed_valid_samples = self.consumed_valid_samples
args.do_train = self.do_train
args.do_valid = self.do_valid
args.do_test = self.do_test

def state_dict(self) -> dict[str, torch.Tensor]:
"""Serializes the training state into a dictionary of tensors.

Expand All @@ -38,7 +66,7 @@ def state_dict(self) -> dict[str, torch.Tensor]:
# for both the state dict and the dataclass attribute. 'iteration' is more consistent with
# Megatron-LM, but using 'step' for the state dict will allow pre-unification Bridge checkpoints
# to work without issue when loading in Megatron-LM after unification.
# The same applies for 'floating_point_operations_so_far' (Megatron-Bridge) vs
# The same applies for 'floating_point_operations_so_far' (Megatron-Bridge) vs
# 'num_floating_point_operations_so_far' (Megatron-LM).
"step": torch.tensor(self.iteration, dtype=torch.int64),
"consumed_train_samples": torch.tensor(self.consumed_train_samples, dtype=torch.int64),
Expand All @@ -62,7 +90,42 @@ def load_state_dict(self, state_dict: dict[str, torch.Tensor]) -> None:
self.consumed_train_samples = state_dict["consumed_train_samples"].item()
self.skipped_train_samples = state_dict["skipped_train_samples"].item()
self.consumed_valid_samples = state_dict["consumed_valid_samples"].item()
self.num_floating_point_operations_so_far = state_dict["floating_point_operations_so_far"].item()
self.num_floating_point_operations_so_far = state_dict[
"floating_point_operations_so_far"
].item()
self.do_train = state_dict["do_train"].item()
self.do_valid = state_dict["do_valid"].item()
self.do_test = state_dict["do_test"].item()


def save_train_state(train_state: TrainState, filename: str | PathLike[str]) -> None:
"""Write a Bridge-compatible train-state sidecar."""
maybe_msc.torch.save(train_state.state_dict(), filename)


def load_train_state(filename: str | PathLike[str]) -> TrainState:
"""Load a train-state sidecar on rank zero and broadcast it to every rank."""
distributed = torch.distributed.is_initialized()
state_obj: list[dict[str, Any] | None] = [None]

if not distributed or torch.distributed.get_rank() == 0:
try:
state_obj[0] = {
"state_dict": maybe_msc.torch.load(filename, map_location="cpu", weights_only=True)
}
except Exception as error:
state_obj[0] = {"error": f"Unable to load train state file {filename}: {error}"}

if distributed:
torch.distributed.broadcast_object_list(state_obj, src=0)

payload = state_obj[0]
if payload is None or "error" in payload:
message = (
"Train-state broadcast returned no payload" if payload is None else payload["error"]
)
raise RuntimeError(message)

train_state = TrainState()
train_state.load_state_dict(payload["state_dict"])
return train_state
Loading
Loading