From d9ac4146f6b5df0389f7a5f25725ac4b889389ab Mon Sep 17 00:00:00 2001 From: Neil Kale <263453039+kalectory@users.noreply.github.com> Date: Sun, 4 Oct 2026 15:50:09 +0000 Subject: [PATCH] Avoid materializing unused Megatron forward outputs Signed-off-by: Neil Kale <263453039+kalectory@users.noreply.github.com> --- .../workers/megatron/megatron_worker.py | 5 +++ .../test_forward_collection_outputs.py | 32 +++++++++++++++++++ 2 files changed, 37 insertions(+) create mode 100644 tests/backends/skyrl_train/workers/megatron/test_forward_collection_outputs.py diff --git a/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py b/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py index 2c5032c582..23a0390b80 100644 --- a/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py +++ b/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py @@ -988,6 +988,9 @@ def forward( # micro-batching (when `max_tokens_per_microbatch > 0`) is handled inside # `_forward_logprobs`, which also reorders back to the original sample order. log_probs = self._forward_logprobs(data) + # All ranks must finish the forward collectives; only collection ranks return token arrays. + if not self.mesh_rank.is_collection_dp_rank(): + return WorkerOutput() loss_fn_outputs = [{"logprobs": log_probs[i].tolist()} for i in range(log_probs.shape[0])] return WorkerOutput(loss_fn_outputs=loss_fn_outputs, metrics={}) @@ -1866,6 +1869,8 @@ def forward(self, data: TrainingInputBatch) -> WorkerOutput: ``max_tokens_per_microbatch > 0``) is handled inside ``_forward_logprobs``. """ log_probs = self._forward_logprobs(data) + if not self.mesh_rank.is_collection_dp_rank(): + return WorkerOutput() loss_fn_outputs = [{"logprobs": log_probs[i].tolist()} for i in range(log_probs.shape[0])] return WorkerOutput(loss_fn_outputs=loss_fn_outputs, metrics={}) diff --git a/tests/backends/skyrl_train/workers/megatron/test_forward_collection_outputs.py b/tests/backends/skyrl_train/workers/megatron/test_forward_collection_outputs.py new file mode 100644 index 0000000000..553dc4eb54 --- /dev/null +++ b/tests/backends/skyrl_train/workers/megatron/test_forward_collection_outputs.py @@ -0,0 +1,32 @@ +"""Only collection ranks materialize token outputs after the distributed forward.""" + +from unittest.mock import Mock + +import pytest +import torch + +pytest.importorskip("megatron.core") + +from skyrl.backends.skyrl_train.distributed.dispatch import MeshRank +from skyrl.backends.skyrl_train.workers.megatron.megatron_worker import ( + MegatronPolicyWorkerBase, + MegatronRefWorkerBase, +) + + +@pytest.mark.parametrize("worker_cls", [MegatronPolicyWorkerBase, MegatronRefWorkerBase]) +@pytest.mark.parametrize("sp,tp,pp", [(0, 0, 1), (1, 0, 1), (0, 1, 1), (0, 0, 0)]) +def test_forward_keeps_computation_but_only_collects_selected_rank(worker_cls, sp, tp, pp): + worker = worker_cls.__new__(worker_cls) + worker.mesh_rank = MeshRank(dp=0, sp=sp, tp=tp, pp=pp, world_size=8, dp_size=1, pp_size=2) + collects = worker.mesh_rank.is_collection_dp_rank() + # A non-collection rank must not even inspect/materialize its returned tensor. + values = torch.tensor([[-0.5, -1.0]]) if collects else object() + worker._forward_logprobs = Mock(return_value=values) + data = object() + + output = worker.forward(data) + + worker._forward_logprobs.assert_called_once_with(data) + assert output.loss_fn_outputs == ([{"logprobs": [-0.5, -1.0]}] if collects else []) + assert output.metrics == {}