Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
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
59 changes: 41 additions & 18 deletions megatron/core/optimizer/distrib_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -537,6 +537,7 @@ def __init__(
)

self._state_offloader: Optional[OptimizerStateOffloader] = None
self._nvme_state_store = None

# when freezing sub-models we have no real optimizer
# but still need a stub DistributedOptimizer class
Expand Down Expand Up @@ -630,6 +631,15 @@ def __init__(
if self.config.offload_optimizer_states:
self._state_offloader = OptimizerStateOffloader(self)

if self.config.optimizer_state_nvme_dir is not None:
from megatron.core.optimizer.nvme_state_store import NVMeOptimizerStateStore

self._nvme_state_store = NVMeOptimizerStateStore(
self,
self.config.optimizer_state_nvme_dir,
self.config.optimizer_state_nvme_chunk_mb,
)

def _get_model_param_range_map(self, param: torch.nn.Parameter):
"""
Given a model param, get the index sub-range of the param that this
Expand All @@ -656,6 +666,11 @@ def state_dict(self):
optimizer state (e.g., exp_avg, exp_avg_sq) are stored in a separate
checkpoint file by calling 'save_parameter_state()'.
"""
if self._nvme_state_store is not None:
raise RuntimeError(
"Checkpointing optimizer state is not supported with "
"--optimizer-state-nvme-dir yet (state lives on NVMe)."
)
inner_state_dict = self.optimizer.state_dict()
state_dict = {}

Expand Down Expand Up @@ -1247,6 +1262,11 @@ def sharded_state_dict(

Regular state dict parameters are saved on DP rank 0 and loaded on all ranks.
"""
if self._nvme_state_store is not None:
raise RuntimeError(
"Checkpointing optimizer state is not supported with "
"--optimizer-state-nvme-dir yet (state lives on NVMe)."
)
if sharding_type is not None:
log_single_rank(
logger,
Expand Down Expand Up @@ -2486,29 +2506,29 @@ def _copy_main_params_to_model_params(self):
# Utility method for copying group params.
def copy_group_params(shard_main_groups, model_groups):
for shard_main_group, model_group in zip(shard_main_groups, model_groups):
for shard_main_param, model_param in zip(shard_main_group, model_group):
self._copy_main_params_to_model_params_for(zip(shard_main_group, model_group))

param_range_map = self._get_model_param_range_map(model_param)
world_range = param_range_map["gbuf_world_in_bucket"]
# Copy shard groups to model groups.
copy_group_params(self.shard_fp32_from_float16_groups, self.model_float16_groups)
copy_group_params(self.shard_fp32_groups, self.model_fp32_groups)

assert world_range.size == shard_main_param.nelement()
def _copy_main_params_to_model_params_for(self, pairs):
"""Copy (shard_main_param, model_param) pairs into the param buffer."""
for shard_main_param, model_param in pairs:
param_range_map = self._get_model_param_range_map(model_param)
world_range = param_range_map["gbuf_world_in_bucket"]

gbuf_index, _, bucket_id = self.model_param_gbuf_map[model_param]
model_param_buffer = self.buffers[gbuf_index].buckets[bucket_id].param_data
assert world_range.size == shard_main_param.nelement()

shard_model_param = model_param_buffer.view(-1)[
world_range.start : world_range.end
]
gbuf_index, _, bucket_id = self.model_param_gbuf_map[model_param]
model_param_buffer = self.buffers[gbuf_index].buckets[bucket_id].param_data

if is_float8tensor(model_param):
# FP8 params are quantized in the above "quantize_param_shard" function.
continue
else:
shard_model_param.data.copy_(shard_main_param)
shard_model_param = model_param_buffer.view(-1)[world_range.start : world_range.end]

# Copy shard groups to model groups.
copy_group_params(self.shard_fp32_from_float16_groups, self.model_float16_groups)
copy_group_params(self.shard_fp32_groups, self.model_fp32_groups)
if is_float8tensor(model_param):
# FP8 params are quantized in the above "quantize_param_shard" function.
continue
shard_model_param.data.copy_(shard_main_param)

def _copy_main_params_to_param_buffer(self):
"""
Expand Down Expand Up @@ -2634,7 +2654,10 @@ def step_with_ready_grads(self) -> bool:
"""
if self._state_offloader is not None:
self._state_offloader.sync_before_step()
update_successful = super().step_with_ready_grads()
if self._nvme_state_store is not None:
update_successful = self._nvme_state_store.step()
else:
update_successful = super().step_with_ready_grads()

timers = self.config.timers
if timers is not None:
Expand Down
206 changes: 206 additions & 0 deletions megatron/core/optimizer/nvme_state_store.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,206 @@
import atexit
import logging
import os
import shutil
import time
from typing import TYPE_CHECKING, Dict, List, Tuple

import torch

if TYPE_CHECKING:
from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer

logger = logging.getLogger(__name__)

_MOMENT_KEYS = ("exp_avg", "exp_avg_sq")


class _BucketSpec:
"""One DDP bucket's slice of optimizer state and its backing file.

The file holds three equally-sized segments, [main | exp_avg | exp_avg_sq],
each laid out as the bucket's param shards concatenated in group order.
"""

def __init__(self, index: int, path: str, entries: List[Tuple[torch.nn.Parameter, torch.Tensor, int]]):
self.index = index
self.path = path
self.entries = entries # (model_param, shard_main_param, master_group_idx)
self.numel = sum(main.numel() for _, main, _ in entries)
offsets = []
pos = 0
for _, main, _ in entries:
offsets.append(pos)
pos += main.numel()
self.entry_offsets = offsets
self.fd = -1
self.adam = None
self.group_master_indices: List[int] = []
self.main_on_disk = False
self.moments_on_disk = False


class NVMeOptimizerStateStore:
"""Owns residency and I/O of one DistributedOptimizer's state."""

def __init__(self, distrib_optimizer: "DistributedOptimizer", dir_root: str, chunk_mb: int):
self.dist_opt = distrib_optimizer
config = distrib_optimizer.config

assert not config.use_precision_aware_optimizer, (
"NVMe state store requires the non-precision-aware optimizer "
"(fp32 main params held by mcore)."
)
assert not config.optimizer_cpu_offload, "NVMe state store is mutually exclusive with CPU offload."
assert not config.offload_optimizer_states, (
"NVMe state store is mutually exclusive with --offload-optimizer-states."
)
assert not distrib_optimizer.ddp_config.use_megatron_fsdp
assert all(len(g) == 0 for g in distrib_optimizer.model_fp32_groups), (
"NVMe state store only supports pure bf16/fp16 models (no fp32 model params)."
)

rank = torch.distributed.get_rank()
instance = distrib_optimizer.distributed_optimizer_instance_id
self.dir = os.path.join(dir_root, f"rank{rank}", f"opt{instance}")
shutil.rmtree(self.dir, ignore_errors=True)
os.makedirs(self.dir, exist_ok=True)
atexit.register(shutil.rmtree, self.dir, ignore_errors=True)

self._chunk = torch.empty(chunk_mb * 1024 * 1024 // 4, dtype=torch.float32, pin_memory=True)
self._chunk_np = self._chunk.numpy()

self.specs = self._build_specs()
self._build_bucket_optimizers()
for spec in self.specs:
spec.fd = os.open(spec.path, os.O_RDWR | os.O_CREAT, 0o600)
os.posix_fallocate(spec.fd, 0, 3 * spec.numel * 4)

total_gb = sum(3 * s.numel * 4 for s in self.specs) / 1024**3
logger.info(
f"NVMe optimizer state store: {len(self.specs)} buckets, "
f"{total_gb:.1f} GB state at {self.dir}"
)

def _build_specs(self) -> List["_BucketSpec"]:
by_bucket: Dict[Tuple, List] = {}
groups = zip(
self.dist_opt.model_float16_groups, self.dist_opt.shard_fp32_from_float16_groups
)
for group_idx, (model_group, main_group) in enumerate(groups):
for model_param, main_param in zip(model_group, main_group):
assert main_param is not None and main_param.dtype == torch.float32
key = self.dist_opt.model_param_gbuf_map[model_param]
by_bucket.setdefault(key, []).append((model_param, main_param, group_idx))
return [
_BucketSpec(i, os.path.join(self.dir, f"bucket{i:05d}.bin"), entries)
for i, (_, entries) in enumerate(sorted(by_bucket.items(), key=lambda kv: kv[0]))
]

def _build_bucket_optimizers(self) -> None:
from megatron.core.optimizer import Adam

master_groups = self.dist_opt.optimizer.param_groups
for spec in self.specs:
groups = []
spec.group_master_indices = sorted({gi for _, _, gi in spec.entries})
for g_idx in spec.group_master_indices:
group = {k: v for k, v in master_groups[g_idx].items() if k != "params"}
group["params"] = [main for _, main, gi in spec.entries if gi == g_idx]
groups.append(group)
spec.adam = Adam(groups, adam_w_mode=self.dist_opt.config.decoupled_weight_decay)

# ------------------------------------------------------------------ step

@torch.no_grad()
def step(self) -> bool:
t0 = time.monotonic()
read_bytes = written_bytes = 0
for spec in self.specs:
read_bytes += self._load_bucket(spec)
self._sync_hyperparams(spec)
spec.adam.step()
self.dist_opt._copy_main_params_to_model_params_for(
(main, model) for model, main, _ in spec.entries
)
written_bytes += self._store_bucket(spec)
logger.info(
f"NVMe streaming step: {len(self.specs)} buckets, "
f"read {read_bytes / 1024**3:.1f} GB, wrote {written_bytes / 1024**3:.1f} GB "
f"in {time.monotonic() - t0:.1f}s"
)
return True

def _sync_hyperparams(self, spec: "_BucketSpec") -> None:
master_groups = self.dist_opt.optimizer.param_groups
for group, g_idx in zip(spec.adam.param_groups, spec.group_master_indices):
group["lr"] = master_groups[g_idx]["lr"]
group["weight_decay"] = master_groups[g_idx]["weight_decay"]

def _load_bucket(self, spec: "_BucketSpec") -> int:
nbytes = 0
if spec.main_on_disk:
for tensor, offset in self._segment(spec, "main"):
self._materialize(tensor)
self._stream(spec.fd, offset, tensor, to_disk=False)
nbytes += tensor.numel() * 4
if spec.moments_on_disk:
for key in _MOMENT_KEYS:
for tensor, offset in self._segment(spec, key):
self._materialize(tensor)
self._stream(spec.fd, offset, tensor, to_disk=False)
nbytes += tensor.numel() * 4
return nbytes

def _store_bucket(self, spec: "_BucketSpec") -> int:
nbytes = 0
for key in ("main",) + _MOMENT_KEYS:
for tensor, offset in self._segment(spec, key):
self._stream(spec.fd, offset, tensor, to_disk=True)
self._release(tensor)
nbytes += tensor.numel() * 4
spec.main_on_disk = True
spec.moments_on_disk = True
return nbytes

# ------------------------------------------------------- residency & I/O

def _segment(self, spec: "_BucketSpec", key: str):
segment_index = ("main",) + _MOMENT_KEYS
base = segment_index.index(key) * spec.numel * 4
for (_, main, _), entry_offset in zip(spec.entries, spec.entry_offsets):
tensor = main if key == "main" else spec.adam.state[main][key]
yield tensor, base + entry_offset * 4

@staticmethod
def _materialize(tensor: torch.Tensor) -> None:
tensor.untyped_storage().resize_(tensor.numel() * tensor.element_size())

@staticmethod
def _release(tensor: torch.Tensor) -> None:
tensor.untyped_storage().resize_(0)

def _stream(self, fd: int, base_offset: int, tensor: torch.Tensor, *, to_disk: bool) -> None:
flat = tensor.view(-1)
chunk_numel = self._chunk.numel()
pos = 0
while pos < flat.numel():
n = min(chunk_numel, flat.numel() - pos)
byte_offset = base_offset + pos * 4
if to_disk:
self._chunk[:n].copy_(flat[pos : pos + n])
self._rw_full(os.pwritev, fd, byte_offset, self._chunk_np[:n])
else:
self._rw_full(os.preadv, fd, byte_offset, self._chunk_np[:n])
flat[pos : pos + n].copy_(self._chunk[:n])
pos += n

@staticmethod
def _rw_full(op, fd: int, offset: int, array) -> None:
mv = memoryview(array).cast("B")
done = 0
while done < len(mv):
n = op(fd, [mv[done:]], offset + done)
if n <= 0:
raise IOError(f"short {op.__name__} ({n}) on optimizer state file at offset {offset + done}")
done += n
10 changes: 10 additions & 0 deletions megatron/core/optimizer/optimizer_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -333,6 +333,16 @@ class OptimizerConfig:
low_memory_resume: bool = False
"""If True, allocate optimizer states on CPU during checkpoint loading to prevent GPU OOM."""

optimizer_state_nvme_dir: Optional[str] = None
"""
If set, fp32 main params and Adam moments live in per-bucket files under this
node-local directory and are streamed through the GPU bucket-by-bucket during
the optimizer step, bounding GPU residency to one bucket.
"""

optimizer_state_nvme_chunk_mb: int = 256
"""Pinned staging chunk size for NVMe optimizer state streaming."""

################
# Miscellaneous
################
Expand Down
7 changes: 7 additions & 0 deletions megatron/training/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -2110,6 +2110,13 @@ def _add_training_args(parser):
'Only support TE FusedAdam optimizer.'
'Note that this still uses pure GPU optimizer instead of '
'HybridDeviceOptimizer for --optimizer-cpu-offload.')
group.add_argument('--optimizer-state-nvme-dir', type=str, default=None,
help='Stream fp32 main params and Adam moments through per-bucket '
'files under this node-local directory during the optimizer step, '
'bounding GPU residency to one bucket. Checkpointing optimizer '
'state is not supported yet.')
group.add_argument('--optimizer-state-nvme-chunk-mb', type=int, default=256,
help='Pinned staging chunk size for NVMe optimizer state streaming.')
group.add_argument('--dataloader-type', type=str, default=None,
choices=['single', 'cyclic', 'external'],
help='Single pass vs multiple pass data loader')
Expand Down
Loading