diff --git a/miles/backends/fsdp_utils/actor.py b/miles/backends/fsdp_utils/actor.py index f3eab3694..9e29faea6 100644 --- a/miles/backends/fsdp_utils/actor.py +++ b/miles/backends/fsdp_utils/actor.py @@ -30,9 +30,11 @@ from . import checkpoint from .diffusion_update_weight_utils import ( + DiffusionUpdateWeightFromDistributed, DiffusionUpdateWeightFromTensor, DiffusionUpdateWeightFromTensorLoRA, DiffusionUpdateWeightFromTensorLoRAIPC, + DiffusionUpdateWeightLoRADistributed, ) from .ema import EmaShadow from .input_dtype_policy import apply_input_dtype_policy @@ -220,6 +222,12 @@ def init(self, args: Namespace, role: str, with_ref: bool = False) -> int: # ty uprate=self.args.ema_decay_ramp, uphold=self.args.ema_decay_max, flat_steps=self.args.ema_decay_flat_steps, + # Async prefetch uses the EMA from before the concurrent training update. + keep_previous_ema=( + getattr(self.args, "train_async", False) + and self.args.ref_mode == "ema" + and self.args.ema_rollout_policy == "ema" + ), ) # sglang-d now supports /update_weights_from_tensor (PR #20464). @@ -227,6 +235,11 @@ def init(self, args: Namespace, role: str, with_ref: bool = False) -> int: # ty self.weight_updater = None elif self.args.use_lora and self.args.lora_ipc_weight_sync: self.weight_updater = DiffusionUpdateWeightFromTensorLoRAIPC(self.args, self.models) + elif not self.args.colocate: + updater = ( + DiffusionUpdateWeightLoRADistributed if self.args.use_lora else DiffusionUpdateWeightFromDistributed + ) + self.weight_updater = updater(self.args, self.models) elif self.args.use_lora: self.weight_updater = DiffusionUpdateWeightFromTensorLoRA(self.args, self.models) else: @@ -304,10 +317,7 @@ def update_weights(self) -> None: # type: ignore[override] ray.get(self.rollout_manager.clear_num_new_engines.remote()) ema_shadow = self.ema_shadow - if ema_shadow is not None: - delta = ema_shadow.update() - if dist.get_rank() == 0: - logger.info("EMA shadow updated (decay=%.4f step=%d)", delta, ema_shadow.step) + # Publish the current EMA; the previous EMA is only a training reference. rollout_weight_context = ( ema_shadow.swap_in() if ema_shadow is not None and self.args.ema_rollout_policy == "ema" else nullcontext() ) @@ -343,6 +353,10 @@ def train(self, rollout_id: int, rollout_data_ref) -> None: # type: ignore[over if self.args.debug_rollout_only: return self._train_core(rollout_id=rollout_id, rollout_data=rollout_data) + if self.ema_shadow is not None: + delta = self.ema_shadow.update() + if dist.get_rank() == 0: + logger.info("EMA shadow updated (decay=%.4f step=%d)", delta, self.ema_shadow.step) train_metric_utils.log_perf_data_raw( rollout_id=rollout_id, @@ -555,7 +569,8 @@ def _compute_noise_pred() -> torch.Tensor: ref_mode = self.args.ref_mode if ref_mode != "none": if ref_mode == "ema": - ref_ctx = self.ema_shadow.swap_in() + # Match the EMA that generated the prefetched batch in async mode. + ref_ctx = self.ema_shadow.swap_in(use_previous_ema=self.ema_shadow.previous_ema is not None) else: ref_ctx = prepared.model.disable_adapter() with torch.no_grad(), ref_ctx: diff --git a/miles/backends/fsdp_utils/checkpoint.py b/miles/backends/fsdp_utils/checkpoint.py index 4137757b3..a4398e7e2 100644 --- a/miles/backends/fsdp_utils/checkpoint.py +++ b/miles/backends/fsdp_utils/checkpoint.py @@ -256,6 +256,7 @@ def load(actor: Any) -> dict[str, Any] | None: "rng": rng_state, "metadata": metadata, "iteration": target_step, + "checkpoint_dir": checkpoint_dir, } @@ -264,6 +265,14 @@ def finalize_load(actor: Any, checkpoint_payload: dict[str, Any] | None) -> None dist.barrier() return + if actor.ema_shadow is not None: + ema_dir = checkpoint_payload["checkpoint_dir"] / "ema" + if ema_dir.exists(): + dcp.load({"ema": actor.ema_shadow}, checkpoint_id=str(ema_dir)) + else: + logger.warning("Checkpoint has no EMA state; initializing EMA from the loaded model.") + actor.ema_shadow.step = checkpoint_payload["iteration"] + if checkpoint_payload.get("rng") is not None and not actor.args.no_load_rng: rng_state = checkpoint_payload["rng"] if "torch" in rng_state: @@ -315,6 +324,9 @@ def save(actor: Any, iteration: int) -> None: state_dict = {"model_state": model_state} dcp.save(state_dict, checkpoint_id=str(model_dir)) + if actor.ema_shadow is not None: + dcp.save({"ema": actor.ema_shadow}, checkpoint_id=str(checkpoint_dir / "ema")) + # --no-save-optim drops both the optimizer and the LR scheduler. if not actor.args.no_save_optim: allowed_missing = actor.train_pipeline_config.optimizer_state_allowed_missing diff --git a/miles/backends/fsdp_utils/diffusion_update_weight_utils.py b/miles/backends/fsdp_utils/diffusion_update_weight_utils.py index 0f68bf4d3..57fa92a13 100644 --- a/miles/backends/fsdp_utils/diffusion_update_weight_utils.py +++ b/miles/backends/fsdp_utils/diffusion_update_weight_utils.py @@ -2,8 +2,10 @@ import logging import os import re +import socket from argparse import Namespace from collections.abc import Mapping, Sequence +from datetime import timedelta import ray import torch @@ -16,7 +18,8 @@ except ImportError: from sglang.srt.patch_torch import monkey_patch_torch_reductions # type: ignore[import] -from sglang.srt.utils import MultiprocessingSerializer +from sglang.srt.utils import MultiprocessingSerializer, init_custom_process_group +from sglang.srt.utils.network import NetworkAddress try: from sglang.srt.weight_sync.tensor_bucket import FlattenedTensorBucket # type: ignore[import] @@ -33,10 +36,9 @@ from miles.ray.utils import get_physical_gpu_id - logger = logging.getLogger(__name__) -LORA_IPC_WEIGHT_UPDATE_MODE = "lora_merge" +LORA_WEIGHT_UPDATE_MODE = "lora_merge" class PeftLoRAKeyMapper: @@ -444,7 +446,7 @@ def _verify_weight_sync(self, pairs: list[tuple[str, torch.Tensor]], target_modu logger.warning(f"[weight_sync verify v{self.weight_version} cross-engine] " f"all_equal={all_equal} {pretty}") -class DiffusionUpdateWeightFromTensorLoRAIPC(DiffusionUpdateWeightFromTensor): +class DiffusionUpdateWeightLoRA(DiffusionUpdateWeight): """Push only lora_A/lora_B tensors; rollout merges locally via weight_update_mode=lora_merge.""" def _prepare_lora_param(self, param: torch.Tensor) -> torch.Tensor: @@ -459,7 +461,7 @@ def _prepare_lora_param(self, param: torch.Tensor) -> torch.Tensor: def _collect_layer_groups( self, model: torch.nn.Module ) -> tuple[list[list[tuple[str, torch.Tensor]]], list[str], int]: - """Group PEFT LoRA tensors so each layer's A/B pair stays in one IPC bucket. + """Group PEFT LoRA tensors so each layer's A/B pair stays in one transfer bucket. Names stay PEFT/diffusers-shaped (``transformer_blocks.0.attn.to_q.lora_A``). sglang-d's ``lora_merge`` path applies ``param_names_mapping`` and the @@ -483,7 +485,7 @@ def update_weights(self) -> None: self.wait_and_update_bucket_weights( bucket, target_module, - weight_update_mode=LORA_IPC_WEIGHT_UPDATE_MODE, + weight_update_mode=LORA_WEIGHT_UPDATE_MODE, ) num_buckets += 1 bucket = [] @@ -497,7 +499,7 @@ def update_weights(self) -> None: self.wait_and_update_bucket_weights( bucket, target_module, - weight_update_mode=LORA_IPC_WEIGHT_UPDATE_MODE, + weight_update_mode=LORA_WEIGHT_UPDATE_MODE, ) num_buckets += 1 @@ -507,7 +509,7 @@ def update_weights(self) -> None: num_layers = len(layer_groups) sample_layers = [PeftLoRAKeyMapper.layer_prefix(group[0][0]) for group in layer_groups[:3]] logger.info( - "LoRA IPC weight sync v%s [%s]: pushed %d lora tensors, " + "LoRA weight sync v%s [%s]: pushed %d lora tensors, " "%d layer prefixes in %d buckets (unmapped=%d)", self.weight_version, target_module, @@ -518,18 +520,119 @@ def update_weights(self) -> None: ) if sample_layers: logger.info( - "LoRA IPC [%s] sample layer prefixes: %s", + "LoRA weight sync [%s] sample layer prefixes: %s", target_module, sample_layers, ) if unmapped_keys: logger.warning( - "LoRA IPC unmapped PEFT keys [%s] (first 5): %s", + "LoRA weight sync unmapped PEFT keys [%s] (first 5): %s", target_module, unmapped_keys[:5], ) if num_lora_keys == 0: logger.error( - "LoRA IPC [%s]: no lora tensors found in training state_dict", + "LoRA weight sync [%s]: no lora tensors found in training state_dict", target_module, ) + + +class DiffusionUpdateWeightFromTensorLoRAIPC(DiffusionUpdateWeightLoRA, DiffusionUpdateWeightFromTensor): + pass + + +def connect_rollout_engines_from_distributed(rollout_engines, engine_gpu_counts, group_name, timeout): + if len(rollout_engines) != len(engine_gpu_counts) or any(count <= 0 for count in engine_gpu_counts): + raise ValueError("Each engine requires a positive GPU count") + master_address = ray._private.services.get_node_ip_address() + with socket.socket() as sock: + sock.bind(("", 0)) + master_port = sock.getsockname()[1] + world_size = 1 + sum(engine_gpu_counts) + refs = [] + rank_offset = 1 + for engine, count in zip(rollout_engines, engine_gpu_counts, strict=True): + refs.append( + engine.init_weights_update_group.remote( + master_address=master_address, + master_port=master_port, + rank_offset=rank_offset, + world_size=world_size, + group_name=group_name, + backend="nccl", + ) + ) + rank_offset += count + options = dist.ProcessGroupNCCL.Options() + group = init_custom_process_group( + backend="nccl", + init_method=NetworkAddress(master_address, master_port).to_tcp(), + world_size=world_size, + rank=0, + group_name=group_name, + timeout=timeout, + pg_options=options, + ) + # Custom groups span independent worlds and cannot split the default communicator. + options.split_from = None + ray.get(refs) + return group + + +def broadcast_bucket(rollout_engines, group, group_name, named_tensors, target_module, **kwargs): + refs = [ + engine.update_weights_from_distributed.remote( + names=[name for name, _ in named_tensors], + dtypes=[str(tensor.dtype).removeprefix("torch.") for _, tensor in named_tensors], + shapes=[list(tensor.shape) for _, tensor in named_tensors], + group_name=group_name, + target_modules=[target_module], + **kwargs, + ) + for engine in rollout_engines + ] + tensors = [tensor.contiguous() for _, tensor in named_tensors] + handles = [dist.broadcast(tensor, src=0, group=group, async_op=True) for tensor in tensors] + for handle in handles: + handle.wait() + ray.get(refs) + + +class DiffusionUpdateWeightFromDistributed(DiffusionUpdateWeight): + def __init__(self, args, models): + super().__init__(args, models) + self._model_update_group = None + self._group_name = "diffusion-weight-update" + + def connect_rollout_engines(self, rollout_engines, rollout_engine_lock): + if dist.get_rank() != 0: + return + if self._model_update_group is not None: + refs = [engine.destroy_weights_update_group.remote(self._group_name) for engine in rollout_engines] + dist.destroy_process_group(self._model_update_group) + ray.get(refs) + self.rollout_engines = rollout_engines + self._model_update_group = connect_rollout_engines_from_distributed( + rollout_engines=rollout_engines, + engine_gpu_counts=[self.args.rollout_num_gpus_per_engine] * len(rollout_engines), + group_name=self._group_name, + timeout=timedelta(minutes=self.args.distributed_timeout_minutes), + ) + + def update_bucket_weights(self, named_tensors, target_module, weight_version=None, weight_update_mode=None): + if dist.get_rank() != 0: + return + broadcast_bucket( + rollout_engines=self.rollout_engines, + group=self._model_update_group, + group_name=self._group_name, + named_tensors=named_tensors, + target_module=target_module, + weight_update_mode=weight_update_mode, + lora_alpha=self.args.lora_alpha, + lora_rank=self.args.lora_rank, + ) + + +class DiffusionUpdateWeightLoRADistributed(DiffusionUpdateWeightLoRA, DiffusionUpdateWeightFromDistributed): + pass diff --git a/miles/backends/fsdp_utils/ema.py b/miles/backends/fsdp_utils/ema.py index 9120ef3b8..47f8fa513 100644 --- a/miles/backends/fsdp_utils/ema.py +++ b/miles/backends/fsdp_utils/ema.py @@ -15,7 +15,13 @@ def _local(t: torch.Tensor) -> torch.Tensor: class EmaShadow: - """EMA shadow of trainable parameters.""" + """EMA shadow of trainable parameters. + + ``shadow`` holds the current EMA. With ``keep_previous_ema=True``, + ``previous_ema`` preserves the EMA from before the most recent ``update()`` + for the async trainer's prefetched-batch reference. At initialization and + checkpoint restore, both snapshots start from the same weights. + """ def __init__( self, @@ -25,6 +31,7 @@ def __init__( uprate: float = 0.001, uphold: float = 0.5, flat_steps: int = 0, + keep_previous_ema: bool = False, ) -> None: self.decay = float(decay) self.uprate = float(uprate) @@ -37,6 +44,7 @@ def __init__( if not self.params: raise ValueError("EmaShadow: model has no trainable parameters") self.shadow = [_local(p.detach()).clone() for p in self.params] + self.previous_ema = [sh.clone() for sh in self.shadow] if keep_previous_ema else None def decay_at(self, t: int) -> float: if t <= self.flat_steps: @@ -50,24 +58,56 @@ def update(self) -> float: raise RuntimeError("EmaShadow.update called while swapped in") self.step += 1 delta = self.decay_at(self.step) + if self.previous_ema is not None: + for previous_ema, current_ema in zip(self.previous_ema, self.shadow, strict=True): + previous_ema.copy_(current_ema) for live, sh in zip(self.params, self.shadow, strict=True): sh.mul_(delta).add_(_local(live.detach()).to(sh.device), alpha=1.0 - delta) return delta + def state_dict(self) -> dict: + # Preserve FSDP shard metadata so DCP can restore on a different mesh. + shadow = [ + ( + DTensor.from_local( + sh, + device_mesh=param.device_mesh, + placements=param.placements, + shape=param.shape, + stride=param.stride(), + ) + if isinstance(param, DTensor) + else sh + ) + for param, sh in zip(self.params, self.shadow, strict=True) + ] + return {"shadow": shadow, "step": self.step} + + @torch.no_grad() + def load_state_dict(self, state_dict: dict) -> None: + for sh, restored in zip(self.shadow, state_dict["shadow"], strict=True): + sh.copy_(_local(restored)) + self.step = int(state_dict["step"]) + # Resume starts a fresh pipeline: its first two batches use the restored EMA. + if self.previous_ema is not None: + for previous_ema, current_ema in zip(self.previous_ema, self.shadow, strict=True): + previous_ema.copy_(current_ema) + @contextmanager - def swap_in(self): - """Temporarily expose EMA weights as the live parameters.""" - self._swap() + def swap_in(self, use_previous_ema: bool = False): + """Temporarily use current EMA weights, or the snapshot before the last update.""" + buffers = self.previous_ema if use_previous_ema else self.shadow + self._swap(buffers) self._swapped = True try: yield finally: - self._swap() + self._swap(buffers) self._swapped = False @torch.no_grad() - def _swap(self) -> None: - for live, sh in zip(self.params, self.shadow, strict=True): + def _swap(self, buffers: list[torch.Tensor]) -> None: + for live, sh in zip(self.params, buffers, strict=True): live_local = _local(live.data) tmp = live_local.clone() live_local.copy_(sh) diff --git a/miles/backends/sglang_diffusion_utils/sglang_diffusion_engine.py b/miles/backends/sglang_diffusion_utils/sglang_diffusion_engine.py index 2f9cdcb87..366df047f 100644 --- a/miles/backends/sglang_diffusion_utils/sglang_diffusion_engine.py +++ b/miles/backends/sglang_diffusion_utils/sglang_diffusion_engine.py @@ -234,7 +234,7 @@ def health_generate(self, timeout: float = 5.0) -> bool: def update_weights_from_tensor( self, serialized_named_tensors: list[str], - payload_gpu_uuids: list[str], + payload_gpu_uuids: list[str] | None, load_format: str | None = None, target_modules: list[str] | None = None, weight_version: str | None = None, @@ -268,6 +268,49 @@ def update_weights_from_tensor( payload, ) + def init_weights_update_group( + self, master_address, master_port, rank_offset, world_size, group_name, backend="nccl" + ): + return self._make_request( + "init_weights_update_group", + { + "master_address": master_address, + "master_port": master_port, + "rank_offset": rank_offset, + "world_size": world_size, + "group_name": group_name, + "backend": backend, + }, + ) + + def destroy_weights_update_group(self, group_name): + return self._make_request("destroy_weights_update_group", {"group_name": group_name}) + + def update_weights_from_distributed( + self, + names, + dtypes, + shapes, + group_name, + target_modules, + weight_update_mode=None, + lora_alpha=None, + lora_rank=None, + ): + return self._make_request( + "update_weights_from_distributed", + { + "names": names, + "dtypes": dtypes, + "shapes": shapes, + "group_name": group_name, + "target_modules": target_modules, + "weight_update_mode": weight_update_mode, + "lora_alpha": lora_alpha, + "lora_rank": lora_rank, + }, + ) + def get_weights_checksum(self, module_names: list[str] | None = None) -> dict: """Query the live engine for SHA-256 checksums of the named pipeline modules. @@ -339,7 +382,7 @@ def _compute_server_args(args, host, port, nccl_port): if hasattr(args, f"sglang_{attr.name}") and attr.name not in kwargs: kwargs[attr.name] = getattr(args, f"sglang_{attr.name}") - if getattr(args, "use_lora", False) and getattr(args, "lora_ipc_weight_sync", False): + if args.use_lora and (args.lora_ipc_weight_sync or not args.colocate): kwargs["lora_target_modules"] = args.lora_target_modules # dit_precision / vae_precision are PipelineConfig fields, not ServerArgs, so forward them explicitly (only when changed from the class default, to avoid clobbering a subclass override). from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig diff --git a/miles/ray/actor_group.py b/miles/ray/actor_group.py index 8ce4c3100..ebc2fead4 100644 --- a/miles/ray/actor_group.py +++ b/miles/ray/actor_group.py @@ -1,5 +1,3 @@ -import os - import ray from ray.util.placement_group import PlacementGroup from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy @@ -49,9 +47,6 @@ def _allocate_gpus_for_actor(self, pg, num_gpus_per_actor): pg, reordered_bundle_indices, _reordered_gpu_ids = pg env_vars = { - # because sglang will always set NCCL_CUMEM_ENABLE to 0 - # we need also set it to 0 to prevent nccl error. - "NCCL_CUMEM_ENABLE": os.environ.get("NCCL_CUMEM_ENABLE", "0"), "NVTE_FP8_BLOCK_SCALING_FP32_SCALES": "1", **self.args.train_env_vars, } diff --git a/miles/ray/rollout.py b/miles/ray/rollout.py index 2da1f07a4..3b684b962 100644 --- a/miles/ray/rollout.py +++ b/miles/ray/rollout.py @@ -529,7 +529,7 @@ def init_rollout_engines(args, pg, all_rollout_engines): "SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION": "false", "SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_IDLE": "false", } - if args.lora_ipc_weight_sync: + if args.use_lora and (args.lora_ipc_weight_sync or not args.colocate): # Merge in the train forward dtype, not fp32, to cut train/rollout consistency error. env_vars["SGLANG_DIFFUSION_LORA_MERGE_FP32"] = "1" if args.diffusion_forward_dtype == "fp32" else "0" diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index 71747fa21..4471976e7 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -1788,6 +1788,10 @@ def miles_validate_args(args): if args.offload_rollout is None: args.offload_rollout = False + if not args.colocate and not args.train_only and not args.debug_rollout_only: + if args.lora_ipc_weight_sync: + raise ValueError("--lora-ipc-weight-sync requires --colocate: CUDA IPC needs shared train/rollout GPUs") + if args.hps_num_workers <= 0: raise ValueError(f"--hps-num-workers must be positive, got {args.hps_num_workers}") if args.hps_batch_size <= 0: diff --git a/scripts/run_diffusion_nft_krea2_async.py b/scripts/run_diffusion_nft_krea2_async.py new file mode 100644 index 000000000..a963316c8 --- /dev/null +++ b/scripts/run_diffusion_nft_krea2_async.py @@ -0,0 +1,179 @@ +"""Async Krea-2-Raw DiffusionNFT training (OCR by default, PickScore via --reward). + +Same NFT shape as run_diffusion_nft_sd3_pickscore.py: EMA reference (--ref-mode ema), +rollout under pi_old (--ema-rollout-policy ema), deterministic ODE rollout +(noise_level=0, sde_type=ode) with no CFG. Krea-2 specifics: bf16, 1024px, and one +sample per rollout request (the engine's krea2 pipeline has no per-request output +expansion). Rollout debug tensors are collected (--diffusion-debug-mode). + +OCR is the default reward: text rendering improves visibly and its accuracy curve is +steep, so both the metric and the wandb images validate the run. It needs no reward +GPU. --reward pickscore switches to the aesthetic direction on one extra GPU. + +Smoke mode shrinks the batch for checking the pipeline end to end without a real run. + +Full OCR uses 6 training + 2 rollout H200 GPUs with NCCL LoRA sync; smoke uses 2+2. +The one-step async pipeline uses the previous EMA as the reference for each +prefetched batch. Smoke mode runs three rollouts to cover the first updated rollout batch. + +Usage: + python3 scripts/run_diffusion_nft_krea2_async.py + python3 scripts/run_diffusion_nft_krea2_async.py --reward pickscore + MILES_SCRIPT_SMOKE=1 python3 scripts/run_diffusion_nft_krea2_async.py +""" + +import os +from dataclasses import dataclass + +import typer + +import miles.utils.external_utils.command_utils as U + +MODEL = "krea/Krea-2-Raw" +DATASET = "rockdu/miles-diffusion-datasets" +WANDB_PROJECT = "diffusionNFT" + + +@dataclass +class ScriptArgs(U.ExecuteTrainConfig): + num_rollout: int = 0 # 0 picks the smoke/full default + data_dir: str = "/root/datasets" + smoke: bool = False + reward: str = "ocr" # ocr | pickscore + extra_args: str = "" + + +def _use_ocr(args: ScriptArgs) -> bool: + return args.smoke or args.reward == "ocr" + + +def _subset(args: ScriptArgs) -> str: + return "flowgrpo_ocr" if _use_ocr(args) else "flowgrpo_pickscore" + + +def _num_gpus(args: ScriptArgs) -> int: + return 4 if args.smoke else (8 if _use_ocr(args) else 5) + + +def prepare(args: ScriptArgs) -> str: + local_dir = U.hf_download_dataset(DATASET, include=f"{_subset(args)}/**", data_dir=args.data_dir) + return f"{local_dir}/{_subset(args)}" + + +def execute(args: ScriptArgs, data_dir: str) -> None: + run_name = f"diffusion_nft_krea2_{args.reward}_async_{U.create_run_id()}" + num_rollout = args.num_rollout or (3 if args.smoke else 100) + full_ocr = _use_ocr(args) and not args.smoke + + ckpt_args = f"--hf-checkpoint {MODEL} --save {args.output_dir}/{run_name}/ckpt --save-interval 20 " + + rollout_args = ( + "--rollout-function-path miles.rollout.sglang_diffusion_rollout.generate_rollout " + f"--prompt-data {data_dir}/train.jsonl " + "--input-key input " + f"--num-rollout {num_rollout} " + "--num-steps-per-rollout 1 " + "--diffusion-num-steps 10 " + "--diffusion-guidance-scale 1.0 " + "--diffusion-noise-level 0.0 " + "--diffusion-sde-type ode " + "--diffusion-height 1024 " + "--diffusion-width 1024 " + "--diffusion-debug-mode " + "--rollout-microgroup-size 1 " + ) + ( + "--rollout-batch-size 2 --n-samples-per-prompt 2 " + if args.smoke + else f"--rollout-batch-size {6 if full_ocr else 8} --n-samples-per-prompt 8 " + ) + + eval_args = "--diffusion-eval-num-steps 52 --skip-eval-before-train " + if not args.smoke: + eval_args += f"--eval-prompt-data {args.reward}_test {data_dir}/test.jsonl " + eval_args += f"--eval-interval {num_rollout if full_ocr else 30} " + + grpo_args = ( + "--loss-type nft " + "--diffusion-nft-beta 1.0 " + "--diffusion-nft-timestep-fraction 0.99 " + "--advantage-estimator grpo " + "--globalize-reward-std " + ) + + ema_args = ( + "--ref-mode ema " + "--use-ema " + "--ema-rollout-policy ema " + "--ema-decay-init 0.001 " + "--ema-decay-ramp 0.001 " + "--ema-decay-max 0.5 " + "--ema-decay-flat-steps 0 " + ) + + optimizer_args = "--lr 3e-4 --adam-beta2 0.999 --weight-decay 1e-4 --clip-grad 1.0 " + + lora_args = "--use-lora --lora-rank 32 --lora-alpha 64 --lora-init-weights gaussian " + + reward_args = ( + "--rm-type ocr " + if _use_ocr(args) + else ( + "--rm-type pickscore " + "--pickscore-num-workers 1 " + "--pickscore-num-gpus-per-worker 1.0 " + "--pickscore-batch-size 8 " + "--pickscore-processor-path laion/CLIP-ViT-H-14-laion2B-s32B-b79K " + "--pickscore-model-path yuvalkirstain/PickScore_v1 " + ) + ) + + wandb_args = U.get_default_wandb_args( + __file__, run_id=run_name, project=WANDB_PROJECT, wandb_log_num_images=8, wandb_log_image_interval=10 + ) + + sglang_args = ( + "--use-miles-router " + "--sglang-server-concurrency 8 " + "--sglang-dit-precision bf16 " + "--sglang-vae-slicing " + "--update-weight-buffer-size 2147483648 " + ) + + train_backend_args = "--train-backend fsdp --diffusion-forward-dtype bf16 " + + micro_batch_size = 1 if args.smoke else (4 if full_ocr else 2) + perf_args = f"--micro-batch-size {micro_batch_size} " + if not full_ocr: + perf_args += "--gradient-checkpointing " + + misc_args = ( + f"--actor-num-gpus-per-node {6 if full_ocr else 2} " + "--rollout-num-gpus 2 " + "--rollout-num-gpus-per-engine 1 " + f"--num-gpus-per-node {_num_gpus(args)} " + "--deterministic-mode " + ) + + U.execute_train( + train_args=( + f"{ckpt_args} {rollout_args} {eval_args} {grpo_args} {ema_args} " + f"{optimizer_args} {lora_args} {reward_args} {wandb_args} {sglang_args} " + f"{train_backend_args} {perf_args} {misc_args} {args.extra_args}" + ), + num_gpus_per_node=_num_gpus(args), + train_script="train_diffusion_async.py", + config=args, + extra_env_vars={ + "HF_TOKEN": os.environ.get("HF_TOKEN", ""), + }, + ) + + +@U.dataclass_cli +def main(args: ScriptArgs) -> None: + data_dir = prepare(args) + execute(args, data_dir) + + +if __name__ == "__main__": + typer.run(main) diff --git a/tests/fast/backends/fsdp_utils/_ema_checkpoint_worker.py b/tests/fast/backends/fsdp_utils/_ema_checkpoint_worker.py new file mode 100644 index 000000000..20b0050ad --- /dev/null +++ b/tests/fast/backends/fsdp_utils/_ema_checkpoint_worker.py @@ -0,0 +1,25 @@ +"""Write distinct, uneven EMA shards for the CPU resharding test.""" + +import sys +from pathlib import Path + +import torch +import torch.distributed as dist +import torch.distributed.checkpoint as dcp +from torch.distributed.device_mesh import init_device_mesh +from torch.distributed.tensor import Shard, distribute_tensor + +from miles.backends.fsdp_utils.ema import EmaShadow + +if __name__ == "__main__": + dist.init_process_group("gloo") + mesh = init_device_mesh("cpu", (dist.get_world_size(),)) + full = torch.arange(15).reshape(5, 3).float() + param = torch.nn.Parameter(distribute_tensor(full, mesh, [Shard(0)])) + ema = EmaShadow([param], decay=0.5, flat_steps=10) + for _ in range(2): + with torch.no_grad(): + param.add_(1) + ema.update() + dcp.save({"ema": ema}, checkpoint_id=str(Path(sys.argv[1]) / "ema")) + dist.destroy_process_group() diff --git a/tests/fast/backends/fsdp_utils/test_distributed_weight_update.py b/tests/fast/backends/fsdp_utils/test_distributed_weight_update.py new file mode 100644 index 000000000..95c638b55 --- /dev/null +++ b/tests/fast/backends/fsdp_utils/test_distributed_weight_update.py @@ -0,0 +1,88 @@ +from tests.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=15, suite="stage-a-cpu", labels=[]) + +from datetime import timedelta +from types import SimpleNamespace + +import torch + +from miles.backends.fsdp_utils import diffusion_update_weight_utils as update + + +def test_connect_assigns_every_engine_rank_before_waiting(monkeypatch): + calls = [] + engines = [] + for index in range(2): + + def init(index=index, **kwargs): + calls.append((index, kwargs)) + return index + + engines.append(SimpleNamespace(init_weights_update_group=SimpleNamespace(remote=init))) + + monkeypatch.setattr(update.ray._private.services, "get_node_ip_address", lambda: "127.0.0.1") + + def join(**kwargs): + assert len(calls) == 2 + assert kwargs["rank"] == 0 + assert kwargs["world_size"] == 7 + return "group" + + monkeypatch.setattr(update, "init_custom_process_group", join) + monkeypatch.setattr(update.ray, "get", lambda refs: calls.append(("wait", refs))) + group = update.connect_rollout_engines_from_distributed(engines, [2, 4], "update", timedelta(seconds=30)) + assert group == "group" + assert [calls[i][1]["rank_offset"] for i in range(2)] == [1, 3] + assert calls[-1] == ("wait", [0, 1]) + + +def test_broadcast_matches_metadata_and_retains_contiguous_buffers(monkeypatch): + events = [] + payloads = [] + tensors = [("b", torch.arange(6).reshape(2, 3).t()), ("a", torch.ones(2, dtype=torch.bfloat16))] + + def receive(**kwargs): + payloads.append(kwargs) + events.append("rpc") + return len(payloads) + + engines = [SimpleNamespace(update_weights_from_distributed=SimpleNamespace(remote=receive)) for _ in range(2)] + + def broadcast(tensor, src, group, async_op): + assert len(payloads) == 2 + assert tensor.is_contiguous() + index = events.count("broadcast") + torch.testing.assert_close(tensor, tensors[index][1]) + events.append("broadcast") + return SimpleNamespace(wait=lambda: events.append("wait")) + + monkeypatch.setattr(update.dist, "broadcast", broadcast) + monkeypatch.setattr(update.ray, "get", lambda refs: events.append("ack")) + update.broadcast_bucket(engines, "group", "update", tensors, "transformer") + assert payloads[0]["names"] == ["b", "a"] + assert payloads[0]["dtypes"] == ["int64", "bfloat16"] + assert payloads[0]["shapes"] == [[3, 2], [2]] + assert events == ["rpc", "rpc", "broadcast", "broadcast", "wait", "wait", "ack"] + + +def test_reconnect_destroys_old_group_before_joining(monkeypatch): + args = SimpleNamespace(rollout_num_gpus_per_engine=2, distributed_timeout_minutes=1) + updater = update.DiffusionUpdateWeightFromDistributed(args, {}) + updater._model_update_group = "old" + events = [] + engine = SimpleNamespace( + destroy_weights_update_group=SimpleNamespace(remote=lambda name: events.append("remote destroy")) + ) + monkeypatch.setattr(update.dist, "get_rank", lambda: 0) + monkeypatch.setattr(update.dist, "destroy_process_group", lambda group: events.append(("destroy", group))) + monkeypatch.setattr(update.ray, "get", lambda refs: events.append("ack")) + + def connect(**kwargs): + events.append("connect") + return "new" + + monkeypatch.setattr(update, "connect_rollout_engines_from_distributed", connect) + updater.connect_rollout_engines([engine], None) + assert events == ["remote destroy", ("destroy", "old"), "ack", "connect"] + assert updater._model_update_group == "new" diff --git a/tests/fast/backends/fsdp_utils/test_ema_checkpoint.py b/tests/fast/backends/fsdp_utils/test_ema_checkpoint.py new file mode 100644 index 000000000..b2b68e03f --- /dev/null +++ b/tests/fast/backends/fsdp_utils/test_ema_checkpoint.py @@ -0,0 +1,97 @@ +"""EMA checkpoint integration, including restoration of distributed shards.""" + +from tests.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=90, suite="stage-a-cpu", labels=[]) + +import shutil +import subprocess +import sys +from argparse import Namespace + +import pytest +import torch +import torch.distributed as dist +import torch.distributed.checkpoint as dcp + +from miles.backends.fsdp_utils import checkpoint +from miles.backends.fsdp_utils.ema import EmaShadow + + +def make_actor(tmp_path): + model = torch.nn.Linear(3, 5, bias=False) + optimizer = torch.optim.AdamW(model.parameters()) + return Namespace( + model=model, + optimizer=optimizer, + lr_scheduler=torch.optim.lr_scheduler.LambdaLR(optimizer, lambda step: 1.0), + ema_shadow=EmaShadow(model.parameters(), decay=0.5, flat_steps=10, keep_previous_ema=True), + global_step=2, + micro_step=0, + train_pipeline_config=Namespace(optimizer_state_allowed_missing=[]), + args=Namespace( + save=str(tmp_path), + load=str(tmp_path), + ckpt_step=None, + use_lora=False, + no_save_optim=False, + no_load_optim=False, + no_load_rng=True, + start_rollout_id=0, + ), + ) + + +@pytest.mark.parametrize("legacy", [False, True]) +def test_checkpoint_restores_ema_and_restarts_reference(tmp_path, monkeypatch, legacy): + dist.init_process_group("gloo", init_method=f"file://{tmp_path}/rendezvous", rank=0, world_size=1) + monkeypatch.setattr(torch.cuda, "synchronize", lambda: None) + monkeypatch.setattr(torch.cuda, "get_rng_state_all", lambda: []) + try: + original = make_actor(tmp_path) + for _ in range(2): + with torch.no_grad(): + original.model.weight.add_(1) + original.ema_shadow.update() + checkpoint.save(original, iteration=1) + if legacy: + shutil.rmtree(tmp_path / "iter_0000002/ema") + restored = make_actor(tmp_path) + payload = checkpoint.load(restored) + # Actor initializes its EMA from the already-restored live model. + restored.ema_shadow = EmaShadow(restored.model.parameters(), decay=0.5, flat_steps=10, keep_previous_ema=True) + checkpoint.finalize_load(restored, payload) + assert restored.args.start_rollout_id == 2 + assert restored.ema_shadow.step == 2 + expected = original.model.weight if legacy else original.ema_shadow.shadow[0] + torch.testing.assert_close(restored.ema_shadow.shadow[0], expected) + torch.testing.assert_close(restored.ema_shadow.previous_ema[0], expected) + torch.testing.assert_close(restored.model.weight, original.model.weight) + if not legacy: + assert original.ema_shadow.update() == restored.ema_shadow.update() + torch.testing.assert_close(restored.ema_shadow.shadow[0], original.ema_shadow.shadow[0]) + finally: + dist.destroy_process_group() + + +def test_ema_checkpoint_reshards_to_single_process(tmp_path): + subprocess.run( + [ + sys.executable, + "-m", + "torch.distributed.run", + "--standalone", + "--nproc_per_node=2", + "--module", + "tests.fast.backends.fsdp_utils._ema_checkpoint_worker", + str(tmp_path), + ], + check=True, + timeout=180, + ) + param = torch.nn.Parameter(torch.zeros(5, 3)) + restored = EmaShadow([param], keep_previous_ema=True) + dcp.load({"ema": restored}, checkpoint_id=str(tmp_path / "ema")) + assert restored.step == 2 + torch.testing.assert_close(restored.shadow[0], torch.arange(15).reshape(5, 3).float() + 1.25) + torch.testing.assert_close(restored.previous_ema[0], restored.shadow[0]) diff --git a/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py b/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py index 8fe5ea642..6ac239050 100644 --- a/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py +++ b/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py @@ -286,3 +286,20 @@ def test_swap_in_restores_exactly(self): with ema.swap_in(): assert torch.equal(m.weight.detach(), live) assert torch.equal(m.weight.detach(), live + 2.0) + + def test_previous_ema_tracks_pre_update_snapshot(self): + m = self._model() + ema = EmaShadow(m.parameters(), decay=0.5, uprate=0.001, uphold=0.5, flat_steps=10, keep_previous_ema=True) + init = m.weight.detach().clone() + with torch.no_grad(): + m.weight.add_(1.0) + ema.update() + assert torch.equal(ema.previous_ema[0], init) + assert torch.allclose(ema.shadow[0], init + 0.5) + with ema.swap_in(use_previous_ema=True): + assert torch.equal(m.weight.detach(), init) + assert torch.equal(m.weight.detach(), init + 1.0) + + def test_previous_ema_disabled_by_default(self): + ema = EmaShadow(self._model().parameters(), decay=0.1) + assert ema.previous_ema is None diff --git a/tests/fast/rollout/test_async_training.py b/tests/fast/rollout/test_async_training.py new file mode 100644 index 000000000..6967e361d --- /dev/null +++ b/tests/fast/rollout/test_async_training.py @@ -0,0 +1,208 @@ +"""Exercise the async driver with real serial Ray actors and tiny CPU weights.""" + +from tests.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=90, suite="stage-a-cpu", labels=[]) + +import json +import tempfile +import time +from argparse import Namespace +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +import ray +import torch +from train_diffusion_async import train_loop + +from miles.backends.fsdp_utils.ema import EmaShadow +from miles.rollout.data_source import RolloutDataSourceWithBuffer +from miles.utils.ray_utils import Box + + +@ray.remote +class RolloutProbe: + def __init__(self, args, delay): + self.source = RolloutDataSourceWithBuffer(args) + self.source.load(args.start_rollout_id - 1) + self.delay = delay + self.weight = 0.0 + self.events = [] + + def generate(self, rollout_id): + start = time.monotonic() + sample = self.source.get_samples(1)[0][0] + time.sleep(self.delay) + self.events.append((rollout_id, start, time.monotonic())) + return [Box(ray.put(dict(rollout_id=rollout_id, weight=self.weight, sample_index=sample.index)))] + + def save(self, rollout_id): + self.source.save(rollout_id) + + def install(self, weight): + self.weight = weight + + def get_rollout_engines_and_lock(self): + return [], None, 0 + + def get_events(self): + return self.events + + +@ray.remote +class TrainerProbe: + def __init__(self, manager, delay, restored=None): + from miles.backends.fsdp_utils import actor + from miles.utils.timer import Timer + + self.actor = actor + self.delay = delay + self.rollout_manager = manager + self.args = Namespace( + offload_train=False, debug_rollout_only=False, train_only=False, ema_rollout_policy="ema" + ) + self.parallel_state = SimpleNamespace(get_mesh=lambda name: SimpleNamespace(get_local_rank=lambda: 0)) + self.param = torch.nn.Parameter(torch.zeros(1)) + self.ema_shadow = EmaShadow([self.param], decay=0.5, flat_steps=100, keep_previous_ema=True) + if restored is not None: + with torch.no_grad(): + self.param.copy_(restored["param"]) + self.ema_shadow.load_state_dict(restored["ema"]) + self.weight_updater = SimpleNamespace(update_weights=self._capture_weight) + self.records = [] + self.saves = {} + Timer().start("train_wait") + + def _capture_weight(self): + self.published = self.param.item() + + def update_weights(self): + with patch.object(self.actor, "clear_memory"): + self.actor.FSDPTrainRayActor.update_weights(self) + return self.published, self.ema_shadow.step + + def _train_core(self, rollout_id, rollout_data): + start = time.monotonic() + with self.ema_shadow.swap_in(use_previous_ema=True): + reference = self.param.item() + assert rollout_data["rollout_id"] == rollout_id + assert reference == rollout_data["weight"] + time.sleep(self.delay) + with torch.no_grad(): + self.param.add_(1) + self.records.append((rollout_data, start, time.monotonic())) + + def train(self, rollout_id, batch): + with patch.object(self.actor.dist, "get_rank", return_value=0), patch.object( + self.actor.train_metric_utils, "log_perf_data_raw" + ): + self.actor.FSDPTrainRayActor.train(self, rollout_id, batch) + + def save_model(self, rollout_id, force_sync=False): + from copy import deepcopy + + self.saves[rollout_id] = deepcopy({"param": self.param.detach(), "ema": self.ema_shadow.state_dict()}) + + def result(self): + return self.records, self.saves + + +class TrainGroupProbe: + def __init__(self, trainer, manager): + self.trainer = trainer + self.manager = manager + self.updates = [] + + def async_train(self, rollout_id, batch): + return [self.trainer.train.remote(rollout_id, batch)] + + def update_weights(self): + requested_at = time.monotonic() + weight, ema_step = ray.get(self.trainer.update_weights.remote()) + self.updates.append((requested_at, ema_step)) + ray.get(self.manager.install.remote(weight)) + + def save_model(self, rollout_id, force_sync=False): + ray.get(self.trainer.save_model.remote(rollout_id, force_sync=force_sync)) + + +@pytest.fixture(scope="module", autouse=True) +def ray_runtime(): + with tempfile.TemporaryDirectory(prefix="pr232-ray-") as directory: + ray.init(address="local", num_cpus=2, num_gpus=0, include_dashboard=False, _temp_dir=directory) + yield + ray.shutdown() + + +def run_loop(tmp_path, monkeypatch, *, train_delay, rollout_delay, start=0, count=3, restored=None): + prompts = tmp_path / "prompts.jsonl" + prompts.write_text("".join(json.dumps({"input": str(i)}) + "\n" for i in range(16))) + args = Namespace( + start_rollout_id=start, + num_rollout=count, + save_interval=2, + eval_interval=None, + rollout_global_dataset=True, + prompt_data=str(prompts), + input_key="input", + metadata_key="metadata", + rollout_seed=42, + rollout_shuffle=False, + n_samples_per_prompt=1, + save=str(tmp_path / "ckpt"), + load=str(tmp_path / "ckpt") if restored else None, + buffer_filter_path=None, + use_wandb=False, + ) + metrics = [] + monkeypatch.setattr("train_diffusion_async.tracking_utils.log", lambda args, data, **kw: metrics.append(data)) + manager = RolloutProbe.remote(args, rollout_delay) + trainer = TrainerProbe.remote(manager, train_delay, restored) + group = TrainGroupProbe(trainer, manager) + try: + group.update_weights() + train_loop(args, group, manager, None) + records, saves = ray.get(trainer.result.remote()) + events = ray.get(manager.get_events.remote()) + for rid in saves: + saved = torch.load(f"{args.save}/rollout/global_dataset_state_dict_{rid}.pt") + assert saved["sample_index"] == rid + 1 + return records, saves, events, group.updates, metrics + finally: + ray.kill(trainer) + ray.kill(manager) + + +@pytest.mark.parametrize("train_delay,rollout_delay", [(0.05, 0.25), (0.25, 0.05)]) +def test_overlap_reference_cursor_and_update_barrier(tmp_path, monkeypatch, train_delay, rollout_delay): + records, saves, events, updates, metrics = run_loop( + tmp_path, monkeypatch, train_delay=train_delay, rollout_delay=rollout_delay + ) + assert [record[0]["sample_index"] for record in records] == [0, 1, 2] + assert [record[0]["weight"] for record in records] == [0.0, 0.0, 0.5] + assert [step for _, step in updates] == [0, 1, 2, 3] + assert saves[1]["ema"]["step"] == 2 + for i in range(2): + assert max(records[i][1], events[i + 1][1]) < min(records[i][2], events[i + 1][2]) + assert updates[i + 1][0] >= events[i + 1][2] + if rollout_delay > train_delay: + # Rollout 1 saves a checkpoint; its drain wait must not disappear into save(). + assert metrics[1]["perf/drain_wait_time"] > 0.05 + + +def test_resume_rewarms_with_restored_ema(tmp_path, monkeypatch): + _, saves, _, _, _ = run_loop(tmp_path, monkeypatch, train_delay=0, rollout_delay=0) + records, _, _, updates, _ = run_loop( + tmp_path, monkeypatch, train_delay=0, rollout_delay=0, start=2, count=5, restored=saves[1] + ) + assert [r[0]["sample_index"] for r in records] == [2, 3, 4] + assert [r[0]["weight"] for r in records] == [1.25, 1.25, 2.125] + assert [step for _, step in updates] == [2, 3, 4, 5] + + +@pytest.mark.parametrize("count", [0, 1]) +def test_empty_and_single_rollout(tmp_path, monkeypatch, count): + records, _, events, updates, _ = run_loop(tmp_path, monkeypatch, train_delay=0, rollout_delay=0, count=count) + assert len(records) == len(events) == count + assert len(updates) == count + 1 diff --git a/tests/fast/utils/test_lora_args.py b/tests/fast/utils/test_lora_args.py index a3615f779..8dcf3124c 100644 --- a/tests/fast/utils/test_lora_args.py +++ b/tests/fast/utils/test_lora_args.py @@ -18,6 +18,7 @@ def _server_args(**overrides): use_lora=True, lora_ipc_weight_sync=True, lora_target_modules=["to_q", "to_k"], + colocate=True, ) base.update(overrides) return Namespace(**base) @@ -33,3 +34,8 @@ def test_lora_ipc_omitted_when_disabled(self): args = _server_args(lora_ipc_weight_sync=False) kwargs = _compute_server_args(args, "127.0.0.1", 15000, 15001) assert "lora_target_modules" not in kwargs + + def test_lora_disaggregated_uses_resolved_args(self): + args = _server_args(lora_ipc_weight_sync=False, colocate=False) + kwargs = _compute_server_args(args, "127.0.0.1", 15000, 15001) + assert kwargs["lora_target_modules"] == ["to_q", "to_k"] diff --git a/train_diffusion_async.py b/train_diffusion_async.py new file mode 100644 index 000000000..2318bf096 --- /dev/null +++ b/train_diffusion_async.py @@ -0,0 +1,83 @@ +"""One-rollout overlap. Resume discards prefetch and starts from the saved EMA.""" + +import sys +import time + +import ray + +from miles.utils import tracking_utils +from miles.utils.metric_utils import compute_rollout_step +from miles.utils.misc import should_run_periodic_action + + +def train_loop(args, actor_model, rollout_manager, num_rollout_per_epoch): + if args.start_rollout_id >= args.num_rollout: + return + + current_batch = ray.get(rollout_manager.generate.remote(args.start_rollout_id)) + for rollout_id in range(args.start_rollout_id, args.num_rollout): + save_checkpoint = should_run_periodic_action( + rollout_id, args.save_interval, num_rollout_per_epoch, args.num_rollout + ) + # This serial actor saves the current cursor before the next generate advances it. + cursor_save = ( + rollout_manager.save.remote(rollout_id) if save_checkpoint and args.rollout_global_dataset else None + ) + next_future = rollout_manager.generate.remote(rollout_id + 1) if rollout_id + 1 < args.num_rollout else None + + ray.get(actor_model.async_train(rollout_id, current_batch)) + + # Measure the exposed generation wait before checkpoint I/O can hide it. + drain_start = time.monotonic() + if next_future is not None: + current_batch = ray.get(next_future) + drain_wait = time.monotonic() - drain_start + + if save_checkpoint: + if cursor_save is not None: + ray.get(cursor_save) + actor_model.save_model(rollout_id, force_sync=rollout_id == args.num_rollout - 1) + + # No generation is in flight while the engines install the new weights. + actor_model.update_weights() + if should_run_periodic_action(rollout_id, args.eval_interval, num_rollout_per_epoch): + ray.get(rollout_manager.eval.remote(rollout_id)) + tracking_utils.log( + args, + {"perf/drain_wait_time": drain_wait, "rollout/step": compute_rollout_step(args, rollout_id)}, + step_key="rollout/step", + ) + + +def train(args): + from miles.ray.placement_group import create_placement_groups, create_rollout_manager, create_training_models + from miles.utils.logging_utils import configure_logger + from miles.utils.tracking_utils import init_tracking + + configure_logger() + if args.colocate or args.offload_train or args.offload_rollout: + raise ValueError("async training requires separate resident train/rollout GPU pools") + args.train_async = True + + pgs = create_placement_groups(args) + init_tracking(args) + rollout_manager, num_rollout_per_epoch = create_rollout_manager(args, pgs["rollout"]) + actor_model = create_training_models(args, pgs, rollout_manager) + + # Publish initial/restored weights without advancing EMA or its decay schedule. + actor_model.update_weights() + if args.eval_interval is not None: + if args.num_rollout == 0: + ray.get(rollout_manager.eval.remote(rollout_id=0)) + elif not args.skip_eval_before_train: + ray.get(rollout_manager.eval.remote(args.start_rollout_id)) + + train_loop(args, actor_model, rollout_manager, num_rollout_per_epoch) + ray.get(rollout_manager.dispose.remote()) + + +if __name__ == "__main__": + from miles.utils.arguments import parse_args + + sys.stdout.reconfigure(line_buffering=True) + train(parse_args())