diff --git a/miles/ray/rollout/rollout_data_conversion.py b/miles/ray/rollout/rollout_data_conversion.py index 539730bd43..9b9b06c7ec 100644 --- a/miles/ray/rollout/rollout_data_conversion.py +++ b/miles/ray/rollout/rollout_data_conversion.py @@ -1,6 +1,7 @@ import itertools import logging +from miles.utils.multi_lora import is_multi_lora_enabled logger = logging.getLogger(__name__) @@ -8,6 +9,15 @@ 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)) @@ -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. @@ -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 diff --git a/miles/ray/rollout/rollout_manager.py b/miles/ray/rollout/rollout_manager.py index ecaa9f41e5..2907fa4157 100644 --- a/miles/ray/rollout/rollout_manager.py +++ b/miles/ray/rollout/rollout_manager.py @@ -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() diff --git a/miles/ray/rollout/train_data_conversion.py b/miles/ray/rollout/train_data_conversion.py index cf80d7890a..5921fc42cc 100644 --- a/miles/ray/rollout/train_data_conversion.py +++ b/miles/ray/rollout/train_data_conversion.py @@ -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) @@ -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] @@ -96,7 +117,12 @@ 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) @@ -104,6 +130,23 @@ def _post_process_rewards(args, samples: list[Sample] | list[list[Sample]], cust 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: @@ -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): @@ -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 @@ -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 diff --git a/miles/rollout/multi_lora/__init__.py b/miles/rollout/multi_lora/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/miles/rollout/multi_lora/async_rollout.py b/miles/rollout/multi_lora/async_rollout.py new file mode 100644 index 0000000000..21bb36179c --- /dev/null +++ b/miles/rollout/multi_lora/async_rollout.py @@ -0,0 +1,584 @@ +"""Fully-async multi-LoRA rollout: a background producer fills per-adapter buffers; batches are collected +round-robin in ``min_groups_per_dp_split`` multiples without overshooting any adapter's remaining batch.""" + +import asyncio +import itertools +import logging +import threading +import time +from collections import defaultdict, deque +from collections.abc import Callable +from dataclasses import dataclass +from typing import Any + +from miles.ray.multi_lora.controller import AdaptersCache, get_multi_lora_controller +from miles.rollout.base_types import RolloutFnTrainOutput +from miles.rollout.filter_hub.base_types import call_dynamic_filter +from miles.rollout.generate_utils.prefill_logprobs import recompute_samples_rollout_logprobs_via_prefill +from miles.rollout.sglang_rollout import GenerateState, generate_and_rm_group, get_model_url +from miles.utils.async_utils import run +from miles.utils.metric_utils import compute_statistics, dict_add_prefix +from miles.utils.misc import load_function +from miles.utils.multi_lora import EmptyBatchTimeoutError, min_groups_per_dp_split +from miles.utils.tracking_utils import tracking +from miles.utils.types import Sample + +logger = logging.getLogger(__name__) + +GenerateFn = Callable[..., Any] + +# Generate fns may return several samples per rollout; the manager flattens later. +Group = list[Sample | list[Sample]] + + +def iter_group_samples(group: Group): + return itertools.chain.from_iterable(item if isinstance(item, list) else (item,) for item in group) + + +def first_sample(group: Group) -> Sample: + return group[0][0] if isinstance(group[0], list) else group[0] + + +def group_adapter_name(group: Group) -> str | None: + head = first_sample(group) if group else None + return head.adapter.name if head is not None and head.adapter else None + + +def group_sample_count(group: Group) -> int: + return sum(1 for _ in iter_group_samples(group)) + + +# Safety valve, same convention as fully_async's queue.Queue(maxsize=1000): +# never hit in practice, just bounds memory if training stalls entirely. +MAX_BUFFERED_GROUPS = 1000 +EMPTY_BATCH_TIMEOUT_S = 30.0 + + +class GroupBuffer: + """One adapter's FIFO of completed prompt groups; bounded — the oldest group is dropped when full.""" + + def __init__(self) -> None: + self._groups: deque[Group] = deque(maxlen=MAX_BUFFERED_GROUPS) + + def __len__(self) -> int: + return len(self._groups) + + def put(self, group: Group) -> None: + self._groups.append(group) + + def get(self, n_groups: int) -> list[Group]: + """Remove and return the n oldest groups (queue.Queue-style API).""" + return [self._groups.popleft() for _ in range(n_groups)] + + def drop_foreign(self, registration_id: str) -> int: + """Drop groups stamped by a different registration of this adapter + name: an in-flight generation of a retired tenant can land after the + buffer was reset for a same-name re-registration. Unstamped groups + (no adapter view at submission time) are kept. Returns the drop count.""" + if not self._groups: + return 0 + kept: deque[Group] = deque(maxlen=MAX_BUFFERED_GROUPS) + dropped = 0 + for group in self._groups: + stamped = first_sample(group).metadata.get("registration_id") + if stamped is not None and stamped != registration_id: + dropped += 1 + else: + kept.append(group) + self._groups = kept + return dropped + + def drop_stale(self, current_version: int, max_staleness: int | None) -> list[int]: + """Drop groups generated too many weight versions ago; returns the + staleness of each dropped group (for metrics).""" + if max_staleness is None or not self._groups: + return [] + kept: deque[Group] = deque(maxlen=MAX_BUFFERED_GROUPS) + dropped: list[int] = [] + for group in self._groups: + stamped = first_sample(group).metadata.get("slot_version") + staleness = current_version - stamped if stamped is not None else 0 + if stamped is not None and staleness > max_staleness: + for sample in iter_group_samples(group): + sample.reset_for_retry() + dropped.append(staleness) + else: + kept.append(group) + self._groups = kept + return dropped + + +@dataclass +class TrainBatch: + """One train batch: the groups for one train call, with its per-adapter bookkeeping.""" + + groups: list[Group] + group_counts: dict[str, int] # prompt groups per adapter in this batch + step_names: list[str] # adapters whose adapter batch completes -> they step + step_slots: list[int] + + +def remaining_groups(adapter) -> int: + """Groups still needed to complete the adapter's batch.""" + remaining = adapter.config.rollout_batch_size - adapter.accumulated_groups + assert remaining > 0, ( + f"adapter '{adapter.name}' accumulated_groups={adapter.accumulated_groups} >= " + f"rollout_batch_size={adapter.config.rollout_batch_size}; batch accounting drifted" + ) + return remaining + + +async def process_group( + args, group: list[Sample], sampling_params: dict, generate_fn: GenerateFn, data_source +) -> Group | None: + """Generate a group; returns None for aborted groups. The slot version is + stamped at submission time (what the staleness filter compares against).""" + adapter_name = group[0].adapter.name if group and group[0].adapter else None + submission_version: int | None = None + submission_registration: str | None = None + if adapter_name is not None: + adapter = await AdaptersCache().get(adapter_name) + submission_version = adapter.version if adapter is not None else None + submission_registration = adapter.registration_id if adapter is not None else None + + if submission_version is not None: + for s in group: + s.metadata["slot_version"] = submission_version + s.metadata["registration_id"] = submission_registration + + result = await generate_fn(args, group, sampling_params) + + if submission_version is not None: + for s in iter_group_samples(result): + s.metadata["slot_version"] = submission_version + s.metadata["registration_id"] = submission_registration + + if any(s.status == Sample.Status.ABORTED for s in iter_group_samples(result)): + for s in iter_group_samples(result): + s.reset_for_retry() + # Re-queuing is not wired up (the per-adapter source is read-only). + return None + return result + + +class MultiLoRAWorkerMetrics: + """Cross-batch metric state; locked because the producer thread records while the trainer thread drains.""" + + def __init__(self) -> None: + self.lock = threading.Lock() + self.dynamic_filter_drop_counts: dict[str, int] = defaultdict(int) + # Staleness of dropped groups per adapter, drained every batch. + self.staleness_values: dict[str, list[int]] = defaultdict(list) + # Per-adapter shipped-sample values, flushed as step statistics when the adapter steps. + self.step_rewards: dict[str, list[float]] = defaultdict(list) + self.step_response_lens: dict[str, list[float]] = defaultdict(list) + # Per-sample mean engine log prob (rough per-adapter entropy trend). + self.step_log_prob_means: dict[str, list[float]] = defaultdict(list) + # Group outcomes for zero-std rates: shipped group counts and each uniform-reward group's reward. + self.step_group_counts: dict[str, int] = defaultdict(int) + self.step_zero_std_rewards: dict[str, list[float]] = defaultdict(list) + + def record_dynamic_filter_drop(self, reason: str) -> None: + with self.lock: + self.dynamic_filter_drop_counts[reason] += 1 + + def record_stale_drops(self, name: str, staleness_values: list[int]) -> None: + with self.lock: + self.staleness_values[name] += staleness_values + + def pop_stale_drops(self) -> dict[str, list[int]]: + """Drain the staleness values of groups dropped since the last batch.""" + with self.lock: + drained = dict(self.staleness_values) + self.staleness_values.clear() + return drained + + def record_shipped_samples( + self, args, data: list[Group], step_names: list[str], adapters: dict + ) -> dict[str, dict[str, float]]: + """Accumulate shipped rewards/response lengths per adapter; flush whole-adapter-batch statistics + for adapters stepping with this batch. Returns {adapter name: flushed metrics}.""" + with self.lock: + for group in data: + name = group_adapter_name(group) + if name is None: + continue + group_rewards = [] + for sample in iter_group_samples(group): + reward = sample.get_reward_value(args) + group_rewards.append(reward) + self.step_rewards[name].append(reward) + self.step_response_lens[name].append(sample.effective_response_length) + if sample.rollout_log_probs: + self.step_log_prob_means[name].append( + sum(sample.rollout_log_probs) / len(sample.rollout_log_probs) + ) + self.step_group_counts[name] += 1 + if len(group_rewards) > 1 and all(reward == group_rewards[0] for reward in group_rewards): + self.step_zero_std_rewards[name].append(round(group_rewards[0], 1)) + + flushed: dict[str, dict[str, float]] = {} + for name in step_names: + rewards = self.step_rewards.pop(name, []) + response_lens = self.step_response_lens.pop(name, []) + log_prob_means = self.step_log_prob_means.pop(name, []) + total_groups = self.step_group_counts.pop(name, 0) + zero_std_rewards = self.step_zero_std_rewards.pop(name, []) + if not rewards: + continue + expected = adapters[name].config.adapter_global_batch_size + if len(rewards) != expected: + logger.warning( + f"Adapter '{name}' stepped with {len(rewards)} shipped samples, expected " + f"adapter_global_batch_size={expected}; batch accounting drifted" + ) + # Single-segment keys so "{name}/" matches the "{name}/*" glob (server globs one segment). + flushed[name] = { + **dict_add_prefix(compute_statistics(rewards), "raw_reward_"), + **dict_add_prefix(compute_statistics(response_lens), "response_len_"), + } + if log_prob_means: + flushed[name]["log_probs"] = sum(log_prob_means) / len(log_prob_means) + if total_groups: + zero = sum(1 for reward in zero_std_rewards if reward == 0.0) + one = sum(1 for reward in zero_std_rewards if reward == 1.0) + flushed[name]["zero_std_all_zero_percentage"] = zero / total_groups + flushed[name]["zero_std_all_one_percentage"] = one / total_groups + return flushed + + def discard_adapter(self, name: str) -> None: + """Drop a retired adapter's partial step accumulation.""" + with self.lock: + self.step_rewards.pop(name, None) + self.step_response_lens.pop(name, None) + self.step_log_prob_means.pop(name, None) + self.step_group_counts.pop(name, None) + self.step_zero_std_rewards.pop(name, None) + self.staleness_values.pop(name, None) + + def pop_metrics(self) -> dict[str, float]: + with self.lock: + metrics = { + f"rollout/dynamic_filter/drop_{reason}": count + for reason, count in self.dynamic_filter_drop_counts.items() + } + self.dynamic_filter_drop_counts.clear() + return metrics + + +class AsyncMultiLoRAWorker: + """Background producer filling bounded per-adapter completed-group buffers; + the collection loop pops from them via ``get_groups``.""" + + global_worker = None + worker_lock = threading.Lock() + + def __init__(self, args, data_source, generate_fn: GenerateFn, concurrency: int = None) -> None: + self.args = args + self.data_source = data_source + self.generate_fn = generate_fn + self.concurrency = concurrency or args.rollout_batch_size + self.running = True + self.worker_thread: threading.Thread | None = None + self.state = GenerateState(args) + self.dynamic_filter = ( + load_function(args.dynamic_sampling_filter_path) if args.dynamic_sampling_filter_path else None + ) + # Guards the buffers: the producer thread puts while get_groups (trainer side) pops. + self.buffer_lock = threading.Lock() + self.buffers: dict[str, GroupBuffer] = defaultdict(GroupBuffer) + # Round-robin cursor over adapters, persisting across get_groups calls and batches. + self.rotation: deque[str] = deque() + self.metrics = MultiLoRAWorkerMetrics() + # Last seen registration id per adapter name; a change means re-registration -> drop inherited state. + self.registrations: dict[str, str] = {} + # Set when run_loop dies; collect_batch surfaces it instead of a misleading empty-batch timeout. + self.failure: Exception | None = None + + @classmethod + def get_or_create(cls, args, data_source, generate_fn: GenerateFn, concurrency: int = None): + with cls.worker_lock: + if cls.global_worker is None or not cls.global_worker.worker_thread.is_alive(): + cls.global_worker = cls(args, data_source, generate_fn, concurrency) + cls.global_worker.start() + return cls.global_worker + + def start(self) -> None: + self.worker_thread = threading.Thread(target=self.thread_main, daemon=True) + self.worker_thread.start() + + def stop(self) -> None: + self.running = False + if self.worker_thread and self.worker_thread.is_alive(): + self.worker_thread.join(timeout=5) + + @classmethod + def stop_global(cls) -> None: + with cls.worker_lock: + if cls.global_worker is None: + return + cls.global_worker.stop() + cls.global_worker = None + + def thread_main(self) -> None: + asyncio.run(self.run_loop()) + + async def run_loop(self) -> None: + active: set[asyncio.Task] = set() + max_concurrent = self.concurrency + try: + while self.running: + done = {t for t in active if t.done()} + for t in done: + try: + t.result() + except Exception as e: + logger.warning(f"generate task failed: {e}") + active.discard(t) + + while len(active) < max_concurrent and self.running: + samples = self.data_source.get_samples(1) + if not samples: + break + active.add(asyncio.create_task(self.process_and_enqueue(samples[0]))) + + await asyncio.sleep(0.01) + except Exception as e: + # Typically the data source: this stops production for EVERY + # adapter, so record the cause for collect_batch to surface. + self.failure = e + logger.exception("multi-LoRA producer failed; generation is stopped") + finally: + for task in active: + task.cancel() + if active: + await asyncio.gather(*active, return_exceptions=True) + + async def process_and_enqueue(self, group: list[Sample]) -> None: + result = await process_group(self.args, group, self.state.sampling_params, self.generate_fn, self.data_source) + if result is None: + return + + filter_result = call_dynamic_filter(self.dynamic_filter, self.args, result) + if not filter_result.keep: + if filter_result.reason: + self.metrics.record_dynamic_filter_drop(filter_result.reason) + return + + adapter_name = group_adapter_name(result) + if adapter_name is None: + return + with self.buffer_lock: + self.buffers[adapter_name].put(result) + + def queue_size(self) -> int: + with self.buffer_lock: + return sum(len(buffer) for buffer in self.buffers.values()) + + def queue_sizes(self) -> dict[str, int]: + """Buffered (completed, not yet shipped) prompt groups per adapter.""" + with self.buffer_lock: + return {name: len(buffer) for name, buffer in self.buffers.items()} + + def get_groups( + self, snapshot: dict, num_samples: int, group_counts: dict[str, int] + ) -> tuple[list[Group], dict[str, int]]: + """Pop groups round-robin in ``min_groups_per_dp_split`` multiples until ``num_samples`` is covered or + nothing is poppable; returns them with an updated ``group_counts`` copy (prevents adapter overshoot).""" + adapters = {**snapshot["active"], **snapshot["retiring"]} + dp_size = self.args.multi_lora_dp_size + max_staleness = getattr(self.args, "max_weight_staleness", None) + group_counts = dict(group_counts) # updated copy; the argument is not modified + popped: list[Group] = [] + popped_samples = 0 + + with self.buffer_lock: + # Retired adapters: discard their buffered tail and partial reward stats. + for name in list(self.buffers): + if name not in adapters: + self.buffers.pop(name) + self.metrics.discard_adapter(name) + self.registrations.pop(name, None) + + # A re-registered name is a new tenant: drop buffered groups and + # partial stats inherited from the old tenant. + for name, adapter in adapters.items(): + previous = self.registrations.get(name) + if previous is not None and previous != adapter.registration_id: + self.buffers.pop(name, None) + self.metrics.discard_adapter(name) + logger.warning(f"Adapter '{name}' was re-registered; dropped the previous tenant's buffered state") + self.registrations[name] = adapter.registration_id + + # Keep the rotation in sync with live adapters. + self.rotation = deque(name for name in self.rotation if name in adapters) + for name in sorted(set(adapters) - set(self.rotation)): + self.rotation.append(name) + + while popped_samples < num_samples: + made_progress = False + for _ in range(len(self.rotation)): + name = self.rotation[0] + self.rotation.rotate(-1) + adapter = adapters[name] + buffer = self.buffers[name] + if dropped := buffer.drop_stale(adapter.version, max_staleness): + self.metrics.record_stale_drops(name, dropped) + # In-flight stragglers of a retired same-name tenant that + # landed after the re-registration sweep reset the buffer. + if foreign := buffer.drop_foreign(adapter.registration_id): + logger.warning(f"Dropped {foreign} buffered groups from a previous registration of '{name}'") + min_groups_per_pop = min_groups_per_dp_split(adapter.config.n_samples_per_prompt, dp_size) + trainable_groups = len(buffer) // min_groups_per_pop * min_groups_per_pop + remaining_allowed_groups = max(0, remaining_groups(adapter) - group_counts.get(name, 0)) + groups_to_pop = min(min_groups_per_pop, trainable_groups, remaining_allowed_groups) + if groups_to_pop <= 0: + continue + popped.extend(buffer.get(groups_to_pop)) + popped_samples += groups_to_pop * adapter.config.n_samples_per_prompt + group_counts[name] = group_counts.get(name, 0) + groups_to_pop + made_progress = True + break + if not made_progress: + break # a full pass over rotation yielded nothing + return popped, group_counts + + +async def collect_batch(args, worker: AsyncMultiLoRAWorker, snapshot: dict) -> TrainBatch: + """Pop group multiples until the batch reaches ``--global-batch-size`` samples, or it is non-empty and + stalls for ``--multi-lora-max-coalesce-wait-s`` (the target can be unreachable; ship what there is).""" + adapters = {**snapshot["active"], **snapshot["retiring"]} + target_samples = args.global_batch_size + wait_s = getattr(args, "multi_lora_max_coalesce_wait_s", 0.5) + empty_wait_s = getattr(args, "multi_lora_max_empty_wait_s", EMPTY_BATCH_TIMEOUT_S) + + collected: list[Group] = [] + group_counts: dict[str, int] = {} + total_samples = 0 + last_progress = time.time() + last_warning = time.time() + + while total_samples < target_samples: + if worker.failure is not None: + raise RuntimeError( + "multi-LoRA producer thread died; generation is stalled for every adapter" + ) from worker.failure + groups, group_counts = worker.get_groups(snapshot, target_samples - total_samples, group_counts) + if groups: + collected.extend(groups) + total_samples += sum(adapters[group_adapter_name(g)].config.n_samples_per_prompt for g in groups) + last_progress = time.time() + continue + stalled_s = time.time() - last_progress + if collected and stalled_s > wait_s: + break + if not collected and stalled_s > empty_wait_s: + raise EmptyBatchTimeoutError( + "No poppable groups collected before empty timeout; this likely means every live adapter is " + "below min_groups_per_dp_split (or sources are exhausted). " + f"queue={worker.queue_size()} active={sorted(snapshot['active'])} retiring={sorted(snapshot['retiring'])}" + ) + if not collected and time.time() - last_warning > 30: + logger.warning( + "No completed groups for 30s. " + f"queue={worker.queue_size()} active={sorted(snapshot['active'])} " + f"retiring={sorted(snapshot['retiring'])}" + ) + last_warning = time.time() + await asyncio.sleep(0.01) + + step_names = sorted(name for name, count in group_counts.items() if count == remaining_groups(adapters[name])) + return TrainBatch( + groups=collected, + group_counts=group_counts, + step_names=step_names, + step_slots=sorted(adapters[name].slot for name in step_names), + ) + + +async def generate_rollout_multi_lora_async( + args, rollout_id: int, data_source, generate_fn: GenerateFn = generate_and_rm_group +) -> RolloutFnTrainOutput: + """Collect one train batch and record its contents on the controller.""" + assert args.rollout_global_dataset + + state = GenerateState(args) + worker = AsyncMultiLoRAWorker.get_or_create(args, data_source, generate_fn) + start_time = time.time() + queue_sizes = worker.queue_sizes() + + # Driver contract: adapter state only changes between generate calls, so one snapshot serves the collection. + snapshot = await get_multi_lora_controller().snapshot.remote() + assert snapshot["active"] or snapshot["retiring"], "generate called with no live adapters" + + batch = await collect_batch(args, worker, snapshot) + + data = sorted( + batch.groups, + key=lambda group: ( + first_sample(group).adapter.slot if first_sample(group).adapter is not None else -1, + first_sample(group).index, + ), + ) + + # Per-sample adapter batch size (drives loss normalization) and batch-level step + # decision (drives selective optimizer stepping), shipped via sample metadata. + adapters = {**snapshot["active"], **snapshot["retiring"]} + for group in data: + adapter = adapters[group_adapter_name(group)] + for sample in iter_group_samples(group): + sample.metadata["adapter_global_batch_size"] = adapter.config.adapter_global_batch_size + if data: + head = first_sample(data[0]) + head.metadata["step_slots"] = list(batch.step_slots) + head.metadata["step_adapter_names"] = list(batch.step_names) + + await get_multi_lora_controller().record_batch_adapters.remote(rollout_id, batch.group_counts, batch.step_names) + + if (x := args.rollout_sample_filter_path) is not None: + load_function(x)(args, data) + + await recompute_samples_rollout_logprobs_via_prefill( + args, + [s for g in data for s in iter_group_samples(g)], + url=get_model_url(args, "default"), + sampling_params=state.sampling_params, + ) + + # Adapter metrics ride the adapter's own optimizer-step axis ({name}/step); this batch completes step + 1. + for name, step_metrics in worker.metrics.record_shipped_samples(args, data, batch.step_names, adapters).items(): + step_key = f"{name}/step" + log_dict = {step_key: adapters[name].step + 1} + log_dict |= {f"{name}/{key}": value for key, value in step_metrics.items()} + tracking.log(args, log_dict, step_key=step_key) + + stale_drops = worker.metrics.pop_stale_drops() + all_staleness = [staleness for values in stale_drops.values() for staleness in values] + metrics = { + **worker.metrics.pop_metrics(), + "perf/fully_async/queue_length": sum(queue_sizes.values()), + "perf/fully_async/stale_dropped": len(all_staleness), + # {name}/perf/* rides rollout/step; two segments under {name}/ keep these off the step axis. + **{f"{name}/perf/queue_length": size for name, size in queue_sizes.items()}, + **{f"{name}/perf/stale_dropped": len(stale_drops.get(name, [])) for name in adapters}, + "perf/fully_async/batch_wait_time": time.time() - start_time, + "perf/fully_async/batch_n_adapters": len(batch.group_counts), + "perf/fully_async/batch_n_groups": len(data), + "perf/fully_async/batch_n_samples": sum(group_sample_count(group) for group in data), + "perf/fully_async/batch_n_adapters_to_step": len(batch.step_names), + } + if all_staleness: + metrics["perf/fully_async/stale_dropped_avg_staleness"] = sum(all_staleness) / len(all_staleness) + metrics["perf/fully_async/stale_dropped_max_staleness"] = max(all_staleness) + for name, values in stale_drops.items(): + if values: + metrics[f"{name}/perf/stale_dropped_avg_staleness"] = sum(values) / len(values) + metrics[f"{name}/perf/stale_dropped_max_staleness"] = max(values) + + return RolloutFnTrainOutput(samples=data, metrics=metrics) + + +def generate_rollout_multi_lora(args, rollout_id: int, data_source, evaluation: bool = False): + if evaluation: + raise ValueError("Evaluation not supported in multi-LoRA async rollout") + return run(generate_rollout_multi_lora_async(args, rollout_id, data_source)) diff --git a/miles/rollout/multi_lora/data_source.py b/miles/rollout/multi_lora/data_source.py new file mode 100644 index 0000000000..426b7af245 --- /dev/null +++ b/miles/rollout/multi_lora/data_source.py @@ -0,0 +1,137 @@ +"""Round-robin per-adapter data source. Deregistration is step-based and +lives in the controller (``mark_batch_trained``); every adapter gets a +``num_step`` at registration, explicit or derived from ``num_epoch``.""" + +import copy +import logging +from argparse import Namespace +from collections import deque +from concurrent.futures import ThreadPoolExecutor + +import ray + +from miles.ray.multi_lora.controller import get_multi_lora_controller +from miles.rollout.data_source import DataSource, RolloutDataSource +from miles.utils.adapter_config import AdapterRun +from miles.utils.types import AdapterRef, RewardSpec, Sample + +logger = logging.getLogger(__name__) + +MAX_RECONCILE_WORKERS = 16 + + +def fetch_snapshot() -> dict: + return ray.get(get_multi_lora_controller().snapshot.remote()) + + +def sampleable(snapshot: dict) -> dict[str, AdapterRun]: + return {**snapshot["active"], **snapshot["retiring"]} + + +class MultiLoRAAsyncDataSource(DataSource): + def __init__(self, args: Namespace): + self.args = args + self.sources: dict[str, RolloutDataSource] = {} + self.source_queue: deque = deque() + + def reconcile(self, adapters: dict[str, AdapterRun]) -> None: + for name in list(self.sources): + if name not in adapters: + del self.sources[name] + logger.info(f"Removed data source for adapter '{name}'") + pending = [(name, a) for name, a in adapters.items() if name not in self.sources] + if pending: + workers = min(MAX_RECONCILE_WORKERS, len(pending)) + if workers > 1: + with ThreadPoolExecutor(max_workers=workers, thread_name_prefix="mlora-ds") as ex: + built = list(ex.map(lambda na: (na[0], self.create_source(na[1])), pending)) + else: + built = [(name, self.create_source(a)) for name, a in pending] + for name, source in built: + self.sources[name] = source + logger.info(f"Created data source for adapter '{name}'") + # Post-filter dataset length; the controller derives num_step + # from num_epoch for adapters that didn't set it. + ray.get(get_multi_lora_controller().resolve_num_step.remote(name, len(source.dataset))) + self.update_queue(set(adapters)) + + def create_source(self, adapter: AdapterRun) -> RolloutDataSource: + config = adapter.config + adapter_args = copy.copy(self.args) + adapter_args.prompt_data = config.data + adapter_args.input_key = config.input_key or self.args.input_key + adapter_args.label_key = config.label_key or self.args.label_key + adapter_args.metadata_key = config.metadata_key or self.args.metadata_key + adapter_args.save = config.save or self.args.save + adapter_args.load = config.save or self.args.load + adapter_args.n_samples_per_prompt = config.n_samples_per_prompt or self.args.n_samples_per_prompt + adapter_args.start_rollout_id = 0 + return RolloutDataSource(adapter_args) + + def update_queue(self, active_names: set[str]) -> None: + new_queue: deque = deque() + in_queue: set[str] = set() + while self.source_queue: + if (name := self.source_queue.popleft()) in active_names: + new_queue.append(name) + in_queue.add(name) + for name in active_names: + if name not in in_queue: + new_queue.append(name) + self.source_queue = new_queue + + def get_samples(self, num_samples: int = 1) -> list[list[Sample]]: + """Return the next prompt group, round-robined across adapters. + + One rotation of the queue: pull one group from the first adapter that + yields, stamp it, and return. Empty list when no adapter can produce. + """ + assert num_samples == 1, "the async producer dispatches one prompt group at a time" + snapshot = fetch_snapshot() + adapters = sampleable(snapshot) + self.reconcile(adapters) + self.update_queue(set(self.sources)) + + for _ in range(len(self.source_queue)): + name = self.source_queue.popleft() + self.source_queue.append(name) + source = self.sources[name] + groups = source.get_samples(1) + if not groups: + continue + + adapter = adapters[name] + config = adapter.config + ref = AdapterRef(name=name, slot=adapter.slot) + reward_spec = RewardSpec(rm_type=config.rm_type, custom_rm_path=config.custom_rm_path) + for sample in groups[0]: + sample.adapter = ref + sample.reward_spec = reward_spec + sample.metadata = {**config.metadata, **sample.metadata} + + return groups + + return [] + + def add_samples(self, samples: list[list[Sample]]) -> None: + """Recycle retried/aborted groups; drop groups for deregistered adapters.""" + adapters = sampleable(fetch_snapshot()) + self.reconcile(adapters) + for group in samples: + name = group[0].adapter.name if group and group[0].adapter else None + if not name or name not in self.sources or name not in adapters: + continue + self.sources[name].add_samples([group]) + + def save(self, rollout_id): + for source in self.sources.values(): + source.save(rollout_id) + + def load(self, rollout_id=None): + for source in self.sources.values(): + source.load(rollout_id) + + def close(self) -> None: + from miles.rollout.multi_lora.async_rollout import AsyncMultiLoRAWorker + + AsyncMultiLoRAWorker.stop_global() diff --git a/tests/fast/ray/rollout/test_multi_lora_batch_collection.py b/tests/fast/ray/rollout/test_multi_lora_batch_collection.py new file mode 100644 index 0000000000..35cb66b88e --- /dev/null +++ b/tests/fast/ray/rollout/test_multi_lora_batch_collection.py @@ -0,0 +1,290 @@ +"""Unit tests for multi-LoRA batch collection (get_groups + collect_batch): +group-multiple math, adapter batch capping, step stamping, coalesce timeout, +round-robin fairness, retirement, and staleness filtering. No Ray, no engines: +the worker is built bare.""" + +import asyncio +import threading +import time +from collections import defaultdict, deque +from types import SimpleNamespace + +import pytest + +from miles.rollout.multi_lora.async_rollout import ( + AsyncMultiLoRAWorker, + GroupBuffer, + MultiLoRAWorkerMetrics, + collect_batch, + group_adapter_name, +) +from miles.utils.adapter_config import AdapterRun, AdapterRunConfig +from miles.utils.types import AdapterRef, Sample + + +def make_args(**overrides) -> SimpleNamespace: + args = SimpleNamespace( + global_batch_size=16, + multi_lora_dp_size=4, + multi_lora_max_coalesce_wait_s=0.05, + max_weight_staleness=None, + ) + for key, value in overrides.items(): + setattr(args, key, value) + return args + + +def make_worker(args=None) -> AsyncMultiLoRAWorker: + worker = AsyncMultiLoRAWorker.__new__(AsyncMultiLoRAWorker) + worker.args = args or make_args() + worker.buffer_lock = threading.Lock() + worker.buffers = defaultdict(GroupBuffer) + worker.rotation = deque() + worker.dynamic_filter = None + worker.metrics = MultiLoRAWorkerMetrics() + worker.registrations = {} + worker.failure = None + return worker + + +def adapter_run( + name: str, + slot: int, + rollout_batch_size: int = 4, + n_samples_per_prompt: int = 4, + accumulated_groups: int = 0, + version: int = 1, + registration_id: str = "", +) -> AdapterRun: + config = AdapterRunConfig( + data="/d", + rank=8, + alpha=16, + rollout_batch_size=rollout_batch_size, + n_samples_per_prompt=n_samples_per_prompt, + ) + return AdapterRun( + name=name, + config=config, + slot=slot, + version=version, + step=0, + accumulated_groups=accumulated_groups, + registration_id=registration_id, + ) + + +def make_group( + adapter: AdapterRun, slot_version: int | None = None, registration_id: str | None = None +) -> list[Sample]: + samples = [] + for _ in range(adapter.config.n_samples_per_prompt): + sample = Sample(prompt="p", adapter=AdapterRef(adapter.name, adapter.slot)) + if slot_version is not None: + sample.metadata["slot_version"] = slot_version + if registration_id is not None: + sample.metadata["registration_id"] = registration_id + samples.append(sample) + return samples + + +def buffer_groups( + worker, adapter: AdapterRun, count: int, slot_version: int | None = None, registration_id: str | None = None +): + for _ in range(count): + worker.buffers[adapter.name].put(make_group(adapter, slot_version, registration_id)) + + +def snapshot_of(*adapters: AdapterRun, retiring: tuple[AdapterRun, ...] = ()) -> dict: + return { + "active": {a.name: a for a in adapters}, + "retiring": {a.name: a for a in retiring}, + "cleanup": [], + } + + +def collect(worker, snapshot): + return asyncio.run(collect_batch(worker.args, worker, snapshot)) + + +def test_no_pop_until_a_whole_group_multiple_is_buffered(): + # dp=8 with n_samples=4 -> multiple = 2 groups; one buffered group is below the multiple. + worker = make_worker(make_args(multi_lora_dp_size=8)) + a = adapter_run("A", 0, rollout_batch_size=4, n_samples_per_prompt=4) + buffer_groups(worker, a, count=1) + groups, counts = worker.get_groups(snapshot_of(a), 16, {}) + assert (groups, counts) == ([], {}) + + buffer_groups(worker, a, count=1) + groups, counts = worker.get_groups(snapshot_of(a), 16, {}) + assert len(groups) == 2 + assert counts == {"A": 2} + + +def test_reaching_target_stops_collecting(): + worker = make_worker() + a = adapter_run("A", 0, rollout_batch_size=8) # adapter batch: 8 groups + buffer_groups(worker, a, count=5) # 20 samples > 16 target + start = time.monotonic() + batch = collect(worker, snapshot_of(a)) + assert time.monotonic() - start < worker.args.multi_lora_max_coalesce_wait_s # no timeout waited + assert batch.group_counts == {"A": 4} # stops once 16 samples are reached + assert batch.step_names == [] # adapter batch (8 groups) not complete + assert len(worker.buffers["A"]) == 1 + + +def test_below_target_ships_after_no_progress_timeout(): + worker = make_worker() + a = adapter_run("A", 0, rollout_batch_size=8) + buffer_groups(worker, a, count=1) # 4 samples < 16 target + start = time.monotonic() + batch = collect(worker, snapshot_of(a)) + assert time.monotonic() - start >= worker.args.multi_lora_max_coalesce_wait_s + assert batch.group_counts == {"A": 1} + + +def test_collection_capped_at_remaining_groups_and_step_stamped(): + worker = make_worker() + # Adapter batch = 4 groups; 3 already banked -> 1 remaining, despite 4 buffered. + a = adapter_run("A", 0, rollout_batch_size=4, accumulated_groups=3) + buffer_groups(worker, a, count=4) + batch = collect(worker, snapshot_of(a)) + assert batch.group_counts == {"A": 1} + assert batch.step_names == ["A"] + assert batch.step_slots == [0] + assert len(worker.buffers["A"]) == 3 # surplus stays buffered + + +def test_batch_never_overshoots_adapter_batch_across_fetches(): + """Groups arriving after an adapter's remaining groups are already in the + batch must not be popped into the same batch.""" + worker = make_worker() + a = adapter_run("A", 0, rollout_batch_size=2) + buffer_groups(worker, a, count=2) + groups, counts = worker.get_groups(snapshot_of(a), 16, {}) + assert len(groups) == 2 # whole remaining batch + + buffer_groups(worker, a, count=2) # fresh arrivals mid-collection + groups, counts = worker.get_groups(snapshot_of(a), 16, counts) + assert groups == [] + + groups, _counts = worker.get_groups(snapshot_of(a), 16, {}) # next batch may pop them + assert len(groups) == 2 + + +def test_pops_interleave_adapters_round_robin(): + worker = make_worker() + a = adapter_run("A", 0, rollout_batch_size=16) + b = adapter_run("B", 1, rollout_batch_size=16) + buffer_groups(worker, a, count=2) + buffer_groups(worker, b, count=2) + groups, counts = worker.get_groups(snapshot_of(a, b), 16, {}) + assert [group_adapter_name(g) for g in groups] == ["A", "B", "A", "B"] + assert counts == {"A": 2, "B": 2} + groups, counts = worker.get_groups(snapshot_of(a, b), 16, counts) + assert groups == [] # buffers drained + + +def test_cursor_persists_across_batches(): + worker = make_worker(make_args(global_batch_size=8)) + a = adapter_run("A", 0, rollout_batch_size=16) + b = adapter_run("B", 1, rollout_batch_size=16) + buffer_groups(worker, a, count=4) + buffer_groups(worker, b, count=4) + + # 8-sample target = 2 groups per batch; collection interleaves A and B. + batch = collect(worker, snapshot_of(a, b)) + assert batch.group_counts == {"A": 1, "B": 1} + + # The next batch continues from the cursor, not from A again. + batch = collect(worker, snapshot_of(a, b)) + assert batch.group_counts == {"A": 1, "B": 1} + assert len(worker.buffers["A"]) == 2 + assert len(worker.buffers["B"]) == 2 + + +def test_retiring_adapter_remains_selectable_until_retired(): + """RETIRING adapters keep serving until the reconcile sync point (base + deregistration semantics): buffered groups stay poppable.""" + worker = make_worker() + a = adapter_run("A", 0, rollout_batch_size=4) + buffer_groups(worker, a, count=4) + batch = collect(worker, snapshot_of(retiring=(a,))) + assert batch.group_counts == {"A": 4} + assert batch.step_names == ["A"] + + +def test_retired_adapter_buffers_are_discarded(): + """Once an adapter leaves the snapshot (retired at reconcile), its buffered + tail is dropped.""" + worker = make_worker() + a = adapter_run("A", 0, rollout_batch_size=4) + b = adapter_run("B", 1, rollout_batch_size=4) + buffer_groups(worker, a, count=3) + groups, _counts = worker.get_groups(snapshot_of(b), 16, {}) # A gone from snapshot + assert groups == [] + assert "A" not in worker.buffers # tail discarded with the adapter + + +def test_stale_buffered_groups_are_dropped(): + worker = make_worker(make_args(max_weight_staleness=1)) + a = adapter_run("A", 0, rollout_batch_size=4, version=5) + buffer_groups(worker, a, count=2, slot_version=3) # staleness 2 > 1 + buffer_groups(worker, a, count=1, slot_version=5) # fresh + batch = collect(worker, snapshot_of(a)) + assert batch.group_counts == {"A": 1} # only the fresh group ships + + +def test_empty_collection_times_out_instead_of_spinning_forever(): + worker = make_worker(make_args(multi_lora_max_empty_wait_s=0.02)) + a = adapter_run("A", 0, rollout_batch_size=4) + with pytest.raises(RuntimeError, match="No poppable groups collected before empty timeout"): + collect(worker, snapshot_of(a)) + + +def test_re_registered_name_drops_previous_tenant_buffer_and_metrics(): + # A retires while its buffer still holds groups; the driver idles (no + # generate), then the operator re-registers the same name. The new + # tenant's first get_groups must not ship the old tenant's groups nor + # inherit its partial step statistics. + worker = make_worker() + old = adapter_run("A", 0, registration_id="reg-old") + buffer_groups(worker, old, count=2, registration_id="reg-old") + worker.get_groups(snapshot_of(old), 0, {}) # worker has seen the old tenant + worker.metrics.step_rewards["A"].append(1.0) # old tenant's partial step stats + + new = adapter_run("A", 0, registration_id="reg-new") + groups, counts = worker.get_groups(snapshot_of(new), 16, {}) + + assert (groups, counts) == ([], {}) + assert len(worker.buffers["A"]) == 0 + assert "A" not in worker.metrics.step_rewards + + +def test_straggler_group_of_previous_registration_is_dropped(): + # An in-flight generation of the old tenant lands in the buffer after the + # re-registration sweep already reset it; only the new tenant's groups ship. + worker = make_worker() + new = adapter_run("A", 0, registration_id="reg-new") + worker.get_groups(snapshot_of(new), 0, {}) # sweep records the new registration + buffer_groups(worker, new, count=1, registration_id="reg-old") # straggler + buffer_groups(worker, new, count=1, registration_id="reg-new") + + groups, counts = worker.get_groups(snapshot_of(new), 16, {}) + + assert counts == {"A": 1} + assert [s.metadata["registration_id"] for g in groups for s in g] == ["reg-new"] * 4 + + +def test_dead_producer_surfaces_its_cause_instead_of_timing_out(): + # A producer-thread failure (e.g. an adapter whose dataset vanished) stops + # generation for every adapter; collect_batch must raise the recorded cause + # immediately, not wait out the empty-batch timeout. + worker = make_worker(make_args(multi_lora_max_empty_wait_s=30.0)) + worker.failure = RuntimeError("dataset gone") + a = adapter_run("A", 0) + start = time.monotonic() + with pytest.raises(RuntimeError, match="producer thread died") as excinfo: + collect(worker, snapshot_of(a)) + assert time.monotonic() - start < 1.0 # no timeout wait + assert "dataset gone" in repr(excinfo.value.__cause__) diff --git a/tests/fast/ray/rollout/test_multi_lora_process_group.py b/tests/fast/ray/rollout/test_multi_lora_process_group.py new file mode 100644 index 0000000000..b97ff2238a --- /dev/null +++ b/tests/fast/ray/rollout/test_multi_lora_process_group.py @@ -0,0 +1,56 @@ +"""Pins process_group's submission-time slot-version stamping: the staleness +filter compares against the version live when the group was submitted, not +when it completed.""" + +import pytest + +import miles.rollout.multi_lora.async_rollout as mod +from miles.rollout.multi_lora.async_rollout import process_group +from miles.utils.types import AdapterRef, Sample + + +class FakeDataSource: + def __init__(self) -> None: + self.added: list = [] + + def add_samples(self, groups) -> None: + self.added.extend(groups) + + +class FakeAdapterView: + def __init__(self, version: int, registration_id: str = "reg-1") -> None: + self.version = version + self.registration_id = registration_id + + +class FakeAdaptersCache: + def __init__(self, versions: dict[str, int]) -> None: + self.versions = versions + + def bump(self, name: str, to: int) -> None: + self.versions[name] = to + + async def get(self, adapter_name: str) -> FakeAdapterView | None: + version = self.versions.get(adapter_name) + return FakeAdapterView(version) if version is not None else None + + +@pytest.mark.asyncio +async def test_process_group_stamps_submission_version(monkeypatch): + """The stamp is the version live at submission (5), not completion (7).""" + cache = FakeAdaptersCache({"A": 5}) + + async def gen(args, group, sampling_params): + cache.bump("A", 7) # update lands mid-generation + for s in group: + s.status = Sample.Status.COMPLETED + return group + + monkeypatch.setattr(mod, "AdaptersCache", lambda: cache) + + g = [Sample(prompt="p", adapter=AdapterRef("A", 0))] + result = await process_group(None, g, {}, gen, FakeDataSource()) + + assert result is g + assert g[0].metadata["slot_version"] == 5 + assert g[0].metadata["registration_id"] == "reg-1" diff --git a/tests/fast/ray/rollout/test_multi_lora_train_data.py b/tests/fast/ray/rollout/test_multi_lora_train_data.py new file mode 100644 index 0000000000..86f259b3b8 --- /dev/null +++ b/tests/fast/ray/rollout/test_multi_lora_train_data.py @@ -0,0 +1,106 @@ +"""Multi-LoRA train-data pipeline: batch metadata extraction, exact dynamic +batch size, per-adapter batch loss scales, step stamping, and per-group reward +normalization with heterogeneous group sizes.""" + +import pytest + +from tests.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=60, suite="stage-a-cpu") + +from tests.fast.ray.rollout.conftest import make_args, make_sample + +from miles.ray.rollout.rollout_data_conversion import postprocess_rollout_data +from miles.ray.rollout.train_data_conversion import convert_samples_to_train_data +from miles.utils.types import AdapterRef + + +def multi_lora_args(**overrides): + defaults = dict( + multi_lora=True, + use_dynamic_global_batch_size=True, + grpo_std_normalization=True, + ) + defaults.update(overrides) + return make_args(**defaults) + + +def adapter_group( + name: str, + slot: int, + n_samples: int, + adapter_global_batch_size: int, + rewards: list[float], + start_index: int = 0, +): + assert len(rewards) == n_samples + group = [] + for k in range(n_samples): + sample = make_sample(index=start_index + k, reward=rewards[k]) + sample.adapter = AdapterRef(name, slot) + sample.metadata = {"adapter_global_batch_size": adapter_global_batch_size} + group.append(sample) + return group + + +def make_batch(): + """Two adapters, heterogeneous group sizes: A steps this batch, B doesn't.""" + groups = [ + adapter_group("A", 0, 4, 16, [1.0, 0.0, 1.0, 0.0], start_index=0), + adapter_group("A", 0, 4, 16, [1.0, 1.0, 1.0, 1.0], start_index=4), + adapter_group("B", 1, 2, 32, [3.0, 1.0], start_index=8), + ] + groups[0][0].metadata["step_slots"] = [0] + groups[0][0].metadata["step_adapter_names"] = ["A"] + return groups + + +def run_pipeline(dp_size: int = 2): + args = multi_lora_args() + data, metadata = postprocess_rollout_data(args, make_batch(), train_parallel_config={"dp_size": dp_size}) + train_data = convert_samples_to_train_data( + args, + data, + metadata=metadata, + custom_convert_samples_to_train_data_func=None, + custom_reward_post_process_func=None, + ) + return data, metadata, train_data + + +def test_postprocess_extracts_batch_metadata_and_exact_batch_size(): + data, metadata, _ = run_pipeline() + assert metadata["prompt_group_sizes"] == [4, 4, 2] + assert metadata["step_slots"] == [0] + assert metadata["step_adapter_names"] == ["A"] + assert metadata["dynamic_global_batch_size"] == 10 # exact batch size, no trim + assert len(data) == 10 # flattened + assert "step_slots" not in data[0].metadata # lifted out + + +def test_multi_lora_rejects_dp_indivisible_batch(): + args = multi_lora_args() + with pytest.raises(ValueError, match="not divisible by dp_size"): + postprocess_rollout_data(args, make_batch(), train_parallel_config={"dp_size": 4}) + + +def test_step_fields(): + _, _, train_data = run_pipeline() + assert train_data["adapter_slots"] == [0] * 8 + [1] * 2 + assert train_data["step_slots"] == [0] + assert train_data["step_adapter_names"] == ["A"] + # Only A steps: the trainer scales slot 0's accumulated gradient by 1/16. + assert train_data["step_adapter_batch_sizes"] == {0: 16} + assert train_data["prompt_group_sizes"] == [4, 4, 2] + + +def test_rewards_normalize_within_heterogeneous_groups(): + _, _, train_data = run_pipeline() + rewards = train_data["rewards"] + # Group boundaries: [0:4], [4:8], [8:10] — each zero-mean. + for start, end in [(0, 4), (4, 8), (8, 10)]: + assert sum(rewards[start:end]) == pytest.approx(0.0, abs=1e-6) + # Constant group (all 1.0) normalizes to zeros, not NaN. + assert rewards[4:8] == pytest.approx([0.0] * 4) + # Singleton-free std normalization applied to group 1 (n=4, mixed). + assert max(abs(r) for r in rewards[0:4]) > 0.5