diff --git a/examples/multi_lora/README.md b/examples/multi_lora/README.md new file mode 100644 index 0000000000..8423ddae1d --- /dev/null +++ b/examples/multi_lora/README.md @@ -0,0 +1,146 @@ +# Multi-LoRA Training Example (fully-async) + +Train multiple LoRA adapters concurrently against a shared base model, using a +fully-async rollout (continuous producer) + a slot-keyed LoRA page table on the +SGLang engines (in-place upsert, no unload, no drain). + +This example trains two adapters on Qwen3-4B: + +- **gsm8k** — grade-school math, `rm_type: math` +- **dapo_math** — competition math (DAPO-Math-17k), `rm_type: deepscaler` + +## Layout + +``` +run_multi_lora.py # launcher: prepare / train / full-train / serve +service_smoke.py # register/deregister smoke test against the API +adapters/ + gsm8k.yaml + dapo_math.yaml +``` + +The implementation lives in the library: the driver is `train_multi_lora_async.py` +at the repo root (next to `train.py`/`train_async.py`), the rollout fn and data +source are `miles/rollout/multi_lora/`, and the controller is +`miles/ray/multi_lora/` (registry + backend + HTTP API, plus the named Ray +actor pinned to the head node). + +## Design (decoupled per-adapter optimizers) + +- **Controller** (Ray actor + control-plane HTTP API) is the source of truth: + `POST/GET/DELETE /adapter_runs` plus `GET /adapter_runs/state`. The data source + reads it; the trainer reads it. Generation traffic goes straight to the router; + on deregister the controller aborts the adapter's in-flight requests + engine-side by rid prefix (`rid = {adapter}::{uuid}`, set in `generate`). +- **Per-adapter gradient accumulation.** Each adapter has its own batch shape: + `rollout_batch_size` prompt groups per optimizer step, each group holding + `n_samples_per_prompt` responses (`adapter_global_batch_size = + rollout_batch_size x n_samples_per_prompt` samples per step). Completed + prompt groups flow into training continuously in multiples of the + adapter's `min_groups_per_dp_split` (the smallest group count whose samples + split evenly across data-parallel ranks), gradients + accumulate in the DDP buffers across train batches, and an adapter's + optimizer steps exactly when its adapter batch fills — independent of every other + adapter. The controller tracks adapter batch progress (`accumulated_groups`) and commits + it only after a successful train call. +- **Per-slot optimizers.** One Adam per adapter slot under Megatron's + `LayerWiseDistributedOptimizer` (whole-parameter ZeRO-1): per-slot state, + step counts, and gradient clipping; optimizer state sharded across DP ranks; + plain DDP all-reduce (no distributed optimizer) makes cross-batch gradient + retention idempotent. +- **Batch collection.** The collection loop (same shape as fully_async's) + pops groups from the per-adapter buffers round-robin, one + `min_groups_per_dp_split` at a time, capped at each adapter's remaining + batch, until the batch reaches `--global-batch-size` samples or a non-empty + batch makes no progress for `--multi-lora-max-coalesce-wait-s` (the target + can be permanently unreachable, so it trains on whatever is ready) — a + single adapter with a small batch trains alone without waiting for + anyone. Samples enter the gradient buffers with weight 1; at step time the + slot's accumulated gradient is scaled by `1/adapter_global_batch_size` + (a constant known in advance), so an adapter's update is identical to what + it would get training alone. +- **Selective weight sync.** Only adapters whose optimizer stepped are pushed + to the engines (upsert into the slot-keyed page table); only their slot + versions bump, keeping staleness filtering per-adapter accurate. +- Adapters deregister on committed optimizer-step count (`num_step`) in the + controller's train-commit path (`mark_batch_trained`), so stop checks happen + exactly when steps advance. `num_step` is relative to the adapter's + start/resume step. When an adapter doesn't set `num_step`, it is derived + from `num_epoch` (default 1) as `num_epoch x len(dataset) // + rollout_batch_size` once the data source loads the dataset (post-filter + length). The trainer's + `reconcile_adapters` (before each generate) retires it at the next sync + point and cleans up (save ckpt + clear Megatron slot + zero its optimizer + state and retained gradients). The adapter's untrained tail — buffered + groups and any partially accumulated gradients — is discarded. +- **Batch ⊆ loaded property:** `reconcile_adapters` runs before `generate`, so the + batch is fetched with loaded = active; active only shrinks during generate, so every + adapter in the batch is live on the trainer. + +## Provision (once) + +```bash +python examples/multi_lora/run_multi_lora.py prepare +``` + +Downloads `Qwen/Qwen3-4B` (to `/root/models`), `zhuzilin/dapo-math-17k`, and +`zhuzilin/gsm8k` (to `/root/datasets`). + +## Run + +```bash +python examples/multi_lora/run_multi_lora.py train # or: full-train (prepare + train) +``` + +Registers the two adapters from CLI flags and trains until each hits its `num_step`, +then exits. + +## Service mode + +```bash +python examples/multi_lora/run_multi_lora.py serve +``` + +Starts with no adapters and idles; register/deregister at runtime through the +control-plane API (port 8068): + +```bash +python examples/multi_lora/service_smoke.py --api-url http://127.0.0.1:8068 \ + --data /root/datasets/gsm8k/train.parquet --input-key messages --label-key label --rm-type math +``` + +## Multi-LoRA CLI flags + +| Flag | Purpose | +| --- | --- | +| `--multi-lora-n-adapters N` | Max concurrent adapter slots. `0` disables (default); `> 0` enables. | +| `--multi-lora-adapter NAME PATH` | Register an adapter at startup. Repeatable. `PATH` → an `adapter.yaml`. | + +Per-adapter `rank` in `adapter.yaml` must be `<= --lora-rank`. + +## adapter.yaml + +```yaml +rank: 16 +alpha: 16 +rollout_batch_size: 32 # prompt groups per optimizer step (defaults to --rollout-batch-size) +n_samples_per_prompt: 4 # group shape (defaults to --n-samples-per-prompt) +data: /root/datasets/gsm8k/train.parquet +input_key: messages +label_key: label +rm_type: math +num_step: 400 # stop adapter after N optimizer steps + # (default: derived from num_epoch, itself default 1) +# optional: save, num_epoch, custom_rm_path, ... +``` + +The derived `adapter_global_batch_size = rollout_batch_size x +n_samples_per_prompt` is the adapter's samples-per-optimizer-step (the +per-adapter analog of `--global-batch-size`). + +Batch-shape constraints (validated at registration, not at runtime): +`n_samples_per_prompt` must be a divisor or multiple of the trainer's +data-parallel size; `rollout_batch_size` must be a multiple of the adapter's +`min_groups_per_dp_split`; +`adapter_global_batch_size` is capped by +`--multi-lora-max-adapter-global-batch-size` (default 4x `--global-batch-size`). diff --git a/examples/multi_lora/adapters/dapo_math.yaml b/examples/multi_lora/adapters/dapo_math.yaml new file mode 100644 index 0000000000..3a1a3ff8a8 --- /dev/null +++ b/examples/multi_lora/adapters/dapo_math.yaml @@ -0,0 +1,9 @@ +rank: 32 +alpha: 32 +rollout_batch_size: 8 # prompt groups per optimizer step +n_samples_per_prompt: 8 # -> 64 samples per step +data: /root/datasets/dapo-math-17k/dapo-math-17k.jsonl +input_key: prompt +label_key: label +rm_type: deepscaler +num_step: 500 diff --git a/examples/multi_lora/adapters/gsm8k.yaml b/examples/multi_lora/adapters/gsm8k.yaml new file mode 100644 index 0000000000..2290611464 --- /dev/null +++ b/examples/multi_lora/adapters/gsm8k.yaml @@ -0,0 +1,9 @@ +rank: 16 +alpha: 16 +rollout_batch_size: 32 # prompt groups per optimizer step +n_samples_per_prompt: 4 # -> 128 samples per step +data: /root/datasets/gsm8k/train.parquet +input_key: messages +label_key: label +rm_type: math +num_step: 400 diff --git a/examples/multi_lora/run_multi_lora.py b/examples/multi_lora/run_multi_lora.py new file mode 100644 index 0000000000..d58e9e8f77 --- /dev/null +++ b/examples/multi_lora/run_multi_lora.py @@ -0,0 +1,204 @@ +"""Multi-LoRA fully-async GRPO example (Qwen3-4B, disaggregated 4 train + 4 rollout GPUs). + +Trains multiple LoRA adapters concurrently on a shared base model. Two example +adapters ship in ``adapters/``: gsm8k (rm_type=math) and dapo_math +(rm_type=deepscaler); each carries its own rank/alpha, batch shape, dataset, +reward, and ``num_step`` stop condition. The driver is +``train_multi_lora_async.py`` at the repo root; fully-async training forbids +``--colocate`` (generation needs continuous GPU). + +Usage: + python examples/multi_lora/run_multi_lora.py prepare # download Qwen3-4B + both datasets (once per node) + python examples/multi_lora/run_multi_lora.py train # bounded run: registers the two adapters, exits when each hits num_step + python examples/multi_lora/run_multi_lora.py full-train # prepare + train + python examples/multi_lora/run_multi_lora.py serve # service mode: no adapters preloaded, idles for registrations (API on :8068) + +Service mode pairs with the smoke client: + python examples/multi_lora/service_smoke.py --api-url http://127.0.0.1:8068 \\ + --data /root/datasets/gsm8k/train.parquet --input-key messages --label-key label --rm-type math +""" + +from dataclasses import dataclass + +import typer + +import miles.utils.external_utils.command_utils as U + +app = typer.Typer() + +_ADAPTER_DIR = f"{U.repo_base_dir}/examples/multi_lora/adapters" + + +@dataclass +class ScriptArgs(U.ExecuteTrainConfig): + run_id: str = U.create_run_id() + + hf_checkpoint: str | None = None + model_dir: str = "/root/models" + data_dir: str = "/root/datasets" + save_dir: str = "/tmp/multi_lora" + megatron_path: str = "/root/Megatron-LM" + + # Disaggregated split (fully-async forbids colocate). + num_gpus_per_node: int = 8 + actor_num_gpus: int = 4 + rollout_num_gpus: int = 4 + tp: int = 2 + + # LoRA slot pool. Per-adapter rank/alpha come from adapter.yaml, capped by lora_rank. + lora_rank: int = 32 + lora_alpha: int = 32 + lora_dropout: float = 0.0 + target_modules: str = "all-linear" + n_adapters: int = 4 + # Comma-separated adapter names; each resolves to adapters/{name}.yaml (train mode only). + adapters: str = "dapo_math,gsm8k" + + # Global rollout defaults; the per-adapter batch shapes live in the yamls. + num_rollout: int = 50 + rollout_batch_size: int = 32 + n_samples_per_prompt: int = 8 + rollout_max_response_len: int = 4096 + global_batch_size: int = 256 + max_weight_staleness: int = 3 + + # Service mode. + api_port: int = 8068 + + save_interval: int = 5 + enable_wandb: bool = False + extra_args: str = "" + + def __post_init__(self): + if self.hf_checkpoint is None: + self.hf_checkpoint = f"{self.model_dir}/Qwen3-4B" + + +def _prepare_download(args: ScriptArgs): + U.exec_command(f"mkdir -p {args.data_dir} {args.model_dir}") + U.exec_command(f"hf download Qwen/Qwen3-4B --local-dir {args.model_dir}/Qwen3-4B") + U.hf_download_dataset("zhuzilin/dapo-math-17k", data_dir=args.data_dir) + U.hf_download_dataset("zhuzilin/gsm8k", data_dir=args.data_dir) + + +def _train(args: ScriptArgs, service: bool): + mode = "service" if service else "bounded" + print( + f"[run] multi-LoRA ({mode}): {args.actor_num_gpus} train + {args.rollout_num_gpus} rollout GPUs, tp={args.tp}" + ) + + ckpt_args = f"--hf-checkpoint {args.hf_checkpoint} --megatron-to-hf-mode bridge " + + lora_args = ( + f"--lora-rank {args.lora_rank} --lora-alpha {args.lora_alpha} " + f'--lora-dropout {args.lora_dropout} --target-modules "{args.target_modules}" ' + ) + + multi_lora_args = f"--multi-lora-n-adapters {args.n_adapters} --multi-lora-idle-poll-s 5 " + if service: + # No adapters preloaded; the control-plane API accepts registrations at runtime. + multi_lora_args += f"--multi-lora-api-port {args.api_port} " + else: + for name in args.adapters.split(","): + multi_lora_args += f'--multi-lora-adapter "{name}" "{_ADAPTER_DIR}/{name}.yaml" ' + multi_lora_args += "--multi-lora-disable-service-mode " + + # in_place pause + upsert weight push is what lets adapters refresh without + # unloading (an unload would deadlock behind paused in-flight requests). + sync_args = f"--pause-generation-mode in_place --max-weight-staleness {args.max_weight_staleness} --use-tis " + + rollout_args = ( + "--apply-chat-template --rollout-shuffle " + f"--num-rollout {args.num_rollout} " + f"--rollout-batch-size {args.rollout_batch_size} " + f"--n-samples-per-prompt {args.n_samples_per_prompt} " + f"--rollout-max-response-len {args.rollout_max_response_len} " + "--rollout-temperature 1 " + f"--global-batch-size {args.global_batch_size} " + ) + + grpo_args = ( + "--advantage-estimator grpo --kl-loss-coef 0.00 --kl-coef 0.00 " + "--entropy-coef 0.00 --eps-clip 0.2 --eps-clip-high 0.28 " + ) + + optimizer_args = ( + "--optimizer adam --lr 1e-5 --lr-decay-style constant --weight-decay 0.1 " + "--adam-beta1 0.9 --adam-beta2 0.98 " + ) + + perf_args = ( + f"--tensor-model-parallel-size {args.tp} --sequence-parallel " + "--pipeline-model-parallel-size 1 --context-parallel-size 1 " + "--expert-model-parallel-size 1 --expert-tensor-parallel-size 1 " + "--use-dynamic-batch-size --max-tokens-per-gpu 9216 " + ) + + sglang_args = "--rollout-num-gpus-per-engine 1 --sglang-mem-fraction-static 0.8 " + + topology_args = ( + f"--actor-num-nodes 1 --actor-num-gpus-per-node {args.actor_num_gpus} " + f"--rollout-num-gpus {args.rollout_num_gpus} --use-miles-router " + ) + + save_args = f"--save {args.save_dir} --save-interval {args.save_interval} " + + misc_args = ( + "--attention-dropout 0.0 --hidden-dropout 0.0 --accumulate-allreduce-grads-in-fp32 " + "--attention-softmax-in-fp32 --attention-backend flash " + ) + + wandb_args = U.get_default_wandb_args(__file__, run_id=args.run_id) if args.enable_wandb else "" + + train_args = ( + f"{ckpt_args} {lora_args} {multi_lora_args} {sync_args} {rollout_args} {grpo_args} " + f"{optimizer_args} {perf_args} {sglang_args} {topology_args} {save_args} {misc_args} " + f"{wandb_args} {args.extra_args} " + ) + + U.execute_train( + train_args=train_args, + config=args, + num_gpus_per_node=args.num_gpus_per_node, + megatron_model_type="qwen3-4B", + train_script="train_multi_lora_async.py", + megatron_path=args.megatron_path, + ) + + +@app.command() +@U.dataclass_cli +def prepare(args: ScriptArgs): + """Download Qwen3-4B and both task datasets. Run once per node before training.""" + _prepare_download(args) + + +@app.command() +@U.dataclass_cli +def train(args: ScriptArgs): + """Bounded run: register the adapters from adapters/, train until each hits num_step, exit.""" + _train(args, service=False) + + +@app.command() +@U.dataclass_cli +def full_train(args: ScriptArgs): + """Download model + datasets, then run the bounded training.""" + _prepare_download(args) + _train(args, service=False) + + +@app.command() +@U.dataclass_cli +def serve(args: ScriptArgs): + """Service mode: no adapters preloaded; register/deregister via the HTTP API while it idles.""" + _train(args, service=True) + + +@app.callback() +def _callback() -> None: + pass + + +if __name__ == "__main__": + app() diff --git a/examples/multi_lora/service_smoke.py b/examples/multi_lora/service_smoke.py new file mode 100644 index 0000000000..fb6685e34f --- /dev/null +++ b/examples/multi_lora/service_smoke.py @@ -0,0 +1,170 @@ +"""Smoke test for multi-LoRA service mode: register/deregister against a running +trainer, using step counts as the race-free progress signal. + +Usage: python examples/multi_lora/service_smoke.py --api-url http://HOST:8068 \\ + --data /root/datasets/gsm8k/train.parquet --input-key messages --label-key label --rm-type math +""" + +import argparse +import sys +import time + +import httpx + +POLL_INTERVAL_S = 5.0 + + +class SmokeFailure(Exception): + pass + + +class ServiceClient: + def __init__(self, api_url: str, timeout_s: float): + self.api_url = api_url.rstrip("/") + self.timeout_s = timeout_s + self.http = httpx.Client(timeout=30.0) + + def adapters(self, states: set[str] | None = None) -> dict: + response = self.http.get(f"{self.api_url}/adapter_runs") + response.raise_for_status() + wanted_states = states if states is not None else {"ACTIVE"} + return { + status["name"]: { + "slot": status["slot"], + "version": status["version"], + "step": status["step"], + "state": status["state"], + } + for status in response.json()["adapters"] + if status["state"] in wanted_states + } + + def active_adapters(self) -> dict: + return self.adapters(states={"ACTIVE"}) + + def register(self, name: str, config: dict) -> httpx.Response: + return self.http.post(f"{self.api_url}/adapter_runs", json={"name": name, "config": config}) + + def deregister(self, name: str) -> None: + response = self.http.delete(f"{self.api_url}/adapter_runs/{name}") + response.raise_for_status() + + def wait_for(self, description: str, predicate) -> dict: + deadline = time.time() + self.timeout_s + while time.time() < deadline: + try: + adapters = self.active_adapters() + except httpx.HTTPError as e: + print(f" ... api not reachable yet ({e})") + adapters = None + if adapters is not None: + if predicate(adapters): + print(f" ok: {description} (active={adapters})") + return adapters + print(f" waiting for {description} (active={adapters})") + time.sleep(POLL_INTERVAL_S) + raise SmokeFailure(f"timed out after {self.timeout_s}s waiting for: {description}") + + def wait_for_step(self, name: str, min_step: int) -> None: + # Step-triggered deregistration can move an adapter to RETIRING quickly; + # count both ACTIVE and RETIRING for progress waits. + self.wait_for( + f"'{name}' to reach step {min_step}", + lambda _active: ( + (adapters := self.adapters(states={"ACTIVE", "RETIRING"})) + and name in adapters + and adapters[name]["step"] >= min_step + ), + ) + + def register_when_allowed(self, name: str, config: dict) -> None: + """Registration is rejected while a same-named adapter is cleaning up; + retry until the name frees.""" + deadline = time.time() + self.timeout_s + while time.time() < deadline: + response = self.register(name, config) + if response.status_code == 200: + print(f" ok: registered '{name}'") + return + print(f" register '{name}' rejected ({response.status_code}): {response.text[:200]}") + time.sleep(POLL_INTERVAL_S) + raise SmokeFailure(f"timed out registering '{name}'") + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--api-url", required=True, help="controller API listener, e.g. http://host:8068") + parser.add_argument("--data", required=True, help="prompt dataset path for the test adapters") + parser.add_argument("--input-key", default="text") + parser.add_argument("--label-key", default="label") + parser.add_argument("--rm-type", default="math") + parser.add_argument("--rank", type=int, default=16) + parser.add_argument("--alpha", type=int, default=16) + parser.add_argument("--save", default=None, help="per-adapter save dir root override (default: trainer --save)") + parser.add_argument("--steps", type=int, default=2, help="training steps to wait for per phase") + parser.add_argument( + "--num-step-smoke", + type=int, + default=1, + help="num_step used by the auto-deregister smoke adapter", + ) + parser.add_argument("--timeout", type=float, default=1800.0, help="per-phase timeout in seconds") + args = parser.parse_args() + + def config(name: str) -> dict: + cfg = { + "rank": args.rank, + "alpha": args.alpha, + "data": args.data, + "input_key": args.input_key, + "label_key": args.label_key, + "rm_type": args.rm_type, + } + if args.save: + cfg["save"] = f"{args.save}/{name}" + return cfg + + client = ServiceClient(args.api_url, args.timeout) + try: + print("phase 1: api reachable, no active adapters expected") + client.wait_for("api reachable", lambda adapters: True) + + print("phase 2: register smoke_auto with num_step; expect auto-deregister after committed steps") + auto_cfg = config("smoke_auto") + auto_cfg["num_step"] = args.num_step_smoke + client.register_when_allowed("smoke_auto", auto_cfg) + client.wait_for_step("smoke_auto", args.num_step_smoke) + client.wait_for("'smoke_auto' auto-deregistered", lambda adapters: "smoke_auto" not in adapters) + + print("phase 3: register smoke_a; expect promotion + training progress") + client.register_when_allowed("smoke_a", config("smoke_a")) + client.wait_for_step("smoke_a", args.steps) + + print("phase 4: register smoke_b mid-run; both must train") + client.register_when_allowed("smoke_b", config("smoke_b")) + client.wait_for_step("smoke_b", args.steps) + + print("phase 5: deregister smoke_a mid-run; smoke_b must keep training") + step_b = client.active_adapters()["smoke_b"]["step"] + client.deregister("smoke_a") + client.wait_for("'smoke_a' gone from active set", lambda adapters: "smoke_a" not in adapters) + client.wait_for_step("smoke_b", step_b + 1) + + print("phase 6: re-register the name smoke_a (waits out cleanup, reuses slot)") + client.register_when_allowed("smoke_a", config("smoke_a")) + client.wait_for_step("smoke_a", 1) + + print("phase 7: deregister everything; service should drain to idle") + client.deregister("smoke_a") + client.deregister("smoke_b") + client.wait_for("no active adapters", lambda adapters: not adapters) + + print("SMOKE TEST PASSED") + return 0 + except SmokeFailure as failure: + print(f"SMOKE TEST FAILED: {failure}", file=sys.stderr) + return 1 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/train_multi_lora_async.py b/train_multi_lora_async.py new file mode 100644 index 0000000000..5420bdef4d --- /dev/null +++ b/train_multi_lora_async.py @@ -0,0 +1,102 @@ +"""Fully-async multi-LoRA trainer driver.""" + +import asyncio +import logging +from pathlib import Path + +import ray + +from miles.ray.multi_lora.controller import create_multilora_controller, get_multi_lora_controller +from miles.ray.placement_group import create_placement_groups, create_rollout_manager, create_training_models +from miles.utils.adapter_config import parse_adapter_run_yaml +from miles.utils.arguments import parse_args +from miles.utils.audit_utils.process_identity import MainProcessIdentity +from miles.utils.logging_utils import configure_logger +from miles.utils.multi_lora import EmptyBatchTimeoutError, define_new_adapter_metrics +from miles.utils.tracking_utils.tracking import init_tracking + +logger = logging.getLogger(__name__) + + +def _is_empty_batch_timeout(task_error: ray.exceptions.RayTaskError) -> bool: + cause = getattr(task_error, "cause", None) + if isinstance(cause, EmptyBatchTimeoutError): + return True + return isinstance(task_error.as_instanceof_cause(), EmptyBatchTimeoutError) + + +async def main(args): + assert ( + not args.colocate + ), "Colocation is not supported for fully-async training (generation needs continuous GPU; colocate time-shares)." + configure_logger(args, source=MainProcessIdentity()) + + # The multi-LoRA rollout fn / data source / global dataset flags are + # defaulted by miles_validate_args when --multi-lora-n-adapters > 0. + pgs = create_placement_groups(args) + init_tracking(args) + rollout_manager, _num_rollout_per_epoch = create_rollout_manager(args, pgs["rollout"]) + + # Create a controller nclusing MultiLoRAController and MultiLoRAHTTPServer to manage lora + router_ip, router_port = await rollout_manager.get_router_address.remote() + args.sglang_router_ip, args.sglang_router_port = router_ip, router_port + controller = create_multilora_controller(args, f"http://{router_ip}:{router_port}") + await controller.start.remote() + host = await controller.http_host.remote() + api_port = await controller.api_port.remote() + logger.info(f"Multi-LoRA control API listening on http://{host}:{api_port} (head node)") + + actor_model, _ = await create_training_models(args, pgs, rollout_manager) + + # CLI-registered adapters are loaded and pushed by the loop's first + # reconcile + update_weights. + for name, path in args.multi_lora_adapters: + config = parse_adapter_run_yaml(Path(path)) + await controller.register_adapter.remote(name, config) + + rollout_id = 0 + while True: + snapshot = await get_multi_lora_controller().snapshot.remote() + + # handle dynamic metrics in tracking backend + define_new_adapter_metrics(snapshot) + if not (snapshot["pending"] or snapshot["active"] or snapshot["retiring"] or snapshot["cleanup"]): + if not args.multi_lora_service_mode: + logger.info("No adapters; exiting.") + break + logger.info(f"No adapters; sleeping for {args.multi_lora_idle_poll_s}s...") + await asyncio.sleep(args.multi_lora_idle_poll_s) + continue + + # Reconcile + push before generate: the push promotes pending adapters, + # and only then does the data source sample them. The actor pushes only + # stale adapter weights (newly loaded, or stepped by the last batch). + await actor_model.reconcile_adapters() + await actor_model.update_weights() + + # With nothing active, generate would wait forever. + post_update = await get_multi_lora_controller().snapshot.remote() + if not (post_update["active"] or post_update["retiring"]): + continue + + try: + rollout_data = await rollout_manager.generate.remote(rollout_id) + except ray.exceptions.RayTaskError as e: + if _is_empty_batch_timeout(e): + logger.warning(f"Generate timed out with no trainable groups; retrying reconcile/update. {e}") + continue + raise + await actor_model.train(rollout_id, rollout_data) + + # Per-adapter save cadence decided inside save_model. + await actor_model.save_model(rollout_id) + + rollout_id += 1 + + await rollout_manager.dispose.remote() + await controller.stop.remote() + + +if __name__ == "__main__": + args = parse_args() + asyncio.run(main(args))