Skip to content
Draft
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
35 changes: 35 additions & 0 deletions skyrl/backends/skyrl_train/distributed/megatron/grad_sync.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,35 @@
"""Gradient synchronization at SkyRL's optimizer-step boundary."""

from contextlib import ExitStack, contextmanager


@contextmanager
def defer_grad_sync(model_chunks):
"""Accumulate locally, including across multiple pipeline schedule calls.

Megatron's schedules normally enable DDP synchronization for their last
microbatch. SkyRL's optimizer window can contain several schedule calls,
so only ``optim_step`` knows when the accumulated gradients are complete.
Clear the schedule callbacks while owning the DDP no-sync contexts: nesting
the schedule's own DDP no-sync would re-enable hooks when its context exits.
"""
with ExitStack() as stack:
for chunk in model_chunks:
config = chunk.config
for name in ("no_sync_func", "grad_sync_func"):
stack.callback(setattr, config, name, getattr(config, name))
setattr(config, name, None)
stack.enter_context(chunk.no_sync())
yield


def start_deferred_grad_sync(model_chunks):
"""Dispatch async reductions once the whole optimizer window is complete.

Non-overlap DDP dispatches synchronously from ``finalize_model_grads``.
Overlap DDP normally dispatches from backward hooks, which were suppressed
during accumulation, so it needs an explicit start before finalization.
"""
for chunk in model_chunks:
if chunk.ddp_config.overlap_grad_reduce:
chunk.start_grad_sync()
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
from contextlib import nullcontext
from dataclasses import asdict
from functools import partial
from typing import Any, Callable, Dict, List, Optional
Expand All @@ -13,6 +14,10 @@
call_model_with_fused_lm_head,
fused_lm_head_output_processor,
)
from skyrl.backends.skyrl_train.distributed.megatron.grad_sync import (
defer_grad_sync,
start_deferred_grad_sync,
)
from skyrl.backends.skyrl_train.distributed.megatron.megatron_utils import (
get_model_config,
make_batch_generator,
Expand Down Expand Up @@ -234,6 +239,7 @@ def run_pending_grad_sync(self) -> None:
"""
pending = self._pending_grad_sync
self._pending_grad_sync = None
start_deferred_grad_sync(self.actor_module)
finalize_model_grads(self.actor_module, pending["num_tokens"] if pending else None)

def train(self):
Expand Down Expand Up @@ -1156,7 +1162,8 @@ def depad(tensor):
batch_generator = make_batch_generator(micro_batches, vpp_size=len(self.actor_module))

replay_enabled = any(batch["rollout_expert_indices"] is not None for batch in micro_batches)
with router_replay_schedule(replay_enabled):
grad_sync_context = nullcontext() if forward_only else defer_grad_sync(self.actor_module)
with router_replay_schedule(replay_enabled), grad_sync_context:
metrics_list = forward_backward_func(
forward_step_func=forward_step,
data_iterator=batch_generator,
Expand Down
2 changes: 2 additions & 0 deletions skyrl/train/config/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -254,6 +254,8 @@ class FSDPConfig(BaseConfig):
class MegatronDDPConfig(BaseConfig):
grad_reduce_in_fp32: bool = True
overlap_grad_reduce: bool = False
"""Use asynchronous gradient reductions. SkyRL dispatches them at ``optim_step``
after all accumulated forward/backward requests, rather than during backward."""
overlap_param_gather: bool = False
average_in_collective: bool = True

Expand Down
78 changes: 78 additions & 0 deletions tests/backends/skyrl_train/distributed/test_megatron_grad_sync.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,78 @@
"""Optimizer-window ownership of Megatron gradient synchronization."""

from contextlib import contextmanager, nullcontext
from types import SimpleNamespace
from unittest.mock import Mock

import pytest

from skyrl.backends.skyrl_train.distributed.megatron.grad_sync import (
defer_grad_sync,
start_deferred_grad_sync,
)


class DDPChunk:
"""DDP's public no-sync contract with observable synchronization calls."""

def __init__(self, overlap=True):
self.hooks_enabled = True
self.ddp_config = SimpleNamespace(overlap_grad_reduce=overlap)
self.start_grad_sync = Mock()
self.config = SimpleNamespace(no_sync_func=self.no_sync, grad_sync_func=self.start_grad_sync)

@contextmanager
def no_sync(self):
self.hooks_enabled = False
try:
yield
finally:
self.hooks_enabled = True


@pytest.mark.parametrize("request_sizes", [(2,), (1,), (3,), (2, 1), (1, 2, 1)])
@pytest.mark.parametrize("shared_config", [False, True])
def test_gradients_stay_local_until_optimizer_step(request_sizes, shared_config):
chunks = [DDPChunk(), DDPChunk()]
if shared_config:
chunks[1].config = chunks[0].config
original_callbacks = [(c.config.no_sync_func, c.config.grad_sync_func) for c in chunks]

for microbatches in request_sizes:
with defer_grad_sync(chunks):
for chunk in chunks:
# A pipeline schedule exits its no-sync context before its last
# backward. That exit must not re-enable the chunk's DDP hooks.
callback = chunk.config.no_sync_func or nullcontext
with callback():
for _ in range(microbatches - 1):
assert not chunk.hooks_enabled
assert not chunk.hooks_enabled
assert chunk.config.grad_sync_func is None
for chunk, callbacks in zip(chunks, original_callbacks):
assert chunk.hooks_enabled
assert (chunk.config.no_sync_func, chunk.config.grad_sync_func) == callbacks
chunk.start_grad_sync.assert_not_called()

start_deferred_grad_sync(chunks)
for chunk in chunks:
chunk.start_grad_sync.assert_called_once_with()


def test_no_sync_and_callbacks_are_restored_after_exception():
chunks = [DDPChunk(), DDPChunk()]
original_callbacks = [(c.config.no_sync_func, c.config.grad_sync_func) for c in chunks]
with pytest.raises(RuntimeError, match="backward failed"):
with defer_grad_sync(chunks):
raise RuntimeError("backward failed")
for chunk, callbacks in zip(chunks, original_callbacks):
assert chunk.hooks_enabled
assert (chunk.config.no_sync_func, chunk.config.grad_sync_func) == callbacks
chunk.start_grad_sync.assert_not_called()


def test_non_overlap_reduction_is_left_to_finalize():
chunks = [DDPChunk(overlap=False), DDPChunk(overlap=True)]
start_deferred_grad_sync(chunks)
chunks[0].start_grad_sync.assert_not_called()
chunks[1].start_grad_sync.assert_called_once_with()
Original file line number Diff line number Diff line change
@@ -0,0 +1,55 @@
"""Compare asynchronous and synchronous DDP across changing optimizer windows.

Requires four GPUs: two independent DP=2 policy groups, with real Megatron
DDP buffers, distributed optimizers, and SkyRL forward/backward schedules.
"""

import pytest
import ray
import torch

from skyrl.backends.skyrl_train.distributed.dispatch import (
WorkerOutput,
loss_fn_outputs_to_tensor,
)
from tests.backends.skyrl_train.gpu.gpu_ci.megatron.test_megatron_worker import (
get_test_actor_config,
get_test_training_batch,
)
from tests.backends.skyrl_train.gpu.utils import init_worker_with_type


@pytest.mark.megatron
def test_overlap_matches_synchronous_accumulation(ray_init_fixture):
groups = []
for overlap in (False, True):
cfg = get_test_actor_config()
cfg.trainer.strategy = "megatron"
cfg.trainer.placement.colocate_all = False
cfg.trainer.placement.policy_num_gpus_per_node = 2
cfg.trainer.policy.megatron_config.ddp_config.overlap_grad_reduce = overlap
cfg.trainer.train_batch_size = 16
cfg.trainer.policy_mini_batch_size = 16
cfg.trainer.micro_train_batch_size_per_gpu = 2
cfg.generator.n_samples_per_prompt = 1
groups.append(init_worker_with_type("policy", num_gpus_per_node=2, cfg=cfg))

# Two microbatches per rank initially, then fewer, then more. The final
# window also spans two forward_backward requests before one optimizer step.
for step, request_sizes in enumerate(((8,), (4,), (12,), (4, 8))):
for batch_size in request_sizes:
batch = get_test_training_batch(batch_size)
batch.metadata["global_step"] = step
for group in groups:
ray.get(group.async_run_ray_method("mesh", "forward_backward", batch, loss_fn="cross_entropy"))
norms = [ray.get(group.async_run_ray_method("pass_through", "optim_step")) for group in groups]
assert all(norm is not None and norm > 0 for rank_norms in norms for norm in rank_norms)
torch.testing.assert_close(torch.tensor(norms[0]), torch.tensor(norms[1]), rtol=1e-3, atol=1e-5)

batch = get_test_training_batch(4)
logprobs = []
for group in groups:
outputs = ray.get(group.async_run_ray_method("mesh", "forward", batch))
combined = WorkerOutput.cat(group.actor_infos, outputs)
logprobs.append(loss_fn_outputs_to_tensor(combined.loss_fn_outputs, key="logprobs"))
torch.testing.assert_close(logprobs[0], logprobs[1], rtol=1e-3, atol=1e-3)
Loading