From 9db4f52b29ddd4eb5c9bad3b9aeca0b26254d9b2 Mon Sep 17 00:00:00 2001 From: niehen6174 Date: Mon, 27 Jul 2026 12:18:40 +0000 Subject: [PATCH 01/12] refactor(diffusion): actor-owned forward with pluggable prepare/loss formula Keep DiT forward in the FSDP actor and swap convert/prepare/formula hooks for Flow-GRPO vs NFT, instead of dispatching a whole-package loss_fn. --- miles/backends/fsdp_utils/actor.py | 365 ++++++------------ miles/backends/fsdp_utils/lora_ema.py | 128 ++++++ .../backends/fsdp_utils/loss_hub/__init__.py | 26 ++ .../fsdp_utils/loss_hub/advantages.py | 50 +++ miles/backends/fsdp_utils/loss_hub/context.py | 54 +++ miles/backends/fsdp_utils/loss_hub/losses.py | 250 ++++++++++++ miles/backends/fsdp_utils/loss_hub/nft.py | 298 ++++++++++++++ miles/backends/fsdp_utils/metrics.py | 8 + miles/ray/rollout.py | 7 +- miles/utils/arguments.py | 162 +++++++- miles/utils/train_data_utils.py | 34 ++ scripts/run-diffusion-nft-sd3-pickscore.sh | 163 ++++++++ .../backends/fsdp_utils/test_loss_hub_nft.py | 159 ++++++++ 13 files changed, 1451 insertions(+), 253 deletions(-) create mode 100644 miles/backends/fsdp_utils/lora_ema.py create mode 100644 miles/backends/fsdp_utils/loss_hub/__init__.py create mode 100644 miles/backends/fsdp_utils/loss_hub/advantages.py create mode 100644 miles/backends/fsdp_utils/loss_hub/context.py create mode 100644 miles/backends/fsdp_utils/loss_hub/losses.py create mode 100644 miles/backends/fsdp_utils/loss_hub/nft.py create mode 100755 scripts/run-diffusion-nft-sd3-pickscore.sh create mode 100644 tests/fast/backends/fsdp_utils/test_loss_hub_nft.py diff --git a/miles/backends/fsdp_utils/actor.py b/miles/backends/fsdp_utils/actor.py index 378c52be..d3ea0657 100644 --- a/miles/backends/fsdp_utils/actor.py +++ b/miles/backends/fsdp_utils/actor.py @@ -1,7 +1,7 @@ import logging import warnings from argparse import Namespace -from contextlib import contextmanager, nullcontext +from contextlib import contextmanager import ray import torch @@ -24,8 +24,8 @@ from miles.utils.train_data_utils import ( build_microbatch_schedule, scheduler_meta_from_rollout, - stack_train_pair_rollout_debug, validate_same_microbatch_counts_across_train_ranks, + validate_sample_aligned_windows, ) from . import checkpoint from .diffusion_update_weight_utils import ( @@ -33,6 +33,8 @@ DiffusionUpdateWeightFromTensorLoRA, DiffusionUpdateWeightFromTensorLoRAIPC, ) +from .loss_hub import DiffusionLossContext, resolve_loss_formula_fn, resolve_prepare_fn +from .lora_ema import LoraEmaShadow, lora_ema_rollout_policy, lora_ema_shadow_enabled, resolve_lora_ema_kwargs from .lr_scheduler import get_lr_scheduler from .metrics import new_metric_buffer, record_rollout_train_abs_diff from .parallel import create_fsdp_parallel_state @@ -176,6 +178,10 @@ def init(self, args: Namespace, role: str, with_ref: bool = False) -> int: # ty sde_timestep_divisor=self.train_pipeline_config.sde_timestep_divisor, ) + # Replaceable parts (defaults = Flow-GRPO). DiT forward stays in this actor. + self.prepare_fn = resolve_prepare_fn(args) + self.loss_formula_fn = resolve_loss_formula_fn(args) + if args.optimizer == "adam": self.optimizer = torch.optim.AdamW( (p for p in self.model.parameters() if p.requires_grad), @@ -202,6 +208,15 @@ def init(self, args: Namespace, role: str, with_ref: bool = False) -> int: # ty checkpoint_payload = checkpoint.load(self) + # Optional LoRA EMA shadow (pi_old). Loss formulas that set ref_mode="ema" + # consume it via actor reference forward. + self.ema_shadow = None + if lora_ema_shadow_enabled(self.args): + self.ema_shadow = LoraEmaShadow( + (p for m in self.models.values() for p in m.parameters()), + **resolve_lora_ema_kwargs(self.args), + ) + # sglang-d now supports /update_weights_from_tensor (PR #20464). if self.args.debug_train_only: self.weight_updater = None @@ -288,7 +303,17 @@ def update_weights(self) -> None: # type: ignore[override] if dist.get_rank() == 0: ray.get(self.rollout_manager.clear_num_new_engines.remote()) - self.weight_updater.update_weights() + if self.ema_shadow is not None: + delta = self.ema_shadow.update() + if dist.get_rank() == 0: + logger.info("LoRA EMA shadow updated (decay=%.4f step=%d)", delta, self.ema_shadow.step) + if lora_ema_rollout_policy(self.args) == "ema": + with self.ema_shadow.swap_in(): + self.weight_updater.update_weights() + else: + self.weight_updater.update_weights() + else: + self.weight_updater.update_weights() clear_memory() def _log_metrics(self, rollout_id: int, log_dict: dict[str, float], step: int) -> None: @@ -327,11 +352,10 @@ def train(self, rollout_id: int, rollout_data_ref) -> None: # type: ignore[over ) def _train_core(self, rollout_id: int, rollout_data) -> None: - """Diffusion GRPO: ``rollout_data[train_data]`` is a flat list of train-pair dicts. + """Shared train loop: schedule → prepare → DiT forward → loss formula → backward. - Optimizer windows are contiguous groups of train pairs. Within a window, consecutive microbatches of - size ``--micro-batch-size`` drive one forward+backward each; gradients - scale as mean over all train pairs in the window (``loss_chunk / num_local_pairs``). + ``prepare_fn`` / ``loss_formula_fn`` are the replaceable parts (defaults = + Flow-GRPO). DiT forward always runs here. """ device = torch.cuda.current_device() @@ -340,19 +364,9 @@ def _train_core(self, rollout_id: int, rollout_data) -> None: raise ValueError("rollout_data['train_data'] is empty") num_pairs = len(train_pairs) - - # ------------- CFG Scale ------------- - guidance_scale = self.args.diffusion_guidance_scale - true_cfg_scale = self.args.diffusion_true_cfg_scale - cfg_scale = true_cfg_scale if true_cfg_scale is not None else guidance_scale - use_cfg = cfg_scale > 0 - - # ------------- Loss / SDE Parameters ------------- - clip_range = self.args.diffusion_clip_range - noise_level = self.args.diffusion_noise_level num_train_timesteps = self.scheduler.config.num_train_timesteps - # ------------- KL loss ------------- + # ------------- KL precondition (Flow-GRPO formula may request ref forward) ------------- kl_beta = float(self.args.diffusion_kl_beta) if kl_beta > 0 and not self.args.use_lora: raise ValueError( @@ -391,6 +405,23 @@ def _train_core(self, rollout_id: int, rollout_data) -> None: microbatch_schedule=microbatch_schedule, parallel_state=self.parallel_state, ) + if getattr(self.loss_formula_fn, "requires_sample_aligned_windows", False): + validate_sample_aligned_windows( + train_pairs=train_pairs, + microbatch_schedule=microbatch_schedule, + ) + + loss_ctx = DiffusionLossContext( + models=self.models, + model=self.model, + train_pipeline_config=self.train_pipeline_config, + sde_backend=self.sde_backend, + scheduler=self.scheduler, + args=self.args, + forward_dtype=self._forward_dtype, + device=device, + ema_shadow=self.ema_shadow, + ) # ------------- Recompute old log-probs (impl-consistent PPO ratio) ------------- if self.args.diffusion_recompute_old_log_prob: @@ -402,15 +433,9 @@ def _train_core(self, rollout_id: int, rollout_data) -> None: legacy_pad_to_len = self._maybe_legacy_window_pad_len(train_pairs, microbatch_ranges) for pair_lo, pair_hi in microbatch_ranges: self._forward_train_pair_batch( + loss_ctx, train_pairs[pair_lo:pair_hi], - use_cfg=use_cfg, - guidance_scale=guidance_scale, - true_cfg_scale=true_cfg_scale, - clip_range=clip_range, - noise_level=noise_level, - num_train_timesteps=num_train_timesteps, metrics=unused_metrics, - device=device, pad_to_len=legacy_pad_to_len, write_old_log_prob=True, ) @@ -432,16 +457,9 @@ def _train_core(self, rollout_id: int, rollout_data) -> None: for pair_lo, pair_hi in microbatch_ranges: chunk = train_pairs[pair_lo:pair_hi] loss_sum = self._forward_train_pair_batch( + loss_ctx, chunk, - use_cfg=use_cfg, - guidance_scale=guidance_scale, - true_cfg_scale=true_cfg_scale, - clip_range=clip_range, - noise_level=noise_level, - num_train_timesteps=num_train_timesteps, metrics=metrics, - device=device, - kl_beta=kl_beta, pad_to_len=legacy_pad_to_len, old_log_prob_from_new=old_log_prob_from_new, ) @@ -483,237 +501,90 @@ def _maybe_legacy_window_pad_len(self, train_pairs: list, microbatch_ranges: lis conds.append(env.neg_cond_kwargs) return self.train_pipeline_config.maybe_legacy_window_pad_len(conds) + def _resolve_ref_mode(self) -> str | None: + """Whether the actor should run a no-grad reference DiT forward. + + Loss formulas may set ``ref_mode`` to ``"ema"`` / ``"lora_base"``. + Flow-GRPO defaults to ``lora_base`` when ``--diffusion-kl-beta > 0``. + """ + explicit = getattr(self.loss_formula_fn, "ref_mode", None) + if explicit == "ema" and getattr(self.args, "diffusion_nft_ref_mode", "ema") == "base": + explicit = "lora_base" + if explicit in ("ema", "lora_base"): + if explicit == "ema" and self.ema_shadow is None: + return "lora_base" + return explicit + if float(self.args.diffusion_kl_beta) > 0: + return "lora_base" + return None + def _forward_train_pair_batch( self, + ctx: DiffusionLossContext, batch: list, *, - use_cfg: bool, - guidance_scale: float, - true_cfg_scale: float | None, - clip_range: float, - noise_level: float, - num_train_timesteps: int, metrics: MetricBuffer, - device: torch.device, - kl_beta: float = 0.0, pad_to_len: int | None = None, write_old_log_prob: bool = False, old_log_prob_from_new: bool = False, ) -> torch.Tensor | None: - """One DiT forward + PPO loss over ``len(batch)`` train pairs. Returns sum of per-pair losses. + """Shared micro-batch path: prepare → DiT forward → [ref] → loss formula. - With ``write_old_log_prob``, the computed log-prob overwrites each pair's - ``log_prob_old`` and no loss is returned (caller wraps in ``no_grad``). + Only ``prepare_fn`` / ``loss_formula_fn`` are swappable; DiT forward stays here. + """ + if not batch: + raise ValueError("_forward_train_pair_batch received empty batch") - With ``old_log_prob_from_new``, the PPO ratio uses this forward's detached - log-prob as old (valid only under pre-update weights).""" - forward_dtype = self._forward_dtype + prepared = self.prepare_fn(ctx, batch, pad_to_len=pad_to_len) train_pipeline_config = self.train_pipeline_config - bsz = len(batch) - - def _stack(key): - return torch.stack([pair[key] for pair in batch]).to(device=device, dtype=torch.float32) - - latents_microbatch = _stack("latent") # (bsz, *latent_dims) - next_latents_microbatch = _stack("next_latent") # (bsz, *latent_dims) - timesteps_microbatch = _stack("timestep") # (bsz,) -- per-pair timestep is scalar - next_timesteps_microbatch = _stack("next_timestep") # (bsz,) -- next rollout timestep (0 at terminal) - log_prob_old_microbatch = _stack("log_prob_old") # (bsz,) -- per-pair log_prob is scalar - - advantage = torch.tensor( # (bsz,) - [float(pair["advantage"]) for pair in batch], - device=device, - dtype=torch.float32, - ) - advantage = torch.clamp(advantage, -self.args.diffusion_adv_clip_max, self.args.diffusion_adv_clip_max) - - if len(self.models) == 1: - component, model = next(iter(self.models.items())) - else: - components = { - train_pipeline_config.component_for_timestep(t, num_train_timesteps) - for t in timesteps_microbatch.tolist() - } - # to prevent mixing denoising phases in a single micro-batch - # Just in case when some customized step strategy is used that - # may violate the assumption of one phase per micro-batch, we raise an error here - if len(components) > 1: - raise ValueError( - f"Micro-batch mixes denoising phases {sorted(components)}; set " - "--micro-batch-size 1 so each forward is phase-pure (one DiT, one CFG scale)." - ) - component = components.pop() - model = self.models[component] - guidance_scale = train_pipeline_config.select_guidance_scale( - float(timesteps_microbatch[0]), - num_train_timesteps, - guidance_scale, - self.args.diffusion_guidance_scale_2, - ) - - # sgl-d's Qwen DiT divides timestep by num_train_timesteps inside - # forward; diffusers' does not. SD3 already expects raw timesteps. - if train_pipeline_config.needs_timestep_scaling: - timesteps_for_model = timesteps_microbatch / float(num_train_timesteps) - else: - timesteps_for_model = timesteps_microbatch - - pos_list = [ - train_pipeline_config.prepare_cond_kwargs(batch[i]["denoising_env"].pos_cond_kwargs, device) - for i in range(bsz) - ] - neg_list = ( - [ - train_pipeline_config.prepare_cond_kwargs(batch[i]["denoising_env"].neg_cond_kwargs, device) - for i in range(bsz) - ] - if use_cfg - else None - ) + forward_dtype = self._forward_dtype - # Collate cond once, up front. With CFG batching, pos+neg must share one - # padded width and go through a single joint forward, so build that joint cond - # directly; otherwise build pos (and neg) separately. (A single-sample - # timestep-stacked micro-batch is just collate of bsz copies of one sample -- - # bitwise-equivalent to the old expand_cond_for_timestep_batch path; the - # all-True mask qwen adds is a verified forward no-op, see - # tests/manual/check_mask_equivalence.py.) - cfg_batching = use_cfg and bool(self.args.fsdp_cfg_batching) - joint_cond = None - pos_cond_microbatch = None - neg_cond_microbatch = None - if cfg_batching: - joint_cond = _cast_cond_to_dtype( - train_pipeline_config.collate_cond_for_sample_batch( - pos_list + neg_list, device, pad_to_len=pad_to_len - ), - forward_dtype, - ) - else: - pos_cond_microbatch = _cast_cond_to_dtype( - train_pipeline_config.collate_cond_for_sample_batch(pos_list, device, pad_to_len=pad_to_len), - forward_dtype, + latents_input = prepared.latents.to(forward_dtype) + timesteps_input = prepared.timesteps_for_model.to(forward_dtype) + + def _compute_noise_pred() -> torch.Tensor: + return train_pipeline_config.compute_noise_pred( + model=prepared.model, + latents_input=latents_input, + timesteps_input=timesteps_input, + pos_cond=prepared.pos_cond, + neg_cond=prepared.neg_cond, + joint_cond=prepared.joint_cond, + use_cfg=prepared.use_cfg, + cfg_batching=prepared.cfg_batching, + guidance_scale=prepared.guidance_scale, + true_cfg_scale=prepared.true_cfg_scale, ) - if use_cfg and neg_list is not None: - neg_cond_microbatch = _cast_cond_to_dtype( - train_pipeline_config.collate_cond_for_sample_batch(neg_list, device, pad_to_len=pad_to_len), - forward_dtype, - ) - - # Cast inputs explicitly: FSDP MixedPrecisionPolicy casts params but - # leaves fp32 inputs, which would run first matmul at higher precision - # than rollout → systematic noise_pred drift. - latents_input = latents_microbatch.to(forward_dtype) - timesteps_input = timesteps_for_model.to(forward_dtype) - - def _compute_noise_pred(disable_adapter: bool = False) -> torch.Tensor: - adapter_ctx = model.disable_adapter() if disable_adapter else nullcontext() - with adapter_ctx: - return train_pipeline_config.compute_noise_pred( - model=model, - latents_input=latents_input, - timesteps_input=timesteps_input, - pos_cond=pos_cond_microbatch, - neg_cond=neg_cond_microbatch, - joint_cond=joint_cond, - use_cfg=use_cfg, - cfg_batching=cfg_batching, - guidance_scale=guidance_scale, - true_cfg_scale=true_cfg_scale, - ) - noise_pred_microbatch = _compute_noise_pred() - - _, log_prob_new_microbatch, prev_sample_mean_new, std_dev_t_new = self.sde_backend.sde_step_logprob( - noise_pred_microbatch.float(), - timesteps_microbatch, - next_timesteps_microbatch, - latents_microbatch.float(), - prev_sample=next_latents_microbatch.float(), - noise_level=noise_level, + new_pred = _compute_noise_pred() + + ref_pred = None + ref_mode = self._resolve_ref_mode() + if ref_mode is not None: + if ref_mode == "ema": + if self.ema_shadow is None: + raise ValueError("ref_mode=ema requires --lora-ema-shadow") + ref_ctx = self.ema_shadow.swap_in() + elif ref_mode == "lora_base": + if not hasattr(prepared.model, "disable_adapter"): + raise ValueError("ref_mode=lora_base requires a LoRA model exposing disable_adapter()") + ref_ctx = prepared.model.disable_adapter() + else: + raise ValueError(f"unknown ref_mode {ref_mode!r}") + with torch.no_grad(), ref_ctx: + ref_pred = _compute_noise_pred().detach() + + return self.loss_formula_fn( + ctx, + batch, + prepared, + new_pred=new_pred, + ref_pred=ref_pred, + metrics=metrics, + write_old_log_prob=write_old_log_prob, + old_log_prob_from_new=old_log_prob_from_new, ) - if write_old_log_prob: - for pair, log_prob in zip(batch, log_prob_new_microbatch, strict=True): - pair["log_prob_old"] = log_prob.cpu() - return None - - log_prob_new = log_prob_new_microbatch # (bsz,) -- sde_step_with_logprob means over non-batch dims - log_prob_old = log_prob_new.detach() if old_log_prob_from_new else log_prob_old_microbatch # (bsz,) - ratio = torch.exp(log_prob_new - log_prob_old) # (bsz,) - unclipped = -advantage * ratio - clipped = -advantage * torch.clamp(ratio, 1.0 - clip_range, 1.0 + clip_range) - per_pair_loss = torch.maximum(unclipped, clipped) - loss_sum = per_pair_loss.sum() - - # ------------- KL loss (vs LoRA base model as reference) ------------- - kl_sum = loss_sum.new_zeros(()) - if kl_beta > 0: - with torch.no_grad(): - ref_noise_pred_microbatch = _compute_noise_pred(disable_adapter=True) - # TODO: unify sde_step_with_logprob with rollout and trainer forward paths. - _, _, prev_sample_mean_ref, _ = self.sde_backend.sde_step_logprob( - ref_noise_pred_microbatch.float(), - timesteps_microbatch, - next_timesteps_microbatch, - latents_microbatch.float(), - prev_sample=next_latents_microbatch.float(), - noise_level=noise_level, - ) - kl_per_pair = ((prev_sample_mean_new - prev_sample_mean_ref) ** 2).mean( - dim=tuple(range(1, prev_sample_mean_new.ndim)), - keepdim=True, - ) / (2 * std_dev_t_new**2) - loss_sum = loss_sum + kl_beta * kl_per_pair.sum() - kl_sum = kl_per_pair.sum() - - with torch.no_grad(): - # Sums, not means: the flush divides by the globally summed pair count. - metrics.emit_mean("loss", total=loss_sum, count=bsz) - metrics.emit_mean("policy_loss", total=per_pair_loss.sum(), count=bsz) - metrics.emit_mean("kl_loss", total=kl_sum, count=bsz) - metrics.emit_mean("loss_abs_mean", total=per_pair_loss.abs().sum(), count=bsz) - metrics.emit_mean("adv_abs_mean", total=advantage.abs().sum(), count=bsz) - metrics.emit_mean("ratio_abs_minus_1", total=(ratio - 1.0).abs().sum(), count=bsz) - metrics.emit_mean("approx_kl", total=0.5 * ((log_prob_new - log_prob_old) ** 2).sum(), count=bsz) - metrics.emit_mean("clipfrac", total=(torch.abs(ratio - 1.0) > clip_range).float().sum(), count=bsz) - # Single-pair probes: one observation per micro-batch, not a batch statistic. - metrics.emit_mean("log_prob_new_idx_0", total=log_prob_new[0], count=1) - metrics.emit_mean("log_prob_old_idx_0", total=log_prob_old[0], count=1) - log_prob_abs_diff_sum = torch.abs(log_prob_new - log_prob_old).sum() - metrics.emit_mean("log_prob_mean_abs_diff", total=log_prob_abs_diff_sum, count=bsz) - if len(self.models) > 1: - metrics.emit_mean(f"log_prob_mean_abs_diff_{component}", total=log_prob_abs_diff_sum, count=bsz) - - # model_output_* checks the train forward reproduces the rollout forward -- the only - # model-dependent consistency metric (std_dev/prev_sample_mean are deterministic - # functions of it). Matches the legacy actor metric name. - rollout_model_output = stack_train_pair_rollout_debug(batch, "rollout_step_model_output") - if rollout_model_output is not None: - record_rollout_train_abs_diff( - metrics, - "model_output", - noise_pred_microbatch.float(), - rollout_model_output.to(device=device, dtype=torch.float32), - component=component if len(self.models) > 1 else None, - ) - - return loss_sum - - -def _cast_cond_to_dtype(cond: dict, dtype: torch.dtype) -> dict: - """Cast floating-point tensors to the model's compute dtype; leave bool - masks / int / list / scalar values untouched. The bool - encoder_hidden_states_mask must NOT be cast. - """ - out: dict = {} - for k, v in cond.items(): - if isinstance(v, torch.Tensor) and v.dtype.is_floating_point: - out[k] = v.to(dtype) - else: - out[k] = v - return out - @torch.no_grad() def move_torch_optimizer(optimizer, device): diff --git a/miles/backends/fsdp_utils/lora_ema.py b/miles/backends/fsdp_utils/lora_ema.py new file mode 100644 index 00000000..c0b678cc --- /dev/null +++ b/miles/backends/fsdp_utils/lora_ema.py @@ -0,0 +1,128 @@ +"""Generic LoRA EMA shadow for diffusion (and future FSDP) training. + +Algorithms that need a slow-moving reference / sampling policy ``pi_old`` share +``LoraEmaShadow``: trainable LoRA weights plus EMA buffers and ``swap_in()`` +for temporary in-place weight exchange. + +Lifecycle (actor / weight sync):: + + ema.update() + with ema.swap_in(): + weight_updater.update_weights() + +Loss-side reference forward (via ``DiffusionLossContext.ema_shadow``):: + + with torch.no_grad(), ctx.ema_shadow.swap_in(): + old_pred = forward(...) + +Works with FSDP2 DTensor shards (per-rank local swap) and colocate CPU offload. + +Checkpointing (intentionally not wired yet) +------------------------------------------- +``shadow`` / ``step`` are **not** saved or restored by ``fsdp_utils.checkpoint``. +On resume the actor rebuilds EMA from the loaded LoRA weights, so ``pi_old`` +cold-starts (decay schedule restarts at step 0). Fine for single-shot runs; +wrong for mid-run resume that must match UniRL's slow ``pi_old``. + +Wiring it later is non-trivial: buffers are per-rank plain clones (not in the +FSDP/DCP model state), must stay aligned with the trainable-param order, and +must not be saved while ``swap_in()`` is active. Prefer a side file such as +``iter_*/lora_ema.pt`` over stuffing into the DCP model dict. +""" + +from __future__ import annotations + +from argparse import Namespace +from contextlib import contextmanager +from typing import Iterable + +import torch +import torch.nn as nn + + +def _local(t: torch.Tensor) -> torch.Tensor: + """Local shard of a (possibly DTensor) tensor; EMA/swap is per-rank, no comm.""" + return t._local_tensor if hasattr(t, "_local_tensor") else t + + +def lora_ema_shadow_enabled(args: Namespace) -> bool: + """True when a LoRA EMA shadow should be constructed.""" + return bool(getattr(args, "lora_ema_shadow", False)) + + +def lora_ema_rollout_policy(args: Namespace) -> str: + """Which LoRA weights to push to rollout engines after each rollout ('live' or 'ema').""" + return getattr(args, "lora_ema_rollout_policy", "live") + + +def resolve_lora_ema_kwargs(args: Namespace) -> dict[str, float | int]: + """Read normalized ``lora_ema_*`` fields from ``args`` (see ``miles_validate_args``).""" + return { + "decay": float(getattr(args, "lora_ema_decay", 0.001)), + "uprate": float(getattr(args, "lora_ema_uprate", 0.001)), + "uphold": float(getattr(args, "lora_ema_uphold", 0.5)), + "flat_steps": int(getattr(args, "lora_ema_flat_steps", 0)), + } + + +class LoraEmaShadow: + """EMA shadow of trainable (LoRA) parameters. + + Not part of the FSDP checkpoint payload today (see module docstring). + """ + + def __init__( + self, + parameters: Iterable[nn.Parameter], + *, + decay: float = 0.001, + uprate: float = 0.001, + uphold: float = 0.5, + flat_steps: int = 0, + ) -> None: + self.decay = float(decay) + self.uprate = float(uprate) + self.uphold = float(uphold) + self.flat_steps = int(flat_steps) + self.step = 0 + self._swapped = False + + self.params = [p for p in parameters if p.requires_grad] + if not self.params: + raise ValueError("LoraEmaShadow: model has no trainable parameters") + self.shadow = [_local(p.detach()).clone() for p in self.params] + + def decay_at(self, t: int) -> float: + if t <= self.flat_steps: + return self.decay + return float(min((t - self.flat_steps) * self.uprate, self.uphold)) + + @torch.no_grad() + def update(self) -> float: + """theta_old <- delta * theta_old + (1 - delta) * theta.""" + if self._swapped: + raise RuntimeError("LoraEmaShadow.update called while swapped in") + self.step += 1 + delta = self.decay_at(self.step) + for live, sh in zip(self.params, self.shadow): + sh.mul_(delta).add_(_local(live.detach()).to(sh.device), alpha=1.0 - delta) + return delta + + @contextmanager + def swap_in(self): + """Temporarily expose EMA weights as the live parameters.""" + self._swap() + self._swapped = True + try: + yield + finally: + self._swap() + self._swapped = False + + @torch.no_grad() + def _swap(self) -> None: + for live, sh in zip(self.params, self.shadow): + live_local = _local(live.data) + tmp = live_local.clone() + live_local.copy_(sh) + sh.copy_(tmp) diff --git a/miles/backends/fsdp_utils/loss_hub/__init__.py b/miles/backends/fsdp_utils/loss_hub/__init__.py new file mode 100644 index 00000000..668e2171 --- /dev/null +++ b/miles/backends/fsdp_utils/loss_hub/__init__.py @@ -0,0 +1,26 @@ +"""Diffusion train hooks: prepare + loss formula (actor owns DiT forward). + +Swap pieces via: + ``--custom-prepare-train-batch-path`` + ``--custom-loss-function-path`` (formula only; receives new_pred / ref_pred) + ``--custom-convert-samples-to-train-data-path`` +""" + +from miles.backends.fsdp_utils.loss_hub.advantages import grpo_normalize_rewards +from miles.backends.fsdp_utils.loss_hub.context import DiffusionLossContext, PreparedBatch +from miles.backends.fsdp_utils.loss_hub.losses import ( + flow_grpo_loss_formula, + prepare_flow_grpo_batch, + resolve_loss_formula_fn, + resolve_prepare_fn, +) + +__all__ = [ + "DiffusionLossContext", + "PreparedBatch", + "flow_grpo_loss_formula", + "grpo_normalize_rewards", + "prepare_flow_grpo_batch", + "resolve_loss_formula_fn", + "resolve_prepare_fn", +] diff --git a/miles/backends/fsdp_utils/loss_hub/advantages.py b/miles/backends/fsdp_utils/loss_hub/advantages.py new file mode 100644 index 00000000..4fed8912 --- /dev/null +++ b/miles/backends/fsdp_utils/loss_hub/advantages.py @@ -0,0 +1,50 @@ +"""Reward → train-signal helpers for diffusion (customization building blocks). + +Default GRPO group normalization. Override with +``--custom-reward-post-process-path`` pointing at a function with the same +signature as ``grpo_normalize_rewards``. +""" + +from __future__ import annotations + +from argparse import Namespace + +import torch + +from miles.utils.types import Sample + + +def grpo_normalize_rewards( + args: Namespace, + samples: list[Sample] | list[list[Sample]], +) -> tuple[list[float], list[float]]: + """Group-relative reward normalization used by Flow-GRPO. + + Returns ``(raw_rewards, normalized_rewards)``. Normalized values are used + as per-sample advantages when building train pairs. + + ``--globalize-reward-mean`` / ``--globalize-reward-std`` are orthogonal. + flow_grpo pickscore_qwenimage uses per-prompt mean + global std + (``PerPromptStatTracker`` with ``global_std=True``), which is + ``--globalize-reward-std`` alone. + """ + raw_rewards = [sample.get_reward_value(args) for sample in samples] + + rewards_flat = torch.tensor(raw_rewards, dtype=torch.float) + rewards = rewards_flat.view(-1, args.n_samples_per_prompt) + + if args.globalize_reward_mean: + mean = rewards_flat.mean() + else: + mean = rewards.mean(dim=-1, keepdim=True) + rewards = rewards - mean + + if args.grpo_std_normalization: + if args.globalize_reward_std: + std = rewards_flat.std() + else: + std = rewards.std(dim=-1, keepdim=True) + # matches flow_grpo's `+ 1e-4` in both stat_tracking branches + rewards = rewards / (std + 1e-4) + + return raw_rewards, rewards.flatten().tolist() diff --git a/miles/backends/fsdp_utils/loss_hub/context.py b/miles/backends/fsdp_utils/loss_hub/context.py new file mode 100644 index 00000000..036de6c4 --- /dev/null +++ b/miles/backends/fsdp_utils/loss_hub/context.py @@ -0,0 +1,54 @@ +"""Shared handles passed into diffusion prepare / loss-formula hooks.""" + +from __future__ import annotations + +from argparse import Namespace +from dataclasses import dataclass, field +from typing import Any + +import torch +import torch.nn as nn + + +@dataclass +class DiffusionLossContext: + """Train-side handles for prepare / loss-formula callables. + + Owned by the FSDP actor; kept free of Ray / optim internals so hooks stay + unit-testable and swappable via ``--custom-*-path``. + """ + + models: dict[str, torch.nn.Module] + model: torch.nn.Module + train_pipeline_config: Any + sde_backend: Any + scheduler: Any + args: Namespace + forward_dtype: torch.dtype + device: torch.device + # Optional LoRA EMA shadow handle; owned by actor (see ``fsdp_utils.lora_ema``). + ema_shadow: Any = None + + +@dataclass +class PreparedBatch: + """Actor-owned DiT forward inputs produced by a prepare hook. + + ``extras`` carries algorithm-specific tensors for the loss formula + (e.g. ``next_latents`` / ``log_prob_old`` for Flow-GRPO, ``x0`` for NFT). + """ + + latents: torch.Tensor + timesteps: torch.Tensor + timesteps_for_model: torch.Tensor + model: nn.Module + component_name: str + guidance_scale: float + use_cfg: bool + cfg_batching: bool + true_cfg_scale: float | None + pos_cond: dict | None + neg_cond: dict | None + joint_cond: dict | None + advantage: torch.Tensor + extras: dict[str, Any] = field(default_factory=dict) diff --git a/miles/backends/fsdp_utils/loss_hub/losses.py b/miles/backends/fsdp_utils/loss_hub/losses.py new file mode 100644 index 00000000..0e6b20d9 --- /dev/null +++ b/miles/backends/fsdp_utils/loss_hub/losses.py @@ -0,0 +1,250 @@ +"""Default Flow-GRPO prepare + loss formula (actor owns DiT forward). + +Custom algorithms swap ``--custom-prepare-train-batch-path`` and/or +``--custom-loss-function-path`` (formula only: receives ``new_pred`` / ``ref_pred``). +""" + +from __future__ import annotations + +from argparse import Namespace +from typing import Callable + +import torch + +from miles.backends.fsdp_utils.loss_hub.context import DiffusionLossContext, PreparedBatch +from miles.backends.fsdp_utils.metrics import record_rollout_train_abs_diff +from miles.utils.metric_buffer import MetricBuffer +from miles.utils.misc import load_function +from miles.utils.train_data_utils import stack_train_pair_rollout_debug + +PrepareFn = Callable[..., PreparedBatch] +LossFormulaFn = Callable[..., torch.Tensor | None] + + +def _cast_cond_to_dtype(cond: dict, dtype: torch.dtype) -> dict: + out = {} + for k, v in cond.items(): + if isinstance(v, torch.Tensor) and v.dtype.is_floating_point: + out[k] = v.to(dtype=dtype) + else: + out[k] = v + return out + + +def _stack_pair_field(batch: list[dict], key: str, device: torch.device) -> torch.Tensor: + return torch.stack([pair[key] for pair in batch]).to(device=device, dtype=torch.float32) + + +def prepare_flow_grpo_batch( + ctx: DiffusionLossContext, + batch: list[dict], + *, + pad_to_len: int | None = None, +) -> PreparedBatch: + """Stack SDE-pair fields and build CFG conditioning (guidance from args).""" + args = ctx.args + device = ctx.device + config = ctx.train_pipeline_config + num_train_timesteps = int(ctx.scheduler.config.num_train_timesteps) + bsz = len(batch) + + latents = _stack_pair_field(batch, "latent", device) + next_latents = _stack_pair_field(batch, "next_latent", device) + timesteps = _stack_pair_field(batch, "timestep", device) + next_timesteps = _stack_pair_field(batch, "next_timestep", device) + log_prob_old = _stack_pair_field(batch, "log_prob_old", device) + advantage = torch.tensor( + [float(pair["advantage"]) for pair in batch], + device=device, + dtype=torch.float32, + ) + advantage = torch.clamp(advantage, -args.diffusion_adv_clip_max, args.diffusion_adv_clip_max) + + guidance_scale = args.diffusion_guidance_scale + true_cfg_scale = args.diffusion_true_cfg_scale + cfg_scale = true_cfg_scale if true_cfg_scale is not None else guidance_scale + use_cfg = cfg_scale > 0 + + if len(ctx.models) == 1: + component_name, model = next(iter(ctx.models.items())) + else: + components = {config.component_for_timestep(t, num_train_timesteps) for t in timesteps.tolist()} + if len(components) > 1: + raise ValueError( + f"Micro-batch mixes denoising phases {sorted(components)}; set " + "--micro-batch-size 1 so each forward is phase-pure (one DiT, one CFG scale)." + ) + component_name = components.pop() + model = ctx.models[component_name] + guidance_scale = config.select_guidance_scale( + float(timesteps[0]), + num_train_timesteps, + guidance_scale, + args.diffusion_guidance_scale_2, + ) + + if config.needs_timestep_scaling: + timesteps_for_model = timesteps / float(num_train_timesteps) + else: + timesteps_for_model = timesteps + + pos_list = [config.prepare_cond_kwargs(batch[i]["denoising_env"].pos_cond_kwargs, device) for i in range(bsz)] + neg_list = ( + [config.prepare_cond_kwargs(batch[i]["denoising_env"].neg_cond_kwargs, device) for i in range(bsz)] + if use_cfg + else None + ) + cfg_batching = use_cfg and bool(args.fsdp_cfg_batching) + joint_cond = pos_cond = neg_cond = None + if cfg_batching: + joint_cond = _cast_cond_to_dtype( + config.collate_cond_for_sample_batch(pos_list + neg_list, device, pad_to_len=pad_to_len), + ctx.forward_dtype, + ) + else: + pos_cond = _cast_cond_to_dtype( + config.collate_cond_for_sample_batch(pos_list, device, pad_to_len=pad_to_len), + ctx.forward_dtype, + ) + if use_cfg and neg_list is not None: + neg_cond = _cast_cond_to_dtype( + config.collate_cond_for_sample_batch(neg_list, device, pad_to_len=pad_to_len), + ctx.forward_dtype, + ) + + return PreparedBatch( + latents=latents, + timesteps=timesteps, + timesteps_for_model=timesteps_for_model, + model=model, + component_name=component_name, + guidance_scale=guidance_scale, + use_cfg=use_cfg, + cfg_batching=cfg_batching, + true_cfg_scale=true_cfg_scale if use_cfg else None, + pos_cond=pos_cond, + neg_cond=neg_cond, + joint_cond=joint_cond, + advantage=advantage, + extras={ + "next_latents": next_latents, + "next_timesteps": next_timesteps, + "log_prob_old": log_prob_old, + }, + ) + + +def flow_grpo_loss_formula( + ctx: DiffusionLossContext, + batch: list[dict], + prepared: PreparedBatch, + *, + new_pred: torch.Tensor, + ref_pred: torch.Tensor | None, + metrics: MetricBuffer, + write_old_log_prob: bool = False, + old_log_prob_from_new: bool = False, +) -> torch.Tensor | None: + """SDE log-prob + PPO-clip (+ optional KL vs ``ref_pred``). Actor owns DiT forward.""" + args = ctx.args + clip_range = args.diffusion_clip_range + noise_level = args.diffusion_noise_level + kl_beta = float(args.diffusion_kl_beta) + + next_latents = prepared.extras["next_latents"] + next_timesteps = prepared.extras["next_timesteps"] + log_prob_old_rollout = prepared.extras["log_prob_old"] + + _, log_prob_new, prev_sample_mean_new, std_dev_t_new = ctx.sde_backend.sde_step_logprob( + new_pred.float(), + prepared.timesteps, + next_timesteps, + prepared.latents.float(), + prev_sample=next_latents.float(), + noise_level=noise_level, + ) + + if write_old_log_prob: + for pair, log_prob in zip(batch, log_prob_new, strict=True): + pair["log_prob_old"] = log_prob.cpu() + return None + + log_prob_old = log_prob_new.detach() if old_log_prob_from_new else log_prob_old_rollout + ratio = torch.exp(log_prob_new - log_prob_old) + unclipped = -prepared.advantage * ratio + clipped = -prepared.advantage * torch.clamp(ratio, 1.0 - clip_range, 1.0 + clip_range) + per_pair_loss = torch.maximum(unclipped, clipped) + loss_sum = per_pair_loss.sum() + bsz = len(batch) + + kl_sum = loss_sum.new_zeros(()) + if kl_beta > 0: + if ref_pred is None: + raise ValueError("Flow-GRPO KL requires a reference DiT forward (actor ref_mode=lora_base)") + _, _, prev_sample_mean_ref, _ = ctx.sde_backend.sde_step_logprob( + ref_pred.float(), + prepared.timesteps, + next_timesteps, + prepared.latents.float(), + prev_sample=next_latents.float(), + noise_level=noise_level, + ) + kl_per_pair = ((prev_sample_mean_new - prev_sample_mean_ref) ** 2).mean( + dim=tuple(range(1, prev_sample_mean_new.ndim)), + keepdim=True, + ) / (2 * std_dev_t_new**2) + loss_sum = loss_sum + kl_beta * kl_per_pair.sum() + kl_sum = kl_per_pair.sum() + + with torch.no_grad(): + metrics.emit_mean("loss", total=loss_sum, count=bsz) + metrics.emit_mean("policy_loss", total=per_pair_loss.sum(), count=bsz) + metrics.emit_mean("kl_loss", total=kl_sum, count=bsz) + metrics.emit_mean("loss_abs_mean", total=per_pair_loss.abs().sum(), count=bsz) + metrics.emit_mean("adv_abs_mean", total=prepared.advantage.abs().sum(), count=bsz) + metrics.emit_mean("ratio_abs_minus_1", total=(ratio - 1.0).abs().sum(), count=bsz) + metrics.emit_mean("approx_kl", total=0.5 * ((log_prob_new - log_prob_old) ** 2).sum(), count=bsz) + metrics.emit_mean("clipfrac", total=(torch.abs(ratio - 1.0) > clip_range).float().sum(), count=bsz) + metrics.emit_mean("log_prob_new_idx_0", total=log_prob_new[0], count=1) + metrics.emit_mean("log_prob_old_idx_0", total=log_prob_old[0], count=1) + log_prob_abs_diff_sum = torch.abs(log_prob_new - log_prob_old).sum() + metrics.emit_mean("log_prob_mean_abs_diff", total=log_prob_abs_diff_sum, count=bsz) + if len(ctx.models) > 1: + metrics.emit_mean( + f"log_prob_mean_abs_diff_{prepared.component_name}", + total=log_prob_abs_diff_sum, + count=bsz, + ) + + rollout_model_output = stack_train_pair_rollout_debug(batch, "rollout_step_model_output") + if rollout_model_output is not None: + record_rollout_train_abs_diff( + metrics, + "model_output", + new_pred.float(), + rollout_model_output.to(device=ctx.device, dtype=torch.float32), + component=prepared.component_name if len(ctx.models) > 1 else None, + ) + + return loss_sum + + +def resolve_prepare_fn(args: Namespace) -> PrepareFn: + path = getattr(args, "custom_prepare_train_batch_path", None) + if path: + fn = load_function(path) + if fn is None: + raise ValueError(f"Failed to load custom prepare from {path!r}") + return fn + return prepare_flow_grpo_batch + + +def resolve_loss_formula_fn(args: Namespace) -> LossFormulaFn: + """Loss *formula* only — DiT forward stays in the actor.""" + path = getattr(args, "custom_loss_function_path", None) + if path: + fn = load_function(path) + if fn is None: + raise ValueError(f"Failed to load custom loss formula from {path!r}") + return fn + return flow_grpo_loss_formula diff --git a/miles/backends/fsdp_utils/loss_hub/nft.py b/miles/backends/fsdp_utils/loss_hub/nft.py new file mode 100644 index 00000000..b201a12a --- /dev/null +++ b/miles/backends/fsdp_utils/loss_hub/nft.py @@ -0,0 +1,298 @@ +"""DiffusionNFT plugin: custom convert + prepare + loss *formula*. + +Actor still owns DiT forward (+ EMA/LoRA-base reference forward). This module +only swaps the replaceable parts via stock customization hooks. +""" + +from __future__ import annotations + +from argparse import Namespace +from typing import Any + +import torch + +from miles.backends.fsdp_utils.loss_hub.advantages import grpo_normalize_rewards +from miles.backends.fsdp_utils.loss_hub.context import DiffusionLossContext, PreparedBatch +from miles.utils.metric_buffer import MetricBuffer +from miles.utils.types import Sample + +# --------------------------------------------------------------------------- +# Forward-process math +# --------------------------------------------------------------------------- + + +def sample_noise(like: torch.Tensor, *, generator: torch.Generator | None = None) -> torch.Tensor: + return torch.randn(like.shape, device=like.device, dtype=like.dtype, generator=generator) + + +def corrupt(x0: torch.Tensor, t: torch.Tensor, eps: torch.Tensor) -> torch.Tensor: + """Linear flow: ``x_t = (1 - t) x_0 + t ε``.""" + while t.ndim < x0.ndim: + t = t.unsqueeze(-1) + return (1.0 - t) * x0 + t * eps + + +def resolve_nft_sigmas( + sigmas_or_scheduler, + *, + training_timestep_fraction: float = 0.99, +) -> torch.Tensor: + if torch.is_tensor(sigmas_or_scheduler): + ts = sigmas_or_scheduler.detach().float().flatten() + else: + raw = getattr(sigmas_or_scheduler, "sigmas", None) + if raw is None: + raise ValueError("NFT needs scheduler.sigmas (or a sigma tensor)") + ts = raw.detach().float().flatten() + if ts.numel() == 0: + raise ValueError("scheduler.sigmas is empty") + if ts.numel() > 1 and torch.isclose(ts[-1], torch.zeros((), dtype=ts.dtype), atol=1e-8): + ts = ts[:-1] + frac = float(training_timestep_fraction) + if frac < 1.0 and ts.numel() > 1: + keep = max(1, int(ts.numel() * frac)) + ts = ts[:keep] + if ts.numel() == 0: + raise ValueError("No training timesteps left after NFT sigma filtering") + return ts + + +def nft_r_from_advantages(advantages: torch.Tensor, *, adv_clip_max: float) -> torch.Tensor: + clip = float(adv_clip_max) + adv_clipped = torch.clamp(advantages, -clip, clip) + r = (adv_clipped / clip) / 2.0 + 0.5 + return torch.clamp(r, 0.0, 1.0) + + +def nft_branch_losses( + *, + x0: torch.Tensor, + xt: torch.Tensor, + t_exp: torch.Tensor, + new_pred: torch.Tensor, + old_pred: torch.Tensor, + beta: float, + use_adaptive: bool, +) -> tuple[torch.Tensor, torch.Tensor]: + reduce_dims = tuple(range(1, x0.ndim)) + positive_pred = beta * new_pred + (1.0 - beta) * old_pred + negative_pred = (1.0 + beta) * old_pred - beta * new_pred + x0_pos = xt.to(dtype=new_pred.dtype) - t_exp.to(dtype=new_pred.dtype) * positive_pred + x0_neg = xt.to(dtype=new_pred.dtype) - t_exp.to(dtype=new_pred.dtype) * negative_pred + x0_tgt = x0.to(dtype=new_pred.dtype) + if use_adaptive: + with torch.no_grad(): + weight_pos = ( + (x0_pos.detach().double() - x0_tgt.double()).abs().mean(dim=reduce_dims, keepdim=True).clamp(min=1e-5) + ).to(dtype=new_pred.dtype) + weight_neg = ( + (x0_neg.detach().double() - x0_tgt.double()).abs().mean(dim=reduce_dims, keepdim=True).clamp(min=1e-5) + ).to(dtype=new_pred.dtype) + pos_loss = ((x0_pos - x0_tgt) ** 2 / weight_pos).mean(dim=reduce_dims) + neg_loss = ((x0_neg - x0_tgt) ** 2 / weight_neg).mean(dim=reduce_dims) + else: + pos_loss = ((x0_pos - x0_tgt) ** 2).mean(dim=reduce_dims) + neg_loss = ((x0_neg - x0_tgt) ** 2).mean(dim=reduce_dims) + return pos_loss, neg_loss + + +def _cast_cond_to_dtype(cond: dict, dtype: torch.dtype) -> dict: + out = {} + for k, v in cond.items(): + if isinstance(v, torch.Tensor) and v.dtype.is_floating_point: + out[k] = v.to(dtype=dtype) + else: + out[k] = v + return out + + +# --------------------------------------------------------------------------- +# Convert (K-expanded pairs) +# --------------------------------------------------------------------------- + + +def _clean_x0_from_sample(sample: Sample) -> torch.Tensor: + traj = sample.dit_trajectory + if traj is None or traj.latents is None or traj.latents.shape[0] < 1: + raise ValueError( + f"sample {sample.index} missing dit_trajectory.latents; " + "NFT needs the final clean latent x0 from rollout" + ) + return traj.latents[-1].detach().cpu().float() + + +def convert_samples_to_nft_train_data(args: Namespace, samples: list[Sample]) -> dict[str, Any]: + """Expand every sample into K ``(x0, t)`` train pairs (sample-major).""" + raw_rewards, advantages = grpo_normalize_rewards(args, samples) + if not samples: + raise ValueError("NFT convert received empty samples") + first_traj = samples[0].dit_trajectory + if first_traj is None: + raise ValueError("sample 0 missing dit_trajectory") + if first_traj.timesteps is None: + raise ValueError("NFT needs dit_trajectory.timesteps from rollout") + num_train_timesteps = int(getattr(args, "diffusion_num_train_timesteps", 1000) or 1000) + if first_traj.sigmas is not None: + scheduler_sigmas = first_traj.sigmas.detach().cpu().float() + else: + # Match scheduler_meta_from_rollout when sglang omits sigmas (e.g. ODE rollout). + ts = first_traj.timesteps.detach().cpu().float() + scheduler_sigmas = torch.cat([ts / float(num_train_timesteps), ts.new_zeros(1)]) + scheduler_meta = { + "scheduler_timesteps": first_traj.timesteps.detach().cpu().float(), + "scheduler_sigmas": scheduler_sigmas, + } + frac = float(getattr(args, "diffusion_nft_timestep_fraction", 0.99) or 0.99) + shuffle_t = bool(getattr(args, "diffusion_nft_shuffle_timesteps", True)) + sigmas = resolve_nft_sigmas(scheduler_meta["scheduler_sigmas"], training_timestep_fraction=frac) + num_timesteps = int(sigmas.numel()) + + train_data: list[dict[str, Any]] = [] + for sample, adv, raw in zip(samples, advantages, raw_rewards, strict=True): + if sample.denoising_env is None: + raise ValueError(f"sample {sample.index} missing denoising_env") + x0 = _clean_x0_from_sample(sample) + sample_sigmas = sigmas[torch.randperm(num_timesteps)] if shuffle_t else sigmas + for t in sample_sigmas.tolist(): + train_data.append( + { + "x0": x0, + "timestep": float(t), + "denoising_env": sample.denoising_env, + "advantage": float(adv), + "raw_reward": float(raw), + "sample_index": sample.index, + "prompt": sample.prompt, + "nft_num_timesteps": num_timesteps, + } + ) + return {"train_data": train_data, **scheduler_meta} + + +# --------------------------------------------------------------------------- +# Prepare (actor still runs DiT forward on the result) +# --------------------------------------------------------------------------- + + +def prepare_nft_batch( + ctx: DiffusionLossContext, + batch: list[dict], + *, + pad_to_len: int | None = None, +) -> PreparedBatch: + """Corrupt clean x0 at each pair's sigma; CFG-free cond.""" + if len(ctx.models) != 1: + raise ValueError("DiffusionNFT currently supports a single DiT component (SD3)") + device = ctx.device + config = ctx.train_pipeline_config + bsz = len(batch) + x0 = torch.stack([pair["x0"] for pair in batch]).to(device=device, dtype=torch.float32) + t = torch.tensor([float(pair["timestep"]) for pair in batch], device=device, dtype=torch.float32) + advantage = torch.tensor([float(pair["advantage"]) for pair in batch], device=device, dtype=torch.float32) + + component_name, model = next(iter(ctx.models.items())) + pos_list = [config.prepare_cond_kwargs(batch[i]["denoising_env"].pos_cond_kwargs, device) for i in range(bsz)] + pos_cond = _cast_cond_to_dtype( + config.collate_cond_for_sample_batch(pos_list, device, pad_to_len=pad_to_len), + ctx.forward_dtype, + ) + + num_train_timesteps = int(getattr(ctx.scheduler.config, "num_train_timesteps", 1000)) + if config.needs_timestep_scaling: + timesteps_for_model = t.to(dtype=torch.float32) + else: + timesteps_for_model = t * float(num_train_timesteps) + + xt = corrupt(x0, t, sample_noise(x0)) + return PreparedBatch( + latents=xt, + timesteps=t, + timesteps_for_model=timesteps_for_model, + model=model, + component_name=component_name, + guidance_scale=0.0, + use_cfg=False, + cfg_batching=False, + true_cfg_scale=None, + pos_cond=pos_cond, + neg_cond=None, + joint_cond=None, + advantage=advantage, + extras={"x0": x0}, + ) + + +# --------------------------------------------------------------------------- +# Loss formula (receives actor's new_pred / ref_pred) +# --------------------------------------------------------------------------- + + +def nft_loss_formula( + ctx: DiffusionLossContext, + batch: list[dict], + prepared: PreparedBatch, + *, + new_pred: torch.Tensor, + ref_pred: torch.Tensor | None, + metrics: MetricBuffer, + write_old_log_prob: bool = False, + old_log_prob_from_new: bool = False, +) -> torch.Tensor | None: + """Dual-policy x0-MSE. Actor must supply ``ref_pred`` (EMA / LoRA-base).""" + if write_old_log_prob: + return None + if old_log_prob_from_new: + raise ValueError("DiffusionNFT has no PPO log-prob; old_log_prob_from_new is unsupported") + if ref_pred is None: + raise ValueError("NFT loss formula requires a reference prediction from the actor") + + args = ctx.args + beta = float(getattr(args, "diffusion_nft_beta", 1.0) or 1.0) + if beta <= 0: + raise ValueError(f"--diffusion-nft-beta must be > 0, got {beta}") + adv_clip_max = float(getattr(args, "diffusion_nft_adv_clip_max", 5.0) or 5.0) + use_adaptive = bool(getattr(args, "diffusion_nft_adaptive_weight", True)) + + x0 = prepared.extras["x0"] + t = prepared.timesteps + t_exp = t.view(len(batch), *([1] * (x0.ndim - 1))) + r = nft_r_from_advantages(prepared.advantage, adv_clip_max=adv_clip_max) + pos_loss, neg_loss = nft_branch_losses( + x0=x0, + xt=prepared.latents, + t_exp=t_exp, + new_pred=new_pred, + old_pred=ref_pred, + beta=beta, + use_adaptive=use_adaptive, + ) + r_b = r.to(dtype=pos_loss.dtype) + per_pair = (r_b * pos_loss / beta + (1.0 - r_b) * neg_loss / beta) * adv_clip_max + loss_sum = per_pair.sum() + + with torch.no_grad(): + num_timesteps = int(batch[0].get("nft_num_timesteps", 0) or 0) + per_pair_total = per_pair.sum() + bsz = len(batch) + metrics.emit_mean("loss", total=per_pair_total * float(max(num_timesteps, 1)), count=bsz) + metrics.emit_mean("nft_loss", total=per_pair_total * float(max(num_timesteps, 1)), count=bsz) + metrics.emit_mean("nft_loss_per_pair", total=per_pair_total, count=bsz) + metrics.emit_mean("nft_r_mean", total=r.sum(), count=bsz) + metrics.emit_mean("nft_pos_loss", total=pos_loss.sum(), count=bsz) + metrics.emit_mean("nft_neg_loss", total=neg_loss.sum(), count=bsz) + metrics.emit_mean("nft_adv_mean", total=prepared.advantage.sum(), count=bsz) + metrics.emit_mean("nft_t_mean", total=t.sum(), count=bsz) + metrics.emit_mean( + "nft_num_timesteps", + total=torch.tensor(float(num_timesteps), device=ctx.device, dtype=torch.float32), + count=1, + ) + metrics.emit_mean("adv_abs_mean", total=prepared.advantage.abs().sum(), count=bsz) + + return loss_sum + + +# Actor: always run a reference DiT forward for NFT (EMA preferred). +nft_loss_formula.ref_mode = "ema" +# Same-sample K pairs must stay in one optimizer window. +nft_loss_formula.requires_sample_aligned_windows = True diff --git a/miles/backends/fsdp_utils/metrics.py b/miles/backends/fsdp_utils/metrics.py index 080c1b55..66af7ffe 100644 --- a/miles/backends/fsdp_utils/metrics.py +++ b/miles/backends/fsdp_utils/metrics.py @@ -23,6 +23,14 @@ "model_output_max_abs_diff": MetricReduce.MAX, "model_output_rel_max": MetricReduce.MAX, "grad_norm": MetricReduce.REPLICATED, + "nft_loss": MetricReduce.MEAN, + "nft_loss_per_pair": MetricReduce.MEAN, + "nft_r_mean": MetricReduce.MEAN, + "nft_pos_loss": MetricReduce.MEAN, + "nft_neg_loss": MetricReduce.MEAN, + "nft_adv_mean": MetricReduce.MEAN, + "nft_t_mean": MetricReduce.MEAN, + "nft_num_timesteps": MetricReduce.MEAN, } diff --git a/miles/ray/rollout.py b/miles/ray/rollout.py index 967431ee..2590c46e 100644 --- a/miles/ray/rollout.py +++ b/miles/ray/rollout.py @@ -363,9 +363,6 @@ def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sampl """ Convert inference generated samples to training data. """ - if self.custom_convert_samples_to_train_data_func is not None: - return self.custom_convert_samples_to_train_data_func(self.args, samples) - raw_rewards, rewards = self._post_process_rewards(samples) assert len(raw_rewards) == len(samples) @@ -375,6 +372,7 @@ def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sampl norm_t = torch.tensor(rewards, dtype=torch.float) # Emit reward distribution stats (raw + normalized) to stdout + wandb. + # Runs for both default SDE-pair expand and custom converts (e.g. NFT). reward_stats = { **_reward_stats_dict(raw_t, "rollout/reward/raw_"), **_reward_stats_dict(norm_t, "rollout/reward/norm_"), @@ -406,6 +404,9 @@ def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sampl reward_key=self.args.reward_key, ) + if self.custom_convert_samples_to_train_data_func is not None: + return self.custom_convert_samples_to_train_data_func(self.args, samples) + return self.train_data_converter.convert_samples(samples, rewards, raw_rewards) def _log_images( diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index e3a583d1..76e5667c 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -859,13 +859,83 @@ def add_algo_arguments(parser): parser.add_argument( "--loss-type", type=str, - choices=["policy_loss", "sft_loss", "custom_loss"], + choices=["policy_loss", "sft_loss", "custom_loss", "nft", "diffusion_nft"], default="policy_loss", help=( - "Choose loss type, currently support ppo policy_loss or sft_loss, " - "if custom_loss is set, we will use the function path from `--custom-loss-function-path`." + "Train objective shortcut. Diffusion default policy_loss = Flow-GRPO prepare+formula. " + "nft auto-fills convert / prepare / loss-formula custom paths. " + "DiT forward always stays in the FSDP actor." ), ) + parser.add_argument( + "--custom-prepare-train-batch-path", + type=str, + default=None, + help=( + "Dotted path to prepare hook " + "`fn(ctx, batch, *, pad_to_len=None) -> PreparedBatch`. " + "Builds DiT inputs; default is Flow-GRPO SDE-pair stacking." + ), + ) + parser.add_argument( + "--custom-loss-function-path", + type=str, + default=None, + help=( + "Dotted path to a loss *formula* only: " + "`fn(ctx, batch, prepared, *, new_pred, ref_pred, metrics, ...) -> Tensor`. " + "DiT forward is owned by the actor; this hook only computes the objective." + ), + ) + parser.add_argument( + "--diffusion-nft-beta", + type=float, + default=1.0, + help="DiffusionNFT dual-prediction blend coefficient (UniRL beta).", + ) + parser.add_argument( + "--diffusion-nft-adv-clip-max", + type=float, + default=5.0, + help="DiffusionNFT advantage clip before remap to r in [0, 1].", + ) + parser.add_argument( + "--diffusion-nft-adaptive-weight", + action="store_true", + default=True, + help="DiffusionNFT adaptive MSE weight (default on; UniRL use_adaptive_weight).", + ) + parser.add_argument( + "--no-diffusion-nft-adaptive-weight", + action="store_false", + dest="diffusion_nft_adaptive_weight", + help="Disable DiffusionNFT adaptive MSE weight.", + ) + parser.add_argument( + "--diffusion-nft-timestep-fraction", + type=float, + default=0.99, + help="Fraction of rollout schedule sigmas kept for NFT loss (drop terminal 0 first).", + ) + parser.add_argument( + "--diffusion-nft-shuffle-timesteps", + action="store_true", + default=True, + help="Shuffle NFT training timesteps per sample at convert time (default on).", + ) + parser.add_argument( + "--no-diffusion-nft-shuffle-timesteps", + action="store_false", + dest="diffusion_nft_shuffle_timesteps", + help="Disable NFT timestep shuffle.", + ) + parser.add_argument( + "--diffusion-nft-ref-mode", + type=str, + choices=["ema", "base"], + default="ema", + help="NFT reference policy: EMA shadow of the LoRA weights ('ema') or LoRA-base ('base').", + ) parser.add_argument( "--advantage-estimator", type=str, @@ -1068,6 +1138,50 @@ def add_debug_arguments(parser): "(requires matching sglang-d LoRAPipeline support)." ), ) + parser.add_argument( + "--lora-ema-shadow", + action="store_true", + default=False, + help=( + "Maintain an EMA shadow of trainable LoRA weights (pi_old). Custom losses " + "read it via ``ctx.ema_shadow.swap_in()``; combine with " + "--lora-ema-rollout-policy ema to sample under pi_old." + ), + ) + parser.add_argument( + "--lora-ema-rollout-policy", + type=str, + choices=["live", "ema"], + default="live", + help=( + "Which LoRA weights to push to rollout after each rollout_end when " + "--lora-ema-shadow is set: live trainable weights, or EMA shadow (pi_old)." + ), + ) + parser.add_argument( + "--lora-ema-decay", + type=float, + default=0.001, + help="LoRA EMA decay while step <= flat_steps.", + ) + parser.add_argument( + "--lora-ema-uprate", + type=float, + default=0.001, + help="LoRA EMA warmup rate after flat_steps.", + ) + parser.add_argument( + "--lora-ema-uphold", + type=float, + default=0.5, + help="LoRA EMA warmup cap.", + ) + parser.add_argument( + "--lora-ema-flat-steps", + type=int, + default=0, + help="LoRA EMA flat steps before warmup begins.", + ) parser.add_argument( "--diffusion-init-lora-weight", type=str, @@ -1444,6 +1558,45 @@ def miles_validate_args(args): "set --diffusion-model (for per-model defaults) or --lora-target-modules." ) + ema_enabled = bool(getattr(args, "lora_ema_shadow", False)) + if ema_enabled and not args.use_lora: + raise ValueError("--lora-ema-shadow requires --use-lora") + + if getattr(args, "loss_type", None) in ("nft", "diffusion_nft"): + # DiffusionNFT: swap convert + prepare + loss formula; DiT forward stays in actor. + if getattr(args, "custom_loss_function_path", None) is None: + args.custom_loss_function_path = "miles.backends.fsdp_utils.loss_hub.nft.nft_loss_formula" + if getattr(args, "custom_prepare_train_batch_path", None) is None: + args.custom_prepare_train_batch_path = "miles.backends.fsdp_utils.loss_hub.nft.prepare_nft_batch" + if args.custom_convert_samples_to_train_data_path is None: + args.custom_convert_samples_to_train_data_path = ( + "miles.backends.fsdp_utils.loss_hub.nft.convert_samples_to_nft_train_data" + ) + if ( + getattr(args, "diffusion_sde_type", "sde") == "sde" + and float(getattr(args, "diffusion_noise_level", 0.7) or 0.0) == 0.0 + ): + logger.info( + "NFT: auto-setting --diffusion-sde-type ode because noise_level=0 " + "(forward-process needs clean x0 only)." + ) + args.diffusion_sde_type = "ode" + if getattr(args, "diffusion_nft_ref_mode", "ema") == "ema" and not ema_enabled: + logger.warning( + "--loss-type nft with --diffusion-nft-ref-mode ema but without --lora-ema-shadow; " + "NFT loss will fall back to LoRA-base (disable_adapter) as the reference." + ) + beta = float(getattr(args, "diffusion_nft_beta", 1.0) or 1.0) + if beta <= 0: + raise ValueError(f"--diffusion-nft-beta must be > 0, got {beta}") + frac = float(getattr(args, "diffusion_nft_timestep_fraction", 0.99) or 0.99) + if not (0.0 < frac <= 1.0): + raise ValueError(f"--diffusion-nft-timestep-fraction must be in (0, 1], got {frac}") + if args.diffusion_recompute_old_log_prob: + raise ValueError( + "--diffusion-recompute-old-log-prob is only supported for policy_loss / Flow-GRPO, not NFT" + ) + if args.dump_details is not None: args.save_debug_rollout_data = f"{args.dump_details}/rollout_data/{{rollout_id}}.pt" args.save_debug_train_data = f"{args.dump_details}/train_data/{{rollout_id}}_{{rank}}.pt" @@ -1521,6 +1674,9 @@ def miles_validate_args(args): sde_step_backends = { "sde": "miles.backends.fsdp_utils.sde_step_backend.DiffusersSdeStepBackend", "cps": "miles.backends.fsdp_utils.sde_step_backend.CpsSdeStepBackend", + # ODE rollouts (DiffusionNFT) do not score logπ on the train path; Diffusers backend + # is a harmless placeholder so actor init can still construct sde_backend. + "ode": "miles.backends.fsdp_utils.sde_step_backend.DiffusersSdeStepBackend", } if args.diffusion_sde_type not in sde_step_backends: raise ValueError( diff --git a/miles/utils/train_data_utils.py b/miles/utils/train_data_utils.py index 9a663317..2da7e322 100644 --- a/miles/utils/train_data_utils.py +++ b/miles/utils/train_data_utils.py @@ -429,6 +429,40 @@ def reorder_train_pairs_for_tiling( return [train_data[i] for step in schedule for micro_batch in step for i in micro_batch] +def validate_sample_aligned_windows( + *, + train_pairs: list[dict[str, Any]], + microbatch_schedule: list[list[tuple[int, int]]], +) -> None: + """Ensure every optimizer window holds whole samples. + + Objectives whose gradient math sums one sample's pairs within a single + optimizer step (e.g. NFT's K timestep pairs) opt in via the loss-function + attribute ``requires_sample_aligned_windows``. + """ + group_sizes: OrderedDict[Any, int] = OrderedDict() + for pair in train_pairs: + key = pair["sample_index"] + group_sizes[key] = group_sizes.get(key, 0) + 1 + sizes = set(group_sizes.values()) + if len(sizes) > 1: + raise ValueError( + f"loss requires sample-aligned optimizer windows, but samples contribute unequal " + f"pair counts on this rank ({sorted(sizes)}); a sample was likely split across " + "DP ranks — make the sample count divisible by the DP size" + ) + for step_ranges in microbatch_schedule: + window_lo = step_ranges[0][0] + if window_lo == 0: + continue + if train_pairs[window_lo - 1]["sample_index"] == train_pairs[window_lo]["sample_index"]: + raise ValueError( + f"loss requires sample-aligned optimizer windows, but the window starting at " + f"pair {window_lo} splits sample_index={train_pairs[window_lo]['sample_index']}; " + "make samples-per-rank divisible by --num-steps-per-rollout" + ) + + def validate_same_microbatch_counts_across_train_ranks( *, microbatch_schedule: list[list[tuple[int, int]]], diff --git a/scripts/run-diffusion-nft-sd3-pickscore.sh b/scripts/run-diffusion-nft-sd3-pickscore.sh new file mode 100755 index 00000000..7dfdd7cb --- /dev/null +++ b/scripts/run-diffusion-nft-sd3-pickscore.sh @@ -0,0 +1,163 @@ +#!/usr/bin/env bash +# SD3.5 medium + DiffusionNFT + PickScore (UniRL sd3_nft / sd3_nft_100roll parity). +# +# UniRL reference (logs/sd3_nft_100roll.log, 1-GPU trainside smoke override): +# 8 prompts × 8 samples, lr=3e-4, wd=1e-4, LoRA r=32 α=64, +# guidance=1.0 (CFG-free), eta=0 / num_sde_steps=0 (ODE → clean x0), +# NFT beta=1, adv_clip=5, adaptive weight, schedule fraction=0.99, +# LoRA EMA shadow (pi_old) updated + synced at rollout_end (uprate=0.001, uphold=0.5). +# Expected early metrics: reward ~0.73–0.82, train/loss ~280–360. +# +# GPU layout (default CUDA_VISIBLE_DEVICES=4,5,2): +# first 2 = FSDP train + sglang colocate; 3rd = PickScore reward worker. +# +# Usage: +# NUM_ROLLOUT=5 CUDA_VISIBLE_DEVICES=4,5,2 \ +# bash scripts/run-diffusion-nft-sd3-pickscore.sh +# +# Smoke (tiny batch, OCR, 2 GPU — set SMOKE=1): +# SMOKE=1 CUDA_VISIBLE_DEVICES=4,5 bash scripts/run-diffusion-nft-sd3-pickscore.sh + +MILES_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" + +set -euo pipefail + +ROOT_DIR="${MILES_ROOT}" +export CUDA_VISIBLE_DEVICES="${CUDA_VISIBLE_DEVICES:-4,5,2}" +export PYTORCH_CUDA_ALLOC_CONF="${PYTORCH_CUDA_ALLOC_CONF:-expandable_segments:True}" +export HF_TOKEN="${HF_TOKEN:-}" +unset RAY_ADDRESS HF_HUB_OFFLINE TRANSFORMERS_OFFLINE 2>/dev/null || true +rm -f /tmp/ray/session_latest + +SD3_MODEL="${SD3_MODEL:-stabilityai/stable-diffusion-3.5-medium}" +SMOKE="${SMOKE:-0}" +if [[ "${SMOKE}" == "1" ]]; then + NUM_ROLLOUT="${NUM_ROLLOUT:-1}" +else + NUM_ROLLOUT="${NUM_ROLLOUT:-100}" +fi +RUN_NAME="diffusion_nft_sd3_pickscore_$(date +%Y%m%d_%H%M%S)" +SAVE_DIR="${ROOT_DIR}/logs/${RUN_NAME}/ckpt" +mkdir -p "${SAVE_DIR}" + +WANDB_ARGS=() +if [[ -n "${WANDB_API_KEY:-}" ]]; then + WANDB_ARGS+=( + --use-wandb + --wandb-project miles-diffusion-nft + --wandb-group "${RUN_NAME}" + --wandb-key "${WANDB_API_KEY}" + --diffusion-log-images 8 + --diffusion-log-image-interval 10 + --disable-wandb-random-suffix + ) +fi + +DATASETS_DIR="${DATASETS_DIR:-/root/datasets/miles-diffusion-datasets}" +if [[ "${SMOKE}" == "1" ]]; then + hf download --repo-type dataset rockdu/miles-diffusion-datasets \ + --include "flowgrpo_ocr/**" \ + --local-dir "${DATASETS_DIR}" + PROMPT_DATA="${DATASETS_DIR}/flowgrpo_ocr/train.jsonl" + REWARD_ARGS=( + --diffusion-reward ocr:1.0 + --rm-type ocr + ) + # 2-GPU smoke: tiny batch, no dedicated reward GPU. + export CUDA_VISIBLE_DEVICES="${CUDA_VISIBLE_DEVICES:-4,5}" + BATCH_ARGS=( + --rollout-batch-size 2 + --n-samples-per-prompt 2 + --num-rollout "${NUM_ROLLOUT}" + --num-steps-per-rollout 1 + --micro-batch-size 2 + --diffusion-microgroup-size 2 + --actor-num-gpus-per-node 2 + --rollout-num-gpus 2 + --rollout-num-gpus-per-engine 1 + --num-gpus-per-node 2 + ) +else + hf download --repo-type dataset rockdu/miles-diffusion-datasets \ + --include "flowgrpo_pickscore/**" \ + --local-dir "${DATASETS_DIR}" + PROMPT_DATA="${DATASETS_DIR}/flowgrpo_pickscore/train.jsonl" + REWARD_ARGS=( + --diffusion-reward pickscore:1.0 + --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 + ) + # Match UniRL 100-roll override: 8×8 prompts/samples, micro=4. + BATCH_ARGS=( + --rollout-batch-size 8 + --n-samples-per-prompt 8 + --num-rollout "${NUM_ROLLOUT}" + --num-steps-per-rollout 1 + --micro-batch-size 4 + --diffusion-microgroup-size 8 + --actor-num-gpus-per-node 2 + --rollout-num-gpus 2 + --rollout-num-gpus-per-engine 1 + --num-gpus-per-node 3 + --eval-prompt-data pickscore_test "${DATASETS_DIR}/flowgrpo_pickscore/test.jsonl" + --eval-interval 30 + ) +fi + +echo "RUN=${RUN_NAME}" | tee "${ROOT_DIR}/logs/${RUN_NAME}.log" + +python -u "${ROOT_DIR}/train_diffusion.py" \ + --train-backend fsdp \ + --rollout-function-path miles.rollout.sglang_diffusion_rollout.generate_rollout \ + --hf-checkpoint "${SD3_MODEL}" \ + --prompt-data "${PROMPT_DATA}" \ + --input-key input \ + "${BATCH_ARGS[@]}" \ + --gradient-checkpointing \ + --colocate \ + --use-miles-router \ + --sglang-server-concurrency 8 \ + --use-lora \ + --lora-ipc-weight-sync \ + --lora-rank 32 \ + --lora-alpha 64 \ + --diffusion-init-lora-weight gaussian \ + --lr 3e-4 \ + --adam-beta2 0.999 \ + --weight-decay 1e-4 \ + --clip-grad 1.0 \ + --loss-type nft \ + --diffusion-nft-beta 1.0 \ + --diffusion-nft-adv-clip-max 5.0 \ + --diffusion-nft-timestep-fraction 0.99 \ + --diffusion-nft-ref-mode ema \ + --lora-ema-shadow \ + --lora-ema-rollout-policy ema \ + --lora-ema-decay 0.001 \ + --lora-ema-uprate 0.001 \ + --lora-ema-uphold 0.5 \ + --lora-ema-flat-steps 0 \ + --advantage-estimator grpo \ + --globalize-reward-std \ + --diffusion-model "${SD3_MODEL}" \ + "${REWARD_ARGS[@]}" \ + --diffusion-forward-dtype fp16 \ + --sglang-dit-precision fp16 \ + --sglang-vae-slicing \ + --diffusion-num-steps 10 \ + --diffusion-eval-num-steps 50 \ + --update-weight-buffer-size 2147483648 \ + --diffusion-guidance-scale 1.0 \ + --diffusion-noise-level 0.0 \ + --diffusion-sde-type ode \ + --diffusion-height 512 \ + --diffusion-width 512 \ + --save "${SAVE_DIR}" \ + --save-interval 20 \ + --skip-eval-before-train \ + "${WANDB_ARGS[@]}" \ + 2>&1 | tee -a "${ROOT_DIR}/logs/${RUN_NAME}.log" diff --git a/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py b/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py new file mode 100644 index 00000000..5aed343d --- /dev/null +++ b/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py @@ -0,0 +1,159 @@ +"""Smoke tests for DiffusionNFT hooks (prepare + loss formula; actor owns DiT).""" + +from tests.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=30, suite="stage-a-cpu", labels=[]) + +from argparse import Namespace + +import pytest +import torch + +from miles.backends.fsdp_utils.loss_hub.losses import resolve_loss_formula_fn, resolve_prepare_fn +from miles.backends.fsdp_utils.loss_hub.nft import ( + convert_samples_to_nft_train_data, + corrupt, + nft_loss_formula, + nft_r_from_advantages, + prepare_nft_batch, + resolve_nft_sigmas, +) +from miles.backends.fsdp_utils.lora_ema import ( + LoraEmaShadow, + lora_ema_rollout_policy, + lora_ema_shadow_enabled, + resolve_lora_ema_kwargs, +) +from miles.backends.fsdp_utils.metrics import new_metric_buffer +from miles.utils.types import Sample + + +def _args(**overrides): + base = dict( + n_samples_per_prompt=2, + globalize_reward_mean=False, + globalize_reward_std=False, + grpo_std_normalization=True, + reward_key=None, + diffusion_nft_timestep_fraction=1.0, + diffusion_nft_shuffle_timesteps=False, + custom_prepare_train_batch_path=None, + custom_loss_function_path=None, + ) + base.update(overrides) + return Namespace(**base) + + +class TestNftMath: + def test_nft_r_remap(self): + r = nft_r_from_advantages(torch.tensor([-5.0, 0.0, 5.0]), adv_clip_max=5.0) + assert torch.allclose(r, torch.tensor([0.0, 0.5, 1.0])) + + def test_corrupt_linear_flow(self): + x0 = torch.ones(2, 4) + eps = torch.zeros(2, 4) + t = torch.tensor([0.25, 0.75]) + xt = corrupt(x0, t, eps) + assert torch.allclose(xt[0], torch.full((4,), 0.75)) + assert torch.allclose(xt[1], torch.full((4,), 0.25)) + + def test_resolve_sigmas_drops_zero_and_fraction(self): + class _Sched: + sigmas = torch.tensor([1.0, 0.8, 0.6, 0.4, 0.2, 0.0]) + + ts = resolve_nft_sigmas(_Sched(), training_timestep_fraction=0.99) + assert torch.allclose(ts, torch.tensor([1.0, 0.8, 0.6, 0.4])) + + +class TestNftHooks: + def test_convert_expands_k_timestep_pairs(self): + class _Traj: + def __init__(self): + self.timesteps = torch.tensor([999.0, 500.0, 0.0]) + self.sigmas = torch.tensor([1.0, 0.5, 0.0]) + self.latents = torch.zeros(3, 2, 2) + + class _Env: + pos_cond_kwargs = {} + neg_cond_kwargs = None + + samples = [ + Sample(index=0, prompt="a", reward=1.0, dit_trajectory=_Traj(), denoising_env=_Env()), + Sample(index=1, prompt="b", reward=3.0, dit_trajectory=_Traj(), denoising_env=_Env()), + ] + out = convert_samples_to_nft_train_data(_args(), samples) + assert len(out["train_data"]) == 4 + assert {p["timestep"] for p in out["train_data"]} == {1.0, 0.5} + assert out["train_data"][0]["x0"] is out["train_data"][1]["x0"] + + def test_formula_write_old_log_prob_is_noop(self): + metrics = new_metric_buffer(None, torch.device("cpu"), ()) + loss = nft_loss_formula( + None, # type: ignore[arg-type] + [], + None, # type: ignore[arg-type] + new_pred=None, # type: ignore[arg-type] + ref_pred=None, + metrics=metrics, + write_old_log_prob=True, + ) + assert loss is None + + def test_formula_declares_ref_and_window_attrs(self): + assert getattr(nft_loss_formula, "ref_mode", None) == "ema" + assert getattr(nft_loss_formula, "requires_sample_aligned_windows", False) is True + + def test_resolve_defaults_are_flow_grpo(self): + from miles.backends.fsdp_utils.loss_hub.losses import flow_grpo_loss_formula, prepare_flow_grpo_batch + + assert resolve_prepare_fn(_args()) is prepare_flow_grpo_batch + assert resolve_loss_formula_fn(_args()) is flow_grpo_loss_formula + + def test_resolve_custom_paths(self): + args = _args( + custom_prepare_train_batch_path="miles.backends.fsdp_utils.loss_hub.nft.prepare_nft_batch", + custom_loss_function_path="miles.backends.fsdp_utils.loss_hub.nft.nft_loss_formula", + ) + assert resolve_prepare_fn(args) is prepare_nft_batch + assert resolve_loss_formula_fn(args) is nft_loss_formula + + +class TestLoraEmaShadow: + def _model(self): + return torch.nn.Linear(4, 4, bias=False) + + def test_snapshot_and_update(self): + m = self._model() + ema = LoraEmaShadow(m.parameters(), decay=0.5, uprate=0.001, uphold=0.5, flat_steps=10) + init = m.weight.detach().clone() + with torch.no_grad(): + m.weight.add_(1.0) + delta = ema.update() + assert delta == 0.5 + assert torch.allclose(ema.shadow[0], init + 0.5) + + def test_swap_in_restores_exactly(self): + m = self._model() + ema = LoraEmaShadow(m.parameters(), decay=0.1) + live = m.weight.detach().clone() + with torch.no_grad(): + m.weight.add_(2.0) + with ema.swap_in(): + assert torch.equal(m.weight.detach(), live) + assert torch.equal(m.weight.detach(), live + 2.0) + + +class TestLoraEmaArgs: + def test_shadow_enabled(self): + assert lora_ema_shadow_enabled(Namespace(lora_ema_shadow=True)) + assert not lora_ema_shadow_enabled(Namespace()) + + def test_rollout_policy(self): + assert lora_ema_rollout_policy(Namespace(lora_ema_rollout_policy="ema")) == "ema" + assert lora_ema_rollout_policy(Namespace()) == "live" + + def test_resolve_kwargs(self): + kwargs = resolve_lora_ema_kwargs( + Namespace(lora_ema_decay=0.01, lora_ema_uprate=0.02, lora_ema_uphold=0.3, lora_ema_flat_steps=5) + ) + assert kwargs == {"decay": 0.01, "uprate": 0.02, "uphold": 0.3, "flat_steps": 5} From 012dc10f1aca4036bf5ee83160170d0b0ebd11ed Mon Sep 17 00:00:00 2001 From: niehen6174 Date: Mon, 27 Jul 2026 12:34:00 +0000 Subject: [PATCH 02/12] style: fix pre-commit ruff/isort/black on loss_hub hooks --- miles/backends/fsdp_utils/actor.py | 3 ++- miles/backends/fsdp_utils/lora_ema.py | 6 +++--- 2 files changed, 5 insertions(+), 4 deletions(-) diff --git a/miles/backends/fsdp_utils/actor.py b/miles/backends/fsdp_utils/actor.py index d3ea0657..1670b2b1 100644 --- a/miles/backends/fsdp_utils/actor.py +++ b/miles/backends/fsdp_utils/actor.py @@ -27,14 +27,15 @@ validate_same_microbatch_counts_across_train_ranks, validate_sample_aligned_windows, ) + from . import checkpoint from .diffusion_update_weight_utils import ( DiffusionUpdateWeightFromTensor, DiffusionUpdateWeightFromTensorLoRA, DiffusionUpdateWeightFromTensorLoRAIPC, ) -from .loss_hub import DiffusionLossContext, resolve_loss_formula_fn, resolve_prepare_fn from .lora_ema import LoraEmaShadow, lora_ema_rollout_policy, lora_ema_shadow_enabled, resolve_lora_ema_kwargs +from .loss_hub import DiffusionLossContext, resolve_loss_formula_fn, resolve_prepare_fn from .lr_scheduler import get_lr_scheduler from .metrics import new_metric_buffer, record_rollout_train_abs_diff from .parallel import create_fsdp_parallel_state diff --git a/miles/backends/fsdp_utils/lora_ema.py b/miles/backends/fsdp_utils/lora_ema.py index c0b678cc..df15d0b8 100644 --- a/miles/backends/fsdp_utils/lora_ema.py +++ b/miles/backends/fsdp_utils/lora_ema.py @@ -33,8 +33,8 @@ from __future__ import annotations from argparse import Namespace +from collections.abc import Iterable from contextlib import contextmanager -from typing import Iterable import torch import torch.nn as nn @@ -104,7 +104,7 @@ def update(self) -> float: raise RuntimeError("LoraEmaShadow.update called while swapped in") self.step += 1 delta = self.decay_at(self.step) - for live, sh in zip(self.params, self.shadow): + 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 @@ -121,7 +121,7 @@ def swap_in(self): @torch.no_grad() def _swap(self) -> None: - for live, sh in zip(self.params, self.shadow): + for live, sh in zip(self.params, self.shadow, strict=True): live_local = _local(live.data) tmp = live_local.clone() live_local.copy_(sh) From cd073be77b1e2b2e7d2b38eebd3f9ee34da44a54 Mon Sep 17 00:00:00 2001 From: niehen6174 Date: Tue, 28 Jul 2026 07:19:38 +0000 Subject: [PATCH 03/12] style: pre-commit fixes after rebase onto main --- miles/backends/fsdp_utils/actor.py | 2 +- miles/backends/fsdp_utils/loss_hub/losses.py | 2 +- tests/fast/backends/fsdp_utils/test_loss_hub_nft.py | 13 ++++++------- 3 files changed, 8 insertions(+), 9 deletions(-) diff --git a/miles/backends/fsdp_utils/actor.py b/miles/backends/fsdp_utils/actor.py index 1670b2b1..a52f519c 100644 --- a/miles/backends/fsdp_utils/actor.py +++ b/miles/backends/fsdp_utils/actor.py @@ -37,7 +37,7 @@ from .lora_ema import LoraEmaShadow, lora_ema_rollout_policy, lora_ema_shadow_enabled, resolve_lora_ema_kwargs from .loss_hub import DiffusionLossContext, resolve_loss_formula_fn, resolve_prepare_fn from .lr_scheduler import get_lr_scheduler -from .metrics import new_metric_buffer, record_rollout_train_abs_diff +from .metrics import new_metric_buffer from .parallel import create_fsdp_parallel_state from .sequence_parallel.plan import apply_sequence_parallel diff --git a/miles/backends/fsdp_utils/loss_hub/losses.py b/miles/backends/fsdp_utils/loss_hub/losses.py index 0e6b20d9..efdb3c65 100644 --- a/miles/backends/fsdp_utils/loss_hub/losses.py +++ b/miles/backends/fsdp_utils/loss_hub/losses.py @@ -7,7 +7,7 @@ from __future__ import annotations from argparse import Namespace -from typing import Callable +from collections.abc import Callable import torch 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 5aed343d..315bf4bc 100644 --- a/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py +++ b/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py @@ -6,9 +6,14 @@ from argparse import Namespace -import pytest import torch +from miles.backends.fsdp_utils.lora_ema import ( + LoraEmaShadow, + lora_ema_rollout_policy, + lora_ema_shadow_enabled, + resolve_lora_ema_kwargs, +) from miles.backends.fsdp_utils.loss_hub.losses import resolve_loss_formula_fn, resolve_prepare_fn from miles.backends.fsdp_utils.loss_hub.nft import ( convert_samples_to_nft_train_data, @@ -18,12 +23,6 @@ prepare_nft_batch, resolve_nft_sigmas, ) -from miles.backends.fsdp_utils.lora_ema import ( - LoraEmaShadow, - lora_ema_rollout_policy, - lora_ema_shadow_enabled, - resolve_lora_ema_kwargs, -) from miles.backends.fsdp_utils.metrics import new_metric_buffer from miles.utils.types import Sample From 12c3078abad78092fc59acb892b9485b9015ec6a Mon Sep 17 00:00:00 2001 From: niehen6174 Date: Tue, 28 Jul 2026 13:05:33 +0000 Subject: [PATCH 04/12] refactor(loss_hub): extract prepare hooks into prepare.py Separate DiT forward input preparation from loss formulas so Flow-GRPO and NFT prepare paths live in one module while losses.py stays formula-only. --- .../backends/fsdp_utils/loss_hub/__init__.py | 9 +- miles/backends/fsdp_utils/loss_hub/losses.py | 130 +----------- miles/backends/fsdp_utils/loss_hub/nft.py | 71 +------ miles/backends/fsdp_utils/loss_hub/prepare.py | 189 ++++++++++++++++++ miles/utils/arguments.py | 2 +- .../backends/fsdp_utils/test_loss_hub_nft.py | 8 +- 6 files changed, 203 insertions(+), 206 deletions(-) create mode 100644 miles/backends/fsdp_utils/loss_hub/prepare.py diff --git a/miles/backends/fsdp_utils/loss_hub/__init__.py b/miles/backends/fsdp_utils/loss_hub/__init__.py index 668e2171..eaf1ef2c 100644 --- a/miles/backends/fsdp_utils/loss_hub/__init__.py +++ b/miles/backends/fsdp_utils/loss_hub/__init__.py @@ -8,12 +8,8 @@ from miles.backends.fsdp_utils.loss_hub.advantages import grpo_normalize_rewards from miles.backends.fsdp_utils.loss_hub.context import DiffusionLossContext, PreparedBatch -from miles.backends.fsdp_utils.loss_hub.losses import ( - flow_grpo_loss_formula, - prepare_flow_grpo_batch, - resolve_loss_formula_fn, - resolve_prepare_fn, -) +from miles.backends.fsdp_utils.loss_hub.losses import flow_grpo_loss_formula, resolve_loss_formula_fn +from miles.backends.fsdp_utils.loss_hub.prepare import prepare_flow_grpo_batch, prepare_nft_batch, resolve_prepare_fn __all__ = [ "DiffusionLossContext", @@ -21,6 +17,7 @@ "flow_grpo_loss_formula", "grpo_normalize_rewards", "prepare_flow_grpo_batch", + "prepare_nft_batch", "resolve_loss_formula_fn", "resolve_prepare_fn", ] diff --git a/miles/backends/fsdp_utils/loss_hub/losses.py b/miles/backends/fsdp_utils/loss_hub/losses.py index efdb3c65..50b38cc3 100644 --- a/miles/backends/fsdp_utils/loss_hub/losses.py +++ b/miles/backends/fsdp_utils/loss_hub/losses.py @@ -1,7 +1,7 @@ -"""Default Flow-GRPO prepare + loss formula (actor owns DiT forward). +"""Default Flow-GRPO loss formula (actor owns DiT forward). -Custom algorithms swap ``--custom-prepare-train-batch-path`` and/or -``--custom-loss-function-path`` (formula only: receives ``new_pred`` / ``ref_pred``). +Custom algorithms swap ``--custom-loss-function-path`` (formula only: receives +``new_pred`` / ``ref_pred``). Batch preparation lives in ``prepare.py``. """ from __future__ import annotations @@ -17,123 +17,9 @@ from miles.utils.misc import load_function from miles.utils.train_data_utils import stack_train_pair_rollout_debug -PrepareFn = Callable[..., PreparedBatch] LossFormulaFn = Callable[..., torch.Tensor | None] -def _cast_cond_to_dtype(cond: dict, dtype: torch.dtype) -> dict: - out = {} - for k, v in cond.items(): - if isinstance(v, torch.Tensor) and v.dtype.is_floating_point: - out[k] = v.to(dtype=dtype) - else: - out[k] = v - return out - - -def _stack_pair_field(batch: list[dict], key: str, device: torch.device) -> torch.Tensor: - return torch.stack([pair[key] for pair in batch]).to(device=device, dtype=torch.float32) - - -def prepare_flow_grpo_batch( - ctx: DiffusionLossContext, - batch: list[dict], - *, - pad_to_len: int | None = None, -) -> PreparedBatch: - """Stack SDE-pair fields and build CFG conditioning (guidance from args).""" - args = ctx.args - device = ctx.device - config = ctx.train_pipeline_config - num_train_timesteps = int(ctx.scheduler.config.num_train_timesteps) - bsz = len(batch) - - latents = _stack_pair_field(batch, "latent", device) - next_latents = _stack_pair_field(batch, "next_latent", device) - timesteps = _stack_pair_field(batch, "timestep", device) - next_timesteps = _stack_pair_field(batch, "next_timestep", device) - log_prob_old = _stack_pair_field(batch, "log_prob_old", device) - advantage = torch.tensor( - [float(pair["advantage"]) for pair in batch], - device=device, - dtype=torch.float32, - ) - advantage = torch.clamp(advantage, -args.diffusion_adv_clip_max, args.diffusion_adv_clip_max) - - guidance_scale = args.diffusion_guidance_scale - true_cfg_scale = args.diffusion_true_cfg_scale - cfg_scale = true_cfg_scale if true_cfg_scale is not None else guidance_scale - use_cfg = cfg_scale > 0 - - if len(ctx.models) == 1: - component_name, model = next(iter(ctx.models.items())) - else: - components = {config.component_for_timestep(t, num_train_timesteps) for t in timesteps.tolist()} - if len(components) > 1: - raise ValueError( - f"Micro-batch mixes denoising phases {sorted(components)}; set " - "--micro-batch-size 1 so each forward is phase-pure (one DiT, one CFG scale)." - ) - component_name = components.pop() - model = ctx.models[component_name] - guidance_scale = config.select_guidance_scale( - float(timesteps[0]), - num_train_timesteps, - guidance_scale, - args.diffusion_guidance_scale_2, - ) - - if config.needs_timestep_scaling: - timesteps_for_model = timesteps / float(num_train_timesteps) - else: - timesteps_for_model = timesteps - - pos_list = [config.prepare_cond_kwargs(batch[i]["denoising_env"].pos_cond_kwargs, device) for i in range(bsz)] - neg_list = ( - [config.prepare_cond_kwargs(batch[i]["denoising_env"].neg_cond_kwargs, device) for i in range(bsz)] - if use_cfg - else None - ) - cfg_batching = use_cfg and bool(args.fsdp_cfg_batching) - joint_cond = pos_cond = neg_cond = None - if cfg_batching: - joint_cond = _cast_cond_to_dtype( - config.collate_cond_for_sample_batch(pos_list + neg_list, device, pad_to_len=pad_to_len), - ctx.forward_dtype, - ) - else: - pos_cond = _cast_cond_to_dtype( - config.collate_cond_for_sample_batch(pos_list, device, pad_to_len=pad_to_len), - ctx.forward_dtype, - ) - if use_cfg and neg_list is not None: - neg_cond = _cast_cond_to_dtype( - config.collate_cond_for_sample_batch(neg_list, device, pad_to_len=pad_to_len), - ctx.forward_dtype, - ) - - return PreparedBatch( - latents=latents, - timesteps=timesteps, - timesteps_for_model=timesteps_for_model, - model=model, - component_name=component_name, - guidance_scale=guidance_scale, - use_cfg=use_cfg, - cfg_batching=cfg_batching, - true_cfg_scale=true_cfg_scale if use_cfg else None, - pos_cond=pos_cond, - neg_cond=neg_cond, - joint_cond=joint_cond, - advantage=advantage, - extras={ - "next_latents": next_latents, - "next_timesteps": next_timesteps, - "log_prob_old": log_prob_old, - }, - ) - - def flow_grpo_loss_formula( ctx: DiffusionLossContext, batch: list[dict], @@ -229,16 +115,6 @@ def flow_grpo_loss_formula( return loss_sum -def resolve_prepare_fn(args: Namespace) -> PrepareFn: - path = getattr(args, "custom_prepare_train_batch_path", None) - if path: - fn = load_function(path) - if fn is None: - raise ValueError(f"Failed to load custom prepare from {path!r}") - return fn - return prepare_flow_grpo_batch - - def resolve_loss_formula_fn(args: Namespace) -> LossFormulaFn: """Loss *formula* only — DiT forward stays in the actor.""" path = getattr(args, "custom_loss_function_path", None) diff --git a/miles/backends/fsdp_utils/loss_hub/nft.py b/miles/backends/fsdp_utils/loss_hub/nft.py index b201a12a..7987c7aa 100644 --- a/miles/backends/fsdp_utils/loss_hub/nft.py +++ b/miles/backends/fsdp_utils/loss_hub/nft.py @@ -1,7 +1,7 @@ -"""DiffusionNFT plugin: custom convert + prepare + loss *formula*. +"""DiffusionNFT plugin: custom convert + loss *formula*. -Actor still owns DiT forward (+ EMA/LoRA-base reference forward). This module -only swaps the replaceable parts via stock customization hooks. +Prepare hook lives in ``prepare.py`` (``prepare_nft_batch``). Actor still owns +DiT forward (+ EMA/LoRA-base reference forward). """ from __future__ import annotations @@ -17,7 +17,7 @@ from miles.utils.types import Sample # --------------------------------------------------------------------------- -# Forward-process math +# Forward-process math (used by prepare_nft_batch in prepare.py) # --------------------------------------------------------------------------- @@ -96,16 +96,6 @@ def nft_branch_losses( return pos_loss, neg_loss -def _cast_cond_to_dtype(cond: dict, dtype: torch.dtype) -> dict: - out = {} - for k, v in cond.items(): - if isinstance(v, torch.Tensor) and v.dtype.is_floating_point: - out[k] = v.to(dtype=dtype) - else: - out[k] = v - return out - - # --------------------------------------------------------------------------- # Convert (K-expanded pairs) # --------------------------------------------------------------------------- @@ -169,59 +159,6 @@ def convert_samples_to_nft_train_data(args: Namespace, samples: list[Sample]) -> return {"train_data": train_data, **scheduler_meta} -# --------------------------------------------------------------------------- -# Prepare (actor still runs DiT forward on the result) -# --------------------------------------------------------------------------- - - -def prepare_nft_batch( - ctx: DiffusionLossContext, - batch: list[dict], - *, - pad_to_len: int | None = None, -) -> PreparedBatch: - """Corrupt clean x0 at each pair's sigma; CFG-free cond.""" - if len(ctx.models) != 1: - raise ValueError("DiffusionNFT currently supports a single DiT component (SD3)") - device = ctx.device - config = ctx.train_pipeline_config - bsz = len(batch) - x0 = torch.stack([pair["x0"] for pair in batch]).to(device=device, dtype=torch.float32) - t = torch.tensor([float(pair["timestep"]) for pair in batch], device=device, dtype=torch.float32) - advantage = torch.tensor([float(pair["advantage"]) for pair in batch], device=device, dtype=torch.float32) - - component_name, model = next(iter(ctx.models.items())) - pos_list = [config.prepare_cond_kwargs(batch[i]["denoising_env"].pos_cond_kwargs, device) for i in range(bsz)] - pos_cond = _cast_cond_to_dtype( - config.collate_cond_for_sample_batch(pos_list, device, pad_to_len=pad_to_len), - ctx.forward_dtype, - ) - - num_train_timesteps = int(getattr(ctx.scheduler.config, "num_train_timesteps", 1000)) - if config.needs_timestep_scaling: - timesteps_for_model = t.to(dtype=torch.float32) - else: - timesteps_for_model = t * float(num_train_timesteps) - - xt = corrupt(x0, t, sample_noise(x0)) - return PreparedBatch( - latents=xt, - timesteps=t, - timesteps_for_model=timesteps_for_model, - model=model, - component_name=component_name, - guidance_scale=0.0, - use_cfg=False, - cfg_batching=False, - true_cfg_scale=None, - pos_cond=pos_cond, - neg_cond=None, - joint_cond=None, - advantage=advantage, - extras={"x0": x0}, - ) - - # --------------------------------------------------------------------------- # Loss formula (receives actor's new_pred / ref_pred) # --------------------------------------------------------------------------- diff --git a/miles/backends/fsdp_utils/loss_hub/prepare.py b/miles/backends/fsdp_utils/loss_hub/prepare.py new file mode 100644 index 00000000..f6c3fb33 --- /dev/null +++ b/miles/backends/fsdp_utils/loss_hub/prepare.py @@ -0,0 +1,189 @@ +"""DiT forward input preparation hooks (actor pipeline step before forward). + +Custom algorithms swap ``--custom-prepare-train-batch-path``; loss formulas live +in ``losses.py`` / algorithm plugins (e.g. ``nft.py``). +""" + +from __future__ import annotations + +from argparse import Namespace +from collections.abc import Callable + +import torch + +from miles.backends.fsdp_utils.loss_hub.context import DiffusionLossContext, PreparedBatch +from miles.backends.fsdp_utils.loss_hub.nft import corrupt, sample_noise +from miles.utils.misc import load_function + +PrepareFn = Callable[..., PreparedBatch] + + +def _cast_cond_to_dtype(cond: dict, dtype: torch.dtype) -> dict: + out = {} + for k, v in cond.items(): + if isinstance(v, torch.Tensor) and v.dtype.is_floating_point: + out[k] = v.to(dtype=dtype) + else: + out[k] = v + return out + + +def _stack_pair_field(batch: list[dict], key: str, device: torch.device) -> torch.Tensor: + return torch.stack([pair[key] for pair in batch]).to(device=device, dtype=torch.float32) + + +def prepare_flow_grpo_batch( + ctx: DiffusionLossContext, + batch: list[dict], + *, + pad_to_len: int | None = None, +) -> PreparedBatch: + """Stack SDE-pair fields and build CFG conditioning (guidance from args).""" + args = ctx.args + device = ctx.device + config = ctx.train_pipeline_config + num_train_timesteps = int(ctx.scheduler.config.num_train_timesteps) + bsz = len(batch) + + latents = _stack_pair_field(batch, "latent", device) + next_latents = _stack_pair_field(batch, "next_latent", device) + timesteps = _stack_pair_field(batch, "timestep", device) + next_timesteps = _stack_pair_field(batch, "next_timestep", device) + log_prob_old = _stack_pair_field(batch, "log_prob_old", device) + advantage = torch.tensor( + [float(pair["advantage"]) for pair in batch], + device=device, + dtype=torch.float32, + ) + advantage = torch.clamp(advantage, -args.diffusion_adv_clip_max, args.diffusion_adv_clip_max) + + guidance_scale = args.diffusion_guidance_scale + true_cfg_scale = args.diffusion_true_cfg_scale + cfg_scale = true_cfg_scale if true_cfg_scale is not None else guidance_scale + use_cfg = cfg_scale > 0 + + if len(ctx.models) == 1: + component_name, model = next(iter(ctx.models.items())) + else: + components = {config.component_for_timestep(t, num_train_timesteps) for t in timesteps.tolist()} + if len(components) > 1: + raise ValueError( + f"Micro-batch mixes denoising phases {sorted(components)}; set " + "--micro-batch-size 1 so each forward is phase-pure (one DiT, one CFG scale)." + ) + component_name = components.pop() + model = ctx.models[component_name] + guidance_scale = config.select_guidance_scale( + float(timesteps[0]), + num_train_timesteps, + guidance_scale, + args.diffusion_guidance_scale_2, + ) + + if config.needs_timestep_scaling: + timesteps_for_model = timesteps / float(num_train_timesteps) + else: + timesteps_for_model = timesteps + + pos_list = [config.prepare_cond_kwargs(batch[i]["denoising_env"].pos_cond_kwargs, device) for i in range(bsz)] + neg_list = ( + [config.prepare_cond_kwargs(batch[i]["denoising_env"].neg_cond_kwargs, device) for i in range(bsz)] + if use_cfg + else None + ) + cfg_batching = use_cfg and bool(args.fsdp_cfg_batching) + joint_cond = pos_cond = neg_cond = None + if cfg_batching: + joint_cond = _cast_cond_to_dtype( + config.collate_cond_for_sample_batch(pos_list + neg_list, device, pad_to_len=pad_to_len), + ctx.forward_dtype, + ) + else: + pos_cond = _cast_cond_to_dtype( + config.collate_cond_for_sample_batch(pos_list, device, pad_to_len=pad_to_len), + ctx.forward_dtype, + ) + if use_cfg and neg_list is not None: + neg_cond = _cast_cond_to_dtype( + config.collate_cond_for_sample_batch(neg_list, device, pad_to_len=pad_to_len), + ctx.forward_dtype, + ) + + return PreparedBatch( + latents=latents, + timesteps=timesteps, + timesteps_for_model=timesteps_for_model, + model=model, + component_name=component_name, + guidance_scale=guidance_scale, + use_cfg=use_cfg, + cfg_batching=cfg_batching, + true_cfg_scale=true_cfg_scale if use_cfg else None, + pos_cond=pos_cond, + neg_cond=neg_cond, + joint_cond=joint_cond, + advantage=advantage, + extras={ + "next_latents": next_latents, + "next_timesteps": next_timesteps, + "log_prob_old": log_prob_old, + }, + ) + + +def prepare_nft_batch( + ctx: DiffusionLossContext, + batch: list[dict], + *, + pad_to_len: int | None = None, +) -> PreparedBatch: + """Corrupt clean x0 at each pair's sigma; CFG-free cond.""" + if len(ctx.models) != 1: + raise ValueError("DiffusionNFT currently supports a single DiT component (SD3)") + device = ctx.device + config = ctx.train_pipeline_config + bsz = len(batch) + x0 = torch.stack([pair["x0"] for pair in batch]).to(device=device, dtype=torch.float32) + t = torch.tensor([float(pair["timestep"]) for pair in batch], device=device, dtype=torch.float32) + advantage = torch.tensor([float(pair["advantage"]) for pair in batch], device=device, dtype=torch.float32) + + component_name, model = next(iter(ctx.models.items())) + pos_list = [config.prepare_cond_kwargs(batch[i]["denoising_env"].pos_cond_kwargs, device) for i in range(bsz)] + pos_cond = _cast_cond_to_dtype( + config.collate_cond_for_sample_batch(pos_list, device, pad_to_len=pad_to_len), + ctx.forward_dtype, + ) + + num_train_timesteps = int(getattr(ctx.scheduler.config, "num_train_timesteps", 1000)) + if config.needs_timestep_scaling: + timesteps_for_model = t.to(dtype=torch.float32) + else: + timesteps_for_model = t * float(num_train_timesteps) + + xt = corrupt(x0, t, sample_noise(x0)) + return PreparedBatch( + latents=xt, + timesteps=t, + timesteps_for_model=timesteps_for_model, + model=model, + component_name=component_name, + guidance_scale=0.0, + use_cfg=False, + cfg_batching=False, + true_cfg_scale=None, + pos_cond=pos_cond, + neg_cond=None, + joint_cond=None, + advantage=advantage, + extras={"x0": x0}, + ) + + +def resolve_prepare_fn(args: Namespace) -> PrepareFn: + path = getattr(args, "custom_prepare_train_batch_path", None) + if path: + fn = load_function(path) + if fn is None: + raise ValueError(f"Failed to load custom prepare from {path!r}") + return fn + return prepare_flow_grpo_batch diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index 76e5667c..31881cf1 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -1567,7 +1567,7 @@ def miles_validate_args(args): if getattr(args, "custom_loss_function_path", None) is None: args.custom_loss_function_path = "miles.backends.fsdp_utils.loss_hub.nft.nft_loss_formula" if getattr(args, "custom_prepare_train_batch_path", None) is None: - args.custom_prepare_train_batch_path = "miles.backends.fsdp_utils.loss_hub.nft.prepare_nft_batch" + args.custom_prepare_train_batch_path = "miles.backends.fsdp_utils.loss_hub.prepare.prepare_nft_batch" if args.custom_convert_samples_to_train_data_path is None: args.custom_convert_samples_to_train_data_path = ( "miles.backends.fsdp_utils.loss_hub.nft.convert_samples_to_nft_train_data" 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 315bf4bc..3334c63f 100644 --- a/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py +++ b/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py @@ -14,15 +14,15 @@ lora_ema_shadow_enabled, resolve_lora_ema_kwargs, ) -from miles.backends.fsdp_utils.loss_hub.losses import resolve_loss_formula_fn, resolve_prepare_fn +from miles.backends.fsdp_utils.loss_hub.losses import flow_grpo_loss_formula, resolve_loss_formula_fn from miles.backends.fsdp_utils.loss_hub.nft import ( convert_samples_to_nft_train_data, corrupt, nft_loss_formula, nft_r_from_advantages, - prepare_nft_batch, resolve_nft_sigmas, ) +from miles.backends.fsdp_utils.loss_hub.prepare import prepare_flow_grpo_batch, prepare_nft_batch, resolve_prepare_fn from miles.backends.fsdp_utils.metrics import new_metric_buffer from miles.utils.types import Sample @@ -103,14 +103,12 @@ def test_formula_declares_ref_and_window_attrs(self): assert getattr(nft_loss_formula, "requires_sample_aligned_windows", False) is True def test_resolve_defaults_are_flow_grpo(self): - from miles.backends.fsdp_utils.loss_hub.losses import flow_grpo_loss_formula, prepare_flow_grpo_batch - assert resolve_prepare_fn(_args()) is prepare_flow_grpo_batch assert resolve_loss_formula_fn(_args()) is flow_grpo_loss_formula def test_resolve_custom_paths(self): args = _args( - custom_prepare_train_batch_path="miles.backends.fsdp_utils.loss_hub.nft.prepare_nft_batch", + custom_prepare_train_batch_path="miles.backends.fsdp_utils.loss_hub.prepare.prepare_nft_batch", custom_loss_function_path="miles.backends.fsdp_utils.loss_hub.nft.nft_loss_formula", ) assert resolve_prepare_fn(args) is prepare_nft_batch From 0c748a398e8751e40c671a7bb9da701fc9a53a14 Mon Sep 17 00:00:00 2001 From: niehen6174 Date: Wed, 29 Jul 2026 12:18:25 +0000 Subject: [PATCH 05/12] refactor(diffusion): split NFT convert/reward hooks and expose --ref-mode Keep custom_convert as a full override, route NFT through a dedicated converter after post_process_rewards, and resolve reference selection in arguments.py so the actor only consumes args.ref_mode. --- miles/backends/fsdp_utils/actor.py | 49 ++----- miles/backends/fsdp_utils/lora_ema.py | 17 +-- .../backends/fsdp_utils/loss_hub/__init__.py | 5 +- .../fsdp_utils/loss_hub/advantages.py | 7 +- miles/backends/fsdp_utils/loss_hub/losses.py | 6 +- miles/backends/fsdp_utils/loss_hub/nft.py | 123 ++++++++++-------- miles/ray/rollout.py | 40 ++---- miles/utils/arguments.py | 68 +++++++--- miles/utils/train_data_utils.py | 9 ++ scripts/run-diffusion-nft-sd3-pickscore.sh | 2 +- .../backends/fsdp_utils/test_loss_hub_nft.py | 29 ++--- tests/fast/utils/test_ref_mode_and_convert.py | 116 +++++++++++++++++ 12 files changed, 307 insertions(+), 164 deletions(-) create mode 100644 tests/fast/utils/test_ref_mode_and_convert.py diff --git a/miles/backends/fsdp_utils/actor.py b/miles/backends/fsdp_utils/actor.py index a52f519c..060cb8b4 100644 --- a/miles/backends/fsdp_utils/actor.py +++ b/miles/backends/fsdp_utils/actor.py @@ -34,7 +34,7 @@ DiffusionUpdateWeightFromTensorLoRA, DiffusionUpdateWeightFromTensorLoRAIPC, ) -from .lora_ema import LoraEmaShadow, lora_ema_rollout_policy, lora_ema_shadow_enabled, resolve_lora_ema_kwargs +from .lora_ema import LoraEmaShadow, resolve_lora_ema_kwargs from .loss_hub import DiffusionLossContext, resolve_loss_formula_fn, resolve_prepare_fn from .lr_scheduler import get_lr_scheduler from .metrics import new_metric_buffer @@ -209,10 +209,10 @@ def init(self, args: Namespace, role: str, with_ref: bool = False) -> int: # ty checkpoint_payload = checkpoint.load(self) - # Optional LoRA EMA shadow (pi_old). Loss formulas that set ref_mode="ema" - # consume it via actor reference forward. + # Optional LoRA EMA shadow (pi_old). Enabled via --lora-ema-shadow (parsed in arguments.py). + # Consumed when --ref-mode ema runs the no-grad reference DiT forward. self.ema_shadow = None - if lora_ema_shadow_enabled(self.args): + if getattr(self.args, "lora_ema_shadow", False): self.ema_shadow = LoraEmaShadow( (p for m in self.models.values() for p in m.parameters()), **resolve_lora_ema_kwargs(self.args), @@ -308,7 +308,7 @@ def update_weights(self) -> None: # type: ignore[override] delta = self.ema_shadow.update() if dist.get_rank() == 0: logger.info("LoRA EMA shadow updated (decay=%.4f step=%d)", delta, self.ema_shadow.step) - if lora_ema_rollout_policy(self.args) == "ema": + if getattr(self.args, "lora_ema_rollout_policy", "live") == "ema": with self.ema_shadow.swap_in(): self.weight_updater.update_weights() else: @@ -367,14 +367,14 @@ def _train_core(self, rollout_id: int, rollout_data) -> None: num_pairs = len(train_pairs) num_train_timesteps = self.scheduler.config.num_train_timesteps - # ------------- KL precondition (Flow-GRPO formula may request ref forward) ------------- - kl_beta = float(self.args.diffusion_kl_beta) - if kl_beta > 0 and not self.args.use_lora: - raise ValueError( - "--diffusion-kl-beta currently requires --use-lora so the base model can be used as reference." + # ------------- Reference forward precondition (--ref-mode resolved in arguments.py) ------------- + ref_mode = getattr(self.args, "ref_mode", "none") + if ref_mode == "lora_base" and not all(hasattr(m, "disable_adapter") for m in self.models.values()): + raise RuntimeError( + "--ref-mode lora_base requires PEFT models exposing disable_adapter() after FSDP wrapping." ) - if kl_beta > 0 and not all(hasattr(m, "disable_adapter") for m in self.models.values()): - raise RuntimeError("Diffusion KL requires PEFT models exposing disable_adapter() after FSDP wrapping.") + if ref_mode == "ema" and self.ema_shadow is None: + raise RuntimeError("--ref-mode ema requires a constructed LoRA EMA shadow") # ------------- Rollout Scheduler Metadata ------------- scheduler_timesteps, scheduler_sigmas = scheduler_meta_from_rollout( @@ -502,23 +502,6 @@ def _maybe_legacy_window_pad_len(self, train_pairs: list, microbatch_ranges: lis conds.append(env.neg_cond_kwargs) return self.train_pipeline_config.maybe_legacy_window_pad_len(conds) - def _resolve_ref_mode(self) -> str | None: - """Whether the actor should run a no-grad reference DiT forward. - - Loss formulas may set ``ref_mode`` to ``"ema"`` / ``"lora_base"``. - Flow-GRPO defaults to ``lora_base`` when ``--diffusion-kl-beta > 0``. - """ - explicit = getattr(self.loss_formula_fn, "ref_mode", None) - if explicit == "ema" and getattr(self.args, "diffusion_nft_ref_mode", "ema") == "base": - explicit = "lora_base" - if explicit in ("ema", "lora_base"): - if explicit == "ema" and self.ema_shadow is None: - return "lora_base" - return explicit - if float(self.args.diffusion_kl_beta) > 0: - return "lora_base" - return None - def _forward_train_pair_batch( self, ctx: DiffusionLossContext, @@ -560,15 +543,11 @@ def _compute_noise_pred() -> torch.Tensor: new_pred = _compute_noise_pred() ref_pred = None - ref_mode = self._resolve_ref_mode() - if ref_mode is not None: + ref_mode = getattr(self.args, "ref_mode", "none") + if ref_mode not in (None, "none"): if ref_mode == "ema": - if self.ema_shadow is None: - raise ValueError("ref_mode=ema requires --lora-ema-shadow") ref_ctx = self.ema_shadow.swap_in() elif ref_mode == "lora_base": - if not hasattr(prepared.model, "disable_adapter"): - raise ValueError("ref_mode=lora_base requires a LoRA model exposing disable_adapter()") ref_ctx = prepared.model.disable_adapter() else: raise ValueError(f"unknown ref_mode {ref_mode!r}") diff --git a/miles/backends/fsdp_utils/lora_ema.py b/miles/backends/fsdp_utils/lora_ema.py index df15d0b8..463a9c63 100644 --- a/miles/backends/fsdp_utils/lora_ema.py +++ b/miles/backends/fsdp_utils/lora_ema.py @@ -45,18 +45,13 @@ def _local(t: torch.Tensor) -> torch.Tensor: return t._local_tensor if hasattr(t, "_local_tensor") else t -def lora_ema_shadow_enabled(args: Namespace) -> bool: - """True when a LoRA EMA shadow should be constructed.""" - return bool(getattr(args, "lora_ema_shadow", False)) - - -def lora_ema_rollout_policy(args: Namespace) -> str: - """Which LoRA weights to push to rollout engines after each rollout ('live' or 'ema').""" - return getattr(args, "lora_ema_rollout_policy", "live") - - def resolve_lora_ema_kwargs(args: Namespace) -> dict[str, float | int]: - """Read normalized ``lora_ema_*`` fields from ``args`` (see ``miles_validate_args``).""" + """Read normalized ``lora_ema_*`` fields from ``args`` (see ``miles_validate_args``). + + Enablement (``args.lora_ema_shadow``) and rollout policy + (``args.lora_ema_rollout_policy``) are plain args — inferred/validated in + ``arguments.py``, not re-wrapped here. + """ return { "decay": float(getattr(args, "lora_ema_decay", 0.001)), "uprate": float(getattr(args, "lora_ema_uprate", 0.001)), diff --git a/miles/backends/fsdp_utils/loss_hub/__init__.py b/miles/backends/fsdp_utils/loss_hub/__init__.py index eaf1ef2c..9cec641f 100644 --- a/miles/backends/fsdp_utils/loss_hub/__init__.py +++ b/miles/backends/fsdp_utils/loss_hub/__init__.py @@ -3,16 +3,19 @@ Swap pieces via: ``--custom-prepare-train-batch-path`` ``--custom-loss-function-path`` (formula only; receives new_pred / ref_pred) - ``--custom-convert-samples-to-train-data-path`` + ``--custom-reward-post-process-path`` (advantage normalisation) + ``--loss-type nft`` selects ``NftTrainDataConverter`` (not a full convert override) """ from miles.backends.fsdp_utils.loss_hub.advantages import grpo_normalize_rewards from miles.backends.fsdp_utils.loss_hub.context import DiffusionLossContext, PreparedBatch from miles.backends.fsdp_utils.loss_hub.losses import flow_grpo_loss_formula, resolve_loss_formula_fn +from miles.backends.fsdp_utils.loss_hub.nft import NftTrainDataConverter from miles.backends.fsdp_utils.loss_hub.prepare import prepare_flow_grpo_batch, prepare_nft_batch, resolve_prepare_fn __all__ = [ "DiffusionLossContext", + "NftTrainDataConverter", "PreparedBatch", "flow_grpo_loss_formula", "grpo_normalize_rewards", diff --git a/miles/backends/fsdp_utils/loss_hub/advantages.py b/miles/backends/fsdp_utils/loss_hub/advantages.py index 4fed8912..64b8053c 100644 --- a/miles/backends/fsdp_utils/loss_hub/advantages.py +++ b/miles/backends/fsdp_utils/loss_hub/advantages.py @@ -1,8 +1,9 @@ -"""Reward → train-signal helpers for diffusion (customization building blocks). +"""Reward → train-signal helpers for diffusion. -Default GRPO group normalization. Override with +``grpo_normalize_rewards`` is the default implementation of +``RolloutManager._post_process_rewards``. Override with ``--custom-reward-post-process-path`` pointing at a function with the same -signature as ``grpo_normalize_rewards``. +signature. """ from __future__ import annotations diff --git a/miles/backends/fsdp_utils/loss_hub/losses.py b/miles/backends/fsdp_utils/loss_hub/losses.py index 50b38cc3..84875916 100644 --- a/miles/backends/fsdp_utils/loss_hub/losses.py +++ b/miles/backends/fsdp_utils/loss_hub/losses.py @@ -116,7 +116,11 @@ def flow_grpo_loss_formula( def resolve_loss_formula_fn(args: Namespace) -> LossFormulaFn: - """Loss *formula* only — DiT forward stays in the actor.""" + """Loss *formula* only — DiT forward stays in the actor. + + Custom path defaults (e.g. NFT) are assigned in ``arguments.py``. When the + path is unset, Flow-GRPO is the default implementation. + """ path = getattr(args, "custom_loss_function_path", None) if path: fn = load_function(path) diff --git a/miles/backends/fsdp_utils/loss_hub/nft.py b/miles/backends/fsdp_utils/loss_hub/nft.py index 7987c7aa..ac5ecc9f 100644 --- a/miles/backends/fsdp_utils/loss_hub/nft.py +++ b/miles/backends/fsdp_utils/loss_hub/nft.py @@ -1,7 +1,11 @@ -"""DiffusionNFT plugin: custom convert + loss *formula*. +"""DiffusionNFT plugin: train-data converter + loss *formula*. Prepare hook lives in ``prepare.py`` (``prepare_nft_batch``). Actor still owns -DiT forward (+ EMA/LoRA-base reference forward). +DiT forward (+ EMA/LoRA-base reference forward via ``--ref-mode``). + +Reward / advantage normalisation stays in ``RolloutManager._post_process_rewards`` +(default: ``grpo_normalize_rewards``). This module only expands samples into +K ``(x0, t)`` train pairs. """ from __future__ import annotations @@ -11,7 +15,6 @@ import torch -from miles.backends.fsdp_utils.loss_hub.advantages import grpo_normalize_rewards from miles.backends.fsdp_utils.loss_hub.context import DiffusionLossContext, PreparedBatch from miles.utils.metric_buffer import MetricBuffer from miles.utils.types import Sample @@ -97,7 +100,7 @@ def nft_branch_losses( # --------------------------------------------------------------------------- -# Convert (K-expanded pairs) +# Converter (K-expanded pairs; rewards already post-processed by rollout) # --------------------------------------------------------------------------- @@ -111,52 +114,71 @@ def _clean_x0_from_sample(sample: Sample) -> torch.Tensor: return traj.latents[-1].detach().cpu().float() -def convert_samples_to_nft_train_data(args: Namespace, samples: list[Sample]) -> dict[str, Any]: - """Expand every sample into K ``(x0, t)`` train pairs (sample-major).""" - raw_rewards, advantages = grpo_normalize_rewards(args, samples) - if not samples: - raise ValueError("NFT convert received empty samples") - first_traj = samples[0].dit_trajectory - if first_traj is None: - raise ValueError("sample 0 missing dit_trajectory") - if first_traj.timesteps is None: - raise ValueError("NFT needs dit_trajectory.timesteps from rollout") - num_train_timesteps = int(getattr(args, "diffusion_num_train_timesteps", 1000) or 1000) - if first_traj.sigmas is not None: - scheduler_sigmas = first_traj.sigmas.detach().cpu().float() - else: - # Match scheduler_meta_from_rollout when sglang omits sigmas (e.g. ODE rollout). - ts = first_traj.timesteps.detach().cpu().float() - scheduler_sigmas = torch.cat([ts / float(num_train_timesteps), ts.new_zeros(1)]) - scheduler_meta = { - "scheduler_timesteps": first_traj.timesteps.detach().cpu().float(), - "scheduler_sigmas": scheduler_sigmas, - } - frac = float(getattr(args, "diffusion_nft_timestep_fraction", 0.99) or 0.99) - shuffle_t = bool(getattr(args, "diffusion_nft_shuffle_timesteps", True)) - sigmas = resolve_nft_sigmas(scheduler_meta["scheduler_sigmas"], training_timestep_fraction=frac) - num_timesteps = int(sigmas.numel()) - - train_data: list[dict[str, Any]] = [] - for sample, adv, raw in zip(samples, advantages, raw_rewards, strict=True): - if sample.denoising_env is None: - raise ValueError(f"sample {sample.index} missing denoising_env") - x0 = _clean_x0_from_sample(sample) - sample_sigmas = sigmas[torch.randperm(num_timesteps)] if shuffle_t else sigmas - for t in sample_sigmas.tolist(): - train_data.append( - { - "x0": x0, - "timestep": float(t), - "denoising_env": sample.denoising_env, - "advantage": float(adv), - "raw_reward": float(raw), - "sample_index": sample.index, - "prompt": sample.prompt, - "nft_num_timesteps": num_timesteps, - } +class NftTrainDataConverter: + """Expand rollout samples into K ``(x0, t)`` train pairs (sample-major). + + Same interface as ``RolloutTrainDataConverter``: receives already-normalized + rewards from ``_post_process_rewards``. Selected by ``--loss-type nft``. + """ + + def __init__(self, args: Namespace) -> None: + self.args = args + + def convert_samples( + self, + samples: list[Sample], + rewards: list[float], + raw_rewards: list[float], + ) -> dict[str, Any]: + if not samples: + raise ValueError("NFT convert received empty samples") + if len(samples) != len(rewards) or len(samples) != len(raw_rewards): + raise ValueError( + f"NFT convert length mismatch: samples={len(samples)} " + f"rewards={len(rewards)} raw_rewards={len(raw_rewards)}" ) - return {"train_data": train_data, **scheduler_meta} + first_traj = samples[0].dit_trajectory + if first_traj is None: + raise ValueError("sample 0 missing dit_trajectory") + if first_traj.timesteps is None: + raise ValueError("NFT needs dit_trajectory.timesteps from rollout") + args = self.args + num_train_timesteps = int(getattr(args, "diffusion_num_train_timesteps", 1000) or 1000) + if first_traj.sigmas is not None: + scheduler_sigmas = first_traj.sigmas.detach().cpu().float() + else: + # Match scheduler_meta_from_rollout when sglang omits sigmas (e.g. ODE rollout). + ts = first_traj.timesteps.detach().cpu().float() + scheduler_sigmas = torch.cat([ts / float(num_train_timesteps), ts.new_zeros(1)]) + scheduler_meta = { + "scheduler_timesteps": first_traj.timesteps.detach().cpu().float(), + "scheduler_sigmas": scheduler_sigmas, + } + frac = float(getattr(args, "diffusion_nft_timestep_fraction", 0.99) or 0.99) + shuffle_t = bool(getattr(args, "diffusion_nft_shuffle_timesteps", True)) + sigmas = resolve_nft_sigmas(scheduler_meta["scheduler_sigmas"], training_timestep_fraction=frac) + num_timesteps = int(sigmas.numel()) + + train_data: list[dict[str, Any]] = [] + for sample, adv, raw in zip(samples, rewards, raw_rewards, strict=True): + if sample.denoising_env is None: + raise ValueError(f"sample {sample.index} missing denoising_env") + x0 = _clean_x0_from_sample(sample) + sample_sigmas = sigmas[torch.randperm(num_timesteps)] if shuffle_t else sigmas + for t in sample_sigmas.tolist(): + train_data.append( + { + "x0": x0, + "timestep": float(t), + "denoising_env": sample.denoising_env, + "advantage": float(adv), + "raw_reward": float(raw), + "sample_index": sample.index, + "prompt": sample.prompt, + "nft_num_timesteps": num_timesteps, + } + ) + return {"train_data": train_data, **scheduler_meta} # --------------------------------------------------------------------------- @@ -229,7 +251,6 @@ def nft_loss_formula( return loss_sum -# Actor: always run a reference DiT forward for NFT (EMA preferred). -nft_loss_formula.ref_mode = "ema" # Same-sample K pairs must stay in one optimizer window. +# Reference forward is selected via --ref-mode (auto-filled for --loss-type nft). nft_loss_formula.requires_sample_aligned_windows = True diff --git a/miles/ray/rollout.py b/miles/ray/rollout.py index 2590c46e..8f1a68b3 100644 --- a/miles/ray/rollout.py +++ b/miles/ray/rollout.py @@ -26,7 +26,7 @@ from miles.utils.ray_utils import Box from miles.utils.timer import timer from miles.utils.tracking_utils import init_tracking -from miles.utils.train_data_utils import RolloutTrainDataConverter, TrainDataDPSplitter, reorder_train_pairs_for_tiling +from miles.utils.train_data_utils import TrainDataDPSplitter, reorder_train_pairs_for_tiling, resolve_train_data_converter from miles.utils.train_metric_utils import log_perf_data_raw from miles.utils.types import Sample @@ -77,7 +77,9 @@ def __init__(self, args, pg): if self.args.custom_convert_samples_to_train_data_path is not None else None ) - self.train_data_converter = RolloutTrainDataConverter() + # Algorithm-specific converters share the default reward post-process path. + # Full convert overrides use --custom-convert-samples-to-train-data-path instead. + self.train_data_converter = resolve_train_data_converter(self.args) self.train_data_dp_splitter = TrainDataDPSplitter() logger.info(f"import {self.args.rollout_function_path} as generate_rollout function.") logger.info(f"import {self.args.eval_function_path} as eval_generate_rollout function.") @@ -335,34 +337,19 @@ def _post_process_rewards(self, samples: list[Sample] | list[list[Sample]]): if self.custom_reward_post_process_func is not None: return self.custom_reward_post_process_func(self.args, samples) - raw_rewards = [sample.get_reward_value(self.args) for sample in samples] + from miles.backends.fsdp_utils.loss_hub.advantages import grpo_normalize_rewards - # --globalize-reward-mean / --globalize-reward-std are orthogonal. flow_grpo - # pickscore_qwenimage uses per-prompt mean + global std (PerPromptStatTracker - # with global_std=True), which is --globalize-reward-std alone. - rewards_flat = torch.tensor(raw_rewards, dtype=torch.float) - rewards = rewards_flat.view(-1, self.args.n_samples_per_prompt) - - if self.args.globalize_reward_mean: - mean = rewards_flat.mean() - else: - mean = rewards.mean(dim=-1, keepdim=True) - rewards = rewards - mean - - if self.args.grpo_std_normalization: - if self.args.globalize_reward_std: - std = rewards_flat.std() - else: - std = rewards.std(dim=-1, keepdim=True) - # matches flow_grpo's `+ 1e-4` in both stat_tracking branches - rewards = rewards / (std + 1e-4) - - return raw_rewards, rewards.flatten().tolist() + return grpo_normalize_rewards(self.args, samples) def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sample]]): """ Convert inference generated samples to training data. """ + # Full override (Miles-LLM style): covers this entire method, including + # reward post-process. Algorithm converters (e.g. NFT) use the path below. + if self.custom_convert_samples_to_train_data_func is not None: + return self.custom_convert_samples_to_train_data_func(self.args, samples) + raw_rewards, rewards = self._post_process_rewards(samples) assert len(raw_rewards) == len(samples) @@ -372,7 +359,7 @@ def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sampl norm_t = torch.tensor(rewards, dtype=torch.float) # Emit reward distribution stats (raw + normalized) to stdout + wandb. - # Runs for both default SDE-pair expand and custom converts (e.g. NFT). + # Runs for both default SDE-pair expand and NFT converter. reward_stats = { **_reward_stats_dict(raw_t, "rollout/reward/raw_"), **_reward_stats_dict(norm_t, "rollout/reward/norm_"), @@ -404,9 +391,6 @@ def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sampl reward_key=self.args.reward_key, ) - if self.custom_convert_samples_to_train_data_func is not None: - return self.custom_convert_samples_to_train_data_func(self.args, samples) - return self.train_data_converter.convert_samples(samples, rewards, raw_rewards) def _log_images( diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index 31881cf1..2797a56c 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -145,8 +145,19 @@ def add_train_arguments(parser): type=float, default=0.0, help=( - "Reference KL coefficient for diffusion GRPO. When > 0 with LoRA, " - "the trainer disables the LoRA adapter to compute the base-model reference." + "Reference KL coefficient for diffusion GRPO. When > 0, enables a " + "reference DiT forward (see --ref-mode; default lora_base)." + ), + ) + parser.add_argument( + "--ref-mode", + type=str, + choices=["none", "lora_base", "ema"], + default=None, + help=( + "Which reference weights to use for the no-grad DiT forward. " + "Auto: lora_base when --diffusion-kl-beta > 0; for --loss-type nft, " + "ema (or lora_base fallback). Explicit values skip auto inference." ), ) parser.add_argument( @@ -863,7 +874,7 @@ def add_algo_arguments(parser): default="policy_loss", help=( "Train objective shortcut. Diffusion default policy_loss = Flow-GRPO prepare+formula. " - "nft auto-fills convert / prepare / loss-formula custom paths. " + "nft auto-fills prepare / loss-formula paths and selects the NFT train-data converter. " "DiT forward always stays in the FSDP actor." ), ) @@ -934,7 +945,10 @@ def add_algo_arguments(parser): type=str, choices=["ema", "base"], default="ema", - help="NFT reference policy: EMA shadow of the LoRA weights ('ema') or LoRA-base ('base').", + help=( + "Deprecated alias for NFT auto --ref-mode: 'ema' or 'base' (lora_base). " + "Prefer --ref-mode directly." + ), ) parser.add_argument( "--advantage-estimator", @@ -1477,6 +1491,34 @@ def _resolve_eval_datasets(args) -> list[EvalDatasetConfig]: return eval_datasets +def resolve_and_validate_ref_mode(args, *, is_nft: bool, ema_enabled: bool) -> None: + """Infer and validate ``args.ref_mode`` (fast-fail; actor only consumes the result).""" + if getattr(args, "ref_mode", None) is None: + if is_nft: + nft_pref = getattr(args, "diffusion_nft_ref_mode", "ema") + if nft_pref == "ema" and not ema_enabled: + logger.warning( + "--loss-type nft prefers EMA ref but --lora-ema-shadow is off; " + "falling back to --ref-mode lora_base." + ) + args.ref_mode = "lora_base" + else: + args.ref_mode = "ema" if nft_pref == "ema" else "lora_base" + elif float(getattr(args, "diffusion_kl_beta", 0.0) or 0.0) > 0: + args.ref_mode = "lora_base" + else: + args.ref_mode = "none" + + if is_nft and args.ref_mode == "none": + raise ValueError("--loss-type nft requires a reference model; set --ref-mode ema or lora_base") + if args.ref_mode == "ema" and not ema_enabled: + raise ValueError("--ref-mode ema requires --lora-ema-shadow") + if args.ref_mode == "lora_base" and not args.use_lora: + raise ValueError("--ref-mode lora_base requires --use-lora") + if float(getattr(args, "diffusion_kl_beta", 0.0) or 0.0) > 0 and args.ref_mode == "none": + raise ValueError("--diffusion-kl-beta > 0 requires a reference model; set --ref-mode lora_base or ema") + + def miles_validate_args(args): args.eval_datasets = _resolve_eval_datasets(args) @@ -1562,16 +1604,14 @@ def miles_validate_args(args): if ema_enabled and not args.use_lora: raise ValueError("--lora-ema-shadow requires --use-lora") - if getattr(args, "loss_type", None) in ("nft", "diffusion_nft"): - # DiffusionNFT: swap convert + prepare + loss formula; DiT forward stays in actor. + is_nft = getattr(args, "loss_type", None) in ("nft", "diffusion_nft") + if is_nft: + # DiffusionNFT: swap prepare + loss formula; NFT converter is selected in RolloutManager. + # Do not overwrite --custom-convert-samples-to-train-data-path (full-convert override). if getattr(args, "custom_loss_function_path", None) is None: args.custom_loss_function_path = "miles.backends.fsdp_utils.loss_hub.nft.nft_loss_formula" if getattr(args, "custom_prepare_train_batch_path", None) is None: args.custom_prepare_train_batch_path = "miles.backends.fsdp_utils.loss_hub.prepare.prepare_nft_batch" - if args.custom_convert_samples_to_train_data_path is None: - args.custom_convert_samples_to_train_data_path = ( - "miles.backends.fsdp_utils.loss_hub.nft.convert_samples_to_nft_train_data" - ) if ( getattr(args, "diffusion_sde_type", "sde") == "sde" and float(getattr(args, "diffusion_noise_level", 0.7) or 0.0) == 0.0 @@ -1581,11 +1621,6 @@ def miles_validate_args(args): "(forward-process needs clean x0 only)." ) args.diffusion_sde_type = "ode" - if getattr(args, "diffusion_nft_ref_mode", "ema") == "ema" and not ema_enabled: - logger.warning( - "--loss-type nft with --diffusion-nft-ref-mode ema but without --lora-ema-shadow; " - "NFT loss will fall back to LoRA-base (disable_adapter) as the reference." - ) beta = float(getattr(args, "diffusion_nft_beta", 1.0) or 1.0) if beta <= 0: raise ValueError(f"--diffusion-nft-beta must be > 0, got {beta}") @@ -1597,6 +1632,9 @@ def miles_validate_args(args): "--diffusion-recompute-old-log-prob is only supported for policy_loss / Flow-GRPO, not NFT" ) + # Resolve --ref-mode once here (fast-fail); actor only consumes the final value. + resolve_and_validate_ref_mode(args, is_nft=is_nft, ema_enabled=ema_enabled) + if args.dump_details is not None: args.save_debug_rollout_data = f"{args.dump_details}/rollout_data/{{rollout_id}}.pt" args.save_debug_train_data = f"{args.dump_details}/train_data/{{rollout_id}}_{{rank}}.pt" diff --git a/miles/utils/train_data_utils.py b/miles/utils/train_data_utils.py index 2da7e322..6ac8dec7 100644 --- a/miles/utils/train_data_utils.py +++ b/miles/utils/train_data_utils.py @@ -41,6 +41,15 @@ def scheduler_meta_from_rollout( return timesteps, sigmas +def resolve_train_data_converter(args): + """Pick the algorithm train-data converter (not a full ``custom_convert`` override).""" + if getattr(args, "loss_type", None) in ("nft", "diffusion_nft"): + from miles.backends.fsdp_utils.loss_hub.nft import NftTrainDataConverter + + return NftTrainDataConverter(args) + return RolloutTrainDataConverter() + + class RolloutTrainDataConverter: """Convert rollout samples into the flat train-pair payload.""" diff --git a/scripts/run-diffusion-nft-sd3-pickscore.sh b/scripts/run-diffusion-nft-sd3-pickscore.sh index 7dfdd7cb..9bc8e482 100755 --- a/scripts/run-diffusion-nft-sd3-pickscore.sh +++ b/scripts/run-diffusion-nft-sd3-pickscore.sh @@ -134,7 +134,7 @@ python -u "${ROOT_DIR}/train_diffusion.py" \ --diffusion-nft-beta 1.0 \ --diffusion-nft-adv-clip-max 5.0 \ --diffusion-nft-timestep-fraction 0.99 \ - --diffusion-nft-ref-mode ema \ + --ref-mode ema \ --lora-ema-shadow \ --lora-ema-rollout-policy ema \ --lora-ema-decay 0.001 \ 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 3334c63f..0e4c2a83 100644 --- a/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py +++ b/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py @@ -8,15 +8,11 @@ import torch -from miles.backends.fsdp_utils.lora_ema import ( - LoraEmaShadow, - lora_ema_rollout_policy, - lora_ema_shadow_enabled, - resolve_lora_ema_kwargs, -) +from miles.backends.fsdp_utils.lora_ema import LoraEmaShadow, resolve_lora_ema_kwargs +from miles.backends.fsdp_utils.loss_hub.advantages import grpo_normalize_rewards from miles.backends.fsdp_utils.loss_hub.losses import flow_grpo_loss_formula, resolve_loss_formula_fn from miles.backends.fsdp_utils.loss_hub.nft import ( - convert_samples_to_nft_train_data, + NftTrainDataConverter, corrupt, nft_loss_formula, nft_r_from_advantages, @@ -80,10 +76,15 @@ class _Env: Sample(index=0, prompt="a", reward=1.0, dit_trajectory=_Traj(), denoising_env=_Env()), Sample(index=1, prompt="b", reward=3.0, dit_trajectory=_Traj(), denoising_env=_Env()), ] - out = convert_samples_to_nft_train_data(_args(), samples) + args = _args() + raw_rewards, rewards = grpo_normalize_rewards(args, samples) + out = NftTrainDataConverter(args).convert_samples(samples, rewards, raw_rewards) assert len(out["train_data"]) == 4 assert {p["timestep"] for p in out["train_data"]} == {1.0, 0.5} assert out["train_data"][0]["x0"] is out["train_data"][1]["x0"] + # Advantages come from post-process, not from the converter. + assert out["train_data"][0]["advantage"] == rewards[0] + assert out["train_data"][2]["advantage"] == rewards[1] def test_formula_write_old_log_prob_is_noop(self): metrics = new_metric_buffer(None, torch.device("cpu"), ()) @@ -98,9 +99,9 @@ def test_formula_write_old_log_prob_is_noop(self): ) assert loss is None - def test_formula_declares_ref_and_window_attrs(self): - assert getattr(nft_loss_formula, "ref_mode", None) == "ema" + def test_formula_declares_window_attr(self): assert getattr(nft_loss_formula, "requires_sample_aligned_windows", False) is True + assert not hasattr(nft_loss_formula, "ref_mode") def test_resolve_defaults_are_flow_grpo(self): assert resolve_prepare_fn(_args()) is prepare_flow_grpo_batch @@ -141,14 +142,6 @@ def test_swap_in_restores_exactly(self): class TestLoraEmaArgs: - def test_shadow_enabled(self): - assert lora_ema_shadow_enabled(Namespace(lora_ema_shadow=True)) - assert not lora_ema_shadow_enabled(Namespace()) - - def test_rollout_policy(self): - assert lora_ema_rollout_policy(Namespace(lora_ema_rollout_policy="ema")) == "ema" - assert lora_ema_rollout_policy(Namespace()) == "live" - def test_resolve_kwargs(self): kwargs = resolve_lora_ema_kwargs( Namespace(lora_ema_decay=0.01, lora_ema_uprate=0.02, lora_ema_uphold=0.3, lora_ema_flat_steps=5) diff --git a/tests/fast/utils/test_ref_mode_and_convert.py b/tests/fast/utils/test_ref_mode_and_convert.py new file mode 100644 index 00000000..0b63c9aa --- /dev/null +++ b/tests/fast/utils/test_ref_mode_and_convert.py @@ -0,0 +1,116 @@ +"""CPU tests for --ref-mode resolution and train-data convert layering.""" + +from tests.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=30, suite="stage-a-cpu", labels=[]) + +from argparse import Namespace +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +from miles.backends.fsdp_utils.loss_hub.nft import NftTrainDataConverter +from miles.utils.arguments import resolve_and_validate_ref_mode +from miles.utils.train_data_utils import RolloutTrainDataConverter, resolve_train_data_converter + + +def _ref_args(**overrides): + base = dict( + ref_mode=None, + use_lora=True, + lora_ema_shadow=False, + diffusion_kl_beta=0.0, + diffusion_nft_ref_mode="ema", + loss_type="policy_loss", + ) + base.update(overrides) + return Namespace(**base) + + +class TestResolveAndValidateRefMode: + def test_default_none(self): + args = _ref_args() + resolve_and_validate_ref_mode(args, is_nft=False, ema_enabled=False) + assert args.ref_mode == "none" + + def test_kl_beta_auto_lora_base(self): + args = _ref_args(diffusion_kl_beta=0.01) + resolve_and_validate_ref_mode(args, is_nft=False, ema_enabled=False) + assert args.ref_mode == "lora_base" + + def test_nft_auto_ema_when_shadow_on(self): + args = _ref_args(loss_type="nft", lora_ema_shadow=True) + resolve_and_validate_ref_mode(args, is_nft=True, ema_enabled=True) + assert args.ref_mode == "ema" + + def test_nft_auto_fallback_without_shadow(self): + args = _ref_args(loss_type="nft", lora_ema_shadow=False) + resolve_and_validate_ref_mode(args, is_nft=True, ema_enabled=False) + assert args.ref_mode == "lora_base" + + def test_nft_rejects_explicit_none(self): + args = _ref_args(ref_mode="none", loss_type="nft") + with pytest.raises(ValueError, match="nft requires a reference model"): + resolve_and_validate_ref_mode(args, is_nft=True, ema_enabled=False) + + def test_explicit_ema_requires_shadow(self): + args = _ref_args(ref_mode="ema") + with pytest.raises(ValueError, match="requires --lora-ema-shadow"): + resolve_and_validate_ref_mode(args, is_nft=False, ema_enabled=False) + + def test_kl_beta_rejects_none(self): + args = _ref_args(ref_mode="none", diffusion_kl_beta=0.1) + with pytest.raises(ValueError, match="diffusion-kl-beta"): + resolve_and_validate_ref_mode(args, is_nft=False, ema_enabled=False) + + def test_nft_alias_base_maps_to_lora_base(self): + args = _ref_args(diffusion_nft_ref_mode="base") + resolve_and_validate_ref_mode(args, is_nft=True, ema_enabled=False) + assert args.ref_mode == "lora_base" + + +class TestResolveTrainDataConverter: + def test_default_is_rollout_converter(self): + conv = resolve_train_data_converter(Namespace(loss_type="policy_loss")) + assert isinstance(conv, RolloutTrainDataConverter) + + def test_nft_selects_nft_converter(self): + conv = resolve_train_data_converter( + Namespace( + loss_type="nft", + diffusion_nft_timestep_fraction=1.0, + diffusion_nft_shuffle_timesteps=False, + ) + ) + assert isinstance(conv, NftTrainDataConverter) + + +class TestConvertSamplesLayering: + def test_custom_convert_short_circuits_before_post_process(self): + from miles.ray.rollout import RolloutManager + + cls = RolloutManager.__ray_actor_class__ + called = {"post": False, "custom": False} + + def custom_convert(args, samples): + called["custom"] = True + return {"train_data": []} + + fake = SimpleNamespace( + custom_convert_samples_to_train_data_func=custom_convert, + args=Namespace(loss_type="policy_loss"), + rollout_id=0, + train_data_converter=MagicMock(), + ) + + def boom_post_process(self, samples): + called["post"] = True + raise AssertionError("post_process should be skipped for full custom_convert") + + fake._post_process_rewards = boom_post_process.__get__(fake, type(fake)) + out = cls._convert_samples_to_train_data(fake, []) + assert out == {"train_data": []} + assert called["custom"] is True + assert called["post"] is False + fake.train_data_converter.convert_samples.assert_not_called() From a357d7263d66f069ec0a4208f037678a05de9d2f Mon Sep 17 00:00:00 2001 From: niehen6174 Date: Wed, 29 Jul 2026 12:19:07 +0000 Subject: [PATCH 06/12] test: drop low-signal ref_mode/convert unit tests --- tests/fast/utils/test_ref_mode_and_convert.py | 116 ------------------ 1 file changed, 116 deletions(-) delete mode 100644 tests/fast/utils/test_ref_mode_and_convert.py diff --git a/tests/fast/utils/test_ref_mode_and_convert.py b/tests/fast/utils/test_ref_mode_and_convert.py deleted file mode 100644 index 0b63c9aa..00000000 --- a/tests/fast/utils/test_ref_mode_and_convert.py +++ /dev/null @@ -1,116 +0,0 @@ -"""CPU tests for --ref-mode resolution and train-data convert layering.""" - -from tests.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=30, suite="stage-a-cpu", labels=[]) - -from argparse import Namespace -from types import SimpleNamespace -from unittest.mock import MagicMock - -import pytest - -from miles.backends.fsdp_utils.loss_hub.nft import NftTrainDataConverter -from miles.utils.arguments import resolve_and_validate_ref_mode -from miles.utils.train_data_utils import RolloutTrainDataConverter, resolve_train_data_converter - - -def _ref_args(**overrides): - base = dict( - ref_mode=None, - use_lora=True, - lora_ema_shadow=False, - diffusion_kl_beta=0.0, - diffusion_nft_ref_mode="ema", - loss_type="policy_loss", - ) - base.update(overrides) - return Namespace(**base) - - -class TestResolveAndValidateRefMode: - def test_default_none(self): - args = _ref_args() - resolve_and_validate_ref_mode(args, is_nft=False, ema_enabled=False) - assert args.ref_mode == "none" - - def test_kl_beta_auto_lora_base(self): - args = _ref_args(diffusion_kl_beta=0.01) - resolve_and_validate_ref_mode(args, is_nft=False, ema_enabled=False) - assert args.ref_mode == "lora_base" - - def test_nft_auto_ema_when_shadow_on(self): - args = _ref_args(loss_type="nft", lora_ema_shadow=True) - resolve_and_validate_ref_mode(args, is_nft=True, ema_enabled=True) - assert args.ref_mode == "ema" - - def test_nft_auto_fallback_without_shadow(self): - args = _ref_args(loss_type="nft", lora_ema_shadow=False) - resolve_and_validate_ref_mode(args, is_nft=True, ema_enabled=False) - assert args.ref_mode == "lora_base" - - def test_nft_rejects_explicit_none(self): - args = _ref_args(ref_mode="none", loss_type="nft") - with pytest.raises(ValueError, match="nft requires a reference model"): - resolve_and_validate_ref_mode(args, is_nft=True, ema_enabled=False) - - def test_explicit_ema_requires_shadow(self): - args = _ref_args(ref_mode="ema") - with pytest.raises(ValueError, match="requires --lora-ema-shadow"): - resolve_and_validate_ref_mode(args, is_nft=False, ema_enabled=False) - - def test_kl_beta_rejects_none(self): - args = _ref_args(ref_mode="none", diffusion_kl_beta=0.1) - with pytest.raises(ValueError, match="diffusion-kl-beta"): - resolve_and_validate_ref_mode(args, is_nft=False, ema_enabled=False) - - def test_nft_alias_base_maps_to_lora_base(self): - args = _ref_args(diffusion_nft_ref_mode="base") - resolve_and_validate_ref_mode(args, is_nft=True, ema_enabled=False) - assert args.ref_mode == "lora_base" - - -class TestResolveTrainDataConverter: - def test_default_is_rollout_converter(self): - conv = resolve_train_data_converter(Namespace(loss_type="policy_loss")) - assert isinstance(conv, RolloutTrainDataConverter) - - def test_nft_selects_nft_converter(self): - conv = resolve_train_data_converter( - Namespace( - loss_type="nft", - diffusion_nft_timestep_fraction=1.0, - diffusion_nft_shuffle_timesteps=False, - ) - ) - assert isinstance(conv, NftTrainDataConverter) - - -class TestConvertSamplesLayering: - def test_custom_convert_short_circuits_before_post_process(self): - from miles.ray.rollout import RolloutManager - - cls = RolloutManager.__ray_actor_class__ - called = {"post": False, "custom": False} - - def custom_convert(args, samples): - called["custom"] = True - return {"train_data": []} - - fake = SimpleNamespace( - custom_convert_samples_to_train_data_func=custom_convert, - args=Namespace(loss_type="policy_loss"), - rollout_id=0, - train_data_converter=MagicMock(), - ) - - def boom_post_process(self, samples): - called["post"] = True - raise AssertionError("post_process should be skipped for full custom_convert") - - fake._post_process_rewards = boom_post_process.__get__(fake, type(fake)) - out = cls._convert_samples_to_train_data(fake, []) - assert out == {"train_data": []} - assert called["custom"] is True - assert called["post"] is False - fake.train_data_converter.convert_samples.assert_not_called() From a4fcea0ffe327c6841c8d49654681a2dc5f7a6f7 Mon Sep 17 00:00:00 2001 From: niehen6174 Date: Wed, 29 Jul 2026 12:26:12 +0000 Subject: [PATCH 07/12] refactor(fsdp): generalize LoRA EMA to trainable-param EmaShadow MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Rename lora_ema → ema / --ema-shadow and drop the --use-lora gate so full finetune can share the same pi_old path as LoRA. --- miles/backends/fsdp_utils/actor.py | 16 ++++---- .../fsdp_utils/{lora_ema.py => ema.py} | 35 ++++++++-------- miles/backends/fsdp_utils/loss_hub/context.py | 2 +- miles/utils/arguments.py | 41 +++++++++---------- scripts/run-diffusion-nft-sd3-pickscore.sh | 14 +++---- .../backends/fsdp_utils/test_loss_hub_nft.py | 14 +++---- 6 files changed, 58 insertions(+), 64 deletions(-) rename miles/backends/fsdp_utils/{lora_ema.py => ema.py} (73%) diff --git a/miles/backends/fsdp_utils/actor.py b/miles/backends/fsdp_utils/actor.py index 060cb8b4..f089d3d2 100644 --- a/miles/backends/fsdp_utils/actor.py +++ b/miles/backends/fsdp_utils/actor.py @@ -34,7 +34,7 @@ DiffusionUpdateWeightFromTensorLoRA, DiffusionUpdateWeightFromTensorLoRAIPC, ) -from .lora_ema import LoraEmaShadow, resolve_lora_ema_kwargs +from .ema import EmaShadow, resolve_ema_kwargs from .loss_hub import DiffusionLossContext, resolve_loss_formula_fn, resolve_prepare_fn from .lr_scheduler import get_lr_scheduler from .metrics import new_metric_buffer @@ -209,13 +209,13 @@ def init(self, args: Namespace, role: str, with_ref: bool = False) -> int: # ty checkpoint_payload = checkpoint.load(self) - # Optional LoRA EMA shadow (pi_old). Enabled via --lora-ema-shadow (parsed in arguments.py). + # Optional EMA shadow of trainable params (pi_old). Enabled via --ema-shadow. # Consumed when --ref-mode ema runs the no-grad reference DiT forward. self.ema_shadow = None - if getattr(self.args, "lora_ema_shadow", False): - self.ema_shadow = LoraEmaShadow( + if getattr(self.args, "ema_shadow", False): + self.ema_shadow = EmaShadow( (p for m in self.models.values() for p in m.parameters()), - **resolve_lora_ema_kwargs(self.args), + **resolve_ema_kwargs(self.args), ) # sglang-d now supports /update_weights_from_tensor (PR #20464). @@ -307,8 +307,8 @@ def update_weights(self) -> None: # type: ignore[override] if self.ema_shadow is not None: delta = self.ema_shadow.update() if dist.get_rank() == 0: - logger.info("LoRA EMA shadow updated (decay=%.4f step=%d)", delta, self.ema_shadow.step) - if getattr(self.args, "lora_ema_rollout_policy", "live") == "ema": + logger.info("EMA shadow updated (decay=%.4f step=%d)", delta, self.ema_shadow.step) + if getattr(self.args, "ema_rollout_policy", "live") == "ema": with self.ema_shadow.swap_in(): self.weight_updater.update_weights() else: @@ -374,7 +374,7 @@ def _train_core(self, rollout_id: int, rollout_data) -> None: "--ref-mode lora_base requires PEFT models exposing disable_adapter() after FSDP wrapping." ) if ref_mode == "ema" and self.ema_shadow is None: - raise RuntimeError("--ref-mode ema requires a constructed LoRA EMA shadow") + raise RuntimeError("--ref-mode ema requires a constructed EMA shadow") # ------------- Rollout Scheduler Metadata ------------- scheduler_timesteps, scheduler_sigmas = scheduler_meta_from_rollout( diff --git a/miles/backends/fsdp_utils/lora_ema.py b/miles/backends/fsdp_utils/ema.py similarity index 73% rename from miles/backends/fsdp_utils/lora_ema.py rename to miles/backends/fsdp_utils/ema.py index 463a9c63..7977533f 100644 --- a/miles/backends/fsdp_utils/lora_ema.py +++ b/miles/backends/fsdp_utils/ema.py @@ -1,8 +1,8 @@ -"""Generic LoRA EMA shadow for diffusion (and future FSDP) training. +"""EMA shadow of trainable parameters for diffusion FSDP training. Algorithms that need a slow-moving reference / sampling policy ``pi_old`` share -``LoraEmaShadow``: trainable LoRA weights plus EMA buffers and ``swap_in()`` -for temporary in-place weight exchange. +``EmaShadow``: tracks all ``requires_grad`` parameters (LoRA adapters or full +finetune), plus ``swap_in()`` for temporary in-place weight exchange. Lifecycle (actor / weight sync):: @@ -20,14 +20,14 @@ Checkpointing (intentionally not wired yet) ------------------------------------------- ``shadow`` / ``step`` are **not** saved or restored by ``fsdp_utils.checkpoint``. -On resume the actor rebuilds EMA from the loaded LoRA weights, so ``pi_old`` +On resume the actor rebuilds EMA from the loaded trainable weights, so ``pi_old`` cold-starts (decay schedule restarts at step 0). Fine for single-shot runs; wrong for mid-run resume that must match UniRL's slow ``pi_old``. Wiring it later is non-trivial: buffers are per-rank plain clones (not in the FSDP/DCP model state), must stay aligned with the trainable-param order, and must not be saved while ``swap_in()`` is active. Prefer a side file such as -``iter_*/lora_ema.pt`` over stuffing into the DCP model dict. +``iter_*/ema.pt`` over stuffing into the DCP model dict. """ from __future__ import annotations @@ -45,23 +45,22 @@ def _local(t: torch.Tensor) -> torch.Tensor: return t._local_tensor if hasattr(t, "_local_tensor") else t -def resolve_lora_ema_kwargs(args: Namespace) -> dict[str, float | int]: - """Read normalized ``lora_ema_*`` fields from ``args`` (see ``miles_validate_args``). +def resolve_ema_kwargs(args: Namespace) -> dict[str, float | int]: + """Read normalized ``ema_*`` fields from ``args`` (see ``miles_validate_args``). - Enablement (``args.lora_ema_shadow``) and rollout policy - (``args.lora_ema_rollout_policy``) are plain args — inferred/validated in - ``arguments.py``, not re-wrapped here. + Enablement (``args.ema_shadow``) and rollout policy (``args.ema_rollout_policy``) + are plain args — inferred/validated in ``arguments.py``, not re-wrapped here. """ return { - "decay": float(getattr(args, "lora_ema_decay", 0.001)), - "uprate": float(getattr(args, "lora_ema_uprate", 0.001)), - "uphold": float(getattr(args, "lora_ema_uphold", 0.5)), - "flat_steps": int(getattr(args, "lora_ema_flat_steps", 0)), + "decay": float(getattr(args, "ema_decay", 0.001)), + "uprate": float(getattr(args, "ema_uprate", 0.001)), + "uphold": float(getattr(args, "ema_uphold", 0.5)), + "flat_steps": int(getattr(args, "ema_flat_steps", 0)), } -class LoraEmaShadow: - """EMA shadow of trainable (LoRA) parameters. +class EmaShadow: + """EMA shadow of trainable parameters (LoRA or full finetune). Not part of the FSDP checkpoint payload today (see module docstring). """ @@ -84,7 +83,7 @@ def __init__( self.params = [p for p in parameters if p.requires_grad] if not self.params: - raise ValueError("LoraEmaShadow: model has no trainable parameters") + raise ValueError("EmaShadow: model has no trainable parameters") self.shadow = [_local(p.detach()).clone() for p in self.params] def decay_at(self, t: int) -> float: @@ -96,7 +95,7 @@ def decay_at(self, t: int) -> float: def update(self) -> float: """theta_old <- delta * theta_old + (1 - delta) * theta.""" if self._swapped: - raise RuntimeError("LoraEmaShadow.update called while swapped in") + raise RuntimeError("EmaShadow.update called while swapped in") self.step += 1 delta = self.decay_at(self.step) for live, sh in zip(self.params, self.shadow, strict=True): diff --git a/miles/backends/fsdp_utils/loss_hub/context.py b/miles/backends/fsdp_utils/loss_hub/context.py index 036de6c4..6c240186 100644 --- a/miles/backends/fsdp_utils/loss_hub/context.py +++ b/miles/backends/fsdp_utils/loss_hub/context.py @@ -26,7 +26,7 @@ class DiffusionLossContext: args: Namespace forward_dtype: torch.dtype device: torch.device - # Optional LoRA EMA shadow handle; owned by actor (see ``fsdp_utils.lora_ema``). + # Optional EMA shadow handle; owned by actor (see ``fsdp_utils.ema``). ema_shadow: Any = None diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index 2797a56c..b4f3929f 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -157,7 +157,7 @@ def add_train_arguments(parser): help=( "Which reference weights to use for the no-grad DiT forward. " "Auto: lora_base when --diffusion-kl-beta > 0; for --loss-type nft, " - "ema (or lora_base fallback). Explicit values skip auto inference." + "ema if --ema-shadow else lora_base. Explicit values skip auto inference." ), ) parser.add_argument( @@ -1153,48 +1153,48 @@ def add_debug_arguments(parser): ), ) parser.add_argument( - "--lora-ema-shadow", + "--ema-shadow", action="store_true", default=False, help=( - "Maintain an EMA shadow of trainable LoRA weights (pi_old). Custom losses " - "read it via ``ctx.ema_shadow.swap_in()``; combine with " - "--lora-ema-rollout-policy ema to sample under pi_old." + "Maintain an EMA shadow of trainable weights (pi_old; LoRA or full finetune). " + "Consumed when --ref-mode ema; combine with --ema-rollout-policy ema to " + "sample under pi_old." ), ) parser.add_argument( - "--lora-ema-rollout-policy", + "--ema-rollout-policy", type=str, choices=["live", "ema"], default="live", help=( - "Which LoRA weights to push to rollout after each rollout_end when " - "--lora-ema-shadow is set: live trainable weights, or EMA shadow (pi_old)." + "Which trainable weights to push to rollout after each rollout_end when " + "--ema-shadow is set: live weights, or EMA shadow (pi_old)." ), ) parser.add_argument( - "--lora-ema-decay", + "--ema-decay", type=float, default=0.001, - help="LoRA EMA decay while step <= flat_steps.", + help="EMA decay while step <= flat_steps.", ) parser.add_argument( - "--lora-ema-uprate", + "--ema-uprate", type=float, default=0.001, - help="LoRA EMA warmup rate after flat_steps.", + help="EMA warmup rate after flat_steps.", ) parser.add_argument( - "--lora-ema-uphold", + "--ema-uphold", type=float, default=0.5, - help="LoRA EMA warmup cap.", + help="EMA warmup cap.", ) parser.add_argument( - "--lora-ema-flat-steps", + "--ema-flat-steps", type=int, default=0, - help="LoRA EMA flat steps before warmup begins.", + help="EMA flat steps before warmup begins.", ) parser.add_argument( "--diffusion-init-lora-weight", @@ -1498,8 +1498,7 @@ def resolve_and_validate_ref_mode(args, *, is_nft: bool, ema_enabled: bool) -> N nft_pref = getattr(args, "diffusion_nft_ref_mode", "ema") if nft_pref == "ema" and not ema_enabled: logger.warning( - "--loss-type nft prefers EMA ref but --lora-ema-shadow is off; " - "falling back to --ref-mode lora_base." + "--loss-type nft prefers EMA ref but --ema-shadow is off; " "falling back to --ref-mode lora_base." ) args.ref_mode = "lora_base" else: @@ -1512,7 +1511,7 @@ def resolve_and_validate_ref_mode(args, *, is_nft: bool, ema_enabled: bool) -> N if is_nft and args.ref_mode == "none": raise ValueError("--loss-type nft requires a reference model; set --ref-mode ema or lora_base") if args.ref_mode == "ema" and not ema_enabled: - raise ValueError("--ref-mode ema requires --lora-ema-shadow") + raise ValueError("--ref-mode ema requires --ema-shadow") if args.ref_mode == "lora_base" and not args.use_lora: raise ValueError("--ref-mode lora_base requires --use-lora") if float(getattr(args, "diffusion_kl_beta", 0.0) or 0.0) > 0 and args.ref_mode == "none": @@ -1600,9 +1599,7 @@ def miles_validate_args(args): "set --diffusion-model (for per-model defaults) or --lora-target-modules." ) - ema_enabled = bool(getattr(args, "lora_ema_shadow", False)) - if ema_enabled and not args.use_lora: - raise ValueError("--lora-ema-shadow requires --use-lora") + ema_enabled = bool(getattr(args, "ema_shadow", False)) is_nft = getattr(args, "loss_type", None) in ("nft", "diffusion_nft") if is_nft: diff --git a/scripts/run-diffusion-nft-sd3-pickscore.sh b/scripts/run-diffusion-nft-sd3-pickscore.sh index 9bc8e482..095b053d 100755 --- a/scripts/run-diffusion-nft-sd3-pickscore.sh +++ b/scripts/run-diffusion-nft-sd3-pickscore.sh @@ -5,7 +5,7 @@ # 8 prompts × 8 samples, lr=3e-4, wd=1e-4, LoRA r=32 α=64, # guidance=1.0 (CFG-free), eta=0 / num_sde_steps=0 (ODE → clean x0), # NFT beta=1, adv_clip=5, adaptive weight, schedule fraction=0.99, -# LoRA EMA shadow (pi_old) updated + synced at rollout_end (uprate=0.001, uphold=0.5). +# EMA shadow (pi_old) updated + synced at rollout_end (uprate=0.001, uphold=0.5). # Expected early metrics: reward ~0.73–0.82, train/loss ~280–360. # # GPU layout (default CUDA_VISIBLE_DEVICES=4,5,2): @@ -135,12 +135,12 @@ python -u "${ROOT_DIR}/train_diffusion.py" \ --diffusion-nft-adv-clip-max 5.0 \ --diffusion-nft-timestep-fraction 0.99 \ --ref-mode ema \ - --lora-ema-shadow \ - --lora-ema-rollout-policy ema \ - --lora-ema-decay 0.001 \ - --lora-ema-uprate 0.001 \ - --lora-ema-uphold 0.5 \ - --lora-ema-flat-steps 0 \ + --ema-shadow \ + --ema-rollout-policy ema \ + --ema-decay 0.001 \ + --ema-uprate 0.001 \ + --ema-uphold 0.5 \ + --ema-flat-steps 0 \ --advantage-estimator grpo \ --globalize-reward-std \ --diffusion-model "${SD3_MODEL}" \ 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 0e4c2a83..958722e3 100644 --- a/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py +++ b/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py @@ -8,7 +8,7 @@ import torch -from miles.backends.fsdp_utils.lora_ema import LoraEmaShadow, resolve_lora_ema_kwargs +from miles.backends.fsdp_utils.ema import EmaShadow, resolve_ema_kwargs from miles.backends.fsdp_utils.loss_hub.advantages import grpo_normalize_rewards from miles.backends.fsdp_utils.loss_hub.losses import flow_grpo_loss_formula, resolve_loss_formula_fn from miles.backends.fsdp_utils.loss_hub.nft import ( @@ -116,13 +116,13 @@ def test_resolve_custom_paths(self): assert resolve_loss_formula_fn(args) is nft_loss_formula -class TestLoraEmaShadow: +class TestEmaShadow: def _model(self): return torch.nn.Linear(4, 4, bias=False) def test_snapshot_and_update(self): m = self._model() - ema = LoraEmaShadow(m.parameters(), decay=0.5, uprate=0.001, uphold=0.5, flat_steps=10) + ema = EmaShadow(m.parameters(), decay=0.5, uprate=0.001, uphold=0.5, flat_steps=10) init = m.weight.detach().clone() with torch.no_grad(): m.weight.add_(1.0) @@ -132,7 +132,7 @@ def test_snapshot_and_update(self): def test_swap_in_restores_exactly(self): m = self._model() - ema = LoraEmaShadow(m.parameters(), decay=0.1) + ema = EmaShadow(m.parameters(), decay=0.1) live = m.weight.detach().clone() with torch.no_grad(): m.weight.add_(2.0) @@ -141,9 +141,7 @@ def test_swap_in_restores_exactly(self): assert torch.equal(m.weight.detach(), live + 2.0) -class TestLoraEmaArgs: +class TestEmaArgs: def test_resolve_kwargs(self): - kwargs = resolve_lora_ema_kwargs( - Namespace(lora_ema_decay=0.01, lora_ema_uprate=0.02, lora_ema_uphold=0.3, lora_ema_flat_steps=5) - ) + kwargs = resolve_ema_kwargs(Namespace(ema_decay=0.01, ema_uprate=0.02, ema_uphold=0.3, ema_flat_steps=5)) assert kwargs == {"decay": 0.01, "uprate": 0.02, "uphold": 0.3, "flat_steps": 5} From 0dca4a5f5cd7a001c191c379ea6ad5a03acbec8e Mon Sep 17 00:00:00 2001 From: niehen6174 Date: Wed, 29 Jul 2026 12:32:02 +0000 Subject: [PATCH 08/12] style: isort multiline import in rollout.py --- miles/ray/rollout.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/miles/ray/rollout.py b/miles/ray/rollout.py index 8f1a68b3..e6a1e298 100644 --- a/miles/ray/rollout.py +++ b/miles/ray/rollout.py @@ -26,7 +26,11 @@ from miles.utils.ray_utils import Box from miles.utils.timer import timer from miles.utils.tracking_utils import init_tracking -from miles.utils.train_data_utils import TrainDataDPSplitter, reorder_train_pairs_for_tiling, resolve_train_data_converter +from miles.utils.train_data_utils import ( + TrainDataDPSplitter, + reorder_train_pairs_for_tiling, + resolve_train_data_converter, +) from miles.utils.train_metric_utils import log_perf_data_raw from miles.utils.types import Sample From 93df3b11960d8b8c90399c3a74f1c05e4d73178c Mon Sep 17 00:00:00 2001 From: rockdu Date: Thu, 30 Jul 2026 02:59:24 -0700 Subject: [PATCH 09/12] refactor(diffusion): move train-pair expansion into algorithm hooks Keep rollout orchestration generic while allowing NFT to provide its own sample expansion without a converter class. Co-authored-by: Cursor --- .../configs/train_pipeline_config.py | 2 +- .../fsdp_utils/loss_hub/advantages.py | 51 ----- miles/ray/data_conversion_hub/__init__.py | 1 + miles/ray/data_conversion_hub/flow_grpo.py | 175 ++++++++++++++++ miles/ray/data_conversion_hub/nft.py | 105 ++++++++++ miles/ray/rollout.py | 43 +++- miles/utils/arguments.py | 17 +- miles/utils/train_data_utils.py | 189 ------------------ .../backends/fsdp_utils/test_loss_hub_nft.py | 29 +-- tests/fast/utils/test_grouping_parity.py | 11 +- .../utils/test_legacy_tile_grouping_golden.py | 2 +- 11 files changed, 341 insertions(+), 284 deletions(-) delete mode 100644 miles/backends/fsdp_utils/loss_hub/advantages.py create mode 100644 miles/ray/data_conversion_hub/__init__.py create mode 100644 miles/ray/data_conversion_hub/flow_grpo.py create mode 100644 miles/ray/data_conversion_hub/nft.py diff --git a/miles/backends/fsdp_utils/configs/train_pipeline_config.py b/miles/backends/fsdp_utils/configs/train_pipeline_config.py index e325d2ad..c392ea57 100644 --- a/miles/backends/fsdp_utils/configs/train_pipeline_config.py +++ b/miles/backends/fsdp_utils/configs/train_pipeline_config.py @@ -7,7 +7,7 @@ - FSDP preprocess hooks Trajectory unpacking for train-pair construction lives in -``miles.utils.train_data_utils.RolloutTrainDataConverter``. +``miles.ray.data_conversion_hub.flow_grpo.expand_samples_to_train_pairs``. """ from __future__ import annotations diff --git a/miles/backends/fsdp_utils/loss_hub/advantages.py b/miles/backends/fsdp_utils/loss_hub/advantages.py deleted file mode 100644 index 64b8053c..00000000 --- a/miles/backends/fsdp_utils/loss_hub/advantages.py +++ /dev/null @@ -1,51 +0,0 @@ -"""Reward → train-signal helpers for diffusion. - -``grpo_normalize_rewards`` is the default implementation of -``RolloutManager._post_process_rewards``. Override with -``--custom-reward-post-process-path`` pointing at a function with the same -signature. -""" - -from __future__ import annotations - -from argparse import Namespace - -import torch - -from miles.utils.types import Sample - - -def grpo_normalize_rewards( - args: Namespace, - samples: list[Sample] | list[list[Sample]], -) -> tuple[list[float], list[float]]: - """Group-relative reward normalization used by Flow-GRPO. - - Returns ``(raw_rewards, normalized_rewards)``. Normalized values are used - as per-sample advantages when building train pairs. - - ``--globalize-reward-mean`` / ``--globalize-reward-std`` are orthogonal. - flow_grpo pickscore_qwenimage uses per-prompt mean + global std - (``PerPromptStatTracker`` with ``global_std=True``), which is - ``--globalize-reward-std`` alone. - """ - raw_rewards = [sample.get_reward_value(args) for sample in samples] - - rewards_flat = torch.tensor(raw_rewards, dtype=torch.float) - rewards = rewards_flat.view(-1, args.n_samples_per_prompt) - - if args.globalize_reward_mean: - mean = rewards_flat.mean() - else: - mean = rewards.mean(dim=-1, keepdim=True) - rewards = rewards - mean - - if args.grpo_std_normalization: - if args.globalize_reward_std: - std = rewards_flat.std() - else: - std = rewards.std(dim=-1, keepdim=True) - # matches flow_grpo's `+ 1e-4` in both stat_tracking branches - rewards = rewards / (std + 1e-4) - - return raw_rewards, rewards.flatten().tolist() diff --git a/miles/ray/data_conversion_hub/__init__.py b/miles/ray/data_conversion_hub/__init__.py new file mode 100644 index 00000000..61613884 --- /dev/null +++ b/miles/ray/data_conversion_hub/__init__.py @@ -0,0 +1 @@ +"""Rollout-to-train-pair conversion hooks for diffusion algorithms.""" diff --git a/miles/ray/data_conversion_hub/flow_grpo.py b/miles/ray/data_conversion_hub/flow_grpo.py new file mode 100644 index 00000000..00dbbbfd --- /dev/null +++ b/miles/ray/data_conversion_hub/flow_grpo.py @@ -0,0 +1,175 @@ +"""Convert Flow-GRPO rollout samples into flat train-pair payloads.""" + +from typing import Any + +import torch + +from miles.utils.types import RolloutDebugTensors, Sample + + +def expand_samples_to_train_pairs( + args, + samples: list[Sample], + rewards: list[float], + raw_rewards: list[float], +) -> dict[str, Any]: + train_data, scheduler_meta = _expand_samples_to_train_pairs(samples, rewards, raw_rewards) + return {"train_data": train_data, **scheduler_meta} + + +def _expand_samples_to_train_pairs( + samples: list[Sample], + rewards: list[float], + raw_rewards: list[float], +) -> tuple[list[dict[str, Any]], dict[str, torch.Tensor]]: + """Flat train pairs in sample-major order (all pairs for sample 0, then sample 1, ...).""" + device = torch.device("cpu") + train_data: list[dict[str, Any]] = [] + first_traj = samples[0].dit_trajectory + # Scheduler meta is taken from sample 0 and returned once for the whole batch; + # the per-sample loop below verifies every sample actually shares it. + scheduler_meta: dict[str, torch.Tensor] = {"scheduler_timesteps": first_traj.timesteps.detach().cpu().float()} + + if first_traj.sigmas is not None: + scheduler_meta["scheduler_sigmas"] = first_traj.sigmas.detach().cpu().float() + + for sample, rew, raw_r in zip(samples, rewards, raw_rewards, strict=True): + traj, denoising_env, rollout_log_probs = _sample_required_inputs(sample) + # Nail down the shared-scheduler-meta assumption: every sample must carry the + # same timesteps/sigmas as sample 0, since one scheduler_meta is returned for all. + if not torch.equal(traj.timesteps.detach().cpu().float(), scheduler_meta["scheduler_timesteps"]): + raise ValueError( + f"sample {sample.index} has different scheduler_timesteps than sample 0; " + "the converter assumes one shared schedule across the batch" + ) + expected_sigmas = scheduler_meta.get("scheduler_sigmas") + traj_sigmas = None if traj.sigmas is None else traj.sigmas.detach().cpu().float() + if (expected_sigmas is None) != (traj_sigmas is None) or ( + expected_sigmas is not None and not torch.equal(traj_sigmas, expected_sigmas) + ): + raise ValueError( + f"sample {sample.index} has different scheduler_sigmas than sample 0; " + "the converter assumes one shared schedule across the batch" + ) + per_sample_features = _build_per_sample_features( + sample, + reward=rew, + raw_reward=raw_r, + denoising_env=denoising_env, + ) + per_timestep_features, idx = _build_per_timestep_features( + sample, + traj=traj, + rollout_log_probs=rollout_log_probs, + device=device, + ) + pair_debug_steps = None + if sample.rollout_debug_tensors is not None: + pair_debug_steps = _slice_rollout_debug_for_train_pairs(sample.rollout_debug_tensors, sde_idx=idx) + sample_t_steps = int(per_timestep_features["latent"].shape[0]) + if pair_debug_steps is not None and len(pair_debug_steps) != sample_t_steps: + raise ValueError( + f"rollout_debug_tensors step count {len(pair_debug_steps)} != train pairs {sample_t_steps} " + f"(sample_index={sample.index})" + ) + + for t in range(sample_t_steps): + pair: dict[str, Any] = { + **per_sample_features, + **{key: tensor[t].detach().cpu() for key, tensor in per_timestep_features.items()}, + } + if pair_debug_steps is not None: + pair["rollout_debug_tensors"] = pair_debug_steps[t] + train_data.append(pair) + + if not train_data: + raise ValueError("No train pairs were produced from rollout samples") + + return train_data, scheduler_meta + + +def _sample_required_inputs(sample: Sample): + traj = sample.dit_trajectory + denoising_env = sample.denoising_env + rollout_log_probs = sample.rollout_log_probs + if traj is None or traj.timesteps is None or denoising_env is None or rollout_log_probs is None: + raise ValueError("Sample missing dit_trajectory, denoising_env, or rollout_log_probs") + return traj, denoising_env, rollout_log_probs + + +def _build_per_sample_features( + sample: Sample, + *, + reward: float, + raw_reward: float, + denoising_env, +) -> dict[str, Any]: + """Fields shared by every train pair produced from one sample.""" + return { + "advantage": float(reward), + "denoising_env": denoising_env, + "sample_index": sample.index, + "prompt": sample.prompt, + "raw_reward": float(raw_reward), + } + + +def _build_per_timestep_features( + sample: Sample, + *, + traj, + rollout_log_probs: torch.Tensor, + device: torch.device, +) -> tuple[dict[str, torch.Tensor], torch.Tensor]: + """Fields with one row per selected denoising step.""" + all_latents = traj.latents.to(device, dtype=torch.float32) + latents = all_latents[:-1] + next_latents = all_latents[1:] + timesteps = traj.timesteps.to(device, dtype=torch.float32) + # The step after the last has no recorded timestep -> terminal (σ=0, timestep 0), + # so the SDE step reads σ_next from the actual next rollout timestep, not a lookup. + next_timesteps = torch.cat([timesteps[1:], timesteps.new_zeros(1)]) + + sde_idx = (sample.train_metadata or {}).get("sde_step_indices") + assert sde_idx is not None, "SDE step indices are required for training" + idx = torch.as_tensor(sde_idx, dtype=torch.long) + return { + "latent": latents[idx], + "next_latent": next_latents[idx], + "timestep": timesteps[idx], + "next_timestep": next_timesteps[idx], + "log_prob_old": rollout_log_probs[idx], + }, idx + + +def _slice_rollout_debug_for_train_pairs( + dbg: RolloutDebugTensors, + *, + sde_idx: torch.Tensor | None = None, +) -> list[dict[str, torch.Tensor]] | None: + """Slice per-sample rollout debug tensors into one debug payload per train pair.""" + rollout_to_train_pair_fields = { + "rollout_variance_noises": "rollout_step_variance_noise", + "rollout_prev_sample_means": "rollout_step_prev_sample_mean", + "rollout_noise_std_devs": "rollout_step_noise_std_dev", + "rollout_model_outputs": "rollout_step_model_output", + } + indexed: dict[str, torch.Tensor] = {} + for rollout_key, train_pair_key in rollout_to_train_pair_fields.items(): + tensor = getattr(dbg, rollout_key, None) + if tensor is None: + continue + tensor = tensor.detach().cpu() + if sde_idx is not None: + tensor = tensor[sde_idx] + indexed[train_pair_key] = tensor + + if not indexed: + return None + + num_steps = int(next(iter(indexed.values())).shape[0]) + for key, tensor in indexed.items(): + if int(tensor.shape[0]) != num_steps: + raise ValueError(f"Rollout debug field {key!r} has {tensor.shape[0]} steps, expected {num_steps}") + + return [{key: tensor[step] for key, tensor in indexed.items()} for step in range(num_steps)] diff --git a/miles/ray/data_conversion_hub/nft.py b/miles/ray/data_conversion_hub/nft.py new file mode 100644 index 00000000..54208703 --- /dev/null +++ b/miles/ray/data_conversion_hub/nft.py @@ -0,0 +1,105 @@ +"""Convert DiffusionNFT rollout samples into timestep-expanded train pairs.""" + +from argparse import Namespace +from typing import Any + +import torch + +from miles.utils.types import Sample + +# --------------------------------------------------------------------------- +# Converter (K-expanded pairs; rewards already post-processed by rollout) +# --------------------------------------------------------------------------- + + +def _clean_x0_from_sample(sample: Sample) -> torch.Tensor: + traj = sample.dit_trajectory + if traj is None or traj.latents is None or traj.latents.shape[0] < 1: + raise ValueError( + f"sample {sample.index} missing dit_trajectory.latents; " + "NFT needs the final clean latent x0 from rollout" + ) + return traj.latents[-1].detach().cpu().float() + + +# TODO: remove and replace sigmas by rollout results +def resolve_nft_sigmas( + sigmas_or_scheduler, + *, + training_timestep_fraction: float = 0.99, +) -> torch.Tensor: + if torch.is_tensor(sigmas_or_scheduler): + ts = sigmas_or_scheduler.detach().float().flatten() + else: + raw = getattr(sigmas_or_scheduler, "sigmas", None) + if raw is None: + raise ValueError("NFT needs scheduler.sigmas (or a sigma tensor)") + ts = raw.detach().float().flatten() + if ts.numel() == 0: + raise ValueError("scheduler.sigmas is empty") + if ts.numel() > 1 and torch.isclose(ts[-1], torch.zeros((), dtype=ts.dtype), atol=1e-8): + ts = ts[:-1] + frac = float(training_timestep_fraction) + if frac < 1.0 and ts.numel() > 1: + keep = max(1, int(ts.numel() * frac)) + ts = ts[:keep] + if ts.numel() == 0: + raise ValueError("No training timesteps left after NFT sigma filtering") + return ts + + +def expand_samples_to_train_pairs( + args: Namespace, + samples: list[Sample], + rewards: list[float], + raw_rewards: list[float], +) -> dict[str, Any]: + """Expand NFT rollout samples into K ``(x0, t)`` train pairs.""" + if not samples: + raise ValueError("NFT convert received empty samples") + if len(samples) != len(rewards) or len(samples) != len(raw_rewards): + raise ValueError( + f"NFT convert length mismatch: samples={len(samples)} " + f"rewards={len(rewards)} raw_rewards={len(raw_rewards)}" + ) + first_traj = samples[0].dit_trajectory + if first_traj is None: + raise ValueError("sample 0 missing dit_trajectory") + if first_traj.timesteps is None: + raise ValueError("NFT needs dit_trajectory.timesteps from rollout") + num_train_timesteps = int(getattr(args, "diffusion_num_train_timesteps", 1000) or 1000) + if first_traj.sigmas is not None: + scheduler_sigmas = first_traj.sigmas.detach().cpu().float() + else: + # Match scheduler_meta_from_rollout when sglang omits sigmas (e.g. ODE rollout). + ts = first_traj.timesteps.detach().cpu().float() + scheduler_sigmas = torch.cat([ts / float(num_train_timesteps), ts.new_zeros(1)]) + scheduler_meta = { + "scheduler_timesteps": first_traj.timesteps.detach().cpu().float(), + "scheduler_sigmas": scheduler_sigmas, + } + frac = float(getattr(args, "diffusion_nft_timestep_fraction", 0.99) or 0.99) + shuffle_t = bool(getattr(args, "diffusion_nft_shuffle_timesteps", True)) + sigmas = resolve_nft_sigmas(scheduler_meta["scheduler_sigmas"], training_timestep_fraction=frac) + num_timesteps = int(sigmas.numel()) + + train_data: list[dict[str, Any]] = [] + for sample, adv, raw in zip(samples, rewards, raw_rewards, strict=True): + if sample.denoising_env is None: + raise ValueError(f"sample {sample.index} missing denoising_env") + x0 = _clean_x0_from_sample(sample) + sample_sigmas = sigmas[torch.randperm(num_timesteps)] if shuffle_t else sigmas + for t in sample_sigmas.tolist(): + train_data.append( + { + "x0": x0, + "timestep": float(t), + "denoising_env": sample.denoising_env, + "advantage": float(adv), + "raw_reward": float(raw), + "sample_index": sample.index, + "prompt": sample.prompt, + "nft_num_timesteps": num_timesteps, + } + ) + return {"train_data": train_data, **scheduler_meta} diff --git a/miles/ray/rollout.py b/miles/ray/rollout.py index e6a1e298..5b035f15 100644 --- a/miles/ray/rollout.py +++ b/miles/ray/rollout.py @@ -13,6 +13,9 @@ from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS from miles.backends.sglang_diffusion_utils.sglang_diffusion_engine import SGLangDiffusionEngine +from miles.ray.data_conversion_hub.flow_grpo import ( + expand_samples_to_train_pairs as flow_grpo_expand_samples_to_train_pairs, +) from miles.rollout.base_types import call_rollout_fn from miles.rollout.rm_hub.core import set_reward_placement_group from miles.utils import tracking_utils @@ -26,11 +29,7 @@ from miles.utils.ray_utils import Box from miles.utils.timer import timer from miles.utils.tracking_utils import init_tracking -from miles.utils.train_data_utils import ( - TrainDataDPSplitter, - reorder_train_pairs_for_tiling, - resolve_train_data_converter, -) +from miles.utils.train_data_utils import TrainDataDPSplitter, reorder_train_pairs_for_tiling from miles.utils.train_metric_utils import log_perf_data_raw from miles.utils.types import Sample @@ -81,9 +80,11 @@ def __init__(self, args, pg): if self.args.custom_convert_samples_to_train_data_path is not None else None ) - # Algorithm-specific converters share the default reward post-process path. - # Full convert overrides use --custom-convert-samples-to-train-data-path instead. - self.train_data_converter = resolve_train_data_converter(self.args) + self.custom_expand_samples_to_train_pairs_func = ( + load_function(self.args.custom_expand_samples_to_train_pairs_path) + if self.args.custom_expand_samples_to_train_pairs_path is not None + else None + ) self.train_data_dp_splitter = TrainDataDPSplitter() logger.info(f"import {self.args.rollout_function_path} as generate_rollout function.") logger.info(f"import {self.args.eval_function_path} as eval_generate_rollout function.") @@ -341,9 +342,27 @@ def _post_process_rewards(self, samples: list[Sample] | list[list[Sample]]): if self.custom_reward_post_process_func is not None: return self.custom_reward_post_process_func(self.args, samples) - from miles.backends.fsdp_utils.loss_hub.advantages import grpo_normalize_rewards + raw_rewards = [sample.get_reward_value(self.args) for sample in samples] + + # --globalize-reward-mean / --globalize-reward-std are orthogonal. + rewards_flat = torch.tensor(raw_rewards, dtype=torch.float) + rewards = rewards_flat.view(-1, self.args.n_samples_per_prompt) + + if self.args.globalize_reward_mean: + mean = rewards_flat.mean() + else: + mean = rewards.mean(dim=-1, keepdim=True) + rewards = rewards - mean + + if self.args.grpo_std_normalization: + if self.args.globalize_reward_std: + std = rewards_flat.std() + else: + std = rewards.std(dim=-1, keepdim=True) + # matches flow_grpo's `+ 1e-4` in both stat_tracking branches + rewards = rewards / (std + 1e-4) - return grpo_normalize_rewards(self.args, samples) + return raw_rewards, rewards.flatten().tolist() def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sample]]): """ @@ -395,7 +414,9 @@ def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sampl reward_key=self.args.reward_key, ) - return self.train_data_converter.convert_samples(samples, rewards, raw_rewards) + if self.custom_expand_samples_to_train_pairs_func is not None: + return self.custom_expand_samples_to_train_pairs_func(self.args, samples, rewards, raw_rewards) + return flow_grpo_expand_samples_to_train_pairs(self.args, samples, rewards, raw_rewards) def _log_images( self, diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index b4f3929f..72748659 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -580,6 +580,16 @@ def add_rollout_arguments(parser): "`def convert_samples_to_train_data(args, samples) -> dict`. " ), ) + parser.add_argument( + "--custom-expand-samples-to-train-pairs-path", + type=str, + default=None, + help=( + "Expand post-processed rollout samples into train pairs. Signature: " + "`def expand_samples_to_train_pairs(args, samples, rewards, raw_rewards) -> dict`. " + "Defaults to the Flow-GRPO implementation." + ), + ) parser.add_argument( "--rollout-sample-filter-path", type=str, @@ -1603,8 +1613,11 @@ def miles_validate_args(args): is_nft = getattr(args, "loss_type", None) in ("nft", "diffusion_nft") if is_nft: - # DiffusionNFT: swap prepare + loss formula; NFT converter is selected in RolloutManager. - # Do not overwrite --custom-convert-samples-to-train-data-path (full-convert override). + # DiffusionNFT selects its rollout-to-train-pair expansion hook. + if getattr(args, "custom_expand_samples_to_train_pairs_path", None) is None: + args.custom_expand_samples_to_train_pairs_path = ( + "miles.ray.data_conversion_hub.nft.expand_samples_to_train_pairs" + ) if getattr(args, "custom_loss_function_path", None) is None: args.custom_loss_function_path = "miles.backends.fsdp_utils.loss_hub.nft.nft_loss_formula" if getattr(args, "custom_prepare_train_batch_path", None) is None: diff --git a/miles/utils/train_data_utils.py b/miles/utils/train_data_utils.py index 6ac8dec7..d4ea726c 100644 --- a/miles/utils/train_data_utils.py +++ b/miles/utils/train_data_utils.py @@ -5,8 +5,6 @@ import torch import torch.distributed as dist -from miles.utils.types import RolloutDebugTensors, Sample - logger = logging.getLogger(__name__) @@ -41,193 +39,6 @@ def scheduler_meta_from_rollout( return timesteps, sigmas -def resolve_train_data_converter(args): - """Pick the algorithm train-data converter (not a full ``custom_convert`` override).""" - if getattr(args, "loss_type", None) in ("nft", "diffusion_nft"): - from miles.backends.fsdp_utils.loss_hub.nft import NftTrainDataConverter - - return NftTrainDataConverter(args) - return RolloutTrainDataConverter() - - -class RolloutTrainDataConverter: - """Convert rollout samples into the flat train-pair payload.""" - - def convert_samples( - self, - samples: list[Sample], - rewards: list[float], - raw_rewards: list[float], - ) -> dict[str, Any]: - train_data, scheduler_meta = self._expand_samples_to_train_pairs(samples, rewards, raw_rewards) - return {"train_data": train_data, **scheduler_meta} - - def _expand_samples_to_train_pairs( - self, - samples: list[Sample], - rewards: list[float], - raw_rewards: list[float], - ) -> tuple[list[dict[str, Any]], dict[str, torch.Tensor]]: - """Flat train pairs in sample-major order (all pairs for sample 0, then sample 1, ...).""" - device = torch.device("cpu") - train_data: list[dict[str, Any]] = [] - first_traj = samples[0].dit_trajectory - # Scheduler meta is taken from sample 0 and returned once for the whole batch; - # the per-sample loop below verifies every sample actually shares it. - scheduler_meta: dict[str, torch.Tensor] = {"scheduler_timesteps": first_traj.timesteps.detach().cpu().float()} - - if first_traj.sigmas is not None: - scheduler_meta["scheduler_sigmas"] = first_traj.sigmas.detach().cpu().float() - - for sample, rew, raw_r in zip(samples, rewards, raw_rewards, strict=True): - traj, denoising_env, rollout_log_probs = self._sample_required_inputs(sample) - # Nail down the shared-scheduler-meta assumption: every sample must carry the - # same timesteps/sigmas as sample 0, since one scheduler_meta is returned for all. - if not torch.equal(traj.timesteps.detach().cpu().float(), scheduler_meta["scheduler_timesteps"]): - raise ValueError( - f"sample {sample.index} has different scheduler_timesteps than sample 0; " - "the converter assumes one shared schedule across the batch" - ) - expected_sigmas = scheduler_meta.get("scheduler_sigmas") - traj_sigmas = None if traj.sigmas is None else traj.sigmas.detach().cpu().float() - if (expected_sigmas is None) != (traj_sigmas is None) or ( - expected_sigmas is not None and not torch.equal(traj_sigmas, expected_sigmas) - ): - raise ValueError( - f"sample {sample.index} has different scheduler_sigmas than sample 0; " - "the converter assumes one shared schedule across the batch" - ) - # build per-sample features for train pairs - per_sample_features = self._build_per_sample_features( - sample, - reward=rew, - raw_reward=raw_r, - denoising_env=denoising_env, - ) - # build per-timestep features for train pairs - per_timestep_features, idx = self._build_per_timestep_features( - sample, - traj=traj, - rollout_log_probs=rollout_log_probs, - device=device, - ) - # build debug tensors for train pairs - pair_debug_steps = None - if sample.rollout_debug_tensors is not None: - pair_debug_steps = self._slice_rollout_debug_for_train_pairs(sample.rollout_debug_tensors, sde_idx=idx) - # validate debug tensors - sample_t_steps = int(per_timestep_features["latent"].shape[0]) - if pair_debug_steps is not None and len(pair_debug_steps) != sample_t_steps: - raise ValueError( - f"rollout_debug_tensors step count {len(pair_debug_steps)} != train pairs {sample_t_steps} " - f"(sample_index={sample.index})" - ) - - for t in range(sample_t_steps): - pair: dict[str, Any] = { - **per_sample_features, - **{key: tensor[t].detach().cpu() for key, tensor in per_timestep_features.items()}, - } - # attach debug tensors to train pair - if pair_debug_steps is not None: - pair["rollout_debug_tensors"] = pair_debug_steps[t] - train_data.append(pair) - - if not train_data: - raise ValueError("No train pairs were produced from rollout samples") - - return train_data, scheduler_meta - - @staticmethod - def _sample_required_inputs(sample: Sample): - traj = sample.dit_trajectory - denoising_env = sample.denoising_env - rollout_log_probs = sample.rollout_log_probs - if traj is None or traj.timesteps is None or denoising_env is None or rollout_log_probs is None: - raise ValueError("Sample missing dit_trajectory, denoising_env, or rollout_log_probs") - return traj, denoising_env, rollout_log_probs - - @staticmethod - def _build_per_sample_features( - sample: Sample, - *, - reward: float, - raw_reward: float, - denoising_env, - ) -> dict[str, Any]: - """Fields shared by every train pair produced from one sample.""" - return { - "advantage": float(reward), - "denoising_env": denoising_env, - "sample_index": sample.index, - "prompt": sample.prompt, - "raw_reward": float(raw_reward), - } - - @staticmethod - def _build_per_timestep_features( - sample: Sample, - *, - traj, - rollout_log_probs: torch.Tensor, - device: torch.device, - ) -> tuple[dict[str, torch.Tensor], torch.Tensor]: - """Fields with one row per selected denoising step.""" - all_latents = traj.latents.to(device, dtype=torch.float32) - latents = all_latents[:-1] - next_latents = all_latents[1:] - timesteps = traj.timesteps.to(device, dtype=torch.float32) - # The step after the last has no recorded timestep -> terminal (σ=0, timestep 0), - # so the SDE step reads σ_next from the actual next rollout timestep, not a lookup. - next_timesteps = torch.cat([timesteps[1:], timesteps.new_zeros(1)]) - - sde_idx = (sample.train_metadata or {}).get("sde_step_indices") - assert sde_idx is not None, "SDE step indices are required for training" - idx = torch.as_tensor(sde_idx, dtype=torch.long) - return { - "latent": latents[idx], - "next_latent": next_latents[idx], - "timestep": timesteps[idx], - # Carry the actual next-step timestep so the SDE step resolves σ_next from a - # rollout value, never assuming train/rollout scheduler positional alignment. - "next_timestep": next_timesteps[idx], - "log_prob_old": rollout_log_probs[idx], - }, idx - - @staticmethod - def _slice_rollout_debug_for_train_pairs( - dbg: RolloutDebugTensors, - *, - sde_idx: torch.Tensor | None = None, - ) -> list[dict[str, torch.Tensor]] | None: - """Slice per-sample rollout debug tensors into one debug payload per train pair.""" - rollout_to_train_pair_fields = { - "rollout_variance_noises": "rollout_step_variance_noise", - "rollout_prev_sample_means": "rollout_step_prev_sample_mean", - "rollout_noise_std_devs": "rollout_step_noise_std_dev", - "rollout_model_outputs": "rollout_step_model_output", - } - indexed: dict[str, torch.Tensor] = {} - for rollout_key, train_pair_key in rollout_to_train_pair_fields.items(): - tensor = getattr(dbg, rollout_key, None) - if tensor is None: - continue - tensor = tensor.detach().cpu() - if sde_idx is not None: - tensor = tensor[sde_idx] - indexed[train_pair_key] = tensor - - if not indexed: - return None - - num_steps = int(next(iter(indexed.values())).shape[0]) - for key, tensor in indexed.items(): - if int(tensor.shape[0]) != num_steps: - raise ValueError(f"Rollout debug field {key!r} has {tensor.shape[0]} steps, expected {num_steps}") - - return [{key: tensor[step] for key, tensor in indexed.items()} for step in range(num_steps)] - - class TrainDataDPSplitter: """Split flat train-pair payloads across DP ranks. 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 958722e3..0f93f92f 100644 --- a/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py +++ b/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py @@ -9,17 +9,9 @@ import torch from miles.backends.fsdp_utils.ema import EmaShadow, resolve_ema_kwargs -from miles.backends.fsdp_utils.loss_hub.advantages import grpo_normalize_rewards -from miles.backends.fsdp_utils.loss_hub.losses import flow_grpo_loss_formula, resolve_loss_formula_fn -from miles.backends.fsdp_utils.loss_hub.nft import ( - NftTrainDataConverter, - corrupt, - nft_loss_formula, - nft_r_from_advantages, - resolve_nft_sigmas, -) -from miles.backends.fsdp_utils.loss_hub.prepare import prepare_flow_grpo_batch, prepare_nft_batch, resolve_prepare_fn +from miles.backends.fsdp_utils.loss_hub.nft import corrupt, nft_loss_formula, nft_r_from_advantages from miles.backends.fsdp_utils.metrics import new_metric_buffer +from miles.ray.data_conversion_hub.nft import expand_samples_to_train_pairs, resolve_nft_sigmas from miles.utils.types import Sample @@ -77,8 +69,9 @@ class _Env: Sample(index=1, prompt="b", reward=3.0, dit_trajectory=_Traj(), denoising_env=_Env()), ] args = _args() - raw_rewards, rewards = grpo_normalize_rewards(args, samples) - out = NftTrainDataConverter(args).convert_samples(samples, rewards, raw_rewards) + raw_rewards = [1.0, 3.0] + rewards = [-1.0, 1.0] + out = expand_samples_to_train_pairs(args, samples, rewards, raw_rewards) assert len(out["train_data"]) == 4 assert {p["timestep"] for p in out["train_data"]} == {1.0, 0.5} assert out["train_data"][0]["x0"] is out["train_data"][1]["x0"] @@ -103,18 +96,6 @@ def test_formula_declares_window_attr(self): assert getattr(nft_loss_formula, "requires_sample_aligned_windows", False) is True assert not hasattr(nft_loss_formula, "ref_mode") - def test_resolve_defaults_are_flow_grpo(self): - assert resolve_prepare_fn(_args()) is prepare_flow_grpo_batch - assert resolve_loss_formula_fn(_args()) is flow_grpo_loss_formula - - def test_resolve_custom_paths(self): - args = _args( - custom_prepare_train_batch_path="miles.backends.fsdp_utils.loss_hub.prepare.prepare_nft_batch", - custom_loss_function_path="miles.backends.fsdp_utils.loss_hub.nft.nft_loss_formula", - ) - assert resolve_prepare_fn(args) is prepare_nft_batch - assert resolve_loss_formula_fn(args) is nft_loss_formula - class TestEmaShadow: def _model(self): diff --git a/tests/fast/utils/test_grouping_parity.py b/tests/fast/utils/test_grouping_parity.py index ea9097f3..9fbdfda2 100644 --- a/tests/fast/utils/test_grouping_parity.py +++ b/tests/fast/utils/test_grouping_parity.py @@ -31,7 +31,8 @@ import torch -from miles.utils.train_data_utils import RolloutTrainDataConverter, TrainDataDPSplitter +from miles.ray.data_conversion_hub.flow_grpo import expand_samples_to_train_pairs +from miles.utils.train_data_utils import TrainDataDPSplitter # -------------------------------------------------------------------------------------- @@ -115,7 +116,7 @@ def test_l2_converter_pairs_match_direct_indexing(): rewards = [0.1, 0.2, 0.3] raw_rewards = [0.4, 0.5, 0.6] - out = RolloutTrainDataConverter().convert_samples(samples, rewards, raw_rewards) + out = expand_samples_to_train_pairs(None, samples, rewards, raw_rewards) pairs = out["train_data"] # count + sample-major ordering + scheduler meta from the first trajectory @@ -156,7 +157,7 @@ def test_l2_converter_pairs_match_direct_indexing(): def test_l2_converter_sigmas_optional(): T, sde = 4, [0, 2] samples = [_mk_sample(i, T, sde, with_sigmas=False, with_debug=False) for i in range(2)] - out = RolloutTrainDataConverter().convert_samples(samples, [1.0, 2.0], [1.0, 2.0]) + out = expand_samples_to_train_pairs(None, samples, [1.0, 2.0], [1.0, 2.0]) assert "scheduler_sigmas" not in out assert len(out["train_data"]) == 2 * len(sde) @@ -167,7 +168,7 @@ def test_l2_converter_rejects_mismatched_scheduler_timesteps(): samples = [_mk_sample(i, 6, [1, 3, 4]) for i in range(3)] samples[2].dit_trajectory.timesteps = samples[2].dit_trajectory.timesteps + 1.0 # tamper try: - RolloutTrainDataConverter().convert_samples(samples, [0.1, 0.2, 0.3], [0.4, 0.5, 0.6]) + expand_samples_to_train_pairs(None, samples, [0.1, 0.2, 0.3], [0.4, 0.5, 0.6]) except ValueError: pass else: @@ -178,7 +179,7 @@ def test_l2_converter_rejects_mismatched_scheduler_sigmas(): samples = [_mk_sample(i, 4, [0, 2]) for i in range(2)] # with_sigmas=True samples[1].dit_trajectory.sigmas = samples[1].dit_trajectory.sigmas + 1.0 # tamper try: - RolloutTrainDataConverter().convert_samples(samples, [1.0, 2.0], [1.0, 2.0]) + expand_samples_to_train_pairs(None, samples, [1.0, 2.0], [1.0, 2.0]) except ValueError: pass else: diff --git a/tests/fast/utils/test_legacy_tile_grouping_golden.py b/tests/fast/utils/test_legacy_tile_grouping_golden.py index 1666d63a..7d7efa07 100644 --- a/tests/fast/utils/test_legacy_tile_grouping_golden.py +++ b/tests/fast/utils/test_legacy_tile_grouping_golden.py @@ -34,7 +34,7 @@ def _refactored_grouping(cfg: dict): - mb_counts: {(rank, optim_step): num_micro_batches} """ # Flat, sample-major train pairs tagged with their real (sample, sde-step) - # identity — exactly what RolloutTrainDataConverter emits, minus the heavy + # identity — exactly what Flow-GRPO train-pair conversion emits, minus the heavy # per-cell tensors (irrelevant to *grouping*). pairs = [ {"sample_index": s, "sde_step": sde} for s in range(cfg["num_samples"]) for sde in cfg["sde_step_indices"] From c0b5596ff8fb9f21a2b4531445ee8293a817e272 Mon Sep 17 00:00:00 2001 From: rockdu Date: Thu, 30 Jul 2026 02:59:36 -0700 Subject: [PATCH 10/12] refactor(loss_hub): colocate algorithm prepare and loss hooks Make Flow-GRPO the explicit actor fallback and keep NFT-specific preparation and loss behavior together without resolver helpers. Co-authored-by: Cursor --- miles/backends/fsdp_utils/actor.py | 41 +++- .../backends/fsdp_utils/loss_hub/__init__.py | 15 +- .../backends/fsdp_utils/loss_hub/flow_grpo.py | 221 ++++++++++++++++++ miles/backends/fsdp_utils/loss_hub/losses.py | 130 ----------- miles/backends/fsdp_utils/loss_hub/nft.py | 176 +++++--------- miles/backends/fsdp_utils/loss_hub/prepare.py | 189 --------------- .../loss_hub/{context.py => types.py} | 2 +- miles/utils/arguments.py | 4 +- 8 files changed, 320 insertions(+), 458 deletions(-) create mode 100644 miles/backends/fsdp_utils/loss_hub/flow_grpo.py delete mode 100644 miles/backends/fsdp_utils/loss_hub/losses.py delete mode 100644 miles/backends/fsdp_utils/loss_hub/prepare.py rename miles/backends/fsdp_utils/loss_hub/{context.py => types.py} (95%) diff --git a/miles/backends/fsdp_utils/actor.py b/miles/backends/fsdp_utils/actor.py index f089d3d2..88f07d31 100644 --- a/miles/backends/fsdp_utils/actor.py +++ b/miles/backends/fsdp_utils/actor.py @@ -35,7 +35,7 @@ DiffusionUpdateWeightFromTensorLoRAIPC, ) from .ema import EmaShadow, resolve_ema_kwargs -from .loss_hub import DiffusionLossContext, resolve_loss_formula_fn, resolve_prepare_fn +from .loss_hub import DiffusionLossContext, flow_grpo_loss_formula, prepare_flow_grpo_batch from .lr_scheduler import get_lr_scheduler from .metrics import new_metric_buffer from .parallel import create_fsdp_parallel_state @@ -172,16 +172,20 @@ def init(self, args: Namespace, role: str, with_ref: bool = False) -> int: # ty else: self.model = torch.nn.ModuleDict(self.models) - from miles.utils.misc import load_function - self.sde_backend = load_function(args.sde_step_backend_path)( self.scheduler, sde_timestep_divisor=self.train_pipeline_config.sde_timestep_divisor, ) # Replaceable parts (defaults = Flow-GRPO). DiT forward stays in this actor. - self.prepare_fn = resolve_prepare_fn(args) - self.loss_formula_fn = resolve_loss_formula_fn(args) + self.custom_prepare_train_batch_func = ( + load_function(args.custom_prepare_train_batch_path) + if args.custom_prepare_train_batch_path is not None + else None + ) + self.custom_loss_formula_func = ( + load_function(args.custom_loss_function_path) if args.custom_loss_function_path is not None else None + ) if args.optimizer == "adam": self.optimizer = torch.optim.AdamW( @@ -355,8 +359,8 @@ def train(self, rollout_id: int, rollout_data_ref) -> None: # type: ignore[over def _train_core(self, rollout_id: int, rollout_data) -> None: """Shared train loop: schedule → prepare → DiT forward → loss formula → backward. - ``prepare_fn`` / ``loss_formula_fn`` are the replaceable parts (defaults = - Flow-GRPO). DiT forward always runs here. + ``prepare_fn`` and the custom loss hook are replaceable; Flow-GRPO is the + default loss. DiT forward always runs here. """ device = torch.cuda.current_device() @@ -406,7 +410,8 @@ def _train_core(self, rollout_id: int, rollout_data) -> None: microbatch_schedule=microbatch_schedule, parallel_state=self.parallel_state, ) - if getattr(self.loss_formula_fn, "requires_sample_aligned_windows", False): + loss_formula_func = self.custom_loss_formula_func or flow_grpo_loss_formula + if getattr(loss_formula_func, "requires_sample_aligned_windows", False): validate_sample_aligned_windows( train_pairs=train_pairs, microbatch_schedule=microbatch_schedule, @@ -514,12 +519,15 @@ def _forward_train_pair_batch( ) -> torch.Tensor | None: """Shared micro-batch path: prepare → DiT forward → [ref] → loss formula. - Only ``prepare_fn`` / ``loss_formula_fn`` are swappable; DiT forward stays here. + Only ``prepare_fn`` and the custom loss hook are swappable; DiT forward stays here. """ if not batch: raise ValueError("_forward_train_pair_batch received empty batch") - prepared = self.prepare_fn(ctx, batch, pad_to_len=pad_to_len) + if self.custom_prepare_train_batch_func is not None: + prepared = self.custom_prepare_train_batch_func(ctx, batch, pad_to_len=pad_to_len) + else: + prepared = prepare_flow_grpo_batch(ctx, batch, pad_to_len=pad_to_len) train_pipeline_config = self.train_pipeline_config forward_dtype = self._forward_dtype @@ -554,7 +562,18 @@ def _compute_noise_pred() -> torch.Tensor: with torch.no_grad(), ref_ctx: ref_pred = _compute_noise_pred().detach() - return self.loss_formula_fn( + if self.custom_loss_formula_func is not None: + return self.custom_loss_formula_func( + ctx, + batch, + prepared, + new_pred=new_pred, + ref_pred=ref_pred, + metrics=metrics, + write_old_log_prob=write_old_log_prob, + old_log_prob_from_new=old_log_prob_from_new, + ) + return flow_grpo_loss_formula( ctx, batch, prepared, diff --git a/miles/backends/fsdp_utils/loss_hub/__init__.py b/miles/backends/fsdp_utils/loss_hub/__init__.py index 9cec641f..581e2e4d 100644 --- a/miles/backends/fsdp_utils/loss_hub/__init__.py +++ b/miles/backends/fsdp_utils/loss_hub/__init__.py @@ -4,23 +4,18 @@ ``--custom-prepare-train-batch-path`` ``--custom-loss-function-path`` (formula only; receives new_pred / ref_pred) ``--custom-reward-post-process-path`` (advantage normalisation) - ``--loss-type nft`` selects ``NftTrainDataConverter`` (not a full convert override) + ``--loss-type nft`` selects the NFT prepare/loss hooks; its rollout converter + lives in ``miles.ray.data_conversion_hub``. """ -from miles.backends.fsdp_utils.loss_hub.advantages import grpo_normalize_rewards -from miles.backends.fsdp_utils.loss_hub.context import DiffusionLossContext, PreparedBatch -from miles.backends.fsdp_utils.loss_hub.losses import flow_grpo_loss_formula, resolve_loss_formula_fn -from miles.backends.fsdp_utils.loss_hub.nft import NftTrainDataConverter -from miles.backends.fsdp_utils.loss_hub.prepare import prepare_flow_grpo_batch, prepare_nft_batch, resolve_prepare_fn +from miles.backends.fsdp_utils.loss_hub.flow_grpo import flow_grpo_loss_formula, prepare_flow_grpo_batch +from miles.backends.fsdp_utils.loss_hub.nft import prepare_nft_batch +from miles.backends.fsdp_utils.loss_hub.types import DiffusionLossContext, PreparedBatch __all__ = [ "DiffusionLossContext", - "NftTrainDataConverter", "PreparedBatch", "flow_grpo_loss_formula", - "grpo_normalize_rewards", "prepare_flow_grpo_batch", "prepare_nft_batch", - "resolve_loss_formula_fn", - "resolve_prepare_fn", ] diff --git a/miles/backends/fsdp_utils/loss_hub/flow_grpo.py b/miles/backends/fsdp_utils/loss_hub/flow_grpo.py new file mode 100644 index 00000000..5fa01c02 --- /dev/null +++ b/miles/backends/fsdp_utils/loss_hub/flow_grpo.py @@ -0,0 +1,221 @@ +"""Flow-GRPO batch preparation and loss formula. + +The actor owns the DiT forward between these two hooks. +""" + +from __future__ import annotations + +import torch + +from miles.backends.fsdp_utils.loss_hub.types import DiffusionLossContext, PreparedBatch +from miles.backends.fsdp_utils.metrics import record_rollout_train_abs_diff +from miles.utils.metric_buffer import MetricBuffer +from miles.utils.train_data_utils import stack_train_pair_rollout_debug + + +def _cast_cond_to_dtype(cond: dict, dtype: torch.dtype) -> dict: + out = {} + for key, value in cond.items(): + if isinstance(value, torch.Tensor) and value.dtype.is_floating_point: + out[key] = value.to(dtype=dtype) + else: + out[key] = value + return out + + +def _stack_pair_field(batch: list[dict], key: str, device: torch.device) -> torch.Tensor: + return torch.stack([pair[key] for pair in batch]).to(device=device, dtype=torch.float32) + + +def prepare_flow_grpo_batch( + ctx: DiffusionLossContext, + batch: list[dict], + *, + pad_to_len: int | None = None, +) -> PreparedBatch: + """Stack SDE-pair fields and build CFG conditioning (guidance from args).""" + args = ctx.args + device = ctx.device + config = ctx.train_pipeline_config + num_train_timesteps = int(ctx.scheduler.config.num_train_timesteps) + bsz = len(batch) + + latents = _stack_pair_field(batch, "latent", device) + next_latents = _stack_pair_field(batch, "next_latent", device) + timesteps = _stack_pair_field(batch, "timestep", device) + next_timesteps = _stack_pair_field(batch, "next_timestep", device) + log_prob_old = _stack_pair_field(batch, "log_prob_old", device) + advantage = torch.tensor( + [float(pair["advantage"]) for pair in batch], + device=device, + dtype=torch.float32, + ) + advantage = torch.clamp(advantage, -args.diffusion_adv_clip_max, args.diffusion_adv_clip_max) + + guidance_scale = args.diffusion_guidance_scale + true_cfg_scale = args.diffusion_true_cfg_scale + cfg_scale = true_cfg_scale if true_cfg_scale is not None else guidance_scale + use_cfg = cfg_scale > 0 + + if len(ctx.models) == 1: + component_name, model = next(iter(ctx.models.items())) + else: + components = {config.component_for_timestep(t, num_train_timesteps) for t in timesteps.tolist()} + if len(components) > 1: + raise ValueError( + f"Micro-batch mixes denoising phases {sorted(components)}; set " + "--micro-batch-size 1 so each forward is phase-pure (one DiT, one CFG scale)." + ) + component_name = components.pop() + model = ctx.models[component_name] + guidance_scale = config.select_guidance_scale( + float(timesteps[0]), + num_train_timesteps, + guidance_scale, + args.diffusion_guidance_scale_2, + ) + + if config.needs_timestep_scaling: + timesteps_for_model = timesteps / float(num_train_timesteps) + else: + timesteps_for_model = timesteps + + pos_list = [config.prepare_cond_kwargs(batch[i]["denoising_env"].pos_cond_kwargs, device) for i in range(bsz)] + neg_list = ( + [config.prepare_cond_kwargs(batch[i]["denoising_env"].neg_cond_kwargs, device) for i in range(bsz)] + if use_cfg + else None + ) + cfg_batching = use_cfg and bool(args.fsdp_cfg_batching) + joint_cond = pos_cond = neg_cond = None + if cfg_batching: + joint_cond = _cast_cond_to_dtype( + config.collate_cond_for_sample_batch(pos_list + neg_list, device, pad_to_len=pad_to_len), + ctx.forward_dtype, + ) + else: + pos_cond = _cast_cond_to_dtype( + config.collate_cond_for_sample_batch(pos_list, device, pad_to_len=pad_to_len), + ctx.forward_dtype, + ) + if use_cfg and neg_list is not None: + neg_cond = _cast_cond_to_dtype( + config.collate_cond_for_sample_batch(neg_list, device, pad_to_len=pad_to_len), + ctx.forward_dtype, + ) + + return PreparedBatch( + latents=latents, + timesteps=timesteps, + timesteps_for_model=timesteps_for_model, + model=model, + component_name=component_name, + guidance_scale=guidance_scale, + use_cfg=use_cfg, + cfg_batching=cfg_batching, + true_cfg_scale=true_cfg_scale if use_cfg else None, + pos_cond=pos_cond, + neg_cond=neg_cond, + joint_cond=joint_cond, + advantage=advantage, + extras={ + "next_latents": next_latents, + "next_timesteps": next_timesteps, + "log_prob_old": log_prob_old, + }, + ) + + +def flow_grpo_loss_formula( + ctx: DiffusionLossContext, + batch: list[dict], + prepared: PreparedBatch, + *, + new_pred: torch.Tensor, + ref_pred: torch.Tensor | None, + metrics: MetricBuffer, + write_old_log_prob: bool = False, + old_log_prob_from_new: bool = False, +) -> torch.Tensor | None: + """SDE log-prob + PPO-clip (+ optional KL vs ``ref_pred``). Actor owns DiT forward.""" + args = ctx.args + clip_range = args.diffusion_clip_range + noise_level = args.diffusion_noise_level + kl_beta = float(args.diffusion_kl_beta) + + next_latents = prepared.extras["next_latents"] + next_timesteps = prepared.extras["next_timesteps"] + log_prob_old_rollout = prepared.extras["log_prob_old"] + + _, log_prob_new, prev_sample_mean_new, std_dev_t_new = ctx.sde_backend.sde_step_logprob( + new_pred.float(), + prepared.timesteps, + next_timesteps, + prepared.latents.float(), + prev_sample=next_latents.float(), + noise_level=noise_level, + ) + + if write_old_log_prob: + for pair, log_prob in zip(batch, log_prob_new, strict=True): + pair["log_prob_old"] = log_prob.cpu() + return None + + log_prob_old = log_prob_new.detach() if old_log_prob_from_new else log_prob_old_rollout + ratio = torch.exp(log_prob_new - log_prob_old) + unclipped = -prepared.advantage * ratio + clipped = -prepared.advantage * torch.clamp(ratio, 1.0 - clip_range, 1.0 + clip_range) + per_pair_loss = torch.maximum(unclipped, clipped) + loss_sum = per_pair_loss.sum() + bsz = len(batch) + + kl_sum = loss_sum.new_zeros(()) + if kl_beta > 0: + if ref_pred is None: + raise ValueError("Flow-GRPO KL requires a reference DiT forward (actor ref_mode=lora_base)") + _, _, prev_sample_mean_ref, _ = ctx.sde_backend.sde_step_logprob( + ref_pred.float(), + prepared.timesteps, + next_timesteps, + prepared.latents.float(), + prev_sample=next_latents.float(), + noise_level=noise_level, + ) + kl_per_pair = ((prev_sample_mean_new - prev_sample_mean_ref) ** 2).mean( + dim=tuple(range(1, prev_sample_mean_new.ndim)), + keepdim=True, + ) / (2 * std_dev_t_new**2) + loss_sum = loss_sum + kl_beta * kl_per_pair.sum() + kl_sum = kl_per_pair.sum() + + with torch.no_grad(): + metrics.emit_mean("loss", total=loss_sum, count=bsz) + metrics.emit_mean("policy_loss", total=per_pair_loss.sum(), count=bsz) + metrics.emit_mean("kl_loss", total=kl_sum, count=bsz) + metrics.emit_mean("loss_abs_mean", total=per_pair_loss.abs().sum(), count=bsz) + metrics.emit_mean("adv_abs_mean", total=prepared.advantage.abs().sum(), count=bsz) + metrics.emit_mean("ratio_abs_minus_1", total=(ratio - 1.0).abs().sum(), count=bsz) + metrics.emit_mean("approx_kl", total=0.5 * ((log_prob_new - log_prob_old) ** 2).sum(), count=bsz) + metrics.emit_mean("clipfrac", total=(torch.abs(ratio - 1.0) > clip_range).float().sum(), count=bsz) + metrics.emit_mean("log_prob_new_idx_0", total=log_prob_new[0], count=1) + metrics.emit_mean("log_prob_old_idx_0", total=log_prob_old[0], count=1) + log_prob_abs_diff_sum = torch.abs(log_prob_new - log_prob_old).sum() + metrics.emit_mean("log_prob_mean_abs_diff", total=log_prob_abs_diff_sum, count=bsz) + if len(ctx.models) > 1: + metrics.emit_mean( + f"log_prob_mean_abs_diff_{prepared.component_name}", + total=log_prob_abs_diff_sum, + count=bsz, + ) + + rollout_model_output = stack_train_pair_rollout_debug(batch, "rollout_step_model_output") + if rollout_model_output is not None: + record_rollout_train_abs_diff( + metrics, + "model_output", + new_pred.float(), + rollout_model_output.to(device=ctx.device, dtype=torch.float32), + component=prepared.component_name if len(ctx.models) > 1 else None, + ) + + return loss_sum diff --git a/miles/backends/fsdp_utils/loss_hub/losses.py b/miles/backends/fsdp_utils/loss_hub/losses.py deleted file mode 100644 index 84875916..00000000 --- a/miles/backends/fsdp_utils/loss_hub/losses.py +++ /dev/null @@ -1,130 +0,0 @@ -"""Default Flow-GRPO loss formula (actor owns DiT forward). - -Custom algorithms swap ``--custom-loss-function-path`` (formula only: receives -``new_pred`` / ``ref_pred``). Batch preparation lives in ``prepare.py``. -""" - -from __future__ import annotations - -from argparse import Namespace -from collections.abc import Callable - -import torch - -from miles.backends.fsdp_utils.loss_hub.context import DiffusionLossContext, PreparedBatch -from miles.backends.fsdp_utils.metrics import record_rollout_train_abs_diff -from miles.utils.metric_buffer import MetricBuffer -from miles.utils.misc import load_function -from miles.utils.train_data_utils import stack_train_pair_rollout_debug - -LossFormulaFn = Callable[..., torch.Tensor | None] - - -def flow_grpo_loss_formula( - ctx: DiffusionLossContext, - batch: list[dict], - prepared: PreparedBatch, - *, - new_pred: torch.Tensor, - ref_pred: torch.Tensor | None, - metrics: MetricBuffer, - write_old_log_prob: bool = False, - old_log_prob_from_new: bool = False, -) -> torch.Tensor | None: - """SDE log-prob + PPO-clip (+ optional KL vs ``ref_pred``). Actor owns DiT forward.""" - args = ctx.args - clip_range = args.diffusion_clip_range - noise_level = args.diffusion_noise_level - kl_beta = float(args.diffusion_kl_beta) - - next_latents = prepared.extras["next_latents"] - next_timesteps = prepared.extras["next_timesteps"] - log_prob_old_rollout = prepared.extras["log_prob_old"] - - _, log_prob_new, prev_sample_mean_new, std_dev_t_new = ctx.sde_backend.sde_step_logprob( - new_pred.float(), - prepared.timesteps, - next_timesteps, - prepared.latents.float(), - prev_sample=next_latents.float(), - noise_level=noise_level, - ) - - if write_old_log_prob: - for pair, log_prob in zip(batch, log_prob_new, strict=True): - pair["log_prob_old"] = log_prob.cpu() - return None - - log_prob_old = log_prob_new.detach() if old_log_prob_from_new else log_prob_old_rollout - ratio = torch.exp(log_prob_new - log_prob_old) - unclipped = -prepared.advantage * ratio - clipped = -prepared.advantage * torch.clamp(ratio, 1.0 - clip_range, 1.0 + clip_range) - per_pair_loss = torch.maximum(unclipped, clipped) - loss_sum = per_pair_loss.sum() - bsz = len(batch) - - kl_sum = loss_sum.new_zeros(()) - if kl_beta > 0: - if ref_pred is None: - raise ValueError("Flow-GRPO KL requires a reference DiT forward (actor ref_mode=lora_base)") - _, _, prev_sample_mean_ref, _ = ctx.sde_backend.sde_step_logprob( - ref_pred.float(), - prepared.timesteps, - next_timesteps, - prepared.latents.float(), - prev_sample=next_latents.float(), - noise_level=noise_level, - ) - kl_per_pair = ((prev_sample_mean_new - prev_sample_mean_ref) ** 2).mean( - dim=tuple(range(1, prev_sample_mean_new.ndim)), - keepdim=True, - ) / (2 * std_dev_t_new**2) - loss_sum = loss_sum + kl_beta * kl_per_pair.sum() - kl_sum = kl_per_pair.sum() - - with torch.no_grad(): - metrics.emit_mean("loss", total=loss_sum, count=bsz) - metrics.emit_mean("policy_loss", total=per_pair_loss.sum(), count=bsz) - metrics.emit_mean("kl_loss", total=kl_sum, count=bsz) - metrics.emit_mean("loss_abs_mean", total=per_pair_loss.abs().sum(), count=bsz) - metrics.emit_mean("adv_abs_mean", total=prepared.advantage.abs().sum(), count=bsz) - metrics.emit_mean("ratio_abs_minus_1", total=(ratio - 1.0).abs().sum(), count=bsz) - metrics.emit_mean("approx_kl", total=0.5 * ((log_prob_new - log_prob_old) ** 2).sum(), count=bsz) - metrics.emit_mean("clipfrac", total=(torch.abs(ratio - 1.0) > clip_range).float().sum(), count=bsz) - metrics.emit_mean("log_prob_new_idx_0", total=log_prob_new[0], count=1) - metrics.emit_mean("log_prob_old_idx_0", total=log_prob_old[0], count=1) - log_prob_abs_diff_sum = torch.abs(log_prob_new - log_prob_old).sum() - metrics.emit_mean("log_prob_mean_abs_diff", total=log_prob_abs_diff_sum, count=bsz) - if len(ctx.models) > 1: - metrics.emit_mean( - f"log_prob_mean_abs_diff_{prepared.component_name}", - total=log_prob_abs_diff_sum, - count=bsz, - ) - - rollout_model_output = stack_train_pair_rollout_debug(batch, "rollout_step_model_output") - if rollout_model_output is not None: - record_rollout_train_abs_diff( - metrics, - "model_output", - new_pred.float(), - rollout_model_output.to(device=ctx.device, dtype=torch.float32), - component=prepared.component_name if len(ctx.models) > 1 else None, - ) - - return loss_sum - - -def resolve_loss_formula_fn(args: Namespace) -> LossFormulaFn: - """Loss *formula* only — DiT forward stays in the actor. - - Custom path defaults (e.g. NFT) are assigned in ``arguments.py``. When the - path is unset, Flow-GRPO is the default implementation. - """ - path = getattr(args, "custom_loss_function_path", None) - if path: - fn = load_function(path) - if fn is None: - raise ValueError(f"Failed to load custom loss formula from {path!r}") - return fn - return flow_grpo_loss_formula diff --git a/miles/backends/fsdp_utils/loss_hub/nft.py b/miles/backends/fsdp_utils/loss_hub/nft.py index ac5ecc9f..377a2b0b 100644 --- a/miles/backends/fsdp_utils/loss_hub/nft.py +++ b/miles/backends/fsdp_utils/loss_hub/nft.py @@ -1,26 +1,21 @@ -"""DiffusionNFT plugin: train-data converter + loss *formula*. +"""DiffusionNFT batch preparation, training math, and loss formula. -Prepare hook lives in ``prepare.py`` (``prepare_nft_batch``). Actor still owns -DiT forward (+ EMA/LoRA-base reference forward via ``--ref-mode``). +The actor owns the DiT forward (+ EMA/LoRA-base reference forward via +``--ref-mode``) between the prepare and loss hooks. -Reward / advantage normalisation stays in ``RolloutManager._post_process_rewards`` -(default: ``grpo_normalize_rewards``). This module only expands samples into -K ``(x0, t)`` train pairs. +Reward / advantage normalisation stays in ``RolloutManager._post_process_rewards``. +Rollout-to-train conversion lives in ``miles.ray.data_conversion_hub.nft``. """ from __future__ import annotations -from argparse import Namespace -from typing import Any - import torch -from miles.backends.fsdp_utils.loss_hub.context import DiffusionLossContext, PreparedBatch +from miles.backends.fsdp_utils.loss_hub.types import DiffusionLossContext, PreparedBatch from miles.utils.metric_buffer import MetricBuffer -from miles.utils.types import Sample # --------------------------------------------------------------------------- -# Forward-process math (used by prepare_nft_batch in prepare.py) +# Forward-process math and batch preparation # --------------------------------------------------------------------------- @@ -35,29 +30,62 @@ def corrupt(x0: torch.Tensor, t: torch.Tensor, eps: torch.Tensor) -> torch.Tenso return (1.0 - t) * x0 + t * eps -def resolve_nft_sigmas( - sigmas_or_scheduler, +def _cast_cond_to_dtype(cond: dict, dtype: torch.dtype) -> dict: + out = {} + for key, value in cond.items(): + if isinstance(value, torch.Tensor) and value.dtype.is_floating_point: + out[key] = value.to(dtype=dtype) + else: + out[key] = value + return out + + +def prepare_nft_batch( + ctx: DiffusionLossContext, + batch: list[dict], *, - training_timestep_fraction: float = 0.99, -) -> torch.Tensor: - if torch.is_tensor(sigmas_or_scheduler): - ts = sigmas_or_scheduler.detach().float().flatten() + pad_to_len: int | None = None, +) -> PreparedBatch: + """Corrupt clean x0 at each pair's sigma; CFG-free cond.""" + if len(ctx.models) != 1: + raise ValueError("DiffusionNFT currently supports a single DiT component (SD3)") + device = ctx.device + config = ctx.train_pipeline_config + bsz = len(batch) + x0 = torch.stack([pair["x0"] for pair in batch]).to(device=device, dtype=torch.float32) + t = torch.tensor([float(pair["timestep"]) for pair in batch], device=device, dtype=torch.float32) + advantage = torch.tensor([float(pair["advantage"]) for pair in batch], device=device, dtype=torch.float32) + + component_name, model = next(iter(ctx.models.items())) + pos_list = [config.prepare_cond_kwargs(batch[i]["denoising_env"].pos_cond_kwargs, device) for i in range(bsz)] + pos_cond = _cast_cond_to_dtype( + config.collate_cond_for_sample_batch(pos_list, device, pad_to_len=pad_to_len), + ctx.forward_dtype, + ) + + num_train_timesteps = int(getattr(ctx.scheduler.config, "num_train_timesteps", 1000)) + if config.needs_timestep_scaling: + timesteps_for_model = t.to(dtype=torch.float32) else: - raw = getattr(sigmas_or_scheduler, "sigmas", None) - if raw is None: - raise ValueError("NFT needs scheduler.sigmas (or a sigma tensor)") - ts = raw.detach().float().flatten() - if ts.numel() == 0: - raise ValueError("scheduler.sigmas is empty") - if ts.numel() > 1 and torch.isclose(ts[-1], torch.zeros((), dtype=ts.dtype), atol=1e-8): - ts = ts[:-1] - frac = float(training_timestep_fraction) - if frac < 1.0 and ts.numel() > 1: - keep = max(1, int(ts.numel() * frac)) - ts = ts[:keep] - if ts.numel() == 0: - raise ValueError("No training timesteps left after NFT sigma filtering") - return ts + timesteps_for_model = t * float(num_train_timesteps) + + xt = corrupt(x0, t, sample_noise(x0)) + return PreparedBatch( + latents=xt, + timesteps=t, + timesteps_for_model=timesteps_for_model, + model=model, + component_name=component_name, + guidance_scale=0.0, + use_cfg=False, + cfg_batching=False, + true_cfg_scale=None, + pos_cond=pos_cond, + neg_cond=None, + joint_cond=None, + advantage=advantage, + extras={"x0": x0}, + ) def nft_r_from_advantages(advantages: torch.Tensor, *, adv_clip_max: float) -> torch.Tensor: @@ -99,88 +127,6 @@ def nft_branch_losses( return pos_loss, neg_loss -# --------------------------------------------------------------------------- -# Converter (K-expanded pairs; rewards already post-processed by rollout) -# --------------------------------------------------------------------------- - - -def _clean_x0_from_sample(sample: Sample) -> torch.Tensor: - traj = sample.dit_trajectory - if traj is None or traj.latents is None or traj.latents.shape[0] < 1: - raise ValueError( - f"sample {sample.index} missing dit_trajectory.latents; " - "NFT needs the final clean latent x0 from rollout" - ) - return traj.latents[-1].detach().cpu().float() - - -class NftTrainDataConverter: - """Expand rollout samples into K ``(x0, t)`` train pairs (sample-major). - - Same interface as ``RolloutTrainDataConverter``: receives already-normalized - rewards from ``_post_process_rewards``. Selected by ``--loss-type nft``. - """ - - def __init__(self, args: Namespace) -> None: - self.args = args - - def convert_samples( - self, - samples: list[Sample], - rewards: list[float], - raw_rewards: list[float], - ) -> dict[str, Any]: - if not samples: - raise ValueError("NFT convert received empty samples") - if len(samples) != len(rewards) or len(samples) != len(raw_rewards): - raise ValueError( - f"NFT convert length mismatch: samples={len(samples)} " - f"rewards={len(rewards)} raw_rewards={len(raw_rewards)}" - ) - first_traj = samples[0].dit_trajectory - if first_traj is None: - raise ValueError("sample 0 missing dit_trajectory") - if first_traj.timesteps is None: - raise ValueError("NFT needs dit_trajectory.timesteps from rollout") - args = self.args - num_train_timesteps = int(getattr(args, "diffusion_num_train_timesteps", 1000) or 1000) - if first_traj.sigmas is not None: - scheduler_sigmas = first_traj.sigmas.detach().cpu().float() - else: - # Match scheduler_meta_from_rollout when sglang omits sigmas (e.g. ODE rollout). - ts = first_traj.timesteps.detach().cpu().float() - scheduler_sigmas = torch.cat([ts / float(num_train_timesteps), ts.new_zeros(1)]) - scheduler_meta = { - "scheduler_timesteps": first_traj.timesteps.detach().cpu().float(), - "scheduler_sigmas": scheduler_sigmas, - } - frac = float(getattr(args, "diffusion_nft_timestep_fraction", 0.99) or 0.99) - shuffle_t = bool(getattr(args, "diffusion_nft_shuffle_timesteps", True)) - sigmas = resolve_nft_sigmas(scheduler_meta["scheduler_sigmas"], training_timestep_fraction=frac) - num_timesteps = int(sigmas.numel()) - - train_data: list[dict[str, Any]] = [] - for sample, adv, raw in zip(samples, rewards, raw_rewards, strict=True): - if sample.denoising_env is None: - raise ValueError(f"sample {sample.index} missing denoising_env") - x0 = _clean_x0_from_sample(sample) - sample_sigmas = sigmas[torch.randperm(num_timesteps)] if shuffle_t else sigmas - for t in sample_sigmas.tolist(): - train_data.append( - { - "x0": x0, - "timestep": float(t), - "denoising_env": sample.denoising_env, - "advantage": float(adv), - "raw_reward": float(raw), - "sample_index": sample.index, - "prompt": sample.prompt, - "nft_num_timesteps": num_timesteps, - } - ) - return {"train_data": train_data, **scheduler_meta} - - # --------------------------------------------------------------------------- # Loss formula (receives actor's new_pred / ref_pred) # --------------------------------------------------------------------------- diff --git a/miles/backends/fsdp_utils/loss_hub/prepare.py b/miles/backends/fsdp_utils/loss_hub/prepare.py deleted file mode 100644 index f6c3fb33..00000000 --- a/miles/backends/fsdp_utils/loss_hub/prepare.py +++ /dev/null @@ -1,189 +0,0 @@ -"""DiT forward input preparation hooks (actor pipeline step before forward). - -Custom algorithms swap ``--custom-prepare-train-batch-path``; loss formulas live -in ``losses.py`` / algorithm plugins (e.g. ``nft.py``). -""" - -from __future__ import annotations - -from argparse import Namespace -from collections.abc import Callable - -import torch - -from miles.backends.fsdp_utils.loss_hub.context import DiffusionLossContext, PreparedBatch -from miles.backends.fsdp_utils.loss_hub.nft import corrupt, sample_noise -from miles.utils.misc import load_function - -PrepareFn = Callable[..., PreparedBatch] - - -def _cast_cond_to_dtype(cond: dict, dtype: torch.dtype) -> dict: - out = {} - for k, v in cond.items(): - if isinstance(v, torch.Tensor) and v.dtype.is_floating_point: - out[k] = v.to(dtype=dtype) - else: - out[k] = v - return out - - -def _stack_pair_field(batch: list[dict], key: str, device: torch.device) -> torch.Tensor: - return torch.stack([pair[key] for pair in batch]).to(device=device, dtype=torch.float32) - - -def prepare_flow_grpo_batch( - ctx: DiffusionLossContext, - batch: list[dict], - *, - pad_to_len: int | None = None, -) -> PreparedBatch: - """Stack SDE-pair fields and build CFG conditioning (guidance from args).""" - args = ctx.args - device = ctx.device - config = ctx.train_pipeline_config - num_train_timesteps = int(ctx.scheduler.config.num_train_timesteps) - bsz = len(batch) - - latents = _stack_pair_field(batch, "latent", device) - next_latents = _stack_pair_field(batch, "next_latent", device) - timesteps = _stack_pair_field(batch, "timestep", device) - next_timesteps = _stack_pair_field(batch, "next_timestep", device) - log_prob_old = _stack_pair_field(batch, "log_prob_old", device) - advantage = torch.tensor( - [float(pair["advantage"]) for pair in batch], - device=device, - dtype=torch.float32, - ) - advantage = torch.clamp(advantage, -args.diffusion_adv_clip_max, args.diffusion_adv_clip_max) - - guidance_scale = args.diffusion_guidance_scale - true_cfg_scale = args.diffusion_true_cfg_scale - cfg_scale = true_cfg_scale if true_cfg_scale is not None else guidance_scale - use_cfg = cfg_scale > 0 - - if len(ctx.models) == 1: - component_name, model = next(iter(ctx.models.items())) - else: - components = {config.component_for_timestep(t, num_train_timesteps) for t in timesteps.tolist()} - if len(components) > 1: - raise ValueError( - f"Micro-batch mixes denoising phases {sorted(components)}; set " - "--micro-batch-size 1 so each forward is phase-pure (one DiT, one CFG scale)." - ) - component_name = components.pop() - model = ctx.models[component_name] - guidance_scale = config.select_guidance_scale( - float(timesteps[0]), - num_train_timesteps, - guidance_scale, - args.diffusion_guidance_scale_2, - ) - - if config.needs_timestep_scaling: - timesteps_for_model = timesteps / float(num_train_timesteps) - else: - timesteps_for_model = timesteps - - pos_list = [config.prepare_cond_kwargs(batch[i]["denoising_env"].pos_cond_kwargs, device) for i in range(bsz)] - neg_list = ( - [config.prepare_cond_kwargs(batch[i]["denoising_env"].neg_cond_kwargs, device) for i in range(bsz)] - if use_cfg - else None - ) - cfg_batching = use_cfg and bool(args.fsdp_cfg_batching) - joint_cond = pos_cond = neg_cond = None - if cfg_batching: - joint_cond = _cast_cond_to_dtype( - config.collate_cond_for_sample_batch(pos_list + neg_list, device, pad_to_len=pad_to_len), - ctx.forward_dtype, - ) - else: - pos_cond = _cast_cond_to_dtype( - config.collate_cond_for_sample_batch(pos_list, device, pad_to_len=pad_to_len), - ctx.forward_dtype, - ) - if use_cfg and neg_list is not None: - neg_cond = _cast_cond_to_dtype( - config.collate_cond_for_sample_batch(neg_list, device, pad_to_len=pad_to_len), - ctx.forward_dtype, - ) - - return PreparedBatch( - latents=latents, - timesteps=timesteps, - timesteps_for_model=timesteps_for_model, - model=model, - component_name=component_name, - guidance_scale=guidance_scale, - use_cfg=use_cfg, - cfg_batching=cfg_batching, - true_cfg_scale=true_cfg_scale if use_cfg else None, - pos_cond=pos_cond, - neg_cond=neg_cond, - joint_cond=joint_cond, - advantage=advantage, - extras={ - "next_latents": next_latents, - "next_timesteps": next_timesteps, - "log_prob_old": log_prob_old, - }, - ) - - -def prepare_nft_batch( - ctx: DiffusionLossContext, - batch: list[dict], - *, - pad_to_len: int | None = None, -) -> PreparedBatch: - """Corrupt clean x0 at each pair's sigma; CFG-free cond.""" - if len(ctx.models) != 1: - raise ValueError("DiffusionNFT currently supports a single DiT component (SD3)") - device = ctx.device - config = ctx.train_pipeline_config - bsz = len(batch) - x0 = torch.stack([pair["x0"] for pair in batch]).to(device=device, dtype=torch.float32) - t = torch.tensor([float(pair["timestep"]) for pair in batch], device=device, dtype=torch.float32) - advantage = torch.tensor([float(pair["advantage"]) for pair in batch], device=device, dtype=torch.float32) - - component_name, model = next(iter(ctx.models.items())) - pos_list = [config.prepare_cond_kwargs(batch[i]["denoising_env"].pos_cond_kwargs, device) for i in range(bsz)] - pos_cond = _cast_cond_to_dtype( - config.collate_cond_for_sample_batch(pos_list, device, pad_to_len=pad_to_len), - ctx.forward_dtype, - ) - - num_train_timesteps = int(getattr(ctx.scheduler.config, "num_train_timesteps", 1000)) - if config.needs_timestep_scaling: - timesteps_for_model = t.to(dtype=torch.float32) - else: - timesteps_for_model = t * float(num_train_timesteps) - - xt = corrupt(x0, t, sample_noise(x0)) - return PreparedBatch( - latents=xt, - timesteps=t, - timesteps_for_model=timesteps_for_model, - model=model, - component_name=component_name, - guidance_scale=0.0, - use_cfg=False, - cfg_batching=False, - true_cfg_scale=None, - pos_cond=pos_cond, - neg_cond=None, - joint_cond=None, - advantage=advantage, - extras={"x0": x0}, - ) - - -def resolve_prepare_fn(args: Namespace) -> PrepareFn: - path = getattr(args, "custom_prepare_train_batch_path", None) - if path: - fn = load_function(path) - if fn is None: - raise ValueError(f"Failed to load custom prepare from {path!r}") - return fn - return prepare_flow_grpo_batch diff --git a/miles/backends/fsdp_utils/loss_hub/context.py b/miles/backends/fsdp_utils/loss_hub/types.py similarity index 95% rename from miles/backends/fsdp_utils/loss_hub/context.py rename to miles/backends/fsdp_utils/loss_hub/types.py index 6c240186..61bdd822 100644 --- a/miles/backends/fsdp_utils/loss_hub/context.py +++ b/miles/backends/fsdp_utils/loss_hub/types.py @@ -1,4 +1,4 @@ -"""Shared handles passed into diffusion prepare / loss-formula hooks.""" +"""Shared data types passed into diffusion prepare and loss hooks.""" from __future__ import annotations diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index 72748659..168dde60 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -1613,7 +1613,7 @@ def miles_validate_args(args): is_nft = getattr(args, "loss_type", None) in ("nft", "diffusion_nft") if is_nft: - # DiffusionNFT selects its rollout-to-train-pair expansion hook. + # DiffusionNFT swaps the train-pair conversion, preparation, and loss formula. if getattr(args, "custom_expand_samples_to_train_pairs_path", None) is None: args.custom_expand_samples_to_train_pairs_path = ( "miles.ray.data_conversion_hub.nft.expand_samples_to_train_pairs" @@ -1621,7 +1621,7 @@ def miles_validate_args(args): if getattr(args, "custom_loss_function_path", None) is None: args.custom_loss_function_path = "miles.backends.fsdp_utils.loss_hub.nft.nft_loss_formula" if getattr(args, "custom_prepare_train_batch_path", None) is None: - args.custom_prepare_train_batch_path = "miles.backends.fsdp_utils.loss_hub.prepare.prepare_nft_batch" + args.custom_prepare_train_batch_path = "miles.backends.fsdp_utils.loss_hub.nft.prepare_nft_batch" if ( getattr(args, "diffusion_sde_type", "sde") == "sde" and float(getattr(args, "diffusion_noise_level", 0.7) or 0.0) == 0.0 From 153b243049bbdea4b5a07fd1c8c3589db1916823 Mon Sep 17 00:00:00 2001 From: rockdu Date: Thu, 30 Jul 2026 03:17:12 -0700 Subject: [PATCH 11/12] refactor(diffusion): simplify NFT configuration and hooks Co-authored-by: Cursor --- miles/backends/fsdp_utils/actor.py | 40 ++---- miles/backends/fsdp_utils/ema.py | 56 +------- .../backends/fsdp_utils/loss_hub/__init__.py | 10 +- .../backends/fsdp_utils/loss_hub/flow_grpo.py | 5 +- miles/backends/fsdp_utils/loss_hub/nft.py | 36 ++--- miles/backends/fsdp_utils/loss_hub/types.py | 15 +- miles/ray/data_conversion_hub/flow_grpo.py | 7 +- miles/ray/data_conversion_hub/nft.py | 28 ++-- miles/ray/rollout.py | 5 +- miles/utils/arguments.py | 133 +++++++----------- scripts/run-diffusion-nft-sd3-pickscore.sh | 19 +-- .../backends/fsdp_utils/test_loss_hub_nft.py | 18 +-- 12 files changed, 93 insertions(+), 279 deletions(-) diff --git a/miles/backends/fsdp_utils/actor.py b/miles/backends/fsdp_utils/actor.py index 88f07d31..b179f9ab 100644 --- a/miles/backends/fsdp_utils/actor.py +++ b/miles/backends/fsdp_utils/actor.py @@ -34,7 +34,7 @@ DiffusionUpdateWeightFromTensorLoRA, DiffusionUpdateWeightFromTensorLoRAIPC, ) -from .ema import EmaShadow, resolve_ema_kwargs +from .ema import EmaShadow from .loss_hub import DiffusionLossContext, flow_grpo_loss_formula, prepare_flow_grpo_batch from .lr_scheduler import get_lr_scheduler from .metrics import new_metric_buffer @@ -177,7 +177,6 @@ def init(self, args: Namespace, role: str, with_ref: bool = False) -> int: # ty sde_timestep_divisor=self.train_pipeline_config.sde_timestep_divisor, ) - # Replaceable parts (defaults = Flow-GRPO). DiT forward stays in this actor. self.custom_prepare_train_batch_func = ( load_function(args.custom_prepare_train_batch_path) if args.custom_prepare_train_batch_path is not None @@ -213,13 +212,14 @@ def init(self, args: Namespace, role: str, with_ref: bool = False) -> int: # ty checkpoint_payload = checkpoint.load(self) - # Optional EMA shadow of trainable params (pi_old). Enabled via --ema-shadow. - # Consumed when --ref-mode ema runs the no-grad reference DiT forward. self.ema_shadow = None - if getattr(self.args, "ema_shadow", False): + if self.args.ema_shadow: self.ema_shadow = EmaShadow( (p for m in self.models.values() for p in m.parameters()), - **resolve_ema_kwargs(self.args), + decay=self.args.ema_decay, + uprate=self.args.ema_uprate, + uphold=self.args.ema_uphold, + flat_steps=self.args.ema_flat_steps, ) # sglang-d now supports /update_weights_from_tensor (PR #20464). @@ -312,7 +312,7 @@ def update_weights(self) -> None: # type: ignore[override] delta = self.ema_shadow.update() if dist.get_rank() == 0: logger.info("EMA shadow updated (decay=%.4f step=%d)", delta, self.ema_shadow.step) - if getattr(self.args, "ema_rollout_policy", "live") == "ema": + if self.args.ema_rollout_policy == "ema": with self.ema_shadow.swap_in(): self.weight_updater.update_weights() else: @@ -357,11 +357,7 @@ def train(self, rollout_id: int, rollout_data_ref) -> None: # type: ignore[over ) def _train_core(self, rollout_id: int, rollout_data) -> None: - """Shared train loop: schedule → prepare → DiT forward → loss formula → backward. - - ``prepare_fn`` and the custom loss hook are replaceable; Flow-GRPO is the - default loss. DiT forward always runs here. - """ + """Run the shared diffusion training loop.""" device = torch.cuda.current_device() train_pairs: list = rollout_data["train_data"] @@ -371,14 +367,11 @@ def _train_core(self, rollout_id: int, rollout_data) -> None: num_pairs = len(train_pairs) num_train_timesteps = self.scheduler.config.num_train_timesteps - # ------------- Reference forward precondition (--ref-mode resolved in arguments.py) ------------- - ref_mode = getattr(self.args, "ref_mode", "none") + ref_mode = self.args.ref_mode if ref_mode == "lora_base" and not all(hasattr(m, "disable_adapter") for m in self.models.values()): raise RuntimeError( "--ref-mode lora_base requires PEFT models exposing disable_adapter() after FSDP wrapping." ) - if ref_mode == "ema" and self.ema_shadow is None: - raise RuntimeError("--ref-mode ema requires a constructed EMA shadow") # ------------- Rollout Scheduler Metadata ------------- scheduler_timesteps, scheduler_sigmas = scheduler_meta_from_rollout( @@ -419,14 +412,12 @@ def _train_core(self, rollout_id: int, rollout_data) -> None: loss_ctx = DiffusionLossContext( models=self.models, - model=self.model, train_pipeline_config=self.train_pipeline_config, sde_backend=self.sde_backend, scheduler=self.scheduler, args=self.args, forward_dtype=self._forward_dtype, device=device, - ema_shadow=self.ema_shadow, ) # ------------- Recompute old log-probs (impl-consistent PPO ratio) ------------- @@ -517,10 +508,7 @@ def _forward_train_pair_batch( write_old_log_prob: bool = False, old_log_prob_from_new: bool = False, ) -> torch.Tensor | None: - """Shared micro-batch path: prepare → DiT forward → [ref] → loss formula. - - Only ``prepare_fn`` and the custom loss hook are swappable; DiT forward stays here. - """ + """Run one prepared diffusion micro-batch.""" if not batch: raise ValueError("_forward_train_pair_batch received empty batch") @@ -551,14 +539,12 @@ def _compute_noise_pred() -> torch.Tensor: new_pred = _compute_noise_pred() ref_pred = None - ref_mode = getattr(self.args, "ref_mode", "none") - if ref_mode not in (None, "none"): + ref_mode = self.args.ref_mode + if ref_mode != "none": if ref_mode == "ema": ref_ctx = self.ema_shadow.swap_in() - elif ref_mode == "lora_base": - ref_ctx = prepared.model.disable_adapter() else: - raise ValueError(f"unknown ref_mode {ref_mode!r}") + ref_ctx = prepared.model.disable_adapter() with torch.no_grad(), ref_ctx: ref_pred = _compute_noise_pred().detach() diff --git a/miles/backends/fsdp_utils/ema.py b/miles/backends/fsdp_utils/ema.py index 7977533f..9120ef3b 100644 --- a/miles/backends/fsdp_utils/ema.py +++ b/miles/backends/fsdp_utils/ema.py @@ -1,69 +1,21 @@ -"""EMA shadow of trainable parameters for diffusion FSDP training. - -Algorithms that need a slow-moving reference / sampling policy ``pi_old`` share -``EmaShadow``: tracks all ``requires_grad`` parameters (LoRA adapters or full -finetune), plus ``swap_in()`` for temporary in-place weight exchange. - -Lifecycle (actor / weight sync):: - - ema.update() - with ema.swap_in(): - weight_updater.update_weights() - -Loss-side reference forward (via ``DiffusionLossContext.ema_shadow``):: - - with torch.no_grad(), ctx.ema_shadow.swap_in(): - old_pred = forward(...) - -Works with FSDP2 DTensor shards (per-rank local swap) and colocate CPU offload. - -Checkpointing (intentionally not wired yet) -------------------------------------------- -``shadow`` / ``step`` are **not** saved or restored by ``fsdp_utils.checkpoint``. -On resume the actor rebuilds EMA from the loaded trainable weights, so ``pi_old`` -cold-starts (decay schedule restarts at step 0). Fine for single-shot runs; -wrong for mid-run resume that must match UniRL's slow ``pi_old``. - -Wiring it later is non-trivial: buffers are per-rank plain clones (not in the -FSDP/DCP model state), must stay aligned with the trainable-param order, and -must not be saved while ``swap_in()`` is active. Prefer a side file such as -``iter_*/ema.pt`` over stuffing into the DCP model dict. -""" +"""EMA shadow of trainable parameters for diffusion FSDP training.""" from __future__ import annotations -from argparse import Namespace from collections.abc import Iterable from contextlib import contextmanager import torch import torch.nn as nn +from torch.distributed.tensor import DTensor def _local(t: torch.Tensor) -> torch.Tensor: - """Local shard of a (possibly DTensor) tensor; EMA/swap is per-rank, no comm.""" - return t._local_tensor if hasattr(t, "_local_tensor") else t - - -def resolve_ema_kwargs(args: Namespace) -> dict[str, float | int]: - """Read normalized ``ema_*`` fields from ``args`` (see ``miles_validate_args``). - - Enablement (``args.ema_shadow``) and rollout policy (``args.ema_rollout_policy``) - are plain args — inferred/validated in ``arguments.py``, not re-wrapped here. - """ - return { - "decay": float(getattr(args, "ema_decay", 0.001)), - "uprate": float(getattr(args, "ema_uprate", 0.001)), - "uphold": float(getattr(args, "ema_uphold", 0.5)), - "flat_steps": int(getattr(args, "ema_flat_steps", 0)), - } + return t.to_local() if isinstance(t, DTensor) else t class EmaShadow: - """EMA shadow of trainable parameters (LoRA or full finetune). - - Not part of the FSDP checkpoint payload today (see module docstring). - """ + """EMA shadow of trainable parameters.""" def __init__( self, diff --git a/miles/backends/fsdp_utils/loss_hub/__init__.py b/miles/backends/fsdp_utils/loss_hub/__init__.py index 581e2e4d..62b78ca9 100644 --- a/miles/backends/fsdp_utils/loss_hub/__init__.py +++ b/miles/backends/fsdp_utils/loss_hub/__init__.py @@ -1,12 +1,4 @@ -"""Diffusion train hooks: prepare + loss formula (actor owns DiT forward). - -Swap pieces via: - ``--custom-prepare-train-batch-path`` - ``--custom-loss-function-path`` (formula only; receives new_pred / ref_pred) - ``--custom-reward-post-process-path`` (advantage normalisation) - ``--loss-type nft`` selects the NFT prepare/loss hooks; its rollout converter - lives in ``miles.ray.data_conversion_hub``. -""" +"""Diffusion batch preparation and loss hooks.""" from miles.backends.fsdp_utils.loss_hub.flow_grpo import flow_grpo_loss_formula, prepare_flow_grpo_batch from miles.backends.fsdp_utils.loss_hub.nft import prepare_nft_batch diff --git a/miles/backends/fsdp_utils/loss_hub/flow_grpo.py b/miles/backends/fsdp_utils/loss_hub/flow_grpo.py index 5fa01c02..979fc778 100644 --- a/miles/backends/fsdp_utils/loss_hub/flow_grpo.py +++ b/miles/backends/fsdp_utils/loss_hub/flow_grpo.py @@ -1,7 +1,4 @@ -"""Flow-GRPO batch preparation and loss formula. - -The actor owns the DiT forward between these two hooks. -""" +"""Flow-GRPO batch preparation and loss formula.""" from __future__ import annotations diff --git a/miles/backends/fsdp_utils/loss_hub/nft.py b/miles/backends/fsdp_utils/loss_hub/nft.py index 377a2b0b..b388c4dc 100644 --- a/miles/backends/fsdp_utils/loss_hub/nft.py +++ b/miles/backends/fsdp_utils/loss_hub/nft.py @@ -1,11 +1,4 @@ -"""DiffusionNFT batch preparation, training math, and loss formula. - -The actor owns the DiT forward (+ EMA/LoRA-base reference forward via -``--ref-mode``) between the prepare and loss hooks. - -Reward / advantage normalisation stays in ``RolloutManager._post_process_rewards``. -Rollout-to-train conversion lives in ``miles.ray.data_conversion_hub.nft``. -""" +"""DiffusionNFT batch preparation and loss formula.""" from __future__ import annotations @@ -14,10 +7,6 @@ from miles.backends.fsdp_utils.loss_hub.types import DiffusionLossContext, PreparedBatch from miles.utils.metric_buffer import MetricBuffer -# --------------------------------------------------------------------------- -# Forward-process math and batch preparation -# --------------------------------------------------------------------------- - def sample_noise(like: torch.Tensor, *, generator: torch.Generator | None = None) -> torch.Tensor: return torch.randn(like.shape, device=like.device, dtype=like.dtype, generator=generator) @@ -63,7 +52,7 @@ def prepare_nft_batch( ctx.forward_dtype, ) - num_train_timesteps = int(getattr(ctx.scheduler.config, "num_train_timesteps", 1000)) + num_train_timesteps = ctx.scheduler.config.num_train_timesteps if config.needs_timestep_scaling: timesteps_for_model = t.to(dtype=torch.float32) else: @@ -127,11 +116,6 @@ def nft_branch_losses( return pos_loss, neg_loss -# --------------------------------------------------------------------------- -# Loss formula (receives actor's new_pred / ref_pred) -# --------------------------------------------------------------------------- - - def nft_loss_formula( ctx: DiffusionLossContext, batch: list[dict], @@ -152,11 +136,9 @@ def nft_loss_formula( raise ValueError("NFT loss formula requires a reference prediction from the actor") args = ctx.args - beta = float(getattr(args, "diffusion_nft_beta", 1.0) or 1.0) - if beta <= 0: - raise ValueError(f"--diffusion-nft-beta must be > 0, got {beta}") - adv_clip_max = float(getattr(args, "diffusion_nft_adv_clip_max", 5.0) or 5.0) - use_adaptive = bool(getattr(args, "diffusion_nft_adaptive_weight", True)) + beta = args.diffusion_nft_beta + adv_clip_max = args.diffusion_nft_adv_clip_max + use_adaptive = args.diffusion_nft_adaptive_weight x0 = prepared.extras["x0"] t = prepared.timesteps @@ -176,11 +158,11 @@ def nft_loss_formula( loss_sum = per_pair.sum() with torch.no_grad(): - num_timesteps = int(batch[0].get("nft_num_timesteps", 0) or 0) + num_timesteps = batch[0]["nft_num_timesteps"] per_pair_total = per_pair.sum() bsz = len(batch) - metrics.emit_mean("loss", total=per_pair_total * float(max(num_timesteps, 1)), count=bsz) - metrics.emit_mean("nft_loss", total=per_pair_total * float(max(num_timesteps, 1)), count=bsz) + metrics.emit_mean("loss", total=per_pair_total * num_timesteps, count=bsz) + metrics.emit_mean("nft_loss", total=per_pair_total * num_timesteps, count=bsz) metrics.emit_mean("nft_loss_per_pair", total=per_pair_total, count=bsz) metrics.emit_mean("nft_r_mean", total=r.sum(), count=bsz) metrics.emit_mean("nft_pos_loss", total=pos_loss.sum(), count=bsz) @@ -197,6 +179,4 @@ def nft_loss_formula( return loss_sum -# Same-sample K pairs must stay in one optimizer window. -# Reference forward is selected via --ref-mode (auto-filled for --loss-type nft). nft_loss_formula.requires_sample_aligned_windows = True diff --git a/miles/backends/fsdp_utils/loss_hub/types.py b/miles/backends/fsdp_utils/loss_hub/types.py index 61bdd822..33d4f767 100644 --- a/miles/backends/fsdp_utils/loss_hub/types.py +++ b/miles/backends/fsdp_utils/loss_hub/types.py @@ -12,31 +12,20 @@ @dataclass class DiffusionLossContext: - """Train-side handles for prepare / loss-formula callables. - - Owned by the FSDP actor; kept free of Ray / optim internals so hooks stay - unit-testable and swappable via ``--custom-*-path``. - """ + """Train-side handles for prepare and loss hooks.""" models: dict[str, torch.nn.Module] - model: torch.nn.Module train_pipeline_config: Any sde_backend: Any scheduler: Any args: Namespace forward_dtype: torch.dtype device: torch.device - # Optional EMA shadow handle; owned by actor (see ``fsdp_utils.ema``). - ema_shadow: Any = None @dataclass class PreparedBatch: - """Actor-owned DiT forward inputs produced by a prepare hook. - - ``extras`` carries algorithm-specific tensors for the loss formula - (e.g. ``next_latents`` / ``log_prob_old`` for Flow-GRPO, ``x0`` for NFT). - """ + """Actor-owned DiT forward inputs produced by a prepare hook.""" latents: torch.Tensor timesteps: torch.Tensor diff --git a/miles/ray/data_conversion_hub/flow_grpo.py b/miles/ray/data_conversion_hub/flow_grpo.py index 00dbbbfd..b7bfabe7 100644 --- a/miles/ray/data_conversion_hub/flow_grpo.py +++ b/miles/ray/data_conversion_hub/flow_grpo.py @@ -26,8 +26,6 @@ def _expand_samples_to_train_pairs( device = torch.device("cpu") train_data: list[dict[str, Any]] = [] first_traj = samples[0].dit_trajectory - # Scheduler meta is taken from sample 0 and returned once for the whole batch; - # the per-sample loop below verifies every sample actually shares it. scheduler_meta: dict[str, torch.Tensor] = {"scheduler_timesteps": first_traj.timesteps.detach().cpu().float()} if first_traj.sigmas is not None: @@ -35,8 +33,6 @@ def _expand_samples_to_train_pairs( for sample, rew, raw_r in zip(samples, rewards, raw_rewards, strict=True): traj, denoising_env, rollout_log_probs = _sample_required_inputs(sample) - # Nail down the shared-scheduler-meta assumption: every sample must carry the - # same timesteps/sigmas as sample 0, since one scheduler_meta is returned for all. if not torch.equal(traj.timesteps.detach().cpu().float(), scheduler_meta["scheduler_timesteps"]): raise ValueError( f"sample {sample.index} has different scheduler_timesteps than sample 0; " @@ -126,8 +122,7 @@ def _build_per_timestep_features( latents = all_latents[:-1] next_latents = all_latents[1:] timesteps = traj.timesteps.to(device, dtype=torch.float32) - # The step after the last has no recorded timestep -> terminal (σ=0, timestep 0), - # so the SDE step reads σ_next from the actual next rollout timestep, not a lookup. + # The terminal next timestep is zero. next_timesteps = torch.cat([timesteps[1:], timesteps.new_zeros(1)]) sde_idx = (sample.train_metadata or {}).get("sde_step_indices") diff --git a/miles/ray/data_conversion_hub/nft.py b/miles/ray/data_conversion_hub/nft.py index 54208703..1a6e4af8 100644 --- a/miles/ray/data_conversion_hub/nft.py +++ b/miles/ray/data_conversion_hub/nft.py @@ -7,10 +7,6 @@ from miles.utils.types import Sample -# --------------------------------------------------------------------------- -# Converter (K-expanded pairs; rewards already post-processed by rollout) -# --------------------------------------------------------------------------- - def _clean_x0_from_sample(sample: Sample) -> torch.Tensor: traj = sample.dit_trajectory @@ -22,19 +18,12 @@ def _clean_x0_from_sample(sample: Sample) -> torch.Tensor: return traj.latents[-1].detach().cpu().float() -# TODO: remove and replace sigmas by rollout results def resolve_nft_sigmas( - sigmas_or_scheduler, + sigmas: torch.Tensor, *, training_timestep_fraction: float = 0.99, ) -> torch.Tensor: - if torch.is_tensor(sigmas_or_scheduler): - ts = sigmas_or_scheduler.detach().float().flatten() - else: - raw = getattr(sigmas_or_scheduler, "sigmas", None) - if raw is None: - raise ValueError("NFT needs scheduler.sigmas (or a sigma tensor)") - ts = raw.detach().float().flatten() + ts = sigmas.detach().float().flatten() if ts.numel() == 0: raise ValueError("scheduler.sigmas is empty") if ts.numel() > 1 and torch.isclose(ts[-1], torch.zeros((), dtype=ts.dtype), atol=1e-8): @@ -67,20 +56,19 @@ def expand_samples_to_train_pairs( raise ValueError("sample 0 missing dit_trajectory") if first_traj.timesteps is None: raise ValueError("NFT needs dit_trajectory.timesteps from rollout") - num_train_timesteps = int(getattr(args, "diffusion_num_train_timesteps", 1000) or 1000) if first_traj.sigmas is not None: scheduler_sigmas = first_traj.sigmas.detach().cpu().float() else: - # Match scheduler_meta_from_rollout when sglang omits sigmas (e.g. ODE rollout). ts = first_traj.timesteps.detach().cpu().float() - scheduler_sigmas = torch.cat([ts / float(num_train_timesteps), ts.new_zeros(1)]) + scheduler_sigmas = torch.cat([ts / 1000.0, ts.new_zeros(1)]) scheduler_meta = { "scheduler_timesteps": first_traj.timesteps.detach().cpu().float(), "scheduler_sigmas": scheduler_sigmas, } - frac = float(getattr(args, "diffusion_nft_timestep_fraction", 0.99) or 0.99) - shuffle_t = bool(getattr(args, "diffusion_nft_shuffle_timesteps", True)) - sigmas = resolve_nft_sigmas(scheduler_meta["scheduler_sigmas"], training_timestep_fraction=frac) + sigmas = resolve_nft_sigmas( + scheduler_meta["scheduler_sigmas"], + training_timestep_fraction=args.diffusion_nft_timestep_fraction, + ) num_timesteps = int(sigmas.numel()) train_data: list[dict[str, Any]] = [] @@ -88,7 +76,7 @@ def expand_samples_to_train_pairs( if sample.denoising_env is None: raise ValueError(f"sample {sample.index} missing denoising_env") x0 = _clean_x0_from_sample(sample) - sample_sigmas = sigmas[torch.randperm(num_timesteps)] if shuffle_t else sigmas + sample_sigmas = sigmas[torch.randperm(num_timesteps)] if args.diffusion_nft_shuffle_timesteps else sigmas for t in sample_sigmas.tolist(): train_data.append( { diff --git a/miles/ray/rollout.py b/miles/ray/rollout.py index 5b035f15..2b49513a 100644 --- a/miles/ray/rollout.py +++ b/miles/ray/rollout.py @@ -368,8 +368,7 @@ def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sampl """ Convert inference generated samples to training data. """ - # Full override (Miles-LLM style): covers this entire method, including - # reward post-process. Algorithm converters (e.g. NFT) use the path below. + # The full override also replaces reward post-processing. if self.custom_convert_samples_to_train_data_func is not None: return self.custom_convert_samples_to_train_data_func(self.args, samples) @@ -381,8 +380,6 @@ def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sampl raw_t = torch.tensor(raw_rewards, dtype=torch.float) norm_t = torch.tensor(rewards, dtype=torch.float) - # Emit reward distribution stats (raw + normalized) to stdout + wandb. - # Runs for both default SDE-pair expand and NFT converter. reward_stats = { **_reward_stats_dict(raw_t, "rollout/reward/raw_"), **_reward_stats_dict(norm_t, "rollout/reward/norm_"), diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index 168dde60..518446c3 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -108,16 +108,6 @@ def add_train_arguments(parser): default="fsdp", help="The backend for training.", ) - # Diffusion GRPO training (DiffusionFSDPTrainRayActor). - # - # Rollout produces ``train_data``: a flat list of train-pair dicts (one - # dict = one sample × one SDE step), sample-major order before DP split. - # DP ranks receive contiguous pair ranges; each optimizer window covers - # a contiguous slice of pairs. - # - # ``--micro-batch-size`` is the number of train-pair dicts per DiT - # forward/backward (contiguous within the window). Gradients match - # mean loss over all train pairs in the window. parser.add_argument( "--diffusion-clip-range", type=float, @@ -880,7 +870,7 @@ def add_algo_arguments(parser): parser.add_argument( "--loss-type", type=str, - choices=["policy_loss", "sft_loss", "custom_loss", "nft", "diffusion_nft"], + choices=["policy_loss", "sft_loss", "custom_loss", "nft"], default="policy_loss", help=( "Train objective shortcut. Diffusion default policy_loss = Flow-GRPO prepare+formula. " @@ -920,16 +910,11 @@ def add_algo_arguments(parser): default=5.0, help="DiffusionNFT advantage clip before remap to r in [0, 1].", ) - parser.add_argument( - "--diffusion-nft-adaptive-weight", - action="store_true", - default=True, - help="DiffusionNFT adaptive MSE weight (default on; UniRL use_adaptive_weight).", - ) parser.add_argument( "--no-diffusion-nft-adaptive-weight", action="store_false", dest="diffusion_nft_adaptive_weight", + default=True, help="Disable DiffusionNFT adaptive MSE weight.", ) parser.add_argument( @@ -938,28 +923,13 @@ def add_algo_arguments(parser): default=0.99, help="Fraction of rollout schedule sigmas kept for NFT loss (drop terminal 0 first).", ) - parser.add_argument( - "--diffusion-nft-shuffle-timesteps", - action="store_true", - default=True, - help="Shuffle NFT training timesteps per sample at convert time (default on).", - ) parser.add_argument( "--no-diffusion-nft-shuffle-timesteps", action="store_false", dest="diffusion_nft_shuffle_timesteps", + default=True, help="Disable NFT timestep shuffle.", ) - parser.add_argument( - "--diffusion-nft-ref-mode", - type=str, - choices=["ema", "base"], - default="ema", - help=( - "Deprecated alias for NFT auto --ref-mode: 'ema' or 'base' (lora_base). " - "Prefer --ref-mode directly." - ), - ) parser.add_argument( "--advantage-estimator", type=str, @@ -1445,6 +1415,7 @@ def parse_args(add_custom_arguments=None): args.rank = 0 # Primary process rank for wandb initialization args.world_size = args.actor_num_nodes * args.actor_num_gpus_per_node + set_default_diffusion_args(args) miles_validate_args(args) sglang_validate_args(args) @@ -1501,32 +1472,28 @@ def _resolve_eval_datasets(args) -> list[EvalDatasetConfig]: return eval_datasets -def resolve_and_validate_ref_mode(args, *, is_nft: bool, ema_enabled: bool) -> None: - """Infer and validate ``args.ref_mode`` (fast-fail; actor only consumes the result).""" - if getattr(args, "ref_mode", None) is None: +def set_default_diffusion_args(args) -> None: + is_nft = args.loss_type == "nft" + if is_nft: + if args.custom_expand_samples_to_train_pairs_path is None: + args.custom_expand_samples_to_train_pairs_path = ( + "miles.ray.data_conversion_hub.nft.expand_samples_to_train_pairs" + ) + if args.custom_prepare_train_batch_path is None: + args.custom_prepare_train_batch_path = "miles.backends.fsdp_utils.loss_hub.nft.prepare_nft_batch" + if args.custom_loss_function_path is None: + args.custom_loss_function_path = "miles.backends.fsdp_utils.loss_hub.nft.nft_loss_formula" + if args.diffusion_sde_type == "sde" and args.diffusion_noise_level == 0: + args.diffusion_sde_type = "ode" + + if args.ref_mode is None: if is_nft: - nft_pref = getattr(args, "diffusion_nft_ref_mode", "ema") - if nft_pref == "ema" and not ema_enabled: - logger.warning( - "--loss-type nft prefers EMA ref but --ema-shadow is off; " "falling back to --ref-mode lora_base." - ) - args.ref_mode = "lora_base" - else: - args.ref_mode = "ema" if nft_pref == "ema" else "lora_base" - elif float(getattr(args, "diffusion_kl_beta", 0.0) or 0.0) > 0: + args.ref_mode = "ema" if args.ema_shadow else "lora_base" + elif args.diffusion_kl_beta > 0: args.ref_mode = "lora_base" else: args.ref_mode = "none" - if is_nft and args.ref_mode == "none": - raise ValueError("--loss-type nft requires a reference model; set --ref-mode ema or lora_base") - if args.ref_mode == "ema" and not ema_enabled: - raise ValueError("--ref-mode ema requires --ema-shadow") - if args.ref_mode == "lora_base" and not args.use_lora: - raise ValueError("--ref-mode lora_base requires --use-lora") - if float(getattr(args, "diffusion_kl_beta", 0.0) or 0.0) > 0 and args.ref_mode == "none": - raise ValueError("--diffusion-kl-beta > 0 requires a reference model; set --ref-mode lora_base or ema") - def miles_validate_args(args): args.eval_datasets = _resolve_eval_datasets(args) @@ -1600,7 +1567,7 @@ def miles_validate_args(args): if args.use_lora and args.lora_target_modules is None: args.lora_target_modules = list(cfg_cls.lora_target_modules) - if getattr(args, "lora_ipc_weight_sync", False): + if args.lora_ipc_weight_sync: if not args.use_lora: raise ValueError("--lora-ipc-weight-sync requires --use-lora") if not args.lora_target_modules: @@ -1609,41 +1576,40 @@ def miles_validate_args(args): "set --diffusion-model (for per-model defaults) or --lora-target-modules." ) - ema_enabled = bool(getattr(args, "ema_shadow", False)) + if not 0.0 <= args.ema_decay <= 1.0: + raise ValueError(f"--ema-decay must be in [0, 1], got {args.ema_decay}") + if args.ema_uprate < 0.0: + raise ValueError(f"--ema-uprate must be non-negative, got {args.ema_uprate}") + if not 0.0 <= args.ema_uphold <= 1.0: + raise ValueError(f"--ema-uphold must be in [0, 1], got {args.ema_uphold}") + if args.ema_flat_steps < 0: + raise ValueError(f"--ema-flat-steps must be non-negative, got {args.ema_flat_steps}") + if args.ema_rollout_policy == "ema" and not args.ema_shadow: + raise ValueError("--ema-rollout-policy ema requires --ema-shadow") - is_nft = getattr(args, "loss_type", None) in ("nft", "diffusion_nft") + is_nft = args.loss_type == "nft" if is_nft: - # DiffusionNFT swaps the train-pair conversion, preparation, and loss formula. - if getattr(args, "custom_expand_samples_to_train_pairs_path", None) is None: - args.custom_expand_samples_to_train_pairs_path = ( - "miles.ray.data_conversion_hub.nft.expand_samples_to_train_pairs" - ) - if getattr(args, "custom_loss_function_path", None) is None: - args.custom_loss_function_path = "miles.backends.fsdp_utils.loss_hub.nft.nft_loss_formula" - if getattr(args, "custom_prepare_train_batch_path", None) is None: - args.custom_prepare_train_batch_path = "miles.backends.fsdp_utils.loss_hub.nft.prepare_nft_batch" - if ( - getattr(args, "diffusion_sde_type", "sde") == "sde" - and float(getattr(args, "diffusion_noise_level", 0.7) or 0.0) == 0.0 - ): - logger.info( - "NFT: auto-setting --diffusion-sde-type ode because noise_level=0 " - "(forward-process needs clean x0 only)." + if args.diffusion_nft_beta <= 0: + raise ValueError(f"--diffusion-nft-beta must be > 0, got {args.diffusion_nft_beta}") + if args.diffusion_nft_adv_clip_max <= 0: + raise ValueError(f"--diffusion-nft-adv-clip-max must be > 0, got {args.diffusion_nft_adv_clip_max}") + if not 0.0 < args.diffusion_nft_timestep_fraction <= 1.0: + raise ValueError( + f"--diffusion-nft-timestep-fraction must be in (0, 1], got {args.diffusion_nft_timestep_fraction}" ) - args.diffusion_sde_type = "ode" - beta = float(getattr(args, "diffusion_nft_beta", 1.0) or 1.0) - if beta <= 0: - raise ValueError(f"--diffusion-nft-beta must be > 0, got {beta}") - frac = float(getattr(args, "diffusion_nft_timestep_fraction", 0.99) or 0.99) - if not (0.0 < frac <= 1.0): - raise ValueError(f"--diffusion-nft-timestep-fraction must be in (0, 1], got {frac}") if args.diffusion_recompute_old_log_prob: raise ValueError( "--diffusion-recompute-old-log-prob is only supported for policy_loss / Flow-GRPO, not NFT" ) - # Resolve --ref-mode once here (fast-fail); actor only consumes the final value. - resolve_and_validate_ref_mode(args, is_nft=is_nft, ema_enabled=ema_enabled) + if is_nft and args.ref_mode == "none": + raise ValueError("--loss-type nft requires a reference model; set --ref-mode ema or lora_base") + if args.ref_mode == "ema" and not args.ema_shadow: + raise ValueError("--ref-mode ema requires --ema-shadow") + if args.ref_mode == "lora_base" and not args.use_lora: + raise ValueError("--ref-mode lora_base requires --use-lora") + if args.diffusion_kl_beta > 0 and args.ref_mode == "none": + raise ValueError("--diffusion-kl-beta > 0 requires a reference model; set --ref-mode lora_base or ema") if args.dump_details is not None: args.save_debug_rollout_data = f"{args.dump_details}/rollout_data/{{rollout_id}}.pt" @@ -1722,8 +1688,7 @@ def miles_validate_args(args): sde_step_backends = { "sde": "miles.backends.fsdp_utils.sde_step_backend.DiffusersSdeStepBackend", "cps": "miles.backends.fsdp_utils.sde_step_backend.CpsSdeStepBackend", - # ODE rollouts (DiffusionNFT) do not score logπ on the train path; Diffusers backend - # is a harmless placeholder so actor init can still construct sde_backend. + # NFT does not use the placeholder backend. "ode": "miles.backends.fsdp_utils.sde_step_backend.DiffusersSdeStepBackend", } if args.diffusion_sde_type not in sde_step_backends: diff --git a/scripts/run-diffusion-nft-sd3-pickscore.sh b/scripts/run-diffusion-nft-sd3-pickscore.sh index 095b053d..947a071f 100755 --- a/scripts/run-diffusion-nft-sd3-pickscore.sh +++ b/scripts/run-diffusion-nft-sd3-pickscore.sh @@ -1,22 +1,5 @@ #!/usr/bin/env bash -# SD3.5 medium + DiffusionNFT + PickScore (UniRL sd3_nft / sd3_nft_100roll parity). -# -# UniRL reference (logs/sd3_nft_100roll.log, 1-GPU trainside smoke override): -# 8 prompts × 8 samples, lr=3e-4, wd=1e-4, LoRA r=32 α=64, -# guidance=1.0 (CFG-free), eta=0 / num_sde_steps=0 (ODE → clean x0), -# NFT beta=1, adv_clip=5, adaptive weight, schedule fraction=0.99, -# EMA shadow (pi_old) updated + synced at rollout_end (uprate=0.001, uphold=0.5). -# Expected early metrics: reward ~0.73–0.82, train/loss ~280–360. -# -# GPU layout (default CUDA_VISIBLE_DEVICES=4,5,2): -# first 2 = FSDP train + sglang colocate; 3rd = PickScore reward worker. -# -# Usage: -# NUM_ROLLOUT=5 CUDA_VISIBLE_DEVICES=4,5,2 \ -# bash scripts/run-diffusion-nft-sd3-pickscore.sh -# -# Smoke (tiny batch, OCR, 2 GPU — set SMOKE=1): -# SMOKE=1 CUDA_VISIBLE_DEVICES=4,5 bash scripts/run-diffusion-nft-sd3-pickscore.sh +# SD3.5 DiffusionNFT training with PickScore. MILES_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" 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 0f93f92f..162c04bb 100644 --- a/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py +++ b/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py @@ -8,7 +8,7 @@ import torch -from miles.backends.fsdp_utils.ema import EmaShadow, resolve_ema_kwargs +from miles.backends.fsdp_utils.ema import EmaShadow from miles.backends.fsdp_utils.loss_hub.nft import corrupt, nft_loss_formula, nft_r_from_advantages from miles.backends.fsdp_utils.metrics import new_metric_buffer from miles.ray.data_conversion_hub.nft import expand_samples_to_train_pairs, resolve_nft_sigmas @@ -45,10 +45,8 @@ def test_corrupt_linear_flow(self): assert torch.allclose(xt[1], torch.full((4,), 0.25)) def test_resolve_sigmas_drops_zero_and_fraction(self): - class _Sched: - sigmas = torch.tensor([1.0, 0.8, 0.6, 0.4, 0.2, 0.0]) - - ts = resolve_nft_sigmas(_Sched(), training_timestep_fraction=0.99) + sigmas = torch.tensor([1.0, 0.8, 0.6, 0.4, 0.2, 0.0]) + ts = resolve_nft_sigmas(sigmas, training_timestep_fraction=0.99) assert torch.allclose(ts, torch.tensor([1.0, 0.8, 0.6, 0.4])) @@ -75,7 +73,6 @@ class _Env: assert len(out["train_data"]) == 4 assert {p["timestep"] for p in out["train_data"]} == {1.0, 0.5} assert out["train_data"][0]["x0"] is out["train_data"][1]["x0"] - # Advantages come from post-process, not from the converter. assert out["train_data"][0]["advantage"] == rewards[0] assert out["train_data"][2]["advantage"] == rewards[1] @@ -93,8 +90,7 @@ def test_formula_write_old_log_prob_is_noop(self): assert loss is None def test_formula_declares_window_attr(self): - assert getattr(nft_loss_formula, "requires_sample_aligned_windows", False) is True - assert not hasattr(nft_loss_formula, "ref_mode") + assert nft_loss_formula.requires_sample_aligned_windows is True class TestEmaShadow: @@ -120,9 +116,3 @@ 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) - - -class TestEmaArgs: - def test_resolve_kwargs(self): - kwargs = resolve_ema_kwargs(Namespace(ema_decay=0.01, ema_uprate=0.02, ema_uphold=0.3, ema_flat_steps=5)) - assert kwargs == {"decay": 0.01, "uprate": 0.02, "uphold": 0.3, "flat_steps": 5} From b3b91429421e179d37643690d0a3f8ddd4f84ed2 Mon Sep 17 00:00:00 2001 From: rockdu Date: Thu, 30 Jul 2026 03:24:19 -0700 Subject: [PATCH 12/12] refactor(diffusion): tighten NFT hook contracts Co-authored-by: Cursor --- miles/backends/fsdp_utils/actor.py | 25 ++++++++----------- .../backends/fsdp_utils/loss_hub/__init__.py | 2 -- .../backends/fsdp_utils/loss_hub/flow_grpo.py | 17 +++---------- miles/backends/fsdp_utils/loss_hub/nft.py | 22 +++------------- miles/backends/fsdp_utils/loss_hub/utils.py | 8 ++++++ miles/utils/arguments.py | 10 ++++---- .../backends/fsdp_utils/test_loss_hub_nft.py | 19 +------------- 7 files changed, 31 insertions(+), 72 deletions(-) create mode 100644 miles/backends/fsdp_utils/loss_hub/utils.py diff --git a/miles/backends/fsdp_utils/actor.py b/miles/backends/fsdp_utils/actor.py index b179f9ab..4844cbaf 100644 --- a/miles/backends/fsdp_utils/actor.py +++ b/miles/backends/fsdp_utils/actor.py @@ -1,7 +1,7 @@ import logging import warnings from argparse import Namespace -from contextlib import contextmanager +from contextlib import contextmanager, nullcontext import ray import torch @@ -308,16 +308,15 @@ def update_weights(self) -> None: # type: ignore[override] if dist.get_rank() == 0: ray.get(self.rollout_manager.clear_num_new_engines.remote()) - if self.ema_shadow is not None: - delta = self.ema_shadow.update() + 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, self.ema_shadow.step) - if self.args.ema_rollout_policy == "ema": - with self.ema_shadow.swap_in(): - self.weight_updater.update_weights() - else: - self.weight_updater.update_weights() - else: + logger.info("EMA shadow updated (decay=%.4f step=%d)", delta, ema_shadow.step) + rollout_weight_context = ( + ema_shadow.swap_in() if ema_shadow is not None and self.args.ema_rollout_policy == "ema" else nullcontext() + ) + with rollout_weight_context: self.weight_updater.update_weights() clear_memory() @@ -403,8 +402,7 @@ def _train_core(self, rollout_id: int, rollout_data) -> None: microbatch_schedule=microbatch_schedule, parallel_state=self.parallel_state, ) - loss_formula_func = self.custom_loss_formula_func or flow_grpo_loss_formula - if getattr(loss_formula_func, "requires_sample_aligned_windows", False): + if self.args.loss_type == "nft": validate_sample_aligned_windows( train_pairs=train_pairs, microbatch_schedule=microbatch_schedule, @@ -509,9 +507,6 @@ def _forward_train_pair_batch( old_log_prob_from_new: bool = False, ) -> torch.Tensor | None: """Run one prepared diffusion micro-batch.""" - if not batch: - raise ValueError("_forward_train_pair_batch received empty batch") - if self.custom_prepare_train_batch_func is not None: prepared = self.custom_prepare_train_batch_func(ctx, batch, pad_to_len=pad_to_len) else: diff --git a/miles/backends/fsdp_utils/loss_hub/__init__.py b/miles/backends/fsdp_utils/loss_hub/__init__.py index 62b78ca9..ac2f4b9a 100644 --- a/miles/backends/fsdp_utils/loss_hub/__init__.py +++ b/miles/backends/fsdp_utils/loss_hub/__init__.py @@ -1,7 +1,6 @@ """Diffusion batch preparation and loss hooks.""" from miles.backends.fsdp_utils.loss_hub.flow_grpo import flow_grpo_loss_formula, prepare_flow_grpo_batch -from miles.backends.fsdp_utils.loss_hub.nft import prepare_nft_batch from miles.backends.fsdp_utils.loss_hub.types import DiffusionLossContext, PreparedBatch __all__ = [ @@ -9,5 +8,4 @@ "PreparedBatch", "flow_grpo_loss_formula", "prepare_flow_grpo_batch", - "prepare_nft_batch", ] diff --git a/miles/backends/fsdp_utils/loss_hub/flow_grpo.py b/miles/backends/fsdp_utils/loss_hub/flow_grpo.py index 979fc778..00f787b6 100644 --- a/miles/backends/fsdp_utils/loss_hub/flow_grpo.py +++ b/miles/backends/fsdp_utils/loss_hub/flow_grpo.py @@ -5,21 +5,12 @@ import torch from miles.backends.fsdp_utils.loss_hub.types import DiffusionLossContext, PreparedBatch +from miles.backends.fsdp_utils.loss_hub.utils import cast_cond_to_dtype from miles.backends.fsdp_utils.metrics import record_rollout_train_abs_diff from miles.utils.metric_buffer import MetricBuffer from miles.utils.train_data_utils import stack_train_pair_rollout_debug -def _cast_cond_to_dtype(cond: dict, dtype: torch.dtype) -> dict: - out = {} - for key, value in cond.items(): - if isinstance(value, torch.Tensor) and value.dtype.is_floating_point: - out[key] = value.to(dtype=dtype) - else: - out[key] = value - return out - - def _stack_pair_field(batch: list[dict], key: str, device: torch.device) -> torch.Tensor: return torch.stack([pair[key] for pair in batch]).to(device=device, dtype=torch.float32) @@ -86,17 +77,17 @@ def prepare_flow_grpo_batch( cfg_batching = use_cfg and bool(args.fsdp_cfg_batching) joint_cond = pos_cond = neg_cond = None if cfg_batching: - joint_cond = _cast_cond_to_dtype( + joint_cond = cast_cond_to_dtype( config.collate_cond_for_sample_batch(pos_list + neg_list, device, pad_to_len=pad_to_len), ctx.forward_dtype, ) else: - pos_cond = _cast_cond_to_dtype( + pos_cond = cast_cond_to_dtype( config.collate_cond_for_sample_batch(pos_list, device, pad_to_len=pad_to_len), ctx.forward_dtype, ) if use_cfg and neg_list is not None: - neg_cond = _cast_cond_to_dtype( + neg_cond = cast_cond_to_dtype( config.collate_cond_for_sample_batch(neg_list, device, pad_to_len=pad_to_len), ctx.forward_dtype, ) diff --git a/miles/backends/fsdp_utils/loss_hub/nft.py b/miles/backends/fsdp_utils/loss_hub/nft.py index b388c4dc..183bd47f 100644 --- a/miles/backends/fsdp_utils/loss_hub/nft.py +++ b/miles/backends/fsdp_utils/loss_hub/nft.py @@ -5,6 +5,7 @@ import torch from miles.backends.fsdp_utils.loss_hub.types import DiffusionLossContext, PreparedBatch +from miles.backends.fsdp_utils.loss_hub.utils import cast_cond_to_dtype from miles.utils.metric_buffer import MetricBuffer @@ -19,16 +20,6 @@ def corrupt(x0: torch.Tensor, t: torch.Tensor, eps: torch.Tensor) -> torch.Tenso return (1.0 - t) * x0 + t * eps -def _cast_cond_to_dtype(cond: dict, dtype: torch.dtype) -> dict: - out = {} - for key, value in cond.items(): - if isinstance(value, torch.Tensor) and value.dtype.is_floating_point: - out[key] = value.to(dtype=dtype) - else: - out[key] = value - return out - - def prepare_nft_batch( ctx: DiffusionLossContext, batch: list[dict], @@ -47,7 +38,7 @@ def prepare_nft_batch( component_name, model = next(iter(ctx.models.items())) pos_list = [config.prepare_cond_kwargs(batch[i]["denoising_env"].pos_cond_kwargs, device) for i in range(bsz)] - pos_cond = _cast_cond_to_dtype( + pos_cond = cast_cond_to_dtype( config.collate_cond_for_sample_batch(pos_list, device, pad_to_len=pad_to_len), ctx.forward_dtype, ) @@ -126,12 +117,8 @@ def nft_loss_formula( metrics: MetricBuffer, write_old_log_prob: bool = False, old_log_prob_from_new: bool = False, -) -> torch.Tensor | None: +) -> torch.Tensor: """Dual-policy x0-MSE. Actor must supply ``ref_pred`` (EMA / LoRA-base).""" - if write_old_log_prob: - return None - if old_log_prob_from_new: - raise ValueError("DiffusionNFT has no PPO log-prob; old_log_prob_from_new is unsupported") if ref_pred is None: raise ValueError("NFT loss formula requires a reference prediction from the actor") @@ -177,6 +164,3 @@ def nft_loss_formula( metrics.emit_mean("adv_abs_mean", total=prepared.advantage.abs().sum(), count=bsz) return loss_sum - - -nft_loss_formula.requires_sample_aligned_windows = True diff --git a/miles/backends/fsdp_utils/loss_hub/utils.py b/miles/backends/fsdp_utils/loss_hub/utils.py new file mode 100644 index 00000000..3f1d0893 --- /dev/null +++ b/miles/backends/fsdp_utils/loss_hub/utils.py @@ -0,0 +1,8 @@ +import torch + + +def cast_cond_to_dtype(cond: dict, dtype: torch.dtype) -> dict: + return { + key: value.to(dtype=dtype) if isinstance(value, torch.Tensor) and value.dtype.is_floating_point else value + for key, value in cond.items() + } diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index 518446c3..4eae5d48 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -146,8 +146,8 @@ def add_train_arguments(parser): default=None, help=( "Which reference weights to use for the no-grad DiT forward. " - "Auto: lora_base when --diffusion-kl-beta > 0; for --loss-type nft, " - "ema if --ema-shadow else lora_base. Explicit values skip auto inference." + "Auto: lora_base when --diffusion-kl-beta > 0 and ema for --loss-type nft. " + "Explicit values skip auto inference." ), ) parser.add_argument( @@ -1483,12 +1483,10 @@ def set_default_diffusion_args(args) -> None: args.custom_prepare_train_batch_path = "miles.backends.fsdp_utils.loss_hub.nft.prepare_nft_batch" if args.custom_loss_function_path is None: args.custom_loss_function_path = "miles.backends.fsdp_utils.loss_hub.nft.nft_loss_formula" - if args.diffusion_sde_type == "sde" and args.diffusion_noise_level == 0: - args.diffusion_sde_type = "ode" if args.ref_mode is None: if is_nft: - args.ref_mode = "ema" if args.ema_shadow else "lora_base" + args.ref_mode = "ema" elif args.diffusion_kl_beta > 0: args.ref_mode = "lora_base" else: @@ -1589,6 +1587,8 @@ def miles_validate_args(args): is_nft = args.loss_type == "nft" if is_nft: + if args.diffusion_noise_level == 0 and args.diffusion_sde_type != "ode": + raise ValueError("--loss-type nft with --diffusion-noise-level 0 requires --diffusion-sde-type ode") if args.diffusion_nft_beta <= 0: raise ValueError(f"--diffusion-nft-beta must be > 0, got {args.diffusion_nft_beta}") if args.diffusion_nft_adv_clip_max <= 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 162c04bb..ae8da9df 100644 --- a/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py +++ b/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py @@ -9,8 +9,7 @@ import torch from miles.backends.fsdp_utils.ema import EmaShadow -from miles.backends.fsdp_utils.loss_hub.nft import corrupt, nft_loss_formula, nft_r_from_advantages -from miles.backends.fsdp_utils.metrics import new_metric_buffer +from miles.backends.fsdp_utils.loss_hub.nft import corrupt, nft_r_from_advantages from miles.ray.data_conversion_hub.nft import expand_samples_to_train_pairs, resolve_nft_sigmas from miles.utils.types import Sample @@ -76,22 +75,6 @@ class _Env: assert out["train_data"][0]["advantage"] == rewards[0] assert out["train_data"][2]["advantage"] == rewards[1] - def test_formula_write_old_log_prob_is_noop(self): - metrics = new_metric_buffer(None, torch.device("cpu"), ()) - loss = nft_loss_formula( - None, # type: ignore[arg-type] - [], - None, # type: ignore[arg-type] - new_pred=None, # type: ignore[arg-type] - ref_pred=None, - metrics=metrics, - write_old_log_prob=True, - ) - assert loss is None - - def test_formula_declares_window_attr(self): - assert nft_loss_formula.requires_sample_aligned_windows is True - class TestEmaShadow: def _model(self):