diff --git a/miles/backends/fsdp_utils/actor.py b/miles/backends/fsdp_utils/actor.py index 378c52be..4844cbaf 100644 --- a/miles/backends/fsdp_utils/actor.py +++ b/miles/backends/fsdp_utils/actor.py @@ -24,17 +24,20 @@ 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 ( DiffusionUpdateWeightFromTensor, DiffusionUpdateWeightFromTensorLoRA, DiffusionUpdateWeightFromTensorLoRAIPC, ) +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, 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 @@ -169,13 +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, ) + 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( (p for p in self.model.parameters() if p.requires_grad), @@ -202,6 +212,16 @@ def init(self, args: Namespace, role: str, with_ref: bool = False) -> int: # ty checkpoint_payload = checkpoint.load(self) + self.ema_shadow = None + if self.args.ema_shadow: + self.ema_shadow = EmaShadow( + (p for m in self.models.values() for p in m.parameters()), + 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). if self.args.debug_train_only: self.weight_updater = None @@ -288,7 +308,16 @@ 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() + ema_shadow = self.ema_shadow + if ema_shadow is not None: + delta = ema_shadow.update() + if dist.get_rank() == 0: + logger.info("EMA shadow updated (decay=%.4f step=%d)", delta, ema_shadow.step) + 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() def _log_metrics(self, rollout_id: int, log_dict: dict[str, float], step: int) -> None: @@ -327,12 +356,7 @@ 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. - - 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``). - """ + """Run the shared diffusion training loop.""" device = torch.cuda.current_device() train_pairs: list = rollout_data["train_data"] @@ -340,26 +364,13 @@ 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_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." + 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 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.") # ------------- Rollout Scheduler Metadata ------------- scheduler_timesteps, scheduler_sigmas = scheduler_meta_from_rollout( @@ -391,6 +402,21 @@ def _train_core(self, rollout_id: int, rollout_data) -> None: microbatch_schedule=microbatch_schedule, parallel_state=self.parallel_state, ) + if self.args.loss_type == "nft": + validate_sample_aligned_windows( + train_pairs=train_pairs, + microbatch_schedule=microbatch_schedule, + ) + + loss_ctx = DiffusionLossContext( + models=self.models, + 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, + ) # ------------- Recompute old log-probs (impl-consistent PPO ratio) ------------- if self.args.diffusion_recompute_old_log_prob: @@ -402,15 +428,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 +452,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, ) @@ -485,235 +498,73 @@ def _maybe_legacy_window_pad_len(self, train_pairs: list, microbatch_ranges: lis 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. - - 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``). - - 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 + """Run one prepared diffusion micro-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: + prepared = prepare_flow_grpo_batch(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) + forward_dtype = self._forward_dtype - 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, + 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, ) - # 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 - ) - - # 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, + new_pred = _compute_noise_pred() + + ref_pred = None + ref_mode = self.args.ref_mode + if ref_mode != "none": + if ref_mode == "ema": + ref_ctx = self.ema_shadow.swap_in() + else: + ref_ctx = prepared.model.disable_adapter() + with torch.no_grad(), ref_ctx: + ref_pred = _compute_noise_pred().detach() + + 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, ) - 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, - ) - 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, + return flow_grpo_loss_formula( + 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/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/ema.py b/miles/backends/fsdp_utils/ema.py new file mode 100644 index 00000000..9120ef3b --- /dev/null +++ b/miles/backends/fsdp_utils/ema.py @@ -0,0 +1,74 @@ +"""EMA shadow of trainable parameters for diffusion FSDP training.""" + +from __future__ import annotations + +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: + return t.to_local() if isinstance(t, DTensor) else t + + +class EmaShadow: + """EMA shadow of trainable parameters.""" + + 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("EmaShadow: 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("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): + 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, strict=True): + 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..ac2f4b9a --- /dev/null +++ b/miles/backends/fsdp_utils/loss_hub/__init__.py @@ -0,0 +1,11 @@ +"""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.types import DiffusionLossContext, PreparedBatch + +__all__ = [ + "DiffusionLossContext", + "PreparedBatch", + "flow_grpo_loss_formula", + "prepare_flow_grpo_batch", +] 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..00f787b6 --- /dev/null +++ b/miles/backends/fsdp_utils/loss_hub/flow_grpo.py @@ -0,0 +1,209 @@ +"""Flow-GRPO batch preparation and loss formula.""" + +from __future__ import annotations + +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 _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/nft.py b/miles/backends/fsdp_utils/loss_hub/nft.py new file mode 100644 index 00000000..183bd47f --- /dev/null +++ b/miles/backends/fsdp_utils/loss_hub/nft.py @@ -0,0 +1,166 @@ +"""DiffusionNFT batch preparation and loss formula.""" + +from __future__ import annotations + +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 + + +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 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 = ctx.scheduler.config.num_train_timesteps + 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 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 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: + """Dual-policy x0-MSE. Actor must supply ``ref_pred`` (EMA / LoRA-base).""" + if ref_pred is None: + raise ValueError("NFT loss formula requires a reference prediction from the actor") + + args = ctx.args + 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 + 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 = batch[0]["nft_num_timesteps"] + per_pair_total = per_pair.sum() + bsz = len(batch) + 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) + 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 diff --git a/miles/backends/fsdp_utils/loss_hub/types.py b/miles/backends/fsdp_utils/loss_hub/types.py new file mode 100644 index 00000000..33d4f767 --- /dev/null +++ b/miles/backends/fsdp_utils/loss_hub/types.py @@ -0,0 +1,43 @@ +"""Shared data types passed into diffusion prepare and loss 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 and loss hooks.""" + + models: dict[str, torch.nn.Module] + train_pipeline_config: Any + sde_backend: Any + scheduler: Any + args: Namespace + forward_dtype: torch.dtype + device: torch.device + + +@dataclass +class PreparedBatch: + """Actor-owned DiT forward inputs produced by a prepare hook.""" + + 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/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/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/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..b7bfabe7 --- /dev/null +++ b/miles/ray/data_conversion_hub/flow_grpo.py @@ -0,0 +1,170 @@ +"""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: 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) + 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 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") + 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..1a6e4af8 --- /dev/null +++ b/miles/ray/data_conversion_hub/nft.py @@ -0,0 +1,93 @@ +"""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 + + +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 resolve_nft_sigmas( + sigmas: torch.Tensor, + *, + training_timestep_fraction: float = 0.99, +) -> torch.Tensor: + 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): + 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") + if first_traj.sigmas is not None: + scheduler_sigmas = first_traj.sigmas.detach().cpu().float() + else: + ts = first_traj.timesteps.detach().cpu().float() + 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, + } + 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]] = [] + 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 args.diffusion_nft_shuffle_timesteps 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 967431ee..2b49513a 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,7 +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 RolloutTrainDataConverter, TrainDataDPSplitter, reorder_train_pairs_for_tiling +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 @@ -77,7 +80,11 @@ 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() + 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.") @@ -337,9 +344,7 @@ def _post_process_rewards(self, samples: list[Sample] | list[list[Sample]]): raw_rewards = [sample.get_reward_value(self.args) for sample in samples] - # --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. + # --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) @@ -363,6 +368,7 @@ def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sampl """ Convert inference generated samples to training data. """ + # 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) @@ -374,7 +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. reward_stats = { **_reward_stats_dict(raw_t, "rollout/reward/raw_"), **_reward_stats_dict(norm_t, "rollout/reward/norm_"), @@ -406,7 +411,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 e3a583d1..4eae5d48 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, @@ -145,8 +135,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 and ema for --loss-type nft. " + "Explicit values skip auto inference." ), ) parser.add_argument( @@ -569,6 +570,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, @@ -859,13 +870,66 @@ 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"], 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 prepare / loss-formula paths and selects the NFT train-data converter. " + "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( + "--no-diffusion-nft-adaptive-weight", + action="store_false", + dest="diffusion_nft_adaptive_weight", + default=True, + 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( + "--no-diffusion-nft-shuffle-timesteps", + action="store_false", + dest="diffusion_nft_shuffle_timesteps", + default=True, + help="Disable NFT timestep shuffle.", + ) parser.add_argument( "--advantage-estimator", type=str, @@ -1068,6 +1132,50 @@ def add_debug_arguments(parser): "(requires matching sglang-d LoRAPipeline support)." ), ) + parser.add_argument( + "--ema-shadow", + action="store_true", + default=False, + help=( + "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( + "--ema-rollout-policy", + type=str, + choices=["live", "ema"], + default="live", + help=( + "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( + "--ema-decay", + type=float, + default=0.001, + help="EMA decay while step <= flat_steps.", + ) + parser.add_argument( + "--ema-uprate", + type=float, + default=0.001, + help="EMA warmup rate after flat_steps.", + ) + parser.add_argument( + "--ema-uphold", + type=float, + default=0.5, + help="EMA warmup cap.", + ) + parser.add_argument( + "--ema-flat-steps", + type=int, + default=0, + help="EMA flat steps before warmup begins.", + ) parser.add_argument( "--diffusion-init-lora-weight", type=str, @@ -1307,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) @@ -1363,6 +1472,27 @@ def _resolve_eval_datasets(args) -> list[EvalDatasetConfig]: return eval_datasets +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.ref_mode is None: + if is_nft: + args.ref_mode = "ema" + elif args.diffusion_kl_beta > 0: + args.ref_mode = "lora_base" + else: + args.ref_mode = "none" + + def miles_validate_args(args): args.eval_datasets = _resolve_eval_datasets(args) @@ -1435,7 +1565,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: @@ -1444,6 +1574,43 @@ def miles_validate_args(args): "set --diffusion-model (for per-model defaults) or --lora-target-modules." ) + 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 = 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: + 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}" + ) + 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 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" args.save_debug_train_data = f"{args.dump_details}/train_data/{{rollout_id}}_{{rank}}.pt" @@ -1521,6 +1688,8 @@ 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", + # 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: raise ValueError( diff --git a/miles/utils/train_data_utils.py b/miles/utils/train_data_utils.py index 9a663317..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,184 +39,6 @@ def scheduler_meta_from_rollout( return timesteps, sigmas -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. @@ -429,6 +249,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..947a071f --- /dev/null +++ b/scripts/run-diffusion-nft-sd3-pickscore.sh @@ -0,0 +1,146 @@ +#!/usr/bin/env bash +# SD3.5 DiffusionNFT training with PickScore. + +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 \ + --ref-mode ema \ + --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}" \ + "${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..ae8da9df --- /dev/null +++ b/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py @@ -0,0 +1,101 @@ +"""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 torch + +from miles.backends.fsdp_utils.ema import EmaShadow +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 + + +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): + 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])) + + +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()), + ] + args = _args() + 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"] + assert out["train_data"][0]["advantage"] == rewards[0] + assert out["train_data"][2]["advantage"] == rewards[1] + + +class TestEmaShadow: + def _model(self): + return torch.nn.Linear(4, 4, bias=False) + + def test_snapshot_and_update(self): + m = self._model() + 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) + 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 = EmaShadow(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) 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"]