Skip to content
Merged
31 changes: 31 additions & 0 deletions miles/ray/rollout/rollout_data_conversion.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,23 @@
import itertools
import logging

from miles.utils.multi_lora import is_multi_lora_enabled

logger = logging.getLogger(__name__)


def postprocess_rollout_data(args, data, train_parallel_config):
metadata = {}

# Multi-LoRA: record group boundaries (heterogeneous per-adapter group sizes)
# and lift the collection loop's batch-level step decision out of sample metadata,
# both before flattening.
if is_multi_lora_enabled(args) and isinstance(data[0], list):
metadata["prompt_group_sizes"] = [_nested_sample_count(group) for group in data]
head = _first_sample(data[0])
metadata["step_slots"] = list(head.metadata.pop("step_slots", []))
metadata["step_adapter_names"] = list(head.metadata.pop("step_adapter_names", []))

# flatten the data if it is a list of lists
while isinstance(data[0], list):
data = list(itertools.chain.from_iterable(data))
Expand All @@ -34,6 +44,16 @@ def postprocess_rollout_data(args, data, train_parallel_config):
return data, metadata


def _first_sample(group):
return _first_sample(group[0]) if isinstance(group[0], list) else group[0]


def _nested_sample_count(group) -> int:
if not isinstance(group, list):
return 1
return sum(_nested_sample_count(item) for item in group)


def _compute_dynamic_global_batch_size(args, train_parallel_config, num_samples: int) -> int:
"""Calculate dynamic global_batch_size to ensure only one training step.

Expand All @@ -43,6 +63,17 @@ def _compute_dynamic_global_batch_size(args, train_parallel_config, num_samples:
dp_size = train_parallel_config["dp_size"]
original_gbs = args.global_batch_size

if is_multi_lora_enabled(args):
# Batches take groups in multiples of each adapter's
# min_groups_per_dp_split, so this holds by construction; a violation
# means a generate fn's group shape broke the invariant.
if num_samples % dp_size != 0:
raise ValueError(
f"Multi-LoRA batch of {num_samples} samples is not divisible by dp_size={dp_size}; "
"the min_groups_per_dp_split invariant was violated (variable-size generate fn output?)"
)
return num_samples

# Round down to a multiple of dp_size to ensure only one training step
dynamic_gbs = (num_samples // dp_size) * dp_size

Expand Down
5 changes: 5 additions & 0 deletions miles/ray/rollout/rollout_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,7 +98,12 @@ def __init__(self, args, pg):
# -------------------------- lifecycle -----------------------------
# TODO: may have a `async def init` here later

def get_router_address(self) -> tuple[str, int]:
return self.args.sglang_router_ip, self.args.sglang_router_port

def dispose(self):
if (close := getattr(self.data_source, "close", None)) is not None:
close()
event_analyzer.run_analysis_from_args(self.args)
if self._metric_checker is not None:
self._metric_checker.dispose()
Expand Down
60 changes: 58 additions & 2 deletions miles/ray/rollout/train_data_conversion.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,10 @@ def convert_samples_to_train_data(
return f(args, samples)

raw_rewards, rewards = _post_process_rewards(
args, samples, custom_reward_post_process_func=custom_reward_post_process_func
args,
samples,
custom_reward_post_process_func=custom_reward_post_process_func,
prompt_group_sizes=metadata.get("prompt_group_sizes"),
)

assert len(raw_rewards) == len(samples)
Expand Down Expand Up @@ -85,6 +88,24 @@ def convert_samples_to_train_data(
if samples[0].teacher_log_probs is not None:
train_data["teacher_log_probs"] = [sample.teacher_log_probs for sample in samples]

if any(sample.adapter is not None for sample in samples):
assert all(sample.adapter is not None for sample in samples), "Cannot mix adapter and adapter-less samples"
train_data["adapter_slots"] = [sample.adapter.slot for sample in samples]
# Slots whose adapter batch completes with this batch: the trainer scales their
# accumulated gradients by 1/adapter-batch-size and advances the LR schedule.
step_slots = sorted(metadata.get("step_slots", []))
train_data["step_slots"] = step_slots
train_data["step_adapter_names"] = sorted(metadata.get("step_adapter_names", []))
step_slot_set = set(step_slots)
train_data["step_adapter_batch_sizes"] = {
sample.adapter.slot: sample.metadata["adapter_global_batch_size"]
for sample in samples
if sample.adapter.slot in step_slot_set
}

if (prompt_group_sizes := metadata.get("prompt_group_sizes")) is not None:
train_data["prompt_group_sizes"] = prompt_group_sizes

if samples[0].opd_reverse_kl is not None:
train_data["opd_reverse_kl"] = [sample.opd_reverse_kl for sample in samples]

Expand All @@ -96,14 +117,36 @@ def convert_samples_to_train_data(
return train_data


def _post_process_rewards(args, samples: list[Sample] | list[list[Sample]], custom_reward_post_process_func):
def _post_process_rewards(
args,
samples: list[Sample] | list[list[Sample]],
custom_reward_post_process_func,
prompt_group_sizes: list[int] | None = None,
):
if (f := custom_reward_post_process_func) is not None:
return f(args, samples)

raw_rewards = [sample.get_reward_value(args) for sample in samples]
if args.advantage_estimator in ["grpo", "gspo", "reinforce_plus_plus_baseline"] and args.rewards_normalization:
# group norm
rewards = torch.tensor(raw_rewards, dtype=torch.float)
if prompt_group_sizes is not None:
# Multi-LoRA: groups may have heterogeneous sizes (per-adapter
# n_samples_per_prompt), so normalize within explicit boundaries.
assert sum(prompt_group_sizes) == len(
raw_rewards
), f"prompt group sizes sum to {sum(prompt_group_sizes)}, but got {len(raw_rewards)} rewards"
normalized_groups = []
for group_rewards in rewards.split(prompt_group_sizes):
centered = group_rewards - group_rewards.mean()
if (
args.advantage_estimator in ["grpo", "gspo"]
and args.grpo_std_normalization
and group_rewards.numel() > 1
):
centered = centered / (group_rewards.std() + 1e-6)
normalized_groups.append(centered)
return raw_rewards, torch.cat(normalized_groups).tolist()
if rewards.shape[-1] == args.n_samples_per_prompt * args.rollout_batch_size:
rewards = rewards.reshape(-1, args.n_samples_per_prompt)
else:
Expand Down Expand Up @@ -137,6 +180,12 @@ def split_train_data_by_dp_raw(args, data: dict[str, Any], *, dp_size: int) -> l
else:
partitions = [range(i, len(total_lengths), dp_size) for i in range(dp_size)]

# Multi-LoRA: sort partitions by adapter slot so each microbatch is
# contiguous-by-slot (required by the per-adapter token-count math).
adapter_slots = data.get("adapter_slots")
if adapter_slots is not None:
partitions = [sorted(p, key=lambda i: adapter_slots[i]) for p in partitions]

ans = []

for i in range(dp_size):
Expand All @@ -160,6 +209,7 @@ def split_train_data_by_dp_raw(args, data: dict[str, Any], *, dp_size: int) -> l
"opd_reverse_kl",
"seq_witness_ids",
"weight_versions",
"adapter_slots",
]:
if key not in data:
continue
Expand All @@ -170,9 +220,15 @@ def split_train_data_by_dp_raw(args, data: dict[str, Any], *, dp_size: int) -> l
"raw_reward",
"total_lengths",
"dynamic_global_batch_size",
"step_slots",
"step_adapter_names",
"step_adapter_batch_sizes",
"prompt_group_sizes",
]:
if key not in data:
continue
rollout_data[key] = data[key]
if "adapter_slots" in rollout_data:
rollout_data["n_adapters"] = args.multi_lora_n_adapters
ans.append(rollout_data)
return ans
Empty file.
Loading
Loading