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
34 changes: 30 additions & 4 deletions megatron/core/optimizer/distrib_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,7 @@
from .grad_scaler import MegatronGradScaler
from .optimizer import MixedPrecisionOptimizer, _zero_grad_group_helper, param_group_identifier_keys
from .optimizer_config import OptimizerConfig
from .fused_adam_patch import apply_fused_adam_patch, is_patch_applied

logger = getLogger(__name__)

Expand Down Expand Up @@ -335,8 +336,10 @@ def _build_model_and_main_param_groups(
shard_fp32_groups = []
shard_fp32_from_float16_groups = []

model_param_group_index_map = {}

# Allocate (or slice) each group's param shard.
for group_range in opt_group_ranges:
for group_index, group_range in enumerate(opt_group_ranges):

# Params of this group.
model_float16_params_this_group = []
Expand Down Expand Up @@ -449,12 +452,24 @@ def _build_model_and_main_param_groups(
*shard_float16_params_this_group,
]

# Build correct order for later access in models with mixed precision params.
# When a model has both fp32 and bf16/fp16 parameters (e.g., some params kept
# in fp32 via _keep_fp32), the parameter ordering in optimizer groups may differ
# from _build_optimizer_group_ranges. This rebuilds the index map to match the
# actual order used by the optimizer after shard allocation.
for new_order, model_param in enumerate(model_fp32_params_this_group):
model_param_group_index_map[model_param] = (group_index, new_order)
offset = len(model_fp32_params_this_group)
for i, model_param in enumerate(model_float16_params_this_group):
model_param_group_index_map[model_param] = (group_index, offset + i)

return (
model_float16_groups,
model_fp32_groups,
shard_float16_groups,
shard_fp32_groups,
shard_fp32_from_float16_groups,
model_param_group_index_map,
)

def __init__(
Expand Down Expand Up @@ -587,7 +602,7 @@ def __init__(
param.main_param_sharded = True

# Optimizer ranges.
(self.model_param_group_index_map, self.opt_group_ranges) = (
(_, self.opt_group_ranges) = (
self._build_optimizer_group_ranges(self.optimizer.param_groups, self.gbuf_ranges)
)

Expand All @@ -598,6 +613,7 @@ def __init__(
self.shard_float16_groups,
self.shard_fp32_groups,
self.shard_fp32_from_float16_groups,
self.model_param_group_index_map,
) = self._build_model_and_main_param_groups(
self.gbuf_ranges, self.model_param_gbuf_map, self.opt_group_ranges, config
)
Expand Down Expand Up @@ -795,8 +811,12 @@ def make_needed_groups(param_group):

# Allocate dummy tensors.
numel = len(param_range_map["gbuf_world"])
init_shard = lambda dtype=torch.float32: torch.empty(
(numel,), dtype=dtype, device=torch.cuda.current_device()
low_mem_resume = self.config.low_memory_resume
if low_mem_resume and USING_TE_OPTIMIZER and not is_patch_applied():
apply_fused_adam_patch()
init_device = 'cpu' if low_mem_resume else torch.cuda.current_device()
init_shard = lambda dtype=torch.float32, _device=init_device: torch.empty(
(numel,), dtype=dtype, device=_device
)

# For precision_aware_optimizer, the empty tensors should also be
Expand Down Expand Up @@ -888,6 +908,12 @@ def make_needed_groups(param_group):
else:
raise NotImplementedError(f'Unknown sharding_type: {sharding_type}')

if self.config.low_memory_resume:
for state in self.optimizer.state.values():
for k, v in state.items():
if isinstance(v, torch.Tensor) and v.device.type == 'cpu':
state[k] = v.to(torch.cuda.current_device())

def _get_main_param_and_optimizer_states(self, model_param):
"""Return a dict containing the main param and optimizer states corresponding to the input
model_param.
Expand Down
66 changes: 66 additions & 0 deletions megatron/core/optimizer/fused_adam_patch.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,66 @@
"""Monkey-patch for TransformerEngine FusedAdam to support CPU allocation
during low-memory checkpoint resume (--low-memory-resume).

This prevents GPU OOM during checkpoint loading by initially allocating
optimizer states (exp_avg, exp_avg_sq) on CPU, then moving them to GPU
after the checkpoint is fully loaded.

The patch is only applied when low_memory_resume is enabled, so
_patched_initialize_state unconditionally allocates on CPU.
"""

import torch
from logging import getLogger

logger = getLogger(__name__)
_patch_applied = False


def _patched_initialize_state(self, param, state_name, zero_buffer, store_param_remainders=False):
"""Patched _initialize_state: allocate on CPU for low-memory resume."""
from transformer_engine.pytorch.tensor.float8_tensor import Float8Quantizer
from transformer_engine.pytorch.quantized_tensor import QuantizedTensor
import transformer_engine_torch as tex

dtype = self.name_to_dtype_map[state_name]
param_for_empty = param.dequantize() if isinstance(param, QuantizedTensor) else param

device = torch.device('cpu')

if store_param_remainders:
data = torch.zeros(param_for_empty.shape, dtype=torch.int16, device=device)
else:
data = torch.empty(param_for_empty.shape, dtype=dtype, device=device)

if zero_buffer:
data.zero_()

if dtype == torch.uint8:
quantizer = Float8Quantizer(
scale=torch.ones([1], dtype=torch.float32, device=device),
amax=torch.zeros([1], dtype=torch.float32, device=device),
fp8_dtype=tex.DType.kFloat8E4M3,
)
self.state[param][state_name] = quantizer.make_empty(param.shape)
self.state[param][state_name].quantize_(data.float())
else:
self.state[param][state_name] = data

if dtype != torch.float32:
if param not in self._scales:
self._scales[param] = {}
self._scales[param][state_name] = torch.ones([1], dtype=torch.float32, device=device)


def apply_fused_adam_patch():
global _patch_applied
if _patch_applied:
return

from transformer_engine.pytorch.optimizers import FusedAdam
FusedAdam._initialize_state = _patched_initialize_state
_patch_applied = True


def is_patch_applied():
return _patch_applied
3 changes: 3 additions & 0 deletions megatron/core/optimizer/optimizer_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -330,6 +330,9 @@ class OptimizerConfig:
reload them before the next optimizer step.
"""

low_memory_resume: bool = False
"""If True, allocate optimizer states on CPU during checkpoint loading to prevent GPU OOM."""

################
# Miscellaneous
################
Expand Down
3 changes: 3 additions & 0 deletions megatron/core/transformer/mlp.py
Original file line number Diff line number Diff line change
Expand Up @@ -334,6 +334,9 @@ def sh_ten_build_fn(

def sh_ten_merge_fn(sub_state_dict):
with torch.no_grad():
from megatron.training import get_args
if get_args().low_memory_resume:
return torch.cat([t.cpu() for t in sub_state_dict])
try:
return torch.cat(sub_state_dict)
except (RuntimeError, torch.cuda.OutOfMemoryError) as e:
Expand Down
3 changes: 3 additions & 0 deletions megatron/training/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -2196,6 +2196,9 @@ def _add_checkpointing_args(parser):
help='Do not load optimizer when loading checkpoint.')
group.add_argument('--no-load-rng', action='store_true', default=None,
help='Do not load rng state when loading checkpoint.')
group.add_argument('--low-memory-resume', action='store_true', default=False,
help='Allocate optimizer states on CPU during distributed optimizer checkpoint loading '
'to prevent GPU OOM on large peak memory.')
group.add_argument('--use-dist-ckpt', action='store_true',
dest='use_dist_ckpt_deprecated',
help='Deprecated: see --ckpt-format.')
Expand Down