From 6041eab6d97f1bf2a3410e7b395f96eeefe77dfd Mon Sep 17 00:00:00 2001 From: Yusheng Su Date: Tue, 21 Jul 2026 00:31:30 -0700 Subject: [PATCH] =?UTF-8?q?[multi-lora]=201/7:=20utils=20foundation=20?= =?UTF-8?q?=E2=80=94=20sample/adapter=20types,=20adapter=20yaml=20config,?= =?UTF-8?q?=20shared=20helpers,=20CLI=20flags=20and=20validation?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- miles/rollout/generate_utils/sample_utils.py | 2 + miles/utils/adapter_config.py | 93 ++++++++++ miles/utils/arguments.py | 90 ++++++++- miles/utils/multi_lora.py | 184 +++++++++++++++++++ miles/utils/tracking_utils/base.py | 22 +++ miles/utils/tracking_utils/tracking.py | 6 + miles/utils/types.py | 21 +++ tests/fast/utils/test_arguments.py | 80 ++++++++ 8 files changed, 496 insertions(+), 2 deletions(-) create mode 100644 miles/utils/adapter_config.py create mode 100644 miles/utils/multi_lora.py diff --git a/miles/rollout/generate_utils/sample_utils.py b/miles/rollout/generate_utils/sample_utils.py index effbce562b..84703c0df9 100644 --- a/miles/rollout/generate_utils/sample_utils.py +++ b/miles/rollout/generate_utils/sample_utils.py @@ -142,6 +142,8 @@ def _merge_metadata(): metadata=_merge_metadata(), generate_function_path=_merge_equal_value("generate_function_path"), train_metadata=_merge_equal_value("train_metadata"), + adapter=_merge_equal_value("adapter"), + reward_spec=_merge_equal_value("reward_spec"), routing_key=_merge_equal_value("routing_key"), non_generation_time=_merge_equal_value("non_generation_time"), spec_info=_merge_spec_info(a.spec_info, b.spec_info), diff --git a/miles/utils/adapter_config.py b/miles/utils/adapter_config.py new file mode 100644 index 0000000000..6c4d224a86 --- /dev/null +++ b/miles/utils/adapter_config.py @@ -0,0 +1,93 @@ +"""Adapter config parsing for multi-LoRA training. + +``AdapterRunConfig`` carries only static, YAML-sourced configuration; the +mutable slot is owned by the controller and exposed through ``AdapterRun`` +views. +""" + +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + +import yaml + + +@dataclass(frozen=True) +class AdapterRunConfig: + + data: str + + # resolves them to CLI defaults if None (--lora-rank / --lora-alpha) on register. + rank: int | None = None + alpha: int | None = None + + # Prompt groups consumed per optimizer step for this adapter (group units, + # like --rollout-batch-size, which it defaults to). The samples-per-step + # analog of --global-batch-size is derived: adapter_global_batch_size = + # rollout_batch_size * n_samples_per_prompt. + rollout_batch_size: int | None = None + n_samples_per_prompt: int | None = None + + save: str | Path | None = None + + input_key: str = "text" + label_key: str | None = None + metadata_key: str | None = None + + rm_type: str | None = None + custom_rm_path: str | None = None + + # Stop after N optimizer steps; derived from num_epoch (default 1) when absent. + num_step: int | None = None + num_epoch: int | None = None + + metadata: dict[str, Any] = field(default_factory=dict) + + @property + def adapter_global_batch_size(self) -> int: + """Samples per optimizer step (per-adapter analog of --global-batch-size).""" + assert self.rollout_batch_size is not None and self.n_samples_per_prompt is not None + return self.rollout_batch_size * self.n_samples_per_prompt + + +@dataclass(frozen=True) +class AdapterRun: + """Read-only join view of a run's static config and current slot.""" + + name: str + config: AdapterRunConfig + slot: int + version: int = 0 + step: int = 0 + # Committed prompt groups accumulated toward the current optimizer step. + accumulated_groups: int = 0 + # Unique per registration (see AdapterRecord.registration_id): lets the + # rollout worker tell a re-registered name apart from the previous tenant. + registration_id: str = "" + + +def parse_adapter_run_yaml(path: Path) -> AdapterRunConfig: + """Parse a single adapter.yaml file. + + ``rank``, ``alpha`` and ``save`` are optional in the YAML; when absent the + caller (e.g. the multi-LoRA controller) is responsible for resolving them. + """ + with open(path) as f: + raw = yaml.safe_load(f) + + return AdapterRunConfig( + rank=raw.get("rank"), + alpha=raw.get("alpha"), + data=raw["data"], + rollout_batch_size=raw.get("rollout_batch_size"), + n_samples_per_prompt=raw.get("n_samples_per_prompt"), + save=Path(raw["save"]) if raw.get("save", None) else None, + input_key=raw.get("input_key", "text"), + label_key=raw.get("label_key"), + metadata_key=raw.get("metadata_key"), + rm_type=raw.get("rm_type"), + custom_rm_path=raw.get("custom_rm_path"), + num_step=raw.get("num_step"), + num_epoch=raw.get("num_epoch"), + metadata=raw.get("metadata") or {}, + ) diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index 6138f65cd5..4cd8c85a83 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -538,7 +538,7 @@ def add_rollout_arguments(parser): default=512 * 1024**2, help=( "buffer size for update weight, in bytes. " - "This is used for updating weights by chunk and should be useful for MoE models." + "This is used for updating weights by batch and should be useful for MoE models." ), ) parser.add_argument( @@ -1400,6 +1400,84 @@ def add_lora_arguments(parser): "down lora_B shared across experts, expert_dim=1). Matches SGLang " "PR #21466's experts_shared_outer_loras=True serving contract.", ) + parser.add_argument( + "--multi-lora-n-adapters", + type=int, + default=0, + help="Maximum number of concurrent adapter slots for multi-LoRA. Set to 0 to disable multi-LoRA (default: 0)", + ) + parser.add_argument( + "--multi-lora-adapter", + nargs=2, + action="append", + type=str, + dest="multi_lora_adapters", + default=[], + ) + parser.add_argument( + "--multi-lora-idle-poll-s", + type=float, + default=5.0, + help="When no adapter is RUNNING, the trainer polls for new registrations every this many seconds (default: 5.0)", + ) + parser.add_argument( + "--multi-lora-http-server-path", + type=str, + default=None, + help=( + "Dotted path to a MultiLoRAHTTPServer subclass to use for the multi-LoRA " + "controller's HTTP server (default: MultiLoRAHTTPServer)" + ), + ) + parser.add_argument( + "--multi-lora-backend-path", + type=str, + default=None, + help=( + "Dotted path to a MultiLoRABackend subclass for the multi-LoRA controller, " + "e.g. to add custom adapter validation via validate_adapter (default: MultiLoRABackend)" + ), + ) + parser.add_argument( + "--multi-lora-api-port", + type=int, + default=8068, + help="Port for the multi-LoRA controller's control-plane API, served from the head node (default: 8068)", + ) + parser.add_argument( + "--multi-lora-disable-service-mode", + action="store_false", + dest="multi_lora_service_mode", + help="Disable service mode. By default, the trainer waits indefinitely for new adapters. With this flag, it exits after all adapters have been processed.", + ) + parser.add_argument( + "--multi-lora-max-adapter-global-batch-size", + type=int, + default=None, + help=( + "Registration-time upper bound on an adapter's samples per optimizer " + "step (rollout_batch_size x n_samples_per_prompt). Defaults to 4x " + "--global-batch-size." + ), + ) + parser.add_argument( + "--multi-lora-max-coalesce-wait-s", + type=float, + default=0.5, + help=( + "Maximum time ready groups wait for the batch to fill toward " + "--global-batch-size before training starts on what is ready (default: 0.5)." + ), + ) + parser.add_argument( + "--multi-lora-max-empty-wait-s", + type=float, + default=30.0, + help=( + "How long a generate call waits for the first poppable group before " + "failing with an empty-batch timeout (default: 30)." + ), + ) return parser def add_router_arguments(parser): @@ -2532,6 +2610,12 @@ def miles_validate_args(args): "shared-outer" if args.experts_shared_outer_loras else "per-expert", ) + # Sets args.multi_lora, then validates/defaults the multi-LoRA arg surface + # (adapter configs themselves are loaded later by the controller). + from miles.utils.multi_lora import validate_multi_lora_args + + validate_multi_lora_args(args) + assert not (args.kl_coef != 0 and args.kl_loss_coef != 0), "Only one of kl_coef and kl_loss_coef can be set" if args.advantage_estimator in ["reinforce_plus_plus", "reinforce_plus_plus_baseline"]: @@ -2700,7 +2784,9 @@ def miles_validate_args(args): ) args.global_batch_size = global_batch_size - if args.n_samples_per_prompt == 1: + # Multi-LoRA adapters carry their own n_samples_per_prompt; the per-group + # normalization path already skips std for singleton groups. + if args.n_samples_per_prompt == 1 and not args.multi_lora: args.grpo_std_normalization = False logger.info("n_samples_per_prompt is set to 1, grpo_std_normalization will be set to False.") diff --git a/miles/utils/multi_lora.py b/miles/utils/multi_lora.py new file mode 100644 index 0000000000..a195f73169 --- /dev/null +++ b/miles/utils/multi_lora.py @@ -0,0 +1,184 @@ +"""Small multi-LoRA helpers shared across the rollout, trainer, and controller. + +The controller-side machinery (AdapterRegistry, MultiLoRABackend, +MultiLoRAHTTPServer) lives in ``miles/ray/multi_lora/``. +""" + +import logging +import uuid +from typing import Any + +logger = logging.getLogger(__name__) + +__all__ = [ + "EmptyBatchTimeoutError", + "RID_SEPARATOR", + "define_new_adapter_metrics", + "is_multi_lora_enabled", + "make_rid", + "min_groups_per_dp_split", + "parse_adapter", + "slot_lora_name", + "validate_multi_lora_args", +] + + +# Must not appear in adapter names so rid prefix aborts can't cross adapters. +RID_SEPARATOR = "::" + + +class EmptyBatchTimeoutError(RuntimeError): + """No trainable groups arrived before empty-wait timeout.""" + + +def is_multi_lora_enabled(args: Any) -> bool: + return getattr(args, "multi_lora", False) + + +def define_new_adapter_metrics(snapshot: dict) -> None: + """Declare metric axes for new adapters ({name}/* -> {name}/step, {name}/perf/* -> rollout/step); must run + in the primary tracking writer. Already-declared adapters are skipped, so calling every snapshot is free.""" + # lazy import tracking deps + from miles.utils.tracking_utils.tracking import define_step_key_metric_group + + for name in {**snapshot["pending"], **snapshot["active"], **snapshot["retiring"]}: + define_step_key_metric_group(prefix=name, step_key=f"{name}/step") + define_step_key_metric_group(prefix=f"{name}/perf", step_key="rollout/step") + + +def validate_multi_lora_args(args: Any) -> None: + """Set ``args.multi_lora``, then validate and default the multi-LoRA arg + surface. Called from ``miles_validate_args``; a no-op for normal runs.""" + args.multi_lora = getattr(args, "multi_lora_n_adapters", 0) > 0 + if not args.multi_lora: + return + + # Swap in the multi-LoRA rollout fn and data source unless the user pointed these flags elsewhere. + standard_rollout_fns = ( + "miles.rollout.inference_rollout.inference_rollout_common.InferenceRolloutFn", + "miles.rollout.sglang_rollout.generate_rollout", + ) + if args.rollout_function_path in standard_rollout_fns: + args.rollout_function_path = "miles.rollout.multi_lora.async_rollout.generate_rollout_multi_lora" + if args.data_source_path == "miles.rollout.data_source.RolloutDataSourceWithBuffer": + args.data_source_path = "miles.rollout.multi_lora.data_source.MultiLoRAAsyncDataSource" + # The per-adapter data source is inherently global (the controller owns + # what is sampleable); rollout workers must not shard it. + args.rollout_global_dataset = True + assert args.lora_rank > 0, "--lora-rank must be set when --multi-lora-n-adapters > 0" + assert args.target_modules is not None, "--target-modules must be set when --multi-lora-n-adapters > 0" + assert args.train_backend == "megatron", "Multi-LoRA currently requires --train-backend megatron" + assert "muon" not in str(getattr(args, "optimizer", "")).lower(), ( + "Multi-LoRA does not support Muon: per-adapter decoupled stepping is only " + "implemented for Adam-family per-slot optimizers" + ) + assert not args.colocate, ( + "Multi-LoRA requires disaggregated rollout engines: weight sync is only " + "implemented for the distributed path, not the colocated tensor path." + ) + assert ( + not getattr(args, "indep_dp", False) and "train" not in args.ft_components + ), "Multi-LoRA does not support independent-DP training; remove 'train' from --ft-components" + assert not args.offload_train, ( + "Multi-LoRA retains per-adapter gradient accumulation in GPU buffers between " + "train calls; --offload-train would destroy it. Disable offload for multi-LoRA." + ) + assert not getattr(args, "enable_witness", False), ( + "Multi-LoRA runs without the distributed optimizer (per-slot LayerWise " + "optimizers); the witness module assumes use_distributed_optimizer" + ) + assert getattr(args, "sglang_tokenizer_worker_num", 1) == 1, ( + "Multi-LoRA requires --sglang-tokenizer-worker-num 1: each tokenizer " + "worker process holds its own LoRA registry, so per-step adapter " + "upserts resolve against whichever worker the router picks and fail " + "non-deterministically. sglang rejects the upsert at runtime anyway; " + "fail at launch instead of burning GPU time until the first weight push." + ) + assert not args.calculate_per_token_loss, ( + "Multi-LoRA normalizes each sample by its adapter batch " + "(sample-mean); per-token loss normalization would make adapter batch weights " + "depend on batch contents. Drop --calculate-per-token-loss." + ) + assert args.multi_lora_max_coalesce_wait_s >= 0, "--multi-lora-max-coalesce-wait-s must be non-negative" + assert (getattr(args, "optimizer", "adam") or "adam").lower() == "adam", ( + "Multi-LoRA requires --optimizer adam: the per-slot optimizer isolation " + "(build_multi_lora_optimizer, slot retirement state cleanup) only implements " + f"Adam semantics; got --optimizer {args.optimizer}" + ) + from miles.utils.environ import enable_experimental_ft_trainer + + assert not enable_experimental_ft_trainer(), ( + "Multi-LoRA is not supported with MILES_EXPERIMENTAL_FT_TRAINER=1: the v2 " + "train group has no reconcile_adapters and does not return train outcomes" + ) + # --global-batch-size may legitimately be unset (Megatron derives it later); + # leave the adapter cap unset too rather than multiplying None. + if args.multi_lora_max_adapter_global_batch_size is None and getattr(args, "global_batch_size", None) is not None: + args.multi_lora_max_adapter_global_batch_size = 4 * args.global_batch_size + if args.multi_lora_max_adapter_global_batch_size is not None: + assert ( + args.multi_lora_max_adapter_global_batch_size > 0 + ), "--multi-lora-max-adapter-global-batch-size must be positive" + + # Trainer DP size, used to validate adapter batch shapes; guarded for harnesses without megatron args set. + if all( + hasattr(args, name) + for name in ( + "world_size", + "tensor_model_parallel_size", + "pipeline_model_parallel_size", + "context_parallel_size", + ) + ): + from miles.utils.megatron_args_utils import compute_megatron_world_size_except_dp + + model_parallel = compute_megatron_world_size_except_dp(args) + assert ( + args.world_size % model_parallel == 0 + ), f"actor world size {args.world_size} is not divisible by tp*pp*cp {model_parallel}" + args.multi_lora_dp_size = args.world_size // model_parallel + else: + args.multi_lora_dp_size = None + + # Batches are variable-sized; carry the exact sample + # count through rollout conversion instead of trimming to --global-batch-size. + assert not args.disable_rollout_trim_samples, ( + "Multi-LoRA computes the exact dynamic batch size in rollout postprocessing; " + "do not pass --disable-rollout-trim-samples" + ) + args.use_dynamic_global_batch_size = True + args.megatron_to_hf_mode = "bridge" + + +def make_rid(adapter_name: str) -> str: + return f"{adapter_name}{RID_SEPARATOR}{uuid.uuid4().hex}" + + +def parse_adapter(rid: str) -> str: + return rid.rsplit(RID_SEPARATOR, 1)[0] + + +def slot_lora_name(slot: int) -> str: + """Engine-side LoRA adapter name for a controller slot. Weight pushes and + every inference request (rollout and prefill scoring) must agree on this.""" + return f"__miles_slot_{slot}" + + +def min_groups_per_dp_split(n_samples_per_prompt: int, dp_size: int) -> int: + """Minimum prompt-group count that splits cleanly across data-parallel + ranks. + + Train batches only pop groups in multiples of this value, so each popped + slice has a sample count divisible by ``dp_size`` with no trimming. + + Requires ``n_samples_per_prompt`` and ``dp_size`` to divide each other + (one must be a multiple of the other). + """ + larger = max(dp_size, n_samples_per_prompt) + smaller = min(dp_size, n_samples_per_prompt) + if larger % smaller == 0: + return larger // n_samples_per_prompt + raise ValueError( + f"n_samples_per_prompt={n_samples_per_prompt} must be a divisor or a multiple of " + f"the data-parallel size {dp_size} so whole prompt groups can split evenly across ranks" + ) diff --git a/miles/utils/tracking_utils/base.py b/miles/utils/tracking_utils/base.py index f3e3a9dceb..b7a92b1ffc 100644 --- a/miles/utils/tracking_utils/base.py +++ b/miles/utils/tracking_utils/base.py @@ -33,11 +33,19 @@ def log(self, metrics: dict[str, Any], step: int | None = None, **kwargs) -> Non @abstractmethod def finish(self) -> None: ... + def define_step_key_metric_group(self, prefix: str, step_key: str) -> None: + """Declare that ``{prefix}/*`` metrics plot against ``step_key``; no-op for + backends that take the step numerically on every log call.""" + return + # Thin adapters for backwards compatibility to keep wandb_utils and tensorboard_utils untouched. class WandbBackend(TrackingBackend): # Delegates to the existing ``wandb_utils`` helpers. + def __init__(self) -> None: + self._defined_step_key_groups: set[tuple[str, str]] = set() + def init(self, args, *, primary: bool = True, **kwargs) -> None: from . import wandb_utils @@ -51,6 +59,16 @@ def log(self, metrics: dict[str, Any], step: int | None = None, **kwargs) -> Non wandb.log(metrics) + def define_step_key_metric_group(self, prefix: str, step_key: str) -> None: + # Call from the primary tracking process: definitions from secondary shared-mode writers are lost. + if (prefix, step_key) in self._defined_step_key_groups: + return + import wandb + + wandb.define_metric(step_key) + wandb.define_metric(f"{prefix}/*", step_metric=step_key) + self._defined_step_key_groups.add((prefix, step_key)) + def finish(self) -> None: import wandb @@ -140,6 +158,10 @@ def log(self, metrics: dict[str, Any], step: int | None = None, step_key: str | for backend in self._backends: backend.log(metrics, step=step, step_key=step_key) + def define_step_key_metric_group(self, prefix: str, step_key: str) -> None: + for backend in self._backends: + backend.define_step_key_metric_group(prefix, step_key) + def finish(self) -> None: for backend in self._backends: try: diff --git a/miles/utils/tracking_utils/tracking.py b/miles/utils/tracking_utils/tracking.py index 6a75534aef..877f578f96 100644 --- a/miles/utils/tracking_utils/tracking.py +++ b/miles/utils/tracking_utils/tracking.py @@ -28,6 +28,12 @@ def init_tracking(args, primary: bool = True, **kwargs): _manager.init(args, primary=primary, **kwargs) +def define_step_key_metric_group(prefix: str, step_key: str) -> None: + """Declare a metric group plotted against its own step key (e.g. ``{name}/*`` vs ``{name}/step``). + Only wandb acts on this; must be called from the primary tracking process or definitions may be lost.""" + _manager.define_step_key_metric_group(prefix, step_key) + + def log(args, metrics, step_key: str): step = metrics.get(step_key) _manager.log(metrics, step=step, step_key=step_key) diff --git a/miles/utils/types.py b/miles/utils/types.py index 6b50145650..0d89ee092e 100644 --- a/miles/utils/types.py +++ b/miles/utils/types.py @@ -6,6 +6,22 @@ import torch +@dataclass(frozen=True) +class AdapterRef: + """Which LoRA adapter a sample is bound to (training slot routing, inference lora_path); ``None`` = no adapter.""" + + name: str + slot: int + + +@dataclass(frozen=True) +class RewardSpec: + """Per-sample spec of how the response is scored; intentionally decoupled from adapter routing.""" + + rm_type: str | None = None + custom_rm_path: str | None = None + + @dataclass class Sample: """The sample generated""" @@ -52,6 +68,11 @@ class Status(Enum): # metadata used during training, e.g., what loss to use for this sample. train_metadata: dict | None = None + # MultiLoRA: which adapter this sample trains/infers with + adapter: AdapterRef | None = None + # Per-sample reward dispatch override (e.g., per-adapter RM in multi-LoRA) + reward_spec: RewardSpec | None = None + # Per-sample routing key for the router's consistent_hashing policy (sent as X-SMG-Routing-Key) routing_key: str | None = None diff --git a/tests/fast/utils/test_arguments.py b/tests/fast/utils/test_arguments.py index 20a7bfb758..2dbec7238f 100644 --- a/tests/fast/utils/test_arguments.py +++ b/tests/fast/utils/test_arguments.py @@ -172,6 +172,86 @@ def test_custom_megatron_post_save_hook_path_requires_save(): miles_validate_args(args) +class TestMultiLoRAValidation: + def _parse(self, extra): + parser = argparse.ArgumentParser() + get_miles_extra_args_provider()(parser) + return parser.parse_args( + [ + "--multi-lora-n-adapters", + "2", + "--lora-rank", + "8", + "--target-modules", + "linear_qkv", + "--num-rollout", + "1", + ] + + extra + + REQUIRED_ARGS + ) + + def test_rejects_multiple_tokenizer_workers(self): + # Each sglang tokenizer worker holds its own LoRA registry, so per-step + # upserts fail non-deterministically; fail at launch, not first push. + args = self._parse(["--sglang-tokenizer-worker-num", "2"]) + + with pytest.raises(AssertionError, match="sglang-tokenizer-worker-num 1"): + miles_validate_args(args) + + def test_accepts_default_single_tokenizer_worker(self): + args = self._parse([]) + + miles_validate_args(args) + + assert args.multi_lora is True + + def test_defaults_rollout_fn_and_data_source_to_multi_lora(self): + args = self._parse([]) + + miles_validate_args(args) + + assert args.rollout_function_path == "miles.rollout.multi_lora.async_rollout.generate_rollout_multi_lora" + assert args.data_source_path == "miles.rollout.multi_lora.data_source.MultiLoRAAsyncDataSource" + assert args.rollout_global_dataset is True + + def test_keeps_user_supplied_rollout_fn_and_data_source(self): + args = self._parse( + ["--rollout-function-path", "my.custom.rollout_fn", "--data-source-path", "my.custom.DataSource"] + ) + + miles_validate_args(args) + + assert args.rollout_function_path == "my.custom.rollout_fn" + assert args.data_source_path == "my.custom.DataSource" + + def test_empty_wait_is_a_registered_argument(self): + assert self._parse([]).multi_lora_max_empty_wait_s == 30.0 + assert self._parse(["--multi-lora-max-empty-wait-s", "5"]).multi_lora_max_empty_wait_s == 5.0 + + def test_rejects_non_adam_optimizer(self): + # Per-slot optimizer isolation (state init, retirement cleanup, step + # clocks) only implements Adam semantics. Muon has its own dedicated + # rejection; anything else non-Adam trips the generic guard. + args = self._parse([]) + args.optimizer = "muon" + with pytest.raises(AssertionError, match="does not support Muon"): + miles_validate_args(args) + + args = self._parse([]) + args.optimizer = "sgd" + with pytest.raises(AssertionError, match="requires --optimizer adam"): + miles_validate_args(args) + + def test_rejects_experimental_ft_trainer(self, monkeypatch): + # The v2 train group has no reconcile_adapters. + monkeypatch.setenv("MILES_EXPERIMENTAL_FT_TRAINER", "1") + args = self._parse([]) + + with pytest.raises(AssertionError, match="MILES_EXPERIMENTAL_FT_TRAINER"): + miles_validate_args(args) + + class TestResolveFtComponents: def test_disabled_with_no_components_returns_empty_without_warning(self, caplog) -> None: """use_fault_tolerance off and no ft_components yields an empty list and no warning."""