diff --git a/miles/ray/actor_group.py b/miles/ray/actor_group.py index f1436cfafaa..bfbfdf55e85 100644 --- a/miles/ray/actor_group.py +++ b/miles/ray/actor_group.py @@ -119,9 +119,9 @@ async def init(self): "init", self.args, self.role, with_ref=self.with_ref, with_opd_teacher=self.with_opd_teacher ) - async def train(self, rollout_id, rollout_data_ref): + async def train(self, rollout_id, rollout_data_pack): """Do one rollout training""" - await self._broadcast("train", rollout_id, rollout_data_ref) + await self._broadcast("train", rollout_id, rollout_data_pack["data_ref"]) async def save_model(self, rollout_id, force_sync=False): """Save actor model""" diff --git a/miles/ray/rollout/rollout_manager.py b/miles/ray/rollout/rollout_manager.py index c216b440b24..107d771233c 100644 --- a/miles/ray/rollout/rollout_manager.py +++ b/miles/ray/rollout/rollout_manager.py @@ -28,6 +28,7 @@ from miles.utils.logging_utils import configure_logger from miles.utils.metric_checker import MetricChecker from miles.utils.misc import load_function +from miles.utils.ray_utils import Box from miles.utils.tracking_utils import init_tracking logging.getLogger("httpx").setLevel(logging.WARNING) @@ -117,7 +118,12 @@ async def generate(self, rollout_id): custom_convert_samples_to_train_data_func=self.custom_convert_samples_to_train_data_func, custom_reward_post_process_func=self.custom_reward_post_process_func, ) - return split_train_data_by_dp(self.args, data, self.train_parallel_config["dp_size"]) + sample_indices = data.get("sample_indices") + if self.args.delay_split_train_data_by_dp: + data_ref = Box(ray.put(data)) + else: + data_ref = split_train_data_by_dp(self.args, data, self.train_parallel_config["dp_size"]) + return dict(sample_indices=sample_indices, data_ref=data_ref) async def eval(self, rollout_id): if self.args.debug_train_only: diff --git a/miles/ray/rollout/train_data_conversion.py b/miles/ray/rollout/train_data_conversion.py index 65bc8d4b6db..e67087fbb3d 100644 --- a/miles/ray/rollout/train_data_conversion.py +++ b/miles/ray/rollout/train_data_conversion.py @@ -122,6 +122,12 @@ def _post_process_rewards(args, samples: list[Sample] | list[list[Sample]], cust def split_train_data_by_dp(args, data, dp_size): + """Split the train data by data parallel size.""" + rollout_data_list = split_train_data_by_dp_raw(args, data, dp_size=dp_size) + return [Box(ray.put(rollout_data)) for rollout_data in rollout_data_list] + + +def split_train_data_by_dp_raw(args, data: dict[str, Any], *, dp_size: int) -> list[dict[str, Any]]: """Split the train data by data parallel size.""" rollout_data = {} @@ -136,7 +142,7 @@ def split_train_data_by_dp(args, data, dp_size): else: partitions = [range(i, len(total_lengths), dp_size) for i in range(dp_size)] - rollout_data_refs = [] + ans = [] for i in range(dp_size): rollout_data = {} @@ -157,6 +163,7 @@ def split_train_data_by_dp(args, data, dp_size): "prompt", "teacher_log_probs", "opd_reverse_kl", + "seq_witness_ids", "weight_versions", ]: if key not in data: @@ -172,5 +179,5 @@ def split_train_data_by_dp(args, data, dp_size): if key not in data: continue rollout_data[key] = data[key] - rollout_data_refs.append(Box(ray.put(rollout_data))) - return rollout_data_refs + ans.append(rollout_data) + return ans diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index 49d99898c76..7ff03b77878 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -270,6 +270,11 @@ def add_train_arguments(parser): parser.add_argument( "--log-probs-chunk-size", type=int, default=-1, help="Chunk size to compute log probs to save memory" ) + parser.add_argument( + "--delay-split-train-data-by-dp", + action="store_true", + default=False, + ) parser.add_argument( "--allgather-cp", action="store_true", diff --git a/miles/utils/data.py b/miles/utils/data.py index 0fe7ad483fd..21d01db972b 100644 --- a/miles/utils/data.py +++ b/miles/utils/data.py @@ -8,6 +8,8 @@ import numpy as np import ray +from miles.ray.rollout.train_data_conversion import split_train_data_by_dp_raw + try: import pyarrow.parquet as pq except ImportError: @@ -272,9 +274,19 @@ def get_minimum_num_micro_batch_size(total_lengths, max_tokens_per_gpu): return len(batches) -def process_rollout_data(args, rollout_data_ref, dp_rank, dp_size): - assert len(rollout_data_ref) == dp_size - rollout_data = ray.get(rollout_data_ref[dp_rank].inner) +def process_rollout_data( + args, + rollout_data_ref, + dp_rank, + dp_size, +): + if args.delay_split_train_data_by_dp: + raw = ray.get(rollout_data_ref.inner) + raw = split_train_data_by_dp_raw(args, raw, dp_size=dp_size) + rollout_data = raw[dp_rank] + else: + assert len(rollout_data_ref) == dp_size + rollout_data = ray.get(rollout_data_ref[dp_rank].inner) partition = rollout_data.pop("partition") total_lengths = rollout_data["total_lengths"] diff --git a/tests/fast/ray/rollout/real_ray/test_rollout_manager.py b/tests/fast/ray/rollout/real_ray/test_rollout_manager.py index ab7d79b75df..962a53099dc 100644 --- a/tests/fast/ray/rollout/real_ray/test_rollout_manager.py +++ b/tests/fast/ray/rollout/real_ray/test_rollout_manager.py @@ -516,9 +516,12 @@ def fake_rollout_fn(input): assert len(captured) == 1 assert isinstance(captured[0], RolloutFnTrainInput) assert captured[0].rollout_id == 42 + # generate returns {"sample_indices": ..., "data_ref": ...}; # split_train_data_by_dp returns Box(ObjectRef) per dp rank - assert len(result) == 2 - partitions = ray.get([box.inner for box in result]) + assert set(result) == {"sample_indices", "data_ref"} + data_refs = result["data_ref"] + assert len(data_refs) == 2 + partitions = ray.get([box.inner for box in data_refs]) for partition in partitions: assert "tokens" in partition assert "rewards" in partition diff --git a/tests/fast/ray/rollout/test_train_data_conversion.py b/tests/fast/ray/rollout/test_train_data_conversion.py index 4a3da1a2883..31f3b1975db 100644 --- a/tests/fast/ray/rollout/test_train_data_conversion.py +++ b/tests/fast/ray/rollout/test_train_data_conversion.py @@ -1,7 +1,10 @@ from __future__ import annotations +from unittest.mock import MagicMock + import pytest import ray +import torch from hypothesis import HealthCheck, given, settings from hypothesis import strategies as st from tests.fast.ray.rollout.conftest import make_args, make_sample, make_samples_grouped @@ -10,6 +13,7 @@ _post_process_rewards, convert_samples_to_train_data, split_train_data_by_dp, + split_train_data_by_dp_raw, ) from miles.utils.types import Sample @@ -472,3 +476,66 @@ def test_partition_indices_form_a_partition(self): parts = [ray.get(r.inner) for r in refs] all_indices = sorted(i for p in parts for i in p["partition"]) assert all_indices == list(range(n)) + + +class TestSplitTrainDataRaw: + def test_witness_ids_split_across_dp(self) -> None: + tokens = [[1, 2, 3], [4, 5], [6, 7, 8, 9], [10, 11]] + witness_ids = [ + torch.tensor([0, 0, 0]), + torch.tensor([1, 1]), + torch.tensor([2, 2, 2, 2]), + torch.tensor([3, 3]), + ] + + data = { + "tokens": tokens, + "seq_witness_ids": witness_ids, + "response_lengths": [1, 1, 1, 1], + "loss_masks": [[0, 0, 1], [0, 1], [0, 0, 0, 1], [0, 1]], + } + + args = MagicMock() + args.balance_data = False + + result = split_train_data_by_dp_raw(args, data, dp_size=2) + + assert len(result) == 2 + assert "seq_witness_ids" in result[0] + assert "seq_witness_ids" in result[1] + assert len(result[0]["seq_witness_ids"]) == 2 + assert len(result[1]["seq_witness_ids"]) == 2 + + def test_indexer_topk_and_opd_reverse_kl_split_across_dp(self) -> None: + """Keys from the rollout-side split (rollout_indexer_topk, opd_reverse_kl) partition per sample.""" + data = { + "tokens": [[1, 2], [3, 4], [5, 6], [7, 8]], + "response_lengths": [1, 1, 1, 1], + "loss_masks": [[0, 1], [0, 1], [0, 1], [0, 1]], + "rollout_indexer_topk": [torch.tensor([i]) for i in range(4)], + "opd_reverse_kl": [[float(i)] for i in range(4)], + } + + args = MagicMock() + args.balance_data = False + + result = split_train_data_by_dp_raw(args, data, dp_size=2) + + assert len(result) == 2 + for part in result: + assert len(part["rollout_indexer_topk"]) == 2 + assert len(part["opd_reverse_kl"]) == 2 + + def test_no_witness_ids_when_absent(self) -> None: + tokens = [[1, 2], [3, 4]] + data = { + "tokens": tokens, + "response_lengths": [1, 1], + "loss_masks": [[0, 1], [0, 1]], + } + + args = MagicMock() + args.balance_data = False + + result = split_train_data_by_dp_raw(args, data, dp_size=1) + assert "seq_witness_ids" not in result[0]