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
2 changes: 2 additions & 0 deletions slime/backends/megatron_utils/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,3 +40,5 @@ def _patched_forward(self, *args, packed_seq_params=None, **kwargs):
pass

logging.getLogger("megatron").setLevel(logging.WARNING)

from . import megatron_patch # noqa: F401, E402
1 change: 1 addition & 0 deletions slime/backends/megatron_utils/megatron_patch/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
from . import megatron_chunked_grad_coalesce_patch # noqa: F401
Original file line number Diff line number Diff line change
@@ -0,0 +1,146 @@
# Patch _allreduce_non_tensor_model_parallel_grads (and its legacy alias
# _allreduce_layernorm_grads) in megatron.core.distributed.finalize_model_grads
# to coalesce/all_reduce TP-side grads in size-bounded chunks instead of one
# large _flatten_dense_tensors(grads). Lowers the peak contiguous-memory
# allocation during TP-side grad sync, avoiding OOM under allocator
# fragmentation when the combined grad buffer would otherwise be very large.
# SUM/AVG are element-wise, so chunking is mathematically equivalent.
# Chunk size: SLIME_GRAD_COALESCE_CHUNK_BYTES, default 1 GiB.
#
# Cross-compatible across the Megatron-LM versions slime is run against:
# the core_v0.13.0 line (DDP config exposes `use_custom_fsdp`, `_get_main_grad_attr`
# takes `(param, use_custom_fsdp)`, target function takes `(model, config)`) and
# the post-core_v0.15.0rc7 dev line (`use_megatron_fsdp`, single-arg
# `_get_main_grad_attr`, `(model, config, tp_group)`). API differences are
# resolved at runtime — no version-conditional imports.

import inspect
import logging
import os
import sys
import warnings

logger = logging.getLogger(__name__)

try:
import torch
from megatron.core import parallel_state
from megatron.core.distributed.finalize_model_grads import (
_flatten_dense_tensors,
_get_main_grad_attr,
_reshard_if_dtensor,
_unflatten_dense_tensors,
_unshard_if_dtensor,
get_attr_wrapped_model,
)

# post-core_v0.15.0rc7 dev takes (param); core_v0.13.0 line takes
# (param, use_custom_fsdp=False).
_gma_takes_fsdp_arg = len(inspect.signature(_get_main_grad_attr).parameters) >= 2

def _grad_attr(param, fsdp_on):
if _gma_takes_fsdp_arg:
return _get_main_grad_attr(param, fsdp_on)
return _get_main_grad_attr(param)

def _fsdp_flag(ddp_config):
return bool(getattr(ddp_config, "use_megatron_fsdp", False) or getattr(ddp_config, "use_custom_fsdp", False))

_chunk_bytes = int(os.environ.get("SLIME_GRAD_COALESCE_CHUNK_BYTES") or (1 << 30))

def _split_into_chunks(params, grads, target_bytes):
"""Greedy split keeping params/grads aligned. A single grad larger
than target_bytes is placed alone in its own chunk."""
chunks, cur_p, cur_g, cur_b = [], [], [], 0
for p, g in zip(params, grads, strict=False):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

Using strict=True is safer here to ensure that params and grads are perfectly aligned and of equal length. Since Megatron-LM core v0.13+ requires Python 3.10+, strict=True is fully supported and prevents potential silent bugs if the lists ever diverge in length.

Suggested change
for p, g in zip(params, grads, strict=False):
for p, g in zip(params, grads, strict=True):

gb = g.numel() * g.element_size()
if cur_g and cur_b + gb > target_bytes:
chunks.append((cur_p, cur_g))
cur_p, cur_g, cur_b = [], [], 0
cur_p.append(p)
cur_g.append(g)
cur_b += gb
if cur_g:
chunks.append((cur_p, cur_g))
return chunks

def _allreduce_non_tensor_model_parallel_grads(model, config, tp_group=None):
# post-core_v0.15.0rc7 dev passes tp_group; core_v0.13.0 line omits it.
# Default-fill from parallel_state so the same body works for both call sites.
if tp_group is None:
tp_group = parallel_state.get_tensor_model_parallel_group()
if tp_group.size() <= 1:
return

params_sum, grads_sum = [], []
params_avg, grads_avg = [], []
ddp_config = None
for model_chunk in model:
ddp_config = model_chunk.ddp_config
fsdp_on = _fsdp_flag(ddp_config)
for name, param in get_attr_wrapped_model(model_chunk, "named_parameters")():
if not param.requires_grad:
continue
if getattr(param, "average_gradients_across_tp_domain", False):
target_params, target_grads = params_avg, grads_avg
elif (config.sequence_parallel and getattr(param, "sequence_parallel", False)) or (
config.qk_layernorm and ("q_layernorm" in name or "k_layernorm" in name)
):
target_params, target_grads = params_sum, grads_sum
else:
continue

grad_attr = _grad_attr(param, fsdp_on)
grad = getattr(param, grad_attr)
if grad is None:
continue
target_params.append(param)
if fsdp_on and hasattr(grad, "_local_tensor"):
target_grads.append(grad._local_tensor.data)
else:
target_grads.append(_unshard_if_dtensor(grad).data)

for params, grads, op in (
(params_sum, grads_sum, torch.distributed.ReduceOp.SUM),
(params_avg, grads_avg, torch.distributed.ReduceOp.AVG),
):
if not grads:
continue
fsdp_on = _fsdp_flag(ddp_config)
for p_chunk, g_chunk in _split_into_chunks(params, grads, _chunk_bytes):
coalesced = _flatten_dense_tensors(g_chunk)
torch.distributed.all_reduce(coalesced, op=op, group=tp_group)
for param, buf, synced in zip(
p_chunk, g_chunk, _unflatten_dense_tensors(coalesced, g_chunk), strict=False
):
buf.copy_(synced)
grad_attr = _grad_attr(param, fsdp_on)
orig_grad = getattr(param, grad_attr)
if fsdp_on and hasattr(orig_grad, "_local_tensor"):
# buf already aliases orig_grad._local_tensor.data;
# restore original DTensor wrapper (post-rc7 dev semantics).
setattr(param, grad_attr, orig_grad)
else:
setattr(param, grad_attr, _reshard_if_dtensor(buf, orig_grad))
del coalesced
Comment on lines +113 to +125

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

PyTorch Memory Retention Issue

In Python, loop variables (such as synced, param, and buf) remain in the local scope even after the loop finishes.

Because _unflatten_dense_tensors returns views of the coalesced tensor, the loop variable synced (which holds the last view in the loop) keeps a reference to the underlying Storage of coalesced alive. As a result, even though del coalesced is called at the end of the chunk iteration, the GPU memory allocated for coalesced (up to 1 GiB by default) is not freed until synced is overwritten or deleted.

In the next iteration of the outer loop, when the next chunk's coalesced tensor is allocated, the previous chunk's memory is still active in GPU memory. This effectively doubles the peak memory overhead of the coalescing buffer, partially defeating the purpose of chunking.

Solution

Explicitly delete synced, param, and buf along with coalesced at the end of each chunk iteration to ensure the memory is immediately reclaimed by PyTorch's caching allocator. Additionally, we should use strict=True in zip to ensure that the unflattened tensors match the chunk size exactly.

Suggested change
for param, buf, synced in zip(
p_chunk, g_chunk, _unflatten_dense_tensors(coalesced, g_chunk), strict=False
):
buf.copy_(synced)
grad_attr = _grad_attr(param, fsdp_on)
orig_grad = getattr(param, grad_attr)
if fsdp_on and hasattr(orig_grad, "_local_tensor"):
# buf already aliases orig_grad._local_tensor.data;
# restore original DTensor wrapper (post-rc7 dev semantics).
setattr(param, grad_attr, orig_grad)
else:
setattr(param, grad_attr, _reshard_if_dtensor(buf, orig_grad))
del coalesced
for param, buf, synced in zip(
p_chunk, g_chunk, _unflatten_dense_tensors(coalesced, g_chunk), strict=True
):
buf.copy_(synced)
grad_attr = _grad_attr(param, fsdp_on)
orig_grad = getattr(param, grad_attr)
if fsdp_on and hasattr(orig_grad, "_local_tensor"):
# buf already aliases orig_grad._local_tensor.data;
# restore original DTensor wrapper (post-rc7 dev semantics).
setattr(param, grad_attr, orig_grad)
else:
setattr(param, grad_attr, _reshard_if_dtensor(buf, orig_grad))
del coalesced, synced, param, buf


# The parent package re-exports a same-named function, shadowing the
# submodule attribute. Pull the real module out of sys.modules to setattr.
_fmg = sys.modules["megatron.core.distributed.finalize_model_grads"]
_fmg._allreduce_non_tensor_model_parallel_grads = _allreduce_non_tensor_model_parallel_grads
_fmg._allreduce_layernorm_grads = _allreduce_non_tensor_model_parallel_grads

logger.info(
"slime grad coalesce patch applied to "
"megatron.core.distributed.finalize_model_grads."
"_allreduce_non_tensor_model_parallel_grads (chunk=%d MiB)",
_chunk_bytes // (1 << 20),
)

except ImportError as exc:
warnings.warn(
f"slime grad coalesce patch not applied — Megatron import failed ({exc!r}). "
"If this is a Megatron upgrade, the symbol layout may have changed; "
"without this patch, large-model TP grad sync may OOM.",
stacklevel=2,
)
Loading