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
4 changes: 2 additions & 2 deletions miles/ray/actor_group.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"""
Expand Down
8 changes: 7 additions & 1 deletion miles/ray/rollout/rollout_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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:
Expand Down
13 changes: 10 additions & 3 deletions miles/ray/rollout/train_data_conversion.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = {}

Expand All @@ -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 = {}
Expand All @@ -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:
Expand All @@ -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
5 changes: 5 additions & 0 deletions miles/utils/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
18 changes: 15 additions & 3 deletions miles/utils/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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,

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.

I think it'd be better to add type annotation for 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"]
Expand Down
7 changes: 5 additions & 2 deletions tests/fast/ray/rollout/real_ray/test_rollout_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
67 changes: 67 additions & 0 deletions tests/fast/ray/rollout/test_train_data_conversion.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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

Expand Down Expand Up @@ -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]
Loading