From 6041eab6d97f1bf2a3410e7b395f96eeefe77dfd Mon Sep 17 00:00:00 2001 From: Yusheng Su Date: Tue, 21 Jul 2026 00:31:30 -0700 Subject: [PATCH 1/3] =?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.""" From b60b8c30f70480a67e74d69f71d39ebe574265fb Mon Sep 17 00:00:00 2001 From: Yusheng Su Date: Tue, 21 Jul 2026 00:31:30 -0700 Subject: [PATCH 2/3] =?UTF-8?q?[multi-lora]=202/7:=20adapter=20controller?= =?UTF-8?q?=20=E2=80=94=20registry=20state=20machine,=20backend,=20control?= =?UTF-8?q?-plane=20HTTP=20API,=20named=20Ray=20actor?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- miles/ray/multi_lora/__init__.py | 0 miles/ray/multi_lora/backend.py | 188 +++++++++ miles/ray/multi_lora/controller.py | 122 ++++++ miles/ray/multi_lora/http_server.py | 129 ++++++ miles/ray/multi_lora/registry.py | 252 +++++++++++ tests/fast/ray/multi_lora/__init__.py | 0 .../ray/multi_lora/test_controller_backend.py | 392 ++++++++++++++++++ .../ray/multi_lora/test_controller_http.py | 239 +++++++++++ 8 files changed, 1322 insertions(+) create mode 100644 miles/ray/multi_lora/__init__.py create mode 100644 miles/ray/multi_lora/backend.py create mode 100644 miles/ray/multi_lora/controller.py create mode 100644 miles/ray/multi_lora/http_server.py create mode 100644 miles/ray/multi_lora/registry.py create mode 100644 tests/fast/ray/multi_lora/__init__.py create mode 100644 tests/fast/ray/multi_lora/test_controller_backend.py create mode 100644 tests/fast/ray/multi_lora/test_controller_http.py diff --git a/miles/ray/multi_lora/__init__.py b/miles/ray/multi_lora/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/miles/ray/multi_lora/backend.py b/miles/ray/multi_lora/backend.py new file mode 100644 index 0000000000..b8dd0eab27 --- /dev/null +++ b/miles/ray/multi_lora/backend.py @@ -0,0 +1,188 @@ +"""Multi-LoRA backend: the registry plus engine-facing aborts, shared by the +controller Ray actor and the HTTP server. Subclass via +``--multi-lora-backend-path``.""" + +import asyncio +import logging +from dataclasses import replace +from pathlib import Path +from typing import Any + +import httpx + +from miles.ray.multi_lora.registry import AdapterRegistry, AdapterState +from miles.utils.adapter_config import AdapterRunConfig +from miles.utils.multi_lora import RID_SEPARATOR, min_groups_per_dp_split + +logger = logging.getLogger(__name__) + + +class MultiLoRABackend: + """Registry + engine-facing aborts, shared by the Ray actor and HTTP server. + Subclass via --multi-lora-backend-path.""" + + def __init__(self, args: Any, router_url: str) -> None: + self.args = args + self.registry = AdapterRegistry(args.multi_lora_n_adapters) + self.router_url = router_url.rstrip("/") + self.client: httpx.AsyncClient | None = None + + async def init(self) -> None: + self.client = httpx.AsyncClient(timeout=httpx.Timeout(30.0)) + + async def close(self) -> None: + if self.client is not None: + await self.client.aclose() + self.client = None + + async def validate_adapter(self, name: str, config: Any) -> None: + """Override to reject adapter registrations (raise ValueError).""" + + def resolve_adapter_config(self, name: str, config: Any) -> Any: + """Resolve optional adapter-local values against process-wide defaults + and validate the batch shape against the trainer's DP layout. + + All batch-shape constraints are enforced here, at registration, so a + bad config fails immediately instead of crashing an arbitrary later + train batch. + """ + if config is None or not isinstance(config, AdapterRunConfig): + return config + + rank = config.rank if config.rank is not None else getattr(self.args, "lora_rank", 1) + alpha = config.alpha if config.alpha is not None else getattr(self.args, "lora_alpha", rank) + rollout_batch_size = ( + config.rollout_batch_size + if config.rollout_batch_size is not None + else getattr(self.args, "rollout_batch_size", None) + ) + n_samples_per_prompt = ( + config.n_samples_per_prompt + if config.n_samples_per_prompt is not None + else getattr(self.args, "n_samples_per_prompt", 1) + ) + + if type(rank) is not int or rank <= 0: + raise ValueError(f"Adapter '{name}' rank must be a positive integer") + if rank > getattr(self.args, "lora_rank", rank): + raise ValueError(f"Adapter '{name}' rank {rank} exceeds the allocated maximum rank {self.args.lora_rank}") + if alpha is None or alpha <= 0: + raise ValueError(f"Adapter '{name}' must have a positive alpha") + if type(rollout_batch_size) is not int or rollout_batch_size <= 0: + raise ValueError(f"Adapter '{name}' rollout_batch_size must be a positive integer (prompt groups)") + if type(n_samples_per_prompt) is not int or n_samples_per_prompt <= 0: + raise ValueError(f"Adapter '{name}' n_samples_per_prompt must be a positive integer") + if config.num_step is not None and (type(config.num_step) is not int or config.num_step <= 0): + raise ValueError(f"Adapter '{name}' num_step must be a positive integer") + if config.num_epoch is not None and (type(config.num_epoch) is not int or config.num_epoch <= 0): + raise ValueError(f"Adapter '{name}' num_epoch must be a positive integer") + if config.num_step is not None and config.num_epoch is not None: + logger.warning(f"Adapter '{name}' sets both num_step and num_epoch; num_step takes precedence") + + # A bad data path or unresolvable reward config does not fail at this + # API otherwise: the data path kills the shared rollout producer thread + # and an empty reward config burns every generated sample, either way + # stalling ALL adapters behind a misleading empty-batch timeout. + if not Path(config.data).expanduser().exists(): + raise ValueError( + f"Adapter '{name}' data path '{config.data}' does not exist " + "(checked from the controller process, which runs on the head node with the rollout data source)" + ) + if ( + config.custom_rm_path is None + and not (config.rm_type or "").strip() + and getattr(self.args, "custom_rm_path", None) is None + and not (getattr(self.args, "rm_type", None) or "").strip() + ): + raise ValueError( + f"Adapter '{name}' has no reward config: set rm_type or custom_rm_path in the adapter " + "config, or launch with --rm-type / --custom-rm-path" + ) + + adapter_global_batch_size = rollout_batch_size * n_samples_per_prompt + if (max_batch := getattr(self.args, "multi_lora_max_adapter_global_batch_size", None)) is not None: + if adapter_global_batch_size > max_batch: + raise ValueError( + f"Adapter '{name}' consumes {adapter_global_batch_size} samples per step " + f"(rollout_batch_size {rollout_batch_size} x n_samples_per_prompt {n_samples_per_prompt}), " + f"exceeding --multi-lora-max-adapter-global-batch-size {max_batch}" + ) + if (dp_size := getattr(self.args, "multi_lora_dp_size", None)) is not None: + try: + group_multiple = min_groups_per_dp_split(n_samples_per_prompt, dp_size) + except ValueError as e: + raise ValueError(f"Adapter '{name}': {e}") from None + if rollout_batch_size % group_multiple != 0: + raise ValueError( + f"Adapter '{name}' rollout_batch_size {rollout_batch_size} must be a multiple of " + f"its min_groups_per_dp_split ({group_multiple} at dp_size={dp_size}), so the " + f"adapter batch can complete from evenly-splitting takes" + ) + + save = Path(config.save) if config.save is not None else None + if save is None: + if getattr(self.args, "save", None) is None: + raise ValueError(f"Adapter '{name}' has no save dir: set 'save' in the adapter config or pass --save") + save = Path(self.args.save) / "adapters" / name + + return replace( + config, + rank=rank, + alpha=alpha, + rollout_batch_size=rollout_batch_size, + n_samples_per_prompt=n_samples_per_prompt, + save=save, + ) + + async def register(self, name: str, config: Any) -> dict: + config = self.resolve_adapter_config(name, config) + await self.validate_adapter(name, config) + result = self.registry.register(name, config) + resolved = getattr(config, "save", None) + if resolved is not None: + logger.info(f"Adapter '{name}' registered (slot {result['slot']}), checkpoints -> {resolved}") + return result + + async def deregister(self, name: str) -> None: + self.registry.deregister(name) + + async def retire_adapters(self) -> list[str]: + names = self.registry.retire_adapters() + for name in names: + await self.abort_adapter_requests(name) + return names + + async def free_slot(self, name: str) -> int: + """Free the adapter's slot after one final abort round: requests can survive the + ``retire_adapters`` abort (e.g. multi-turn groups), and must not leak to the slot's next tenant.""" + record = self.registry.records.get(name) + if record is not None and record.state is AdapterState.CLEANUP: + await self.abort_adapter_requests(name) + return self.registry.free_slot(name) + + async def worker_urls(self) -> list[str]: + assert self.client is not None + for endpoint, extract in ( + ("/list_workers", lambda body: body["urls"]), + ("/workers", lambda body: [worker["url"] for worker in body["workers"]]), + ): + try: + resp = await self.client.get(f"{self.router_url}{endpoint}") + if resp.status_code == 200: + return extract(resp.json()) + except Exception: + continue + return [] + + async def abort_adapter_requests(self, adapter_name: str) -> None: + prefix = f"{adapter_name}{RID_SEPARATOR}" + urls = await self.worker_urls() + if not urls: + logger.warning(f"Abort for adapter '{adapter_name}': no workers discovered at {self.router_url}") + return + results = await asyncio.gather( + *(self.client.post(f"{url}/abort_request", json={"rid": prefix, "prefix": True}) for url in urls), + return_exceptions=True, + ) + if failures := sum(isinstance(r, Exception) for r in results): + logger.warning(f"Abort for adapter '{adapter_name}': {failures}/{len(results)} posts failed") diff --git a/miles/ray/multi_lora/controller.py b/miles/ray/multi_lora/controller.py new file mode 100644 index 0000000000..7cbff2b5b9 --- /dev/null +++ b/miles/ray/multi_lora/controller.py @@ -0,0 +1,122 @@ +"""Named Ray actor wrapping the multi-LoRA backend + HTTP server.""" + +import time +from functools import cache +from typing import Any + +import ray + +from miles.ray.multi_lora.backend import MultiLoRABackend +from miles.ray.multi_lora.http_server import MultiLoRAHTTPServer +from miles.utils.adapter_config import AdapterRun +from miles.utils.misc import SingletonMeta, get_current_node_ip, load_function +from miles.utils.ray_utils import compute_ray_pin_head_options + +CONTROLLER_NAME = "miles_multi_lora_controller" +CONTROLLER_NAMESPACE = "miles" + + +@cache +def get_multi_lora_controller(): + return ray.get_actor(CONTROLLER_NAME, namespace=CONTROLLER_NAMESPACE) + + +class AdaptersCache(metaclass=SingletonMeta): + """TTL-cached controller snapshot; get/get_all expose the sampleable + projection (active + retiring).""" + + def __init__(self, ttl_s: float = 1.0) -> None: + self.ttl_s = ttl_s + self.snapshot: dict = {"pending": {}, "active": {}, "retiring": {}, "cleanup": []} + self.last_refresh: float | None = None + + async def get_snapshot(self) -> dict: + now = time.monotonic() + if self.last_refresh is None or now - self.last_refresh >= self.ttl_s: + try: + self.snapshot = await get_multi_lora_controller().snapshot.remote() + self.last_refresh = now + except Exception: + pass + return self.snapshot + + async def get_all(self) -> dict[str, "AdapterRun"]: + snapshot = await self.get_snapshot() + return {**snapshot["active"], **snapshot["retiring"]} + + async def get(self, adapter_name: str) -> "AdapterRun | None": + return (await self.get_all()).get(adapter_name) + + +def _load_subclass(path: str | None, base_cls): + if not path: + return base_cls + cls = load_function(path) + assert issubclass(cls, base_cls), f"{path} must point to a {base_cls.__name__} subclass, got {cls}" + return cls + + +@ray.remote(num_cpus=0) +class MultiLoRAController: + def __init__(self, args, router_url: str, host: str = "0.0.0.0") -> None: + backend_cls = _load_subclass(getattr(args, "multi_lora_backend_path", None), MultiLoRABackend) + server_cls = _load_subclass(getattr(args, "multi_lora_http_server_path", None), MultiLoRAHTTPServer) + self.backend = backend_cls(args, router_url) + self.server = server_cls(self.backend, host, api_port=getattr(args, "multi_lora_api_port", 0)) + + async def start(self) -> int: + await self.backend.init() + await self.server.start() + return self.server.actual_api_port + + async def stop(self) -> None: + await self.server.stop() + await self.backend.close() + + async def register_adapter(self, name: str, config: Any) -> dict: + return await self.backend.register(name, config) + + async def deregister_adapter(self, name: str) -> None: + await self.backend.deregister(name) + + async def retire_adapters(self) -> list[str]: + return await self.backend.retire_adapters() + + async def free_slot(self, name: str) -> int: + return await self.backend.free_slot(name) + + def record_weight_update(self, names: list[str]) -> None: + self.backend.registry.record_weight_update(names) + + def record_batch_adapters(self, rollout_id: int, groups: dict[str, int], step_names: list[str]) -> None: + self.backend.registry.record_batch_adapters(rollout_id, groups, step_names) + + def mark_batch_trained(self, rollout_id: int) -> list[str]: + return self.backend.registry.mark_batch_trained(rollout_id) + + def resolve_num_step(self, name: str, dataset_rows: int) -> None: + self.backend.registry.resolve_num_step(name, dataset_rows) + + def set_adapter_step(self, name: str, step: int) -> None: + self.backend.registry.set_step(name, step) + + def adapter_step(self, name: str) -> int: + return self.backend.registry.step_count(name) + + def snapshot(self) -> dict: + return self.backend.registry.snapshot() + + def http_host(self) -> str: + return get_current_node_ip() + + def api_port(self) -> int: + return self.server.actual_api_port + + +def create_multilora_controller(args, router_url: str, host: str = "0.0.0.0"): + # Pinned to the head node so the API sits at a port-forwardable address. + return MultiLoRAController.options( + name=CONTROLLER_NAME, + namespace=CONTROLLER_NAMESPACE, + **compute_ray_pin_head_options(), + ).remote(args, router_url, host) diff --git a/miles/ray/multi_lora/http_server.py b/miles/ray/multi_lora/http_server.py new file mode 100644 index 0000000000..b209142e1e --- /dev/null +++ b/miles/ray/multi_lora/http_server.py @@ -0,0 +1,129 @@ +"""Multi-LoRA control-plane HTTP API over a MultiLoRABackend. + +Subclass via ``--multi-lora-http-server-path`` (override add_routes / +create_app).""" + +import asyncio +from dataclasses import asdict +from pathlib import Path + +import uvicorn +from fastapi import FastAPI, HTTPException, Query, Request +from fastapi.responses import JSONResponse +from pydantic import BaseModel + +from miles.ray.multi_lora.registry import AdapterState +from miles.utils.adapter_config import AdapterRunConfig, parse_adapter_run_yaml + + +class RegisterAdapterRequest(BaseModel): + """Exactly one of ``config`` (inline) or ``yaml_path`` must be set.""" + + name: str + config: AdapterRunConfig | None = None + yaml_path: str | None = None + + +_NAMES_QUERY = Query(default_factory=list) + + +class MultiLoRAHTTPServer: + """Control-plane API over a MultiLoRABackend. Subclass via + --multi-lora-http-server-path (add_routes / create_app).""" + + def __init__(self, backend, host="127.0.0.1", api_port=0): + self.backend = backend + self.host = host + self.api_port = api_port + self.api_server: uvicorn.Server | None = None + self.api_task: asyncio.Task | None = None + + @property + def actual_api_port(self) -> int: + if self.api_server is not None and self.api_server.started: + return self.api_server.servers[0].sockets[0].getsockname()[1] + return self.api_port + + def create_app(self) -> FastAPI: + app = FastAPI(title="Miles Multi-LoRA Controller") + + @app.exception_handler(ValueError) + async def value_error_handler(request: Request, exc: ValueError): + return JSONResponse({"detail": str(exc)}, status_code=400) + + @app.exception_handler(RuntimeError) + async def runtime_error_handler(request: Request, exc: RuntimeError): + status = 409 if "No free adapter slots" in str(exc) else 500 + return JSONResponse({"detail": str(exc)}, status_code=status) + + return app + + def add_routes(self, app: FastAPI) -> None: + app.get("/health")(self.health) + app.get("/adapter_runs")(self.list_adapters) + app.get("/adapter_runs/state")(self.adapter_states) # before /adapter_runs/{name} + app.get("/adapter_runs/{name}")(self.get_adapter) + app.post("/adapter_runs")(self.register_adapter) + app.delete("/adapter_runs/{name}")(self.deregister_adapter) + + async def start(self) -> None: + app = self.create_app() + self.add_routes(app) + config = uvicorn.Config(app, host=self.host, port=self.api_port, log_level="warning", access_log=False) + self.api_server = uvicorn.Server(config) + self.api_task = asyncio.create_task(self.api_server.serve()) + while not self.api_server.started: + if self.api_task.done(): + self.api_task.result() + raise RuntimeError("uvicorn exited before startup completed") + await asyncio.sleep(0.01) + + async def stop(self) -> None: + if self.api_server is not None: + self.api_server.should_exit = True + await self.api_task + self.api_server = self.api_task = None + + async def health(self) -> dict: + return {"status": "healthy"} + + def adapter_statuses(self) -> list[dict]: + registry = self.backend.registry + statuses = [] + for record in registry.records.values(): + flat = asdict(registry.view(record)) + flat |= flat.pop("config") + flat["save"] = str(flat["save"]) + flat["state"] = record.state + if record.state is AdapterState.COMPLETED: + flat["version"] = None + statuses.append(flat) + return statuses + + async def list_adapters(self) -> dict: + return {"adapters": self.adapter_statuses()} + + async def adapter_states(self, names: list[str] = _NAMES_QUERY) -> dict: + return {"states": {name: self.backend.registry.adapter_state(name) for name in names}} + + async def get_adapter(self, name: str) -> dict: + for status in self.adapter_statuses(): + if status["name"] == name: + return status + raise HTTPException(status_code=404, detail=f"Adapter '{name}' not registered") + + async def register_adapter(self, request: RegisterAdapterRequest) -> dict: + if (request.config is None) == (request.yaml_path is None): + raise HTTPException(status_code=400, detail="Exactly one of 'config' or 'yaml_path' must be set") + if request.yaml_path is not None: + config = parse_adapter_run_yaml(Path(request.yaml_path)) + else: + config = request.config + return await self.backend.register(request.name, config) + + async def deregister_adapter(self, name: str) -> dict: + state = self.backend.registry.adapter_state(name) + if state is None: + raise HTTPException(status_code=404, detail=f"Adapter '{name}' not registered") + await self.backend.deregister(name) + return {"status": "ok", "name": name} diff --git a/miles/ray/multi_lora/registry.py b/miles/ray/multi_lora/registry.py new file mode 100644 index 0000000000..4c8723c29d --- /dev/null +++ b/miles/ray/multi_lora/registry.py @@ -0,0 +1,252 @@ +"""Multi-LoRA adapter registry: the controller-owned lifecycle state machine. + +One record per adapter name, walking PENDING -> ACTIVE -> RETIRING -> CLEANUP +-> COMPLETED. Slots are reused across registrations but ``slot_versions`` +never reset, so a (slot, version) pair never recurs. +""" + +import logging +import re +import uuid +from dataclasses import dataclass, field, replace +from enum import Enum +from pathlib import Path +from typing import Any + +from miles.utils.adapter_config import AdapterRun, AdapterRunConfig + +logger = logging.getLogger(__name__) + +VALID_ADAPTER_NAME = re.compile(r"^[A-Za-z0-9._-]+$") + + +class AdapterState(str, Enum): + PENDING = "PENDING" + ACTIVE = "ACTIVE" + RETIRING = "RETIRING" + CLEANUP = "CLEANUP" + COMPLETED = "COMPLETED" + + +# States that hold a slot. +LIVE_STATES = ( + AdapterState.PENDING, + AdapterState.ACTIVE, + AdapterState.RETIRING, + AdapterState.CLEANUP, +) + + +@dataclass +class AdapterRecord: + name: str + slot: int + config: Any + step: int = 0 + # Baseline step for relative num_step stopping (supports checkpoint resume). + start_step: int = 0 + # Committed prompt groups accumulated toward the current optimizer step. + # Only advanced by mark_batch_trained (after a successful train call). + accumulated_groups: int = 0 + state: AdapterState = AdapterState.PENDING + # Unique per registration: a re-registered name is a new tenant, and + # rollout-side state stamped by the previous tenant must not carry over. + registration_id: str = field(default_factory=lambda: uuid.uuid4().hex) + + +MAX_BATCH_RECORDS = 16 +MAX_COMPLETED_RECORDS = 1024 + + +class AdapterRegistry: + """One record per name; ``slot_versions`` never reset, so (slot, version) + never recurs across slot reuse.""" + + def __init__(self, max_adapters: int) -> None: + self.max_adapters = max_adapters + self.free_slots: set[int] = set(range(max_adapters)) + self.slot_versions: list[int] = [0] * max_adapters + self.records: dict[str, AdapterRecord] = {} + self.batch_records: dict[int, dict] = {} + + def in_state(self, *states: AdapterState) -> dict[str, AdapterRecord]: + return {name: r for name, r in self.records.items() if r.state in states} + + def find(self, name: str) -> AdapterRecord | None: + record = self.records.get(name) + return record if record is not None and record.state in LIVE_STATES else None + + def is_active(self, name: str) -> bool: + record = self.records.get(name) + return record is not None and record.state in (AdapterState.ACTIVE, AdapterState.RETIRING) + + def register(self, name: str, config: Any) -> dict: + if not VALID_ADAPTER_NAME.match(name) or name in (".", ".."): + raise ValueError(f"Adapter name '{name}' is invalid: use only letters, digits, '.', '_' and '-'") + if (existing := self.records.get(name)) is not None: + if existing.state in (AdapterState.PENDING, AdapterState.ACTIVE): + raise ValueError(f"Adapter '{name}' already registered") + if existing.state in (AdapterState.RETIRING, AdapterState.CLEANUP): + raise ValueError(f"Adapter '{name}' is still cleaning up; retry shortly") + if (save_dir := getattr(config, "save", None)) is not None: + for record in self.in_state(*LIVE_STATES).values(): + other_save = getattr(record.config, "save", None) + if other_save is not None and Path(other_save).resolve() == Path(save_dir).resolve(): + raise ValueError( + f"Adapter '{name}' save dir '{save_dir}' is already used by adapter '{record.name}'" + ) + if not self.free_slots: + raise RuntimeError(f"No free adapter slots (max {self.max_adapters})") + slot = min(self.free_slots) + self.free_slots.remove(slot) + self.records.pop(name, None) + self.records[name] = AdapterRecord(name=name, slot=slot, config=config) + return {"name": name, "slot": slot} + + def deregister(self, name: str) -> None: + record = self.records.get(name) + if record is not None and record.state in (AdapterState.PENDING, AdapterState.ACTIVE): + record.state = AdapterState.RETIRING + + def retire_adapters(self) -> list[str]: + retired = sorted(self.in_state(AdapterState.RETIRING)) + for name in retired: + self.records[name].state = AdapterState.CLEANUP + return retired + + def free_slot(self, name: str) -> int: + record = self.records.get(name) + if record is None or record.state is not AdapterState.CLEANUP: + return -1 + self.free_slots.add(record.slot) + record.state = AdapterState.COMPLETED + self.records[name] = self.records.pop(name) + completed = self.in_state(AdapterState.COMPLETED) + for oldest in list(completed)[: len(completed) - MAX_COMPLETED_RECORDS]: + self.records.pop(oldest) + return record.slot + + def adapter_state(self, name: str) -> AdapterState | None: + record = self.records.get(name) + if record is None: + return None + if record.state is AdapterState.COMPLETED: + self.records[name] = self.records.pop(name) + return record.state + + def record_weight_update(self, names: list[str]) -> None: + """A weight push landed: bump slot versions, promote PENDING to ACTIVE.""" + for name in names: + record = self.find(name) + if record is None: + continue + self.slot_versions[record.slot] += 1 + if record.state is AdapterState.PENDING: + record.state = AdapterState.ACTIVE + + def record_batch_adapters(self, rollout_id: int, groups: dict[str, int], step_names: list[str]) -> None: + """Register what a train batch contains before it trains. + + ``groups`` maps adapter name -> prompt groups riding in this batch; + ``step_names`` lists adapters whose adapter batch completes with + this batch (decided by the collection loop, which caps per-adapter + contributions at the adapter's remaining groups). + """ + unknown = set(step_names) - set(groups) + assert not unknown, f"step adapters {sorted(unknown)} not present in batch groups" + self.batch_records[rollout_id] = {"groups": dict(groups), "step_names": list(step_names)} + while len(self.batch_records) > MAX_BATCH_RECORDS: + self.batch_records.pop(next(iter(self.batch_records))) + + def mark_batch_trained(self, rollout_id: int) -> list[str]: + """Bank the batch's trained groups and fire steps; returns adapters that stepped. Only place + accumulation/step state advances, so a failed/retried train call leaves the registry untouched.""" + record_entry = self.batch_records.pop(rollout_id, None) + if record_entry is None: + return [] + stepped = [] + reached_num_step = [] + for name, n_groups in record_entry["groups"].items(): + record = self.records.get(name) + if record is None or record.state not in ( + AdapterState.ACTIVE, + AdapterState.RETIRING, + AdapterState.CLEANUP, + ): + continue + record.accumulated_groups += n_groups + if name in record_entry["step_names"]: + target = record.config.rollout_batch_size + if record.accumulated_groups != target: + logger.warning( + f"Adapter '{name}' stepped with accumulated_groups={record.accumulated_groups} " + f"!= rollout_batch_size={target}; adapter batch accounting drifted" + ) + record.step += 1 + record.accumulated_groups = 0 + stepped.append(name) + if ( + getattr(record.config, "num_step", None) is not None + and record.state is AdapterState.ACTIVE + and (record.step - record.start_step) >= record.config.num_step + ): + reached_num_step.append(name) + for name in reached_num_step: + logger.info( + f"Adapter '{name}' reached num_step={self.records[name].config.num_step} " + f"(start_step={self.records[name].start_step}, step={self.records[name].step}), deregistering" + ) + self.deregister(name) + return stepped + + def resolve_num_step(self, name: str, dataset_rows: int) -> None: + """Derive num_step from num_epoch once the data source knows the + post-filter dataset length. No-op when num_step was set explicitly.""" + record = self.find(name) + if record is None or not isinstance(record.config, AdapterRunConfig): + return + if record.config.num_step is not None: + return + num_epoch = record.config.num_epoch or 1 + num_step = max(1, num_epoch * dataset_rows // record.config.rollout_batch_size) + record.config = replace(record.config, num_step=num_step) + logger.info(f"Adapter '{name}': num_epoch={num_epoch} x {dataset_rows} rows -> num_step={num_step}") + + def set_step(self, name: str, step: int) -> None: + if (record := self.find(name)) is not None: + record.step = step + record.start_step = step + + def step_count(self, name: str) -> int: + record = self.find(name) + return record.step if record is not None else 0 + + def view(self, record: AdapterRecord) -> AdapterRun: + return AdapterRun( + name=record.name, + config=record.config, + slot=record.slot, + version=self.slot_versions[record.slot], + step=record.step, + accumulated_groups=record.accumulated_groups, + registration_id=record.registration_id, + ) + + def active_adapters(self) -> dict[str, AdapterRun]: + """Sampleable view: RETIRING keeps serving until retired.""" + return { + name: self.view(record) + for name, record in self.in_state(AdapterState.ACTIVE, AdapterState.RETIRING).items() + } + + def snapshot(self) -> dict: + def views(state: AdapterState) -> dict[str, AdapterRun]: + return {name: self.view(record) for name, record in self.in_state(state).items()} + + return { + "pending": views(AdapterState.PENDING), + "active": views(AdapterState.ACTIVE), + "retiring": views(AdapterState.RETIRING), + "cleanup": list(self.in_state(AdapterState.CLEANUP)), + "completed": list(self.in_state(AdapterState.COMPLETED)), + } diff --git a/tests/fast/ray/multi_lora/__init__.py b/tests/fast/ray/multi_lora/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/tests/fast/ray/multi_lora/test_controller_backend.py b/tests/fast/ray/multi_lora/test_controller_backend.py new file mode 100644 index 0000000000..fe75f985ef --- /dev/null +++ b/tests/fast/ray/multi_lora/test_controller_backend.py @@ -0,0 +1,392 @@ +"""Fast tests for AdapterRegistry + MultiLoRABackend validation +(no Ray, no HTTP I/O, no SGLang, no torch).""" + +from types import SimpleNamespace + +from tests.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=60, suite="stage-a-cpu") + +import pytest + +from miles.ray.multi_lora.backend import MultiLoRABackend +from miles.ray.multi_lora.registry import AdapterRegistry, AdapterState +from miles.utils.adapter_config import AdapterRunConfig +from miles.utils.multi_lora import make_rid, min_groups_per_dp_split, parse_adapter + + +# Registration validates that the data path exists; the test file itself is a +# convenient always-present stand-in. +DATA_FILE = __file__ + + +def make_args(max_adapters: int = 4, save: str | None = None, dp_size: int = 2) -> SimpleNamespace: + return SimpleNamespace( + multi_lora_n_adapters=max_adapters, + save=save, + lora_rank=32, + lora_alpha=32, + rollout_batch_size=16, + n_samples_per_prompt=4, + multi_lora_dp_size=dp_size, + multi_lora_max_adapter_global_batch_size=256, + ) + + +def make_backend(max_adapters: int = 4, save: str | None = None, dp_size: int = 2) -> MultiLoRABackend: + return MultiLoRABackend(make_args(max_adapters, save, dp_size), "http://unused") + + +def make_config(save: str | None = None, **overrides) -> AdapterRunConfig: + kwargs = dict( + rank=8, + alpha=16, + data=DATA_FILE, + rollout_batch_size=4, + n_samples_per_prompt=4, + save=save, + input_key="text", + label_key="label", + rm_type="math", + ) + kwargs.update(overrides) + return AdapterRunConfig(**kwargs) + + +def register_and_promote(registry: AdapterRegistry, name: str, config=None) -> None: + registry.register(name, config) + registry.record_weight_update([name]) + + +def test_rid_roundtrip_preserves_names_with_underscores(): + for name in ["a", "adapter_a", "weird__name", "x_y_z"]: + assert parse_adapter(make_rid(name)) == name + + +def test_register_starts_pending_and_push_promotes(): + registry = AdapterRegistry(max_adapters=4) + result = registry.register("A", config={"rm_type": "x"}) + assert result == {"name": "A", "slot": 0} + assert registry.active_adapters() == {} # pending: not sampleable + + registry.record_weight_update(["A"]) + assert registry.active_adapters()["A"].slot == 0 + view = registry.active_adapters()["A"] + assert view.slot == 0 + assert view.config == {"rm_type": "x"} + assert view.version == 1 + + +def test_snapshot_reports_sets_in_registry_vocabulary(): + registry = AdapterRegistry(max_adapters=4) + register_and_promote(registry, "A") + registry.register("B", None) + snapshot = registry.snapshot() + assert set(snapshot["active"]) == {"A"} + assert set(snapshot["pending"]) == {"B"} + assert snapshot["retiring"] == {} + assert snapshot["cleanup"] == [] + assert set(registry.active_adapters()) == {"A"} # only active adapters are sampleable + + +def test_slot_version_is_monotonic_across_slot_reuse(): + registry = AdapterRegistry(max_adapters=2) + register_and_promote(registry, "A") # slot 0, version 1 + registry.record_weight_update(["A"]) # version 2 + registry.deregister("A") + registry.retire_adapters() + registry.free_slot("A") + + registry.register("A2", None) # reuses slot 0 + assert registry.snapshot()["pending"]["A2"].version == 2 # inherits, not reset + registry.record_weight_update(["A2"]) + assert registry.active_adapters()["A2"].version == 3 + + +def test_record_weight_update_only_touches_reported_names(): + registry = AdapterRegistry(max_adapters=4) + register_and_promote(registry, "A") + register_and_promote(registry, "B") + registry.record_weight_update(["A"]) + assert registry.active_adapters()["A"].version == 2 + assert registry.active_adapters()["B"].version == 1 + + +def test_register_name_rejected_until_cleanup_done(): + registry = AdapterRegistry(max_adapters=4) + register_and_promote(registry, "A") + registry.deregister("A") + with pytest.raises(ValueError, match="cleaning up"): + registry.register("A", None) # retiring + registry.retire_adapters() + with pytest.raises(ValueError, match="cleaning up"): + registry.register("A", None) # cleanup + registry.free_slot("A") + assert registry.register("A", None) == {"name": "A", "slot": 0} + + +def test_deregister_retires_but_keeps_serving_until_demoted(): + registry = AdapterRegistry(max_adapters=4) + register_and_promote(registry, "A") + registry.deregister("A") + assert registry.adapter_state("A") == AdapterState.RETIRING + assert "A" in registry.active_adapters() # still sampleable this iteration + assert "A" in registry.snapshot()["retiring"] + assert registry.retire_adapters() == ["A"] + assert registry.active_adapters() == {} + assert registry.adapter_state("A") == AdapterState.CLEANUP + assert registry.retire_adapters() == [] # idempotent + + +# make_config(): rollout_batch_size=4 groups/step, n_samples_per_prompt=4. + + +def test_mark_batch_trained_accumulates_and_steps_on_completion(): + registry = AdapterRegistry(max_adapters=4) + register_and_promote(registry, "A", make_config()) + register_and_promote(registry, "B", make_config()) + + # Two partial batches accumulate; the third completes the adapter batch. + registry.record_batch_adapters(1, {"A": 1, "B": 2}, step_names=[]) + assert registry.mark_batch_trained(1) == [] + assert registry.records["A"].accumulated_groups == 1 + assert registry.records["B"].accumulated_groups == 2 + + registry.record_batch_adapters(2, {"A": 1}, step_names=[]) + assert registry.mark_batch_trained(2) == [] + assert registry.records["A"].accumulated_groups == 2 + + registry.record_batch_adapters(3, {"A": 2, "B": 2}, step_names=["A", "B"]) + assert registry.mark_batch_trained(3) == ["A", "B"] + assert registry.step_count("A") == 1 + assert registry.step_count("B") == 1 + assert registry.records["A"].accumulated_groups == 0 + assert registry.records["B"].accumulated_groups == 0 + + assert registry.mark_batch_trained(3) == [] # record consumed + + +def test_batch_trained_counts_deregistered_adapter_until_freed(): + registry = AdapterRegistry(max_adapters=4) + register_and_promote(registry, "A", make_config()) + registry.record_batch_adapters(3, {"A": 4}, step_names=["A"]) + registry.deregister("A") # deregistered while its batch is training + assert registry.mark_batch_trained(3) == ["A"] + assert registry.step_count("A") == 1 # final ckpt reads this + registry.retire_adapters() + assert registry.step_count("A") == 1 # cleanup record still holds it + registry.free_slot("A") + assert registry.step_count("A") == 0 + + +def test_set_step_on_resume(): + registry = AdapterRegistry(max_adapters=2) + registry.register("A", make_config()) + registry.set_step("A", 40) + registry.record_batch_adapters(1, {"A": 4}, step_names=["A"]) + registry.record_weight_update(["A"]) + registry.mark_batch_trained(1) + assert registry.step_count("A") == 41 + + +def test_num_step_deregisters_on_committed_steps(): + registry = AdapterRegistry(max_adapters=2) + register_and_promote(registry, "A", make_config(num_step=2)) + registry.record_batch_adapters(1, {"A": 4}, step_names=["A"]) + assert registry.mark_batch_trained(1) == ["A"] + assert registry.adapter_state("A") == AdapterState.ACTIVE + + registry.record_batch_adapters(2, {"A": 4}, step_names=["A"]) + assert registry.mark_batch_trained(2) == ["A"] + assert registry.step_count("A") == 2 + assert registry.adapter_state("A") == AdapterState.RETIRING + + +def test_num_step_is_relative_to_resume_step(): + registry = AdapterRegistry(max_adapters=2) + register_and_promote(registry, "A", make_config(num_step=2)) + registry.set_step("A", 40) + + registry.record_batch_adapters(1, {"A": 4}, step_names=["A"]) + registry.mark_batch_trained(1) + assert registry.step_count("A") == 41 + assert registry.adapter_state("A") == AdapterState.ACTIVE + + registry.record_batch_adapters(2, {"A": 4}, step_names=["A"]) + registry.mark_batch_trained(2) + assert registry.step_count("A") == 42 + assert registry.adapter_state("A") == AdapterState.RETIRING + + +def test_min_groups_per_dp_split(): + assert min_groups_per_dp_split(n_samples_per_prompt=4, dp_size=8) == 2 # divisor + assert min_groups_per_dp_split(n_samples_per_prompt=8, dp_size=8) == 1 # equal + assert min_groups_per_dp_split(n_samples_per_prompt=16, dp_size=8) == 1 # multiple + with pytest.raises(ValueError, match="divisor or a multiple"): + min_groups_per_dp_split(n_samples_per_prompt=6, dp_size=8) + + +@pytest.mark.asyncio +async def test_register_resolves_batch_shape_defaults(tmp_path): + backend = make_backend(save=str(tmp_path)) + await backend.register("A", AdapterRunConfig(data=DATA_FILE, rm_type="math")) + config = backend.registry.records["A"].config + assert config.rollout_batch_size == 16 # <- args.rollout_batch_size + assert config.n_samples_per_prompt == 4 # <- args.n_samples_per_prompt + assert config.rank == 32 and config.alpha == 32 + assert config.adapter_global_batch_size == 64 + + +@pytest.mark.asyncio +async def test_register_rejects_bad_batch_shapes(tmp_path): + backend = make_backend(save=str(tmp_path), dp_size=8) + with pytest.raises(ValueError, match="divisor or a multiple"): + await backend.register("B", make_config(n_samples_per_prompt=6, rollout_batch_size=4)) + with pytest.raises(ValueError, match="min_groups_per_dp_split"): + # dp=8, n_samples=4 -> multiple of 2 groups; 3 groups is not + await backend.register("C", make_config(rollout_batch_size=3)) + with pytest.raises(ValueError, match="exceeding"): + await backend.register("D", make_config(rollout_batch_size=128)) # 512 samples > cap 256 + with pytest.raises(ValueError, match="exceeds the allocated maximum rank"): + await backend.register("E", make_config(rank=64)) + with pytest.raises(ValueError, match="positive integer"): + await backend.register("F", make_config(rollout_batch_size=0)) + with pytest.raises(ValueError, match="num_step must be a positive integer"): + await backend.register("G", make_config(num_step=0)) + with pytest.raises(ValueError, match="num_epoch must be a positive integer"): + await backend.register("H", make_config(num_epoch=0)) + # A valid shape registers fine. + await backend.register("OK", make_config(rollout_batch_size=8)) + + +def test_deregister_holds_slot_until_free_slot(): + registry = AdapterRegistry(max_adapters=2) + register_and_promote(registry, "A") # slot 0 + register_and_promote(registry, "B") # slot 1 + registry.deregister("A") + registry.retire_adapters() + assert not registry.free_slots # slot 0 held until cleanup + with pytest.raises(RuntimeError, match="No free adapter slots"): + registry.register("C", None) + registry.free_slot("A") + assert registry.register("C", None) == {"name": "C", "slot": 0} + + +@pytest.mark.asyncio +async def test_free_slot_reaborts_before_releasing_slot(): + """Requests can survive the single retire-time abort (multi-turn groups + submitting between turns, engine tokenizer-adapter batch misses); free_slot must + fire one more abort round before the slot becomes reusable.""" + backend = make_backend() + aborted: list[str] = [] + + async def record_abort(name: str) -> None: + aborted.append(name) + + backend.abort_adapter_requests = record_abort + + register_and_promote(backend.registry, "A") + await backend.deregister("A") + await backend.retire_adapters() + assert aborted == ["A"] + + assert await backend.free_slot("A") == 0 + assert aborted == ["A", "A"] + assert backend.registry.free_slots == {0, 1, 2, 3} + + +@pytest.mark.asyncio +async def test_free_slot_skips_abort_when_not_in_cleanup(): + backend = make_backend() + aborted: list[str] = [] + + async def record_abort(name: str) -> None: + aborted.append(name) + + backend.abort_adapter_requests = record_abort + + register_and_promote(backend.registry, "A") # ACTIVE, not CLEANUP + assert await backend.free_slot("A") == -1 + assert await backend.free_slot("never-registered") == -1 + assert aborted == [] + + +@pytest.mark.asyncio +async def test_custom_backend_validation_rejects(): + class StrictBackend(MultiLoRABackend): + async def validate_adapter(self, name, config): + if not config: + raise ValueError("adapter config is required") + + backend = StrictBackend(make_args(), "http://unused") + with pytest.raises(ValueError, match="config is required"): + await backend.register("A", None) + assert backend.registry.active_adapters() == {} + + result = await backend.register("A", {"rm_type": "x"}) + assert result == {"name": "A", "slot": 0} + + +def test_register_rejects_unsafe_names(): + registry = AdapterRegistry(max_adapters=4) + for bad in ["a/b", "..", "a::b", "a b", ""]: + with pytest.raises(ValueError, match="invalid"): + registry.register(bad, None) + registry.register("ok-name_1.2", None) + + +def test_register_rejects_duplicate_save_dir(tmp_path): + registry = AdapterRegistry(max_adapters=4) + registry.register("A", make_config(save=tmp_path / "x")) + with pytest.raises(ValueError, match="already used by adapter 'A'"): + registry.register("B", make_config(save=tmp_path / "x")) + registry.register("C", make_config(save=tmp_path / "y")) + + +@pytest.mark.asyncio +async def test_save_dir_defaults_under_save_root(tmp_path): + backend = make_backend(save=str(tmp_path)) + await backend.register("A", make_config()) + saved = backend.registry.records["A"].config.save + assert saved == tmp_path / "adapters" / "A" + + +@pytest.mark.asyncio +async def test_explicit_save_dir_wins_over_root(tmp_path): + backend = make_backend(save=str(tmp_path)) + await backend.register("A", make_config(save=tmp_path / "custom")) + assert backend.registry.records["A"].config.save == tmp_path / "custom" + + +@pytest.mark.asyncio +async def test_register_fails_without_any_save_dir(): + backend = make_backend(save=None) + with pytest.raises(ValueError, match="no save dir"): + await backend.register("A", make_config()) + + +@pytest.mark.asyncio +async def test_register_rejects_missing_data_path(tmp_path): + # A nonexistent data path would otherwise kill the shared rollout producer + # thread at the first get_samples, stalling every adapter. + backend = make_backend(save=str(tmp_path)) + with pytest.raises(ValueError, match="data path"): + await backend.register("A", make_config(data=str(tmp_path / "missing.jsonl"))) + + +@pytest.mark.asyncio +async def test_register_rejects_unresolvable_reward_config(tmp_path): + # No adapter rm_type/custom_rm_path and no process-wide --rm-type: every + # sample would fail reward computation and be dropped. + backend = make_backend(save=str(tmp_path)) + with pytest.raises(ValueError, match="reward config"): + await backend.register("A", make_config(rm_type=None)) + + +@pytest.mark.asyncio +async def test_register_accepts_reward_config_from_global_args(tmp_path): + args = make_args(save=str(tmp_path)) + args.rm_type = "math" + backend = MultiLoRABackend(args, "http://unused") + await backend.register("A", make_config(rm_type=None)) + assert backend.registry.records["A"].config.rm_type is None # resolved at reward time via args diff --git a/tests/fast/ray/multi_lora/test_controller_http.py b/tests/fast/ray/multi_lora/test_controller_http.py new file mode 100644 index 0000000000..b270017597 --- /dev/null +++ b/tests/fast/ray/multi_lora/test_controller_http.py @@ -0,0 +1,239 @@ +"""HTTP tests for the MultiLoRAHTTPServer control plane with a mock router +(no Ray, no SGLang).""" + +import json +from contextlib import asynccontextmanager +from pathlib import Path +from types import SimpleNamespace + +import aiohttp +import pytest +from aiohttp import web + +from tests.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=60, suite="stage-a-cpu") + +from miles.ray.multi_lora.backend import MultiLoRABackend +from miles.ray.multi_lora.http_server import MultiLoRAHTTPServer +from miles.utils.adapter_config import AdapterRunConfig +from miles.utils.multi_lora import RID_SEPARATOR + + +# Registration validates that the data path exists; the test file itself is a +# convenient always-present stand-in. +DATA_FILE = __file__ + + +def minimal_config(name: str) -> dict: + return {"data": DATA_FILE, "rm_type": "math", "save": f"/tmp/adapters/{name}"} + + +class ControllerHarness: + """Running control plane (backend + API listener) against a mock router + that serves /list_workers and records /abort_request posts.""" + + def __init__(self, session: aiohttp.ClientSession, backend: MultiLoRABackend, srv: MultiLoRAHTTPServer): + self.session = session + self.backend = backend + self.srv = srv + self.aborts: list[dict] = [] + + @property + def api_base(self) -> str: + return f"http://127.0.0.1:{self.srv.actual_api_port}" + + async def api_post(self, path: str, payload: dict) -> tuple[int, dict]: + async with self.session.post(f"{self.api_base}{path}", json=payload) as resp: + return resp.status, await resp.json() + + async def api_get(self, path: str) -> tuple[int, dict, dict]: + async with self.session.get(f"{self.api_base}{path}") as resp: + headers = {k.lower(): v for k, v in resp.headers.items()} + return resp.status, await resp.json(), headers + + async def api_delete(self, path: str) -> tuple[int, dict]: + async with self.session.delete(f"{self.api_base}{path}") as resp: + return resp.status, await resp.json() + + async def register(self, name: str) -> tuple[int, dict]: + status, body = await self.api_post("/adapter_runs", {"name": name, "config": minimal_config(name)}) + # Registered adapters start pending; a weight push promotes them. + self.backend.registry.record_weight_update([name]) + return status, body + + async def deregister(self, name: str) -> tuple[int, dict]: + return await self.api_delete(f"/adapter_runs/{name}") + + async def active(self) -> dict: + _, body, _ = await self.api_get("/adapter_runs") + return { + s["name"]: {"slot": s["slot"], "version": s["version"], "step": s["step"]} + for s in body["adapters"] + if s["state"] == "ACTIVE" + } + + +@asynccontextmanager +async def running_controller(server_cls=MultiLoRAHTTPServer): + router_url = "" + harness: ControllerHarness | None = None + + async def router_handler(request): + if request.path == "/list_workers": + return web.json_response({"urls": [router_url]}) + if request.path == "/abort_request": + harness.aborts.append(json.loads(await request.read())) + return web.json_response({}) + return web.json_response({}, status=404) + + app = web.Application() + app.router.add_resource("/{tail:.*}").add_route("*", router_handler) + runner = web.AppRunner(app) + await runner.setup() + site = web.TCPSite(runner, "127.0.0.1", 0) + await site.start() + router_url = f"http://127.0.0.1:{site._server.sockets[0].getsockname()[1]}" + + backend = MultiLoRABackend( + SimpleNamespace( + multi_lora_n_adapters=4, + save=None, + lora_rank=32, + lora_alpha=32, + rollout_batch_size=16, + n_samples_per_prompt=4, + multi_lora_dp_size=2, + multi_lora_max_adapter_global_batch_size=256, + ), + router_url, + ) + srv = server_cls(backend) + await backend.init() + await srv.start() + try: + async with aiohttp.ClientSession() as session: + harness = ControllerHarness(session, backend, srv) + yield harness + finally: + await srv.stop() + await backend.close() + await runner.cleanup() + + +@pytest.mark.asyncio +async def test_register_and_active_view(): + async with running_controller() as ctl: + status, body = await ctl.register("A") + assert status == 200 + assert body["slot"] == 0 + assert await ctl.active() == {"A": {"slot": 0, "version": 1, "step": 0}} + + +@pytest.mark.asyncio +async def test_deregister_marks_and_retire_adapters_aborts(): + """Deregistration only marks; the driver-synced apply performs the + demotion and fans out one prefix abort per worker.""" + async with running_controller() as ctl: + await ctl.register("A") + status, _ = await ctl.deregister("A") + assert status == 200 + assert ctl.aborts == [] # still serving until the sync point + assert "A" in ctl.backend.registry.active_adapters() + + applied = await ctl.backend.retire_adapters() + assert applied == ["A"] + assert ctl.aborts == [{"rid": f"A{RID_SEPARATOR}", "prefix": True}] + assert ctl.backend.registry.active_adapters() == {} + + +@pytest.mark.asyncio +async def test_register_json_config_validates_to_adapter_config(): + """FastAPI validates the JSON body straight into AdapterRunConfig (422 on bad + payloads).""" + async with running_controller() as ctl: + config = { + "rank": 8, + "data": DATA_FILE, + "save": "/tmp/adapters/A", + "rm_type": "math", + } + status, _ = await ctl.api_post("/adapter_runs", {"name": "A", "config": config}) + assert status == 200 + record = ctl.backend.registry.find("A") + assert isinstance(record.config, AdapterRunConfig) + assert record.config.data == DATA_FILE + assert Path(record.config.save) == Path("/tmp/adapters/A") + assert record.config.input_key == "text" # dataclass default + + status, _ = await ctl.api_post("/adapter_runs", {"name": "B", "config": {"rank": 8}}) + assert status == 422 # data is required + + status, _ = await ctl.api_post("/adapter_runs", {"name": "C"}) + assert status == 400 # exactly one of config/yaml_path + + +@pytest.mark.asyncio +async def test_state_endpoint_reports_lifecycle_and_completed(): + """States walk PENDING -> ACTIVE -> RETIRING -> CLEANUP -> COMPLETED; + unknown names report null; COMPLETED is retained after free_slot.""" + async with running_controller() as ctl: + await ctl.api_post("/adapter_runs", {"name": "A", "config": minimal_config("A")}) + + async def state_of(name): + _, body, _ = await ctl.api_get(f"/adapter_runs/state?names={name}") + return body["states"][name] + + assert await state_of("A") == "PENDING" + ctl.backend.registry.record_weight_update(["A"]) + assert await state_of("A") == "ACTIVE" + + await ctl.deregister("A") + assert await state_of("A") == "RETIRING" + await ctl.backend.retire_adapters() + assert await state_of("A") == "CLEANUP" + + ctl.backend.registry.free_slot("A") + assert await state_of("A") == "COMPLETED" + assert await state_of("nope") is None + + # GET by name serves the completed record; DELETE of unknown 404s. + status, body, _ = await ctl.api_get("/adapter_runs/A") + assert status == 200 and body["state"] == "COMPLETED" + status, _ = await ctl.api_delete("/adapter_runs/nope") + assert status == 404 + + # Re-registration reclaims the name; the completed record is dropped. + status, _ = await ctl.api_post( + "/adapter_runs", + {"name": "A", "config": {"data": DATA_FILE, "rm_type": "math", "save": "/tmp/adapters/A2"}}, + ) + assert status == 200 + assert await state_of("A") == "PENDING" + + +@pytest.mark.asyncio +async def test_custom_server_subclass_adds_routes(): + class CustomServer(MultiLoRAHTTPServer): + def create_app(self): + app = super().create_app() + + @app.middleware("http") + async def tag_response(request, call_next): + response = await call_next(request) + response.headers["X-Custom-Server"] = "1" + return response + + return app + + def add_routes(self, app): + super().add_routes(app) + app.get("/custom_status")(self.custom_status) + + async def custom_status(self): + return {"custom": True, "active": sorted(self.backend.registry.active_adapters())} + + async with running_controller(server_cls=CustomServer) as ctl: + _, body, headers = await ctl.api_get("/custom_status") + assert headers.get("x-custom-server") == "1" + assert body == {"custom": True, "active": []} From fad91f81b81b4ee79981a04ad95d6a9193cf6a48 Mon Sep 17 00:00:00 2001 From: Yusheng Su Date: Tue, 21 Jul 2026 00:31:30 -0700 Subject: [PATCH 3/3] =?UTF-8?q?[multi-lora]=203/7:=20trainer=20core=20?= =?UTF-8?q?=E2=80=94=20per-slot=20optimizers,=20per-adapter=20LR=20schedul?= =?UTF-8?q?es,=20slot=20lifecycle,=20batch=20routing=20in=20get=5Fbatch?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- miles/backends/megatron_utils/arguments.py | 3 + .../megatron_utils/bridge_lora_helpers.py | 16 +- miles/backends/megatron_utils/model.py | 101 +++- .../megatron_utils/multi_lora_optimizer.py | 195 ++++++++ .../megatron_utils/multi_lora_scheduler.py | 88 ++++ .../megatron_utils/multi_lora_utils.py | 441 ++++++++++++++++++ miles/backends/training_utils/data.py | 42 +- miles/backends/training_utils/log_utils.py | 6 + miles/backends/training_utils/loss.py | 6 + .../test_multi_lora_scheduler.py | 114 +++++ .../test_multi_lora_slot_cleanup.py | 91 ++++ .../megatron_utils/test_slice_lora_to_rank.py | 48 ++ .../test_get_batch_multi_lora_cp.py | 134 ++++++ 13 files changed, 1255 insertions(+), 30 deletions(-) create mode 100644 miles/backends/megatron_utils/multi_lora_optimizer.py create mode 100644 miles/backends/megatron_utils/multi_lora_scheduler.py create mode 100644 miles/backends/megatron_utils/multi_lora_utils.py create mode 100644 tests/fast/backends/megatron_utils/test_multi_lora_scheduler.py create mode 100644 tests/fast/backends/megatron_utils/test_multi_lora_slot_cleanup.py create mode 100644 tests/fast/backends/megatron_utils/test_slice_lora_to_rank.py create mode 100644 tests/fast/backends/training_utils/test_get_batch_multi_lora_cp.py diff --git a/miles/backends/megatron_utils/arguments.py b/miles/backends/megatron_utils/arguments.py index 668fb6b9ac..06e198b661 100644 --- a/miles/backends/megatron_utils/arguments.py +++ b/miles/backends/megatron_utils/arguments.py @@ -15,6 +15,9 @@ def set_default_megatron_args(args): args.use_distributed_optimizer = (args.optimizer is None or args.optimizer.lower() == "adam") and not getattr( args, "debug_disable_optimizer", False ) + # Multi-LoRA: per-slot LayerWise optimizers require plain DDP all-reduce. + if getattr(args, "multi_lora_n_adapters", 0) > 0: + args.use_distributed_optimizer = False # TODO: maybe change this after megatron has good fp8 support args.bf16 = not args.fp16 # placeholders diff --git a/miles/backends/megatron_utils/bridge_lora_helpers.py b/miles/backends/megatron_utils/bridge_lora_helpers.py index c56b516c7c..9114a14c5b 100644 --- a/miles/backends/megatron_utils/bridge_lora_helpers.py +++ b/miles/backends/megatron_utils/bridge_lora_helpers.py @@ -12,8 +12,9 @@ from megatron.core.utils import get_attr_wrapped_model from miles.utils.hf_config import load_hf_config +from miles.utils.multi_lora import is_multi_lora_enabled -from .lora_utils import create_lora_instance, patch_param_grad_buffer_for_colocate_mode_lora +from .lora_utils import patch_param_grad_buffer_for_colocate_mode_lora @dataclass @@ -111,7 +112,14 @@ def _setup_lora_model_via_bridge(args: Namespace) -> list: provider.dsa_attention_backend = getattr(args, "dsa_attention_backend", "megatron") provider.finalize() - lora = create_lora_instance(args) + if is_multi_lora_enabled(args): + from miles.backends.megatron_utils.multi_lora_utils import create_multi_lora_instance + + lora = create_multi_lora_instance(args) + else: + from .lora_utils import create_lora_instance + + lora = create_lora_instance(args) def apply_lora_hook(model_chunks): transformed = lora(model_chunks, training=True) @@ -129,6 +137,10 @@ def apply_lora_hook(model_chunks): provider.register_pre_wrap_hook(_make_value_model_hook(hidden_size, provider.sequence_parallel)) use_distributed_optimizer = "muon" not in (args.optimizer or "").lower() + if is_multi_lora_enabled(args): + # Per-slot LayerWise optimizers: plain DDP all-reduce keeps full grads on + # every rank (whole-param sharding + retained-gradient idempotency). + use_distributed_optimizer = False ddp_config = DistributedDataParallelConfig( use_distributed_optimizer=use_distributed_optimizer, grad_reduce_in_fp32=args.accumulate_allreduce_grads_in_fp32, diff --git a/miles/backends/megatron_utils/model.py b/miles/backends/megatron_utils/model.py index 61a49fe395..30e70661b9 100644 --- a/miles/backends/megatron_utils/model.py +++ b/miles/backends/megatron_utils/model.py @@ -6,6 +6,7 @@ import math from argparse import Namespace from collections.abc import Callable, Sequence +from contextlib import nullcontext from functools import partial from pathlib import Path @@ -31,6 +32,7 @@ from miles.utils.audit_utils.witness.module import witness_dump_and_clear_stale from miles.utils.dumper_utils import DumperMegatronUtil, DumperPhase from miles.utils.memory_utils import clear_memory +from miles.utils.multi_lora import is_multi_lora_enabled from miles.utils.test_utils.ft_test_actions import FTTestActionActorExecutor from miles.utils.tracking_utils.structured_log import log_structured @@ -137,7 +139,11 @@ def setup_model_and_optimizer( assert not args.moe_use_upcycling assert args.load is not None or args.pretrained_checkpoint is not None - if is_lora_enabled(args) and role == "actor" and args.megatron_to_hf_mode == "bridge": + # Multi-LoRA and single-LoRA (actor, bridge) both build via the bridge helper, + # which picks the adapter type internally. + if is_multi_lora_enabled(args) or ( + is_lora_enabled(args) and role == "actor" and args.megatron_to_hf_mode == "bridge" + ): model = _setup_lora_model_via_bridge(args) else: model = get_model(get_model_provider_func(args, role), ModelType.encoder_or_decoder) @@ -165,6 +171,10 @@ def setup_model_and_optimizer( use_gloo_process_groups=args.enable_gloo_process_groups, layer_wise_distributed_optimizer="dist" in config.optimizer.lower(), ) + elif is_multi_lora_enabled(args): + from miles.backends.megatron_utils.multi_lora_optimizer import build_multi_lora_optimizer + + optimizer = build_multi_lora_optimizer(args, config, model) else: optimizer = get_megatron_optimizer( config=config, @@ -291,6 +301,11 @@ def forward_step( total_lengths = batch["total_lengths"] response_lengths = batch["response_lengths"] + if "adapter_token_counts" in batch: + from megatron.bridge.peft.multi_lora_layers import set_tokens_per_adapter_slot + + set_tokens_per_adapter_slot(model, batch["adapter_token_counts"]) + output_tensor = model( input_ids=tokens, position_ids=None, @@ -355,6 +370,13 @@ def forward_step( return rollout_data +def _zero_grads(model: Sequence[DDP], optimizer: MegatronOptimizer | None, disable_optimizer: bool) -> None: + for model_chunk in model: + model_chunk.zero_grad_buffer() + if not disable_optimizer: + optimizer.zero_grad() + + def train_one_step( args: Namespace, rollout_id: int, @@ -373,6 +395,10 @@ def train_one_step( Runs forward/backward over ``num_microbatches``, applies optimizer step and one scheduler step when gradients are valid. + Multi-LoRA: gradients are retained across train calls (per-adapter + gradient accumulation); only the slots in the batch's ``step_slots`` step, + and only their gradients are zeroed. + Args: args: Runtime arguments. rollout_id: Rollout identifier. @@ -390,12 +416,16 @@ def train_one_step( parallel_state = get_parallel_state() dumper_phase_util = DumperMegatronUtil(args, model, DumperPhase.FWD_BWD, rollout_id=rollout_id) disable_optimizer = args.debug_disable_optimizer or optimizer is None + multi_lora = is_multi_lora_enabled(args) - # Set grad to zero. - for model_chunk in model: - model_chunk.zero_grad_buffer() - if not disable_optimizer: - optimizer.zero_grad() + if multi_lora: + from miles.backends.megatron_utils.multi_lora_optimizer import reset_grad_metadata_keep_grads + + # Retain accumulated per-adapter gradients; reset only the per-iteration + # DDP bookkeeping. Slot grads are zeroed selectively at step time. + reset_grad_metadata_keep_grads(model) + else: + _zero_grads(model, optimizer, disable_optimizer) if args.custom_megatron_before_train_step_hook_path: from miles.utils.misc import load_function @@ -445,6 +475,11 @@ def forward_step(data_iterator: DataIterator, model: GPTModel, return_schedule_p allgather_cp=args.allgather_cp, ) + if "adapter_token_counts" in batch: + from megatron.bridge.peft.multi_lora_layers import set_tokens_per_adapter_slot + + set_tokens_per_adapter_slot(model, batch["adapter_token_counts"]) + from miles.utils.replay_base import all_replay_managers old_stages = [m.stage for m in all_replay_managers] @@ -517,7 +552,7 @@ def forward_step(data_iterator: DataIterator, model: GPTModel, return_schedule_p outcome = TrainStepOutcome.DISCARDED_SHOULD_RETRY valid_step = False - if (not disable_optimizer) and (not getattr(args, "check_for_nan_in_loss_and_grad", True)): + if (not disable_optimizer) and (not multi_lora) and (not getattr(args, "check_for_nan_in_loss_and_grad", True)): found_inf_flag = optimizer.prepare_grads() if found_inf_flag: valid_step = False @@ -542,18 +577,24 @@ def forward_step(data_iterator: DataIterator, model: GPTModel, return_schedule_p dumper_phase_util.finalize(model) if not disable_optimizer and valid_step: - # Update parameters. - update_successful, grad_norm, num_zeros_in_grad = optimizer.step() + if multi_lora: + from miles.backends.megatron_utils.multi_lora_utils import step_stepped_adapter_slots - # Update learning rate. - assert update_successful - opt_param_scheduler.step(increment=args.global_batch_size) + grad_norm = step_stepped_adapter_slots( + args, model, optimizer, data_iterator[0].rollout_data, rollout_id, step_id + ) + else: + # Update parameters. + update_successful, grad_norm, num_zeros_in_grad = optimizer.step() - # release grad - for model_chunk in model: - model_chunk.zero_grad_buffer() - if not disable_optimizer: - optimizer.zero_grad() + # Update learning rate. + assert update_successful + opt_param_scheduler.step(increment=args.global_batch_size) + + # release grad (multi-LoRA retains accumulated grads; stepped slots were + # zeroed selectively inside step_adapter_slots) + if not multi_lora: + _zero_grads(model, optimizer, disable_optimizer) log_structured( logger.info, @@ -909,13 +950,25 @@ def initialize_model_and_optimizer( model, optimizer, opt_param_scheduler = setup_model_and_optimizer(args, role) model[0].role = role clear_memory() - iteration, _ = load_checkpoint( - model, - optimizer, - opt_param_scheduler, - checkpointing_context=checkpointing_context, - skip_load_to_model_and_opt=False, - ) + + multi_lora = is_multi_lora_enabled(args) + if multi_lora: + # Hide adapter params so the bridge's conversion-task walk doesn't see them + # while loading the base checkpoint. + from megatron.bridge.peft.multi_lora_layers import hide_adapters + + load_ctx = hide_adapters(model) + else: + load_ctx = nullcontext() + + with load_ctx: + iteration, _ = load_checkpoint( + model, + optimizer, + opt_param_scheduler, + checkpointing_context=checkpointing_context, + skip_load_to_model_and_opt=False, + ) check_peak_gpu_memory_after_load(args) clear_memory() diff --git a/miles/backends/megatron_utils/multi_lora_optimizer.py b/miles/backends/megatron_utils/multi_lora_optimizer.py new file mode 100644 index 0000000000..1397bd7569 --- /dev/null +++ b/miles/backends/megatron_utils/multi_lora_optimizer.py @@ -0,0 +1,195 @@ +"""Per-slot decoupled Adam optimizers for multi-LoRA, chained under Megatron's LayerWiseDistributedOptimizer; +requires plain DDP all-reduce (use_distributed_optimizer OFF) so cross-batch gradient retention stays idempotent.""" + +import logging +from argparse import Namespace +from collections.abc import Sequence +from contextlib import contextmanager + +import torch +from megatron.core.optimizer import get_megatron_optimizer +from megatron.core.optimizer.clip_grads import clip_grad_by_total_norm_fp32, get_grad_norm_fp32 +from megatron.core.optimizer.layer_wise_optimizer import LayerWiseDistributedOptimizer +from megatron.core.optimizer.optimizer import MegatronOptimizer +from megatron.core.optimizer.optimizer_config import OptimizerConfig +from megatron.core.process_groups_config import ProcessGroupCollection + +logger = logging.getLogger(__name__) + + +def adapter_slot_parameters(model, slot: int) -> list[torch.nn.Parameter]: + """All parameters belonging to one adapter slot, across model chunks.""" + from megatron.bridge.peft.multi_lora_layers import MultiLoRALinear + + parameters = [] + seen = set() + model_chunks = model if isinstance(model, (list, tuple)) else [model] + for model_chunk in model_chunks: + for module in model_chunk.modules(): + if not isinstance(module, MultiLoRALinear): + continue + for param in module.adapters[slot].parameters(): + if id(param) not in seen: + parameters.append(param) + seen.add(id(param)) + return parameters + + +def _adam_init_state_fn(opt, config=None): + for group in opt.param_groups: + for p in group["params"]: + if len(opt.state[p]) == 0: + opt.state[p]["exp_avg"] = torch.zeros_like(p.data) + opt.state[p]["exp_avg_sq"] = torch.zeros_like(p.data) + + +@contextmanager +def _only_slot_trainable(model_chunks, slot_params: list[torch.nn.Parameter]): + """Temporarily freeze every trainable param outside ``slot_params`` so the + stock param-group builder sees exactly one slot (the Muon construction + pattern from megatron's ``get_megatron_muon_optimizer``).""" + slot_ids = {id(p) for p in slot_params} + frozen = [] + for model_chunk in model_chunks: + for param in model_chunk.parameters(): + if param.requires_grad and id(param) not in slot_ids: + param.requires_grad = False + frozen.append(param) + try: + yield + finally: + for param in frozen: + param.requires_grad = True + + +def build_multi_lora_optimizer( + args: Namespace, + config: OptimizerConfig, + model_chunks: Sequence, +) -> MegatronOptimizer: + """Build one Float16-wrapped Adam per adapter slot under a LayerWiseDistributedOptimizer (ChainedOptimizer); + each child's param groups are tagged with ``miles_multi_lora_slot`` and narrowed to this rank's shard.""" + assert not config.use_distributed_optimizer, ( + "multi-LoRA per-slot optimizers require use_distributed_optimizer=False: " + "gradient retention relies on all-reduce idempotency, and LayerWise " + "sharding replaces byte-level ZeRO" + ) + assert not config.fp16, "multi-LoRA per-slot optimizers require bf16 (no dynamic loss scaler)" + assert (config.optimizer or "").lower() == "adam", ( + "multi-LoRA per-slot optimizers only implement Adam semantics (state init, " + f"slot retirement cleanup, step clocks); got optimizer={config.optimizer!r}" + ) + + pg_collection = ProcessGroupCollection.use_mpu_process_groups() + + # Defer bf16 master-weight creation into LayerWise (post-sharding) so fp32 masters exist only for owned params. + reset_bf16 = config.bf16 + config.bf16 = False + + base_optimizers: list = [] + init_fns: list = [] + slot_child_indices: dict[int, list[int]] = {} + try: + for slot in range(args.multi_lora_n_adapters): + slot_params = adapter_slot_parameters(model_chunks, slot) + assert slot_params, f"adapter slot {slot} has no parameters; is this a multi-LoRA model?" + with _only_slot_trainable(model_chunks, slot_params): + chained = get_megatron_optimizer( + config, + list(model_chunks), + use_gloo_process_groups=args.enable_gloo_process_groups, + ) + children = [ + child + for child in chained.chained_optimizers + if getattr(child, "optimizer", None) is not None and child.get_parameters() + ] + assert children, f"adapter slot {slot} produced no optimizer children" + slot_child_indices[slot] = list(range(len(base_optimizers), len(base_optimizers) + len(children))) + for child in children: + for group in child.param_groups: + group["miles_multi_lora_slot"] = slot + base_optimizers.append(child) + init_fns.append(_adam_init_state_fn) + finally: + config.bf16 = reset_bf16 + + optimizer = LayerWiseDistributedOptimizer(base_optimizers, config, pg_collection, init_state_fn_list=init_fns) + + # Params are scattered whole across DP ranks, so per-child norm/clip reductions must span the world. + for child in optimizer.chained_optimizers: + child.grad_stats_parallel_group = None + + optimizer.miles_slot_child_indices = slot_child_indices + logger.info( + f"Built multi-LoRA LayerWise optimizer: {args.multi_lora_n_adapters} slots, " + f"{len(optimizer.chained_optimizers)} chained children" + ) + return optimizer + + +def _slot_children(optimizer, slot: int): + return [optimizer.chained_optimizers[i] for i in optimizer.miles_slot_child_indices[slot]] + + +def reset_grad_metadata_keep_grads(model_chunks) -> None: + """Reset DDP per-iteration grad bookkeeping WITHOUT zeroing grad buffers, so per-adapter accumulation + survives across train batches (replaces ``DistributedDataParallel.zero_grad_buffer``).""" + for model_chunk in model_chunks: + if getattr(model_chunk.config, "cuda_graph_impl", "none") != "transformer_engine": + for param in model_chunk.params_with_grad: + param.grad_added_to_main_grad = False + for bucket_group in model_chunk.bucket_groups + model_chunk.expert_parallel_bucket_groups: + bucket_group.reset() + + +def zero_adapter_slot_grads(model, slot: int) -> None: + """Zero one slot's gradients everywhere they live: the DDP ``main_grad`` buffer views + and any lingering ``grad``/``main_param.grad`` references.""" + for param in adapter_slot_parameters(model, slot): + if (main_grad := getattr(param, "main_grad", None)) is not None: + main_grad.zero_() + param.grad = None + if (main_param := getattr(param, "main_param", None)) is not None: + main_param.grad = None + + +def step_adapter_slots( + optimizer, + model, + step_batch_sizes: dict[int, int], + clip_grad: float, +) -> dict[int, float]: + """Step exactly the slots in ``step_batch_sizes`` (slot -> batch size), retaining all other slots' gradients; + scales each slot's accumulated grad sum by 1/batch_size and returns the grad norm per stepped slot.""" + grad_norms: dict[int, float] = {} + + for slot, batch_size in step_batch_sizes.items(): + children = _slot_children(optimizer, slot) + # Copy accumulated main_grads into the owned masters' grads, then scale the sum to the adapter-batch mean. + for child in children: + child.prepare_grads() + for main_param in child.get_parameters(): + if main_param.grad is not None: + main_param.grad.mul_(1.0 / batch_size) + + # Per-slot grad norm over the slot's children, reduced across the whole world (whole-param DP scatter). + grads_for_norm = [] + slot_params = [] + for child in children: + grads_for_norm += child.get_main_grads_for_grad_norm() + slot_params += child.get_parameters() + slot_norm = get_grad_norm_fp32(grads_for_norm, grad_stats_parallel_group=None) + if clip_grad > 0.0 and slot_params: + clip_grad_by_total_norm_fp32(slot_params, clip_grad, slot_norm, False) + grad_norms[slot] = float(slot_norm) + + for child in children: + child.step_with_ready_grads() + + zero_adapter_slot_grads(model, slot) + + if step_batch_sizes: + optimizer.allgather_params() + + return grad_norms diff --git a/miles/backends/megatron_utils/multi_lora_scheduler.py b/miles/backends/megatron_utils/multi_lora_scheduler.py new file mode 100644 index 0000000000..4f148eb599 --- /dev/null +++ b/miles/backends/megatron_utils/multi_lora_scheduler.py @@ -0,0 +1,88 @@ +"""Per-adapter LR/WD schedules for multi-LoRA: one ``OptimizerParamScheduler`` per adapter slot, positioned by +the adapter's own trained samples. Adapters without a known ``num_step`` warm up, then hold ``--lr`` constant.""" + +import logging +from argparse import Namespace + +from megatron.core.optimizer_param_scheduler import OptimizerParamScheduler + +logger = logging.getLogger(__name__) + + +class _SlotParamGroups: + """Minimal optimizer facade: the scheduler only reads ``.param_groups``.""" + + def __init__(self, param_groups: list[dict]): + self.param_groups = param_groups + + +def build_slot_scheduler(args: Namespace, optimizer, adapter, resume_step: int) -> OptimizerParamScheduler: + """Build the slot's scheduler and position it at the adapter's committed + samples. Rebuilt on every adapter load, so slot reuse starts fresh.""" + from miles.backends.megatron_utils.multi_lora_optimizer import _slot_children + + groups = [group for child in _slot_children(optimizer, adapter.slot) for group in child.param_groups] + samples_per_step = adapter.config.adapter_global_batch_size + num_step = adapter.config.num_step + + decay_steps = num_step * samples_per_step if num_step is not None else None + if args.lr_warmup_fraction is not None and decay_steps is not None: + lr_warmup_steps = args.lr_warmup_fraction * decay_steps + else: + lr_warmup_steps = args.lr_warmup_iters * samples_per_step + if decay_steps is None: + # No horizon: warm up, then hold constant. The decay steps only need + # to satisfy the scheduler's warmup < decay invariant. + lr_decay_style = "constant" + decay_steps = int(lr_warmup_steps) + 1 + else: + lr_decay_style = args.lr_decay_style + + scheduler = OptimizerParamScheduler( + _SlotParamGroups(groups), + init_lr=args.lr_warmup_init, + max_lr=args.lr, + min_lr=args.min_lr, + lr_warmup_steps=lr_warmup_steps, + lr_decay_steps=decay_steps, + lr_decay_style=lr_decay_style, + start_wd=args.start_weight_decay, + end_wd=args.end_weight_decay, + wd_incr_steps=decay_steps, + wd_incr_style=args.weight_decay_incr_style, + use_checkpoint_opt_param_scheduler=False, + override_opt_param_scheduler=False, + wsd_decay_steps=( + args.lr_wsd_decay_iters * samples_per_step + if lr_decay_style == "WSD" and args.lr_wsd_decay_iters is not None + else None + ), + lr_wsd_decay_style=args.lr_wsd_decay_style, + ) + if resume_step: + scheduler.step(increment=resume_step * samples_per_step) + return scheduler + + +def install_slot_scheduler(args: Namespace, optimizer, adapter, resume_step: int) -> None: + """Attach the adapter's scheduler to the optimizer, keyed by slot.""" + if not hasattr(optimizer, "miles_slot_schedulers"): + optimizer.miles_slot_schedulers = {} + optimizer.miles_slot_schedulers[adapter.slot] = build_slot_scheduler(args, optimizer, adapter, resume_step) + + +def drop_slot_scheduler(optimizer, slot: int) -> None: + """Detach a retired slot's scheduler (the next tenant installs its own).""" + getattr(optimizer, "miles_slot_schedulers", {}).pop(slot, None) + + +def step_slot_schedulers(optimizer, step_batch_sizes: dict[int, int]) -> dict[int, float]: + """Advance exactly the stepped slots' schedules by their batch samples. + Returns slot -> new learning rate, for logging.""" + lr_by_slot: dict[int, float] = {} + for slot, batch_size in step_batch_sizes.items(): + scheduler = optimizer.miles_slot_schedulers[slot] + scheduler.step(increment=batch_size) + if scheduler.optimizer.param_groups: # empty on ranks owning none of the slot's params + lr_by_slot[slot] = scheduler.optimizer.param_groups[0]["lr"] + return lr_by_slot diff --git a/miles/backends/megatron_utils/multi_lora_utils.py b/miles/backends/megatron_utils/multi_lora_utils.py new file mode 100644 index 0000000000..707541fa4e --- /dev/null +++ b/miles/backends/megatron_utils/multi_lora_utils.py @@ -0,0 +1,441 @@ +import json +import logging +import os +from argparse import Namespace +from collections.abc import Mapping +from pathlib import Path + +import ray +import torch +import torch.distributed as dist + +from miles.backends.training_utils.parallel import get_parallel_state +from miles.ray.multi_lora.controller import get_multi_lora_controller +from miles.utils.adapter_config import AdapterRun + +logger = logging.getLogger(__name__) + + +def create_multi_lora_instance(args: Namespace): + """Create a MultiLoRA instance from training args.""" + from megatron.bridge.peft.multi_lora import MultiLoRA + + from miles.backends.megatron_utils.lora_utils import convert_target_modules_to_megatron + + lora_type_name = getattr(args, "lora_type", "lora").lower() + if lora_type_name == "canonical_lora": + from megatron.bridge.peft.canonical_lora import CanonicalLoRA + + lora_cls = CanonicalLoRA + else: + from megatron.bridge.peft.lora import LoRA + + lora_cls = LoRA + + return MultiLoRA( + target_modules=convert_target_modules_to_megatron(args.target_modules, lora_type=lora_cls), + n_adapters=args.multi_lora_n_adapters, + dim=args.lora_rank, + alpha=args.lora_alpha, + dropout=getattr(args, "lora_dropout", 0.0), + lora_A_init_method=getattr(args, "lora_A_init_method", "xavier"), + lora_B_init_method=getattr(args, "lora_B_init_method", "zero"), + ) + + +def all_megatron_checkpoints_exist(step_dir: Path, tp_size, pp_size) -> bool: + return all( + (step_dir / f"adapter_megatron_tp{tp}_pp{pp}.pt").exists() for tp in range(tp_size) for pp in range(pp_size) + ) + + +def find_latest_checkpoint(ckpt_dir: Path) -> tuple[Path | None, int]: + if not ckpt_dir.exists(): + return None, 0 + + parallel_state = get_parallel_state() + tp_size = parallel_state.tp.size + pp_size = parallel_state.pp.size + tp_rank = parallel_state.tp.rank + pp_rank = parallel_state.pp.rank + + def get_step(d): + return int(d.name.split("_")[1]) + + step_dirs = sorted( + [d for d in ckpt_dir.iterdir() if d.is_dir() and d.name.startswith("step_")], + key=get_step, + reverse=True, + ) + for step_dir in step_dirs: + step = get_step(step_dir) + if all_megatron_checkpoints_exist(step_dir, tp_size, pp_size): + return step_dir / f"adapter_megatron_tp{tp_rank}_pp{pp_rank}.pt", step + + return None, 0 + + +def zero_optimizer_state_for_adapter(optimizer, model, idx: int) -> None: + from megatron.bridge.peft.multi_lora_layers import MultiLoRALinear, _iter_multi_lora_modules + + target_main_params = set() + for module in _iter_multi_lora_modules(model): + if not isinstance(module, MultiLoRALinear): + continue + adapter = module.adapters[idx] + for param in adapter.parameters(): + main = getattr(param, "main_param", None) + target_main_params.add(id(main if main is not None else param)) + + chained = getattr(optimizer, "chained_optimizers", [optimizer]) + for chained_optimizer in chained: + inner = getattr(chained_optimizer, "optimizer", chained_optimizer) + if inner is None: + continue + # TE/apex FusedAdam tracks the Adam step per param GROUP, not per param; + # reset the retired slot's groups so the next tenant restarts bias correction. + for group in inner.param_groups: + if group.get("miles_multi_lora_slot") == idx and "step" in group: + if isinstance(group["step"], torch.Tensor): + group["step"].zero_() + else: + group["step"] = 0 + for param, state in inner.state.items(): + if id(param) not in target_main_params: + continue + if "exp_avg" in state: + state["exp_avg"].zero_() + if "exp_avg_sq" in state: + state["exp_avg_sq"].zero_() + # Bias correction restarts for the slot's next tenant. + if "step" in state: + if isinstance(state["step"], torch.Tensor): + state["step"].zero_() + else: + state["step"] = 0 + + +def slice_lora_to_rank(hf_name: str, tensor: torch.Tensor, adapter_rank: int) -> torch.Tensor: + if "lora_A" in hf_name and adapter_rank < tensor.shape[0]: + remainder = tensor[adapter_rank:] + assert remainder.abs().max() == 0, ( + f"lora_A padded dims are non-zero: {hf_name}, " + f"max={remainder.abs().max().item():.6e}, shape={tensor.shape}, rank={adapter_rank}" + ) + return tensor[:adapter_rank] + if "lora_B" in hf_name and adapter_rank < tensor.shape[1]: + remainder = tensor[:, adapter_rank:] + assert remainder.abs().max() == 0, ( + f"lora_B padded dims are non-zero: {hf_name}, " + f"max={remainder.abs().max().item():.6e}, shape={tensor.shape}, rank={adapter_rank}" + ) + return tensor[:, :adapter_rank] + return tensor + + +def save_multi_lora_checkpoints( + args, + model, + adapter_steps: Mapping[str, int], + adapters: Mapping[str, AdapterRun], +): + """Save per-adapter checkpoints in two formats per adapter. + + Layout (per adapter):: + + {adapter.save}/checkpoints/step_{iteration}/ + ├── adapter_megatron_tp{tp}_pp{pp}.pt ← per-rank shard, fast resume + ├── adapter_model.safetensors ← gathered HF, inference / external + └── adapter_config.json ← HF PEFT metadata (r, alpha, ...) + """ + from megatron.bridge import AutoBridge + from megatron.bridge.peft.multi_lora_layers import expose_adapter_slot + from safetensors.torch import save_file as save_safetensors + + from miles.backends.megatron_utils.lora_utils import convert_target_modules_to_hf + from miles.utils import megatron_bridge_utils + + parallel_state = get_parallel_state() + tp_rank = parallel_state.tp.rank + pp_rank = parallel_state.pp.rank + # One writer per (tp, pp) shard: LoRA params are replicated across DP AND + # CP, so gate on the combined dp×cp group. Gating on intra_dp alone left + # every CP rank writing the same shard file and racing the os.replace. + is_dp_cp_rank_0 = parallel_state.intra_dp_cp.rank == 0 + is_global_writer = is_dp_cp_rank_0 and tp_rank == 0 and pp_rank == 0 + + target_modules_hf = ( + convert_target_modules_to_hf(list(args.target_modules)) + if args.target_modules + else ["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"] + ) + + bridge = AutoBridge.from_hf_pretrained(args.hf_checkpoint, trust_remote_code=True) + + for adapter_name, adapter in adapters.items(): + config = adapter.config + log_prefix = f"[multilora] ({adapter_name})" + iteration = adapter_steps[adapter_name] + + if config.save is None: + logger.info(f"{log_prefix} skipping checkpoint (no save dir configured)") + continue + + final_dir = config.save / "checkpoints" / f"step_{iteration}" + tmp_dir = config.save / "checkpoints" / f"_tmp_step_{iteration}" + if is_dp_cp_rank_0: + tmp_dir.mkdir(parents=True, exist_ok=True) + if dist.is_initialized(): + dist.barrier() + + with expose_adapter_slot(model, adapter.slot): + # Megatron checkpoints + if is_dp_cp_rank_0: + shard: dict[str, torch.Tensor] = { + name: param.data.cpu() + for batch in model + for name, param in batch.named_parameters() + if ".adapter." in name + } + native_path = tmp_dir / f"adapter_megatron_tp{tp_rank}_pp{pp_rank}.pt" + torch.save(shard, native_path) + logger.info(f"{log_prefix} saved Megatron shard " f"({len(shard)} tensors) to {native_path}") + + hf_state: dict[str, torch.Tensor] = {} + with megatron_bridge_utils.patch_megatron_model(model): + for hf_name, weight, _megatron_name in bridge.export_adapter_weights( + model, + cpu=True, + show_progress=False, + ): + # Slice from the shared --lora-rank down to this adapter's real rank to + # match adapter_config's r; clone() since safetensors rejects aliased views. + hf_state[hf_name] = slice_lora_to_rank(hf_name, weight, config.rank).clone() + + if is_global_writer: + save_safetensors( + hf_state, + str(tmp_dir / "adapter_model.safetensors"), + metadata={"format": "pt"}, + ) + adapter_config_json = { + "peft_type": "LORA", + "r": config.rank, + "lora_alpha": config.alpha, + "target_modules": target_modules_hf, + "lora_dropout": getattr(args, "lora_dropout", 0.0), + "bias": "none", + "task_type": "CAUSAL_LM", + } + with open(tmp_dir / "adapter_config.json", "w") as f: + json.dump(adapter_config_json, f, indent=2) + os.sync() + logger.info(f"{log_prefix} saved HF PEFT to {tmp_dir} " f"({len(hf_state)} tensors)") + + if dist.is_initialized(): + dist.barrier() + + # Write to a temp dir and move into place so readers never see a + # partially written checkpoint. + if is_global_writer: + if final_dir.exists(): + import shutil + + shutil.rmtree(final_dir) + os.replace(tmp_dir, final_dir) + logger.info(f"{log_prefix} promoted checkpoint to {final_dir}") + if dist.is_initialized(): + dist.barrier() + + +def _register_adapter(adapter: AdapterRun, model) -> int: + """Install one adapter on this rank's local model shard. Returns the step + of the checkpoint it resumed from (0 for a fresh adapter).""" + from megatron.bridge.peft.multi_lora_layers import init_adapter_slot, load_adapter + + name = adapter.name + config = adapter.config + slot = adapter.slot + log_prefix = f"[multilora] ({name})" + + step = 0 + if config.save is not None: + ckpt_root = config.save / "checkpoints" + ckpt, step = find_latest_checkpoint(ckpt_root) + else: + ckpt = None + + if ckpt is None: + logger.info(f"{log_prefix} no checkpoint, starting from random init") + step = 0 + else: + state_dict = torch.load(ckpt, map_location="cpu", weights_only=True) + loaded = load_adapter(model, slot, state_dict) + assert loaded > 0, ( + f"{log_prefix} loaded 0 tensors from {ckpt} " + f"(state_dict has {len(state_dict)} entries) — name mismatch?" + ) + logger.info(f"{log_prefix} loaded from {ckpt} ({loaded} tensors)") + + init_adapter_slot(model, slot, rank=config.rank, alpha=config.alpha) + logger.info(f"{log_prefix} installed at slot {slot}") + return step + + +def _deregister_adapter(adapter: AdapterRun, args, model, optimizer) -> None: + """Model-side cleanup for one adapter.""" + from megatron.bridge.peft.multi_lora_layers import clear_adapter_slot + + name = adapter.name + slot = adapter.slot + log_prefix = f"[multilora] ({name})" + + if args.save_interval is not None: + # The controller still holds the step count until free_slot runs. + step = ray.get(get_multi_lora_controller().adapter_step.remote(name)) + save_multi_lora_checkpoints(args, model, {name: step}, {name: adapter}) + logger.info(f"{log_prefix} saved final checkpoint at step {step}") + else: + logger.info(f"{log_prefix} save_interval unset; skipping final checkpoint") + + clear_adapter_slot(model, slot) + logger.info(f"{log_prefix} cleared adapter slot {slot}") + + # Prevent future slot tenants from inheriting optimizer momentum or the + # previous tenant's partially accumulated gradients. + from miles.backends.megatron_utils.multi_lora_optimizer import zero_adapter_slot_grads + + from miles.backends.megatron_utils.multi_lora_scheduler import drop_slot_scheduler + + zero_optimizer_state_for_adapter(optimizer, model, slot) + zero_adapter_slot_grads(model, slot) + drop_slot_scheduler(optimizer, slot) + optimizer.reload_model_params() + logger.info(f"{log_prefix} cleared optimizer state and retained grads for slot {slot}") + + +def load_adapters(args, model, optimizer, adapters) -> int: + """Load adapters into Megatron slots; resumes step counts from checkpoints.""" + from miles.backends.megatron_utils.initialize import is_first_replica_megatron_main_rank + from miles.utils.distributed_utils import get_gloo_group + + if dist.is_initialized(): + dist.barrier(group=get_gloo_group()) + if not adapters: + return 0 + from miles.backends.megatron_utils.multi_lora_scheduler import install_slot_scheduler + + resume_steps: dict[str, int] = {} + for adapter in adapters: + resume_steps[adapter.name] = _register_adapter(adapter, model) + # Per-adapter LR/WD schedule, positioned at the resumed step count. + install_slot_scheduler(args, optimizer, adapter, resume_steps[adapter.name]) + if dist.is_initialized(): + dist.barrier(group=get_gloo_group()) + optimizer.reload_model_params() + if is_first_replica_megatron_main_rank(): + for name, step in resume_steps.items(): + if step > 0: + ray.get(get_multi_lora_controller().set_adapter_step.remote(name, step)) + return len(adapters) + + +def cleanup_adapters(args, model, optimizer, adapters) -> int: + """Save final ckpt + clear Megatron slot, then free_slot on the controller.""" + from miles.backends.megatron_utils.initialize import is_first_replica_megatron_main_rank + from miles.utils.distributed_utils import get_gloo_group + + if dist.is_initialized(): + dist.barrier(group=get_gloo_group()) + if not adapters: + return 0 + for adapter in adapters: + _deregister_adapter(adapter, args, model, optimizer) + if dist.is_initialized(): + dist.barrier(group=get_gloo_group()) + if is_first_replica_megatron_main_rank(): + for adapter in adapters: + ray.get(get_multi_lora_controller().free_slot.remote(adapter.name)) + return len(adapters) + + +def step_stepped_adapter_slots(args, model, optimizer, rollout_data, rollout_id: int, step_id: int) -> float: + """Optimizer-step the slots whose adapter batch completes with this train batch and advance + their per-adapter LR/WD schedules. Returns the max grad norm across stepped slots (0.0 if none).""" + from miles.backends.megatron_utils.multi_lora_optimizer import step_adapter_slots + from miles.backends.megatron_utils.multi_lora_scheduler import step_slot_schedulers + from miles.utils.tracking_utils.structured_log import log_structured + + # slot -> adapter_global_batch_size for adapter batches completing now. + step_batch_sizes = dict(rollout_data.get("step_adapter_batch_sizes", {})) + grad_norms_by_slot = step_adapter_slots( + optimizer, + model, + step_batch_sizes, + clip_grad=args.clip_grad, + ) + + if lr_by_slot := step_slot_schedulers(optimizer, step_batch_sizes): + log_structured( + logger.info, + op="adapter_lr", + rollout=rollout_id, + step=step_id, + **{f"slot_{slot}": lr for slot, lr in lr_by_slot.items()}, + ) + return max(grad_norms_by_slot.values(), default=0.0) + + +def commit_trained_batch(rollout_data, rollout_id: int, pending_push: set) -> None: + """A train call landed: schedule the stepped adapters' engine push and + commit the batch on the controller (main rank only). The stepped set ships + with the train data, identical on all ranks.""" + from miles.backends.megatron_utils.initialize import is_first_replica_megatron_main_rank + + pending_push.update(rollout_data.get("step_adapter_names", [])) + if is_first_replica_megatron_main_rank(): + ray.get(get_multi_lora_controller().mark_batch_trained.remote(rollout_id)) + + +def save_due_adapter_checkpoints(args, model) -> bool: + """Save per-adapter checkpoints for adapters at a save-interval multiple + without a checkpoint on disk. Rank 0 picks and broadcasts, so the + collective export lines up. Returns False when nothing is due.""" + from miles.backends.megatron_utils.initialize import is_first_replica_megatron_main_rank + from miles.utils.distributed_utils import get_gloo_group + + due_buffer = [None] + if is_first_replica_megatron_main_rank() and args.save_interval is not None: + snapshot = ray.get(get_multi_lora_controller().snapshot.remote()) + adapters = {**snapshot["active"], **snapshot["retiring"]} + due_buffer[0] = { + name: adapter + for name, adapter in adapters.items() + if adapter.step > 0 + and adapter.step % args.save_interval == 0 + and adapter.config.save is not None + and not (Path(adapter.config.save) / "checkpoints" / f"step_{adapter.step}").exists() + } + if dist.is_initialized(): + dist.broadcast_object_list(due_buffer, src=0, group=get_gloo_group()) + due_adapters = due_buffer[0] + if not due_adapters: + return False + adapter_steps = {name: adapter.step for name, adapter in due_adapters.items()} + save_multi_lora_checkpoints(args, model, adapter_steps, due_adapters) + return True + + +def select_adapters_to_push(loaded_adapters: dict, pending_push: set, has_new_engines: bool) -> tuple[dict, list]: + """Pick the stale adapters to push (all loaded adapters when engines are new). Returns + (adapters to push keyed by name, names to version-bump — only those whose weights changed).""" + pending = pending_push & set(loaded_adapters) + push_names = set(loaded_adapters) if has_new_engines else pending + return {name: loaded_adapters[name] for name in sorted(push_names)}, sorted(pending) + + +def commit_weight_push(version_update_names: list, is_main_rank: bool) -> None: + """A weight push landed: bump the pushed adapters' slot versions on the + controller (promotes PENDING adapters to ACTIVE).""" + if version_update_names and is_main_rank: + ray.get(get_multi_lora_controller().record_weight_update.remote(version_update_names)) diff --git a/miles/backends/training_utils/data.py b/miles/backends/training_utils/data.py index b5dccb8061..92468eb18b 100644 --- a/miles/backends/training_utils/data.py +++ b/miles/backends/training_utils/data.py @@ -128,7 +128,7 @@ def get_batch( Steps: - Fetch raw fields via iterator. - Save original token tensors under "unconcat_tokens". - - Slice tokens into two chunks for Context Parallelism (CP), concatenate, and pad to a configurable multiple. + - Slice tokens into two batches for Context Parallelism (CP), concatenate, and pad to a configurable multiple. - Build cu_seqlens and `PackedSeqParams` with T-H-D layout (T: sequence length, H: attention heads, D: head dimension). Args: @@ -146,6 +146,10 @@ def get_batch( parallel_state = get_parallel_state() assert "tokens" in keys + # get_batch consumes adapter_slots itself (per-adapter token counts below); + # fetch it here so callers don't have to know. None for non-multi-LoRA runs. + if "adapter_slots" not in keys: + keys = [*keys, "adapter_slots"] batch = data_iterator.get_next(keys) if "dynamic_global_batch_size" in data_iterator.rollout_data: @@ -167,7 +171,9 @@ def get_batch( if qkv_format == "bshd": max_seqlen = batch["max_seq_lens"][0] assert max([t.size(0) for t in tokens]) <= max_seqlen + if allgather_cp: + assert batch.get("adapter_slots") is None, "allgather CP is currently not supported with multi-LoRA: " assert max_seqlen % cp_size == 0, f"max_seqlen {max_seqlen} not divisible by cp_size {cp_size}" local_len = max_seqlen // cp_size start = parallel_state.cp.rank * local_len @@ -176,21 +182,23 @@ def get_batch( ] else: tokens = [slice_with_cp(t, pad_token_id, qkv_format, max_seqlen) for t in tokens] + sample_token_lengths = [t.size(0) for t in tokens] tokens = torch.stack(tokens) elif qkv_format == "thd": cp_rank = parallel_state.cp.rank if allgather_cp: + assert batch.get("adapter_slots") is None, "allgather CP is currently not supported with multi-LoRA: " # DSA mode: concatenate all sequences first, then slice once with CP. - # We also pad the *global* concatenated stream to make per-rank chunks equal. + # We also pad the *global* concatenated stream to make per-rank batches equal. cu_seqlens_list: list[int] = [0] for t in tokens: cu_seqlens_list.append(cu_seqlens_list[-1] + t.size(0)) tokens = torch.cat(tokens, dim=0) - # Pad global stream so (1) divisible by cp_size (equal chunks), + # Pad global stream so (1) divisible by cp_size (equal batches), # (2) divisible by pad_size (reduce fragmentation). global_pad_size = cp_size * pad_size pad = (global_pad_size - tokens.size(0) % global_pad_size) % global_pad_size @@ -202,6 +210,7 @@ def get_batch( tokens = tokens.chunk(cp_size, dim=0)[cp_rank] else: tokens = [slice_with_cp(t, pad_token_id, qkv_format) for t in tokens] + sample_token_lengths = [t.size(0) for t in tokens] cu_seqlens = [0] for t in tokens: @@ -227,6 +236,21 @@ def get_batch( else: raise ValueError(f"Unsupported qkv_format: {qkv_format}") + # Multi-LoRA: compute per-adapter token counts from post-CP per-sample lengths. + # NOTE: allgather CP is currently not supported + adapter_slots = batch.get("adapter_slots") + if adapter_slots is not None: + assert all( + adapter_slots[i] <= adapter_slots[i + 1] for i in range(len(adapter_slots) - 1) + ), f"adapter_slots not sorted in micro-batch: {adapter_slots}" + n_adapters = data_iterator.rollout_data["n_adapters"] + total_tokens = tokens.numel() + counts = torch.zeros(n_adapters, dtype=torch.int32, device=torch.cuda.current_device()) + for slot, length in zip(adapter_slots, sample_token_lengths, strict=True): + counts[slot] += length + counts[adapter_slots[-1]] += total_tokens - counts.sum().item() + batch["adapter_token_counts"] = counts + batch["tokens"] = tokens def _compute_transform_like_token_ids(ids_list: list): @@ -346,7 +370,7 @@ def get_next(self, keys: Sequence[str]) -> dict[str, list[object] | None]: - If `micro_batch_indices` is provided, selects rows according to the current index list for each requested key. - - Otherwise, slices a contiguous window of size `micro_batch_size` starting + - Otherwise, slices a contiguous adapter batch of size `micro_batch_size` starting at the current offset. Returns a dict mapping each key to a list subset (or None if absent). @@ -426,6 +450,12 @@ def _generate_data_iterator(rollout_data, micro_batch_size, micro_batch_indices= return data_iterator if not args.use_dynamic_batch_size: + if "adapter_slots" in rollout_data and num_local_gbs % args.micro_batch_size != 0: + raise ValueError( + "A multi-LoRA local batch must be divisible by --micro-batch-size; " + f"got local_batch_size={num_local_gbs}, micro_batch_size={args.micro_batch_size}. " + "Use --use-dynamic-batch-size or choose compatible adapter batch shapes." + ) num_microbatches = [num_local_gbs // args.micro_batch_size for _ in range(num_steps_per_rollout)] data_iterator = _generate_data_iterator(rollout_data, args.micro_batch_size) else: @@ -463,6 +493,10 @@ def _generate_data_iterator(rollout_data, micro_batch_size, micro_batch_indices= for j in range(num_mbs): for k in range(len(partitions[j])): partitions[j][k] += start + # Multi-LoRA: microbatches must be contiguous-by-slot for the + # grouped GEMM's per-adapter token-count math. + if "adapter_slots" in rollout_data: + partitions[j].sort(key=lambda index: rollout_data["adapter_slots"][index]) micro_batch_indices.extend(partitions) assert len(set(sum(micro_batch_indices, []))) == num_local_samples diff --git a/miles/backends/training_utils/log_utils.py b/miles/backends/training_utils/log_utils.py index 45b073ce9f..a4aa866ed7 100644 --- a/miles/backends/training_utils/log_utils.py +++ b/miles/backends/training_utils/log_utils.py @@ -137,6 +137,12 @@ def log_rollout_data(rollout_id: int, args: Namespace, rollout_data: RolloutBatc "witness_ids", "weight_versions", "metadata", + "n_adapters", + "adapter_slots", + "step_slots", + "step_adapter_names", + "step_adapter_batch_sizes", + "prompt_group_sizes", ]: continue # Upload per sample mean for each rollout value diff --git a/miles/backends/training_utils/loss.py b/miles/backends/training_utils/loss.py index 270fcde3e4..d817d93bf1 100644 --- a/miles/backends/training_utils/loss.py +++ b/miles/backends/training_utils/loss.py @@ -12,6 +12,7 @@ from miles.backends.training_utils.parallel import get_parallel_state from miles.utils.audit_utils.event_logger.logger import get_event_logger, is_event_logger_initialized from miles.utils.audit_utils.event_logger.models import TrainAdvantageComputationEvent +from miles.utils.multi_lora import is_multi_lora_enabled from miles.utils.types import RolloutBatch @@ -154,6 +155,11 @@ def loss_function( # Here we need to divide by cp_size because to cancel the multiply in Megatron. assert args.use_dynamic_global_batch_size == ("dynamic_global_batch_size" in batch) global_batch_size = batch.get("dynamic_global_batch_size", args.global_batch_size) + # Multi-LoRA: samples enter the gradient buffers with weight 1; per-adapter + # normalization (1/adapter_global_batch_size, a constant known in advance) + # is applied to the accumulated slot gradient at optimizer-step time. + if is_multi_lora_enabled(args): + global_batch_size = 1 if not args.calculate_per_token_loss: if apply_megatron_loss_scaling: loss_parallel_size = ( diff --git a/tests/fast/backends/megatron_utils/test_multi_lora_scheduler.py b/tests/fast/backends/megatron_utils/test_multi_lora_scheduler.py new file mode 100644 index 0000000000..be48d7118e --- /dev/null +++ b/tests/fast/backends/megatron_utils/test_multi_lora_scheduler.py @@ -0,0 +1,114 @@ +"""Per-adapter LR schedules: parameters come from the global args, position is per adapter. +Pins two fixes: late loads don't inherit the decayed position; resume rebuilds position from committed steps.""" + +from types import SimpleNamespace + +import pytest + +from miles.backends.megatron_utils.multi_lora_scheduler import install_slot_scheduler, step_slot_schedulers + +LR = 2e-5 + + +def make_args(**overrides) -> SimpleNamespace: + args = SimpleNamespace( + lr=LR, + min_lr=0.0, + lr_warmup_init=0.0, + lr_warmup_fraction=None, + lr_warmup_iters=0, + lr_decay_style="cosine", + start_weight_decay=0.1, + end_weight_decay=0.1, + weight_decay_incr_style="constant", + lr_wsd_decay_iters=None, + lr_wsd_decay_style=None, + ) + for key, value in overrides.items(): + setattr(args, key, value) + return args + + +def make_optimizer(n_slots: int = 2) -> SimpleNamespace: + children = [SimpleNamespace(param_groups=[{"lr": 0.0, "weight_decay": 0.0}]) for _ in range(n_slots)] + return SimpleNamespace( + chained_optimizers=children, + miles_slot_child_indices={slot: [slot] for slot in range(n_slots)}, + ) + + +def make_adapter(slot: int, num_step: int | None, samples_per_step: int = 64) -> SimpleNamespace: + config = SimpleNamespace(num_step=num_step, adapter_global_batch_size=samples_per_step) + return SimpleNamespace(slot=slot, name=f"a{slot}", config=config) + + +def slot_lr(optimizer, slot: int) -> float: + return optimizer.chained_optimizers[slot].param_groups[0]["lr"] + + +def test_decaying_adapter_walks_its_own_cosine_schedule(): + optimizer = make_optimizer() + adapter = make_adapter(slot=0, num_step=10) + install_slot_scheduler(make_args(), optimizer, adapter, resume_step=0) + + assert slot_lr(optimizer, 0) == pytest.approx(LR) # fresh: top of the schedule + + step_slot_schedulers(optimizer, {0: 5 * 64}) # half the horizon + assert slot_lr(optimizer, 0) == pytest.approx(LR / 2) + + step_slot_schedulers(optimizer, {0: 100 * 64}) # far past the horizon + assert slot_lr(optimizer, 0) == pytest.approx(0.0) # clamped at min_lr + + +def test_adapter_without_num_step_holds_constant(): + optimizer = make_optimizer() + install_slot_scheduler(make_args(), optimizer, make_adapter(slot=0, num_step=None), resume_step=0) + + step_slot_schedulers(optimizer, {0: 12345 * 64}) + assert slot_lr(optimizer, 0) == pytest.approx(LR) # no horizon: never decays + + +def test_resume_position_is_deterministic_from_committed_steps(): + stepped = make_optimizer() + install_slot_scheduler(make_args(), stepped, make_adapter(slot=0, num_step=10), resume_step=0) + step_slot_schedulers(stepped, {0: 5 * 64}) + + resumed = make_optimizer() + install_slot_scheduler(make_args(), resumed, make_adapter(slot=0, num_step=10), resume_step=5) + + assert slot_lr(resumed, 0) == pytest.approx(slot_lr(stepped, 0)) + + +def test_only_stepped_slots_advance(): + optimizer = make_optimizer() + args = make_args() + install_slot_scheduler(args, optimizer, make_adapter(slot=0, num_step=10), resume_step=0) + install_slot_scheduler(args, optimizer, make_adapter(slot=1, num_step=10), resume_step=0) + + lr_by_slot = step_slot_schedulers(optimizer, {0: 5 * 64}) + + assert set(lr_by_slot) == {0} + assert slot_lr(optimizer, 0) == pytest.approx(LR / 2) + assert slot_lr(optimizer, 1) == pytest.approx(LR) # co-tenant untouched + + +def test_slot_reuse_installs_a_fresh_schedule(): + optimizer = make_optimizer() + args = make_args() + install_slot_scheduler(args, optimizer, make_adapter(slot=0, num_step=10), resume_step=0) + step_slot_schedulers(optimizer, {0: 5 * 64}) + + install_slot_scheduler(args, optimizer, make_adapter(slot=0, num_step=20), resume_step=0) + assert slot_lr(optimizer, 0) == pytest.approx(LR) # next tenant starts at the top + + +def test_warmup_ramps_from_init_lr(): + optimizer = make_optimizer() + args = make_args(lr_warmup_iters=2) # 2 adapter steps of warmup + install_slot_scheduler(args, optimizer, make_adapter(slot=0, num_step=10), resume_step=0) + + assert slot_lr(optimizer, 0) == pytest.approx(0.0) # init_lr + step_slot_schedulers(optimizer, {0: 64}) + assert slot_lr(optimizer, 0) == pytest.approx(LR / 2) # mid-warmup + step_slot_schedulers(optimizer, {0: 64}) + assert slot_lr(optimizer, 0) == pytest.approx(LR) # warmed up diff --git a/tests/fast/backends/megatron_utils/test_multi_lora_slot_cleanup.py b/tests/fast/backends/megatron_utils/test_multi_lora_slot_cleanup.py new file mode 100644 index 0000000000..93fcd3bd93 --- /dev/null +++ b/tests/fast/backends/megatron_utils/test_multi_lora_slot_cleanup.py @@ -0,0 +1,91 @@ +"""zero_optimizer_state_for_adapter must reset a retired slot's Adam moments and step clock +(group-level FusedAdam or per-param torch AdamW) while leaving co-tenant slots untouched.""" + +import sys +import types +from types import SimpleNamespace + +import pytest +import torch + +from miles.backends.megatron_utils.multi_lora_utils import zero_optimizer_state_for_adapter + +MLL_MODULE = "megatron.bridge.peft.multi_lora_layers" + + +class FakeAdapter: + def __init__(self, params): + self._params = list(params) + + def parameters(self): + return self._params + + +class FakeMultiLoRALinear: + def __init__(self, adapters): + self.adapters = adapters + + +@pytest.fixture() +def rig(monkeypatch): + # Stub the lazily imported bridge module so the test needs no bridge build that ships multi-LoRA. + p0 = torch.nn.Parameter(torch.ones(4)) + p1 = torch.nn.Parameter(torch.ones(4)) + module = FakeMultiLoRALinear({0: FakeAdapter([p0]), 1: FakeAdapter([p1])}) + stub = types.ModuleType(MLL_MODULE) + stub.MultiLoRALinear = FakeMultiLoRALinear + stub._iter_multi_lora_modules = lambda model: [module] + monkeypatch.setitem(sys.modules, MLL_MODULE, stub) + return SimpleNamespace(p0=p0, p1=p1, model=object()) + + +def make_optimizer(groups, state): + inner = SimpleNamespace(param_groups=groups, state=state) + return inner, SimpleNamespace(chained_optimizers=[SimpleNamespace(optimizer=inner)]) + + +def test_group_level_fused_adam_clock_resets_only_for_the_retired_slot(rig): + inner, optimizer = make_optimizer( + groups=[ + {"params": [rig.p0], "miles_multi_lora_slot": 0, "step": 50}, + {"params": [rig.p1], "miles_multi_lora_slot": 1, "step": 50}, + ], + state={ + rig.p0: {"exp_avg": torch.ones(4), "exp_avg_sq": torch.ones(4)}, + rig.p1: {"exp_avg": torch.ones(4), "exp_avg_sq": torch.ones(4)}, + }, + ) + + zero_optimizer_state_for_adapter(optimizer, rig.model, 0) + + assert inner.param_groups[0]["step"] == 0 + assert inner.param_groups[1]["step"] == 50 # co-tenant slot untouched + assert float(inner.state[rig.p0]["exp_avg"].abs().sum()) == 0.0 + assert float(inner.state[rig.p0]["exp_avg_sq"].abs().sum()) == 0.0 + assert float(inner.state[rig.p1]["exp_avg"].abs().sum()) == 4.0 + + +def test_tensor_valued_group_clock_resets_in_place(rig): + step = torch.tensor(50) + inner, optimizer = make_optimizer( + groups=[{"params": [rig.p0], "miles_multi_lora_slot": 0, "step": step}], + state={rig.p0: {"exp_avg": torch.ones(4), "exp_avg_sq": torch.ones(4)}}, + ) + + zero_optimizer_state_for_adapter(optimizer, rig.model, 0) + + assert int(step) == 0 # zeroed in place, no rebinding needed + + +def test_per_param_adamw_fallback_clock_resets(rig): + # torch.optim.AdamW keeps the clock per param; groups carry no "step". + inner, optimizer = make_optimizer( + groups=[{"params": [rig.p0], "miles_multi_lora_slot": 0}], + state={ + rig.p0: {"exp_avg": torch.ones(4), "exp_avg_sq": torch.ones(4), "step": torch.tensor(50.0)}, + }, + ) + + zero_optimizer_state_for_adapter(optimizer, rig.model, 0) + + assert float(inner.state[rig.p0]["step"]) == 0.0 diff --git a/tests/fast/backends/megatron_utils/test_slice_lora_to_rank.py b/tests/fast/backends/megatron_utils/test_slice_lora_to_rank.py new file mode 100644 index 0000000000..6b22665766 --- /dev/null +++ b/tests/fast/backends/megatron_utils/test_slice_lora_to_rank.py @@ -0,0 +1,48 @@ +"""slice_lora_to_rank trims max-rank-padded LoRA tensors to the adapter's real rank; +used by weight-sync and HF PEFT export (PEFT rejects tensors padded past the declared rank).""" + +import pytest +import torch + +from miles.backends.megatron_utils.multi_lora_utils import slice_lora_to_rank + + +def _padded(shape, live_rows=None, live_cols=None): + t = torch.zeros(shape) + if live_rows is not None: + t[:live_rows] = 1.0 + if live_cols is not None: + t[:, :live_cols] = 1.0 + return t + + +def test_lora_a_is_sliced_on_the_rank_dim(): + tensor = _padded((32, 8), live_rows=16) + out = slice_lora_to_rank("base_model.q_proj.lora_A.weight", tensor, 16) + assert out.shape == (16, 8) + assert torch.equal(out, tensor[:16]) + + +def test_lora_b_is_sliced_on_the_rank_dim(): + tensor = _padded((8, 32), live_cols=16) + out = slice_lora_to_rank("base_model.q_proj.lora_B.weight", tensor, 16) + assert out.shape == (8, 16) + assert torch.equal(out, tensor[:, :16]) + + +def test_nonzero_padding_is_rejected(): + # Live values beyond the adapter's rank mean the pad rows were trained — + # slicing would silently drop signal, so it must hard-fail instead. + tensor = torch.ones(32, 8) + with pytest.raises(AssertionError, match="padded dims are non-zero"): + slice_lora_to_rank("x.lora_A.weight", tensor, 16) + + +def test_full_rank_tensor_passes_through(): + tensor = torch.ones(16, 8) + assert slice_lora_to_rank("x.lora_A.weight", tensor, 16) is tensor + + +def test_non_lora_names_pass_through(): + tensor = torch.ones(32, 8) + assert slice_lora_to_rank("x.some_other.weight", tensor, 16) is tensor diff --git a/tests/fast/backends/training_utils/test_get_batch_multi_lora_cp.py b/tests/fast/backends/training_utils/test_get_batch_multi_lora_cp.py new file mode 100644 index 0000000000..cdc4e07ca1 --- /dev/null +++ b/tests/fast/backends/training_utils/test_get_batch_multi_lora_cp.py @@ -0,0 +1,134 @@ +"""CP=2 tests for get_batch's multi-LoRA per-adapter token counts; CUDA is stubbed to run on CPU.""" + +from types import SimpleNamespace + +import pytest +import torch + +import miles.backends.training_utils.cp_utils as cp_utils_mod +import miles.backends.training_utils.data as data_mod +from miles.backends.training_utils.cp_utils import slice_with_cp +from miles.backends.training_utils.data import get_batch + + +def _parallel_state(cp_rank: int, cp_size: int, tp_size: int = 1) -> SimpleNamespace: + return SimpleNamespace( + cp=SimpleNamespace(rank=cp_rank, size=cp_size), + tp=SimpleNamespace(size=tp_size), + ) + + +class _FakeIterator: + def __init__(self, batch: dict, n_adapters: int): + self._batch = batch + self.rollout_data = {"n_adapters": n_adapters} + + def get_next(self, keys): + return {key: self._batch[key] for key in keys} + + +KEYS = ["tokens", "loss_masks", "total_lengths", "response_lengths", "adapter_slots"] + +# 5 samples over 3 of 4 slots (sorted), ragged/odd lengths to exercise the +# per-sample zigzag padding (2 * cp_size chunking). +LENGTHS = [7, 13, 5, 9, 4] +SLOTS = [0, 0, 1, 2, 2] +N_ADAPTERS = 4 + + +def _make_batch(max_seqlen: int | None = None) -> dict: + # Token values start at 1 so zigzag pad (value 0) is distinguishable. + tokens = [torch.arange(1, length + 1, dtype=torch.long) for length in LENGTHS] + response_lengths = [max(1, length // 2) for length in LENGTHS] + batch = { + "tokens": tokens, + "loss_masks": [torch.ones(r, dtype=torch.int) for r in response_lengths], + "total_lengths": list(LENGTHS), + "response_lengths": response_lengths, + "adapter_slots": list(SLOTS), + } + if max_seqlen is not None: + batch["max_seq_lens"] = [max_seqlen] * len(LENGTHS) + return batch + + +@pytest.fixture(autouse=True) +def _stub_cuda(monkeypatch): + monkeypatch.setattr(torch.Tensor, "cuda", lambda self, *args, **kwargs: self, raising=False) + monkeypatch.setattr(torch.cuda, "current_device", lambda: "cpu", raising=False) + + +def _patch_state(monkeypatch, state: SimpleNamespace) -> None: + # data.get_batch and cp_utils.slice_with_cp each resolve the state themselves. + monkeypatch.setattr(data_mod, "get_parallel_state", lambda: state) + monkeypatch.setattr(cp_utils_mod, "get_parallel_state", lambda: state) + + +def _expected_thd_counts(state: SimpleNamespace, local_total: int) -> torch.Tensor: + sliced_lengths = [ + slice_with_cp(t, 0, "thd", parallel_state=state).numel() + for t in (torch.arange(1, length + 1, dtype=torch.long) for length in LENGTHS) + ] + expected = torch.zeros(N_ADAPTERS, dtype=torch.int32) + for slot, sliced_length in zip(SLOTS, sliced_lengths, strict=True): + expected[slot] += sliced_length + stream_pad = local_total - int(sum(sliced_lengths)) + assert stream_pad >= 0, "rank-local stream shorter than the sliced samples" + expected[SLOTS[-1]] += stream_pad + return expected + + +@pytest.mark.parametrize("cp_rank", [0, 1]) +def test_thd_cp2_adapter_token_counts(monkeypatch, cp_rank): + state = _parallel_state(cp_rank, cp_size=2) + _patch_state(monkeypatch, state) + + out = get_batch(_FakeIterator(_make_batch(), N_ADAPTERS), KEYS, pad_multiplier=8, qkv_format="thd") + + counts = out["adapter_token_counts"] + local_total = out["tokens"].numel() + assert counts.dtype == torch.int32 + assert counts.tolist() == _expected_thd_counts(state, local_total).tolist() + assert int(counts.sum()) == local_total, "counts must cover every rank-local token incl. padding" + + +def test_thd_cp2_counts_identical_across_ranks(monkeypatch): + # Zigzag gives each rank the same padded share of every sample, so the + # grouped-GEMM routing counts must not depend on the CP rank. + per_rank = [] + for cp_rank in (0, 1): + state = _parallel_state(cp_rank, cp_size=2) + _patch_state(monkeypatch, state) + out = get_batch(_FakeIterator(_make_batch(), N_ADAPTERS), KEYS, pad_multiplier=8, qkv_format="thd") + per_rank.append(out["adapter_token_counts"].tolist()) + assert per_rank[0] == per_rank[1] + + +def test_unsorted_adapter_slots_rejected(monkeypatch): + _patch_state(monkeypatch, _parallel_state(cp_rank=0, cp_size=2)) + batch = _make_batch() + batch["adapter_slots"] = [2, 0, 1, 0, 2] + with pytest.raises(AssertionError, match="not sorted"): + get_batch(_FakeIterator(batch, N_ADAPTERS), KEYS, pad_multiplier=8, qkv_format="thd") + + +def test_bshd_cp2_tokens_are_single_sliced(monkeypatch): + # Regression: the bshd path used to slice tokens twice under CP>1, zeroing the back half of each rank's tokens. + max_seqlen = 16 # divisible by 2 * cp_size + state = _parallel_state(cp_rank=0, cp_size=2) + _patch_state(monkeypatch, state) + + out = get_batch( + _FakeIterator(_make_batch(max_seqlen=max_seqlen), N_ADAPTERS), + KEYS + ["max_seq_lens"], + pad_multiplier=8, + qkv_format="bshd", + ) + + expected = torch.stack( + [ + slice_with_cp(torch.arange(1, length + 1, dtype=torch.long), 0, "bshd", max_seqlen, parallel_state=state) + for length in LENGTHS + ] + ) + assert torch.equal(out["tokens"], expected)