From 6041eab6d97f1bf2a3410e7b395f96eeefe77dfd Mon Sep 17 00:00:00 2001 From: Yusheng Su Date: Tue, 21 Jul 2026 00:31:30 -0700 Subject: [PATCH 1/2] =?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/2] =?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": []}