diff --git a/miles/backends/fsdp_utils/actor.py b/miles/backends/fsdp_utils/actor.py index e3c3ba46..e00837e5 100644 --- a/miles/backends/fsdp_utils/actor.py +++ b/miles/backends/fsdp_utils/actor.py @@ -381,11 +381,34 @@ def _train_core(self, rollout_id: int, rollout_data) -> None: parallel_state=self.parallel_state, ) + # ------------- Recompute old log-probs (impl-consistent PPO ratio) ------------- + if self.args.diffusion_recompute_old_log_prob: + with timer("recompute_old_log_prob"), torch.no_grad(): + # Skip window 0: its training forward runs on the same pre-update weights and doubles as the recompute. + for microbatch_ranges in microbatch_schedule[1:]: + 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( + 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, + log_stats=defaultdict(list), + device=device, + pad_to_len=legacy_pad_to_len, + write_old_log_prob=True, + ) + # ------------- Forward / Backward ------------- with timer("actor_train"): - for microbatch_ranges in microbatch_schedule: + for optim_step_idx, microbatch_ranges in enumerate(microbatch_schedule): self.optimizer.zero_grad(set_to_none=True) + old_log_prob_from_new = self.args.diffusion_recompute_old_log_prob and optim_step_idx == 0 + num_local_pairs = sum(pair_hi - pair_lo for pair_lo, pair_hi in microbatch_ranges) # LEGACY 2D parity: pad cond to the whole-window width. TODO: remove with legacy 2D path. @@ -407,6 +430,7 @@ def _train_core(self, rollout_id: int, rollout_data) -> None: device=device, kl_beta=kl_beta, pad_to_len=legacy_pad_to_len, + old_log_prob_from_new=old_log_prob_from_new, ) if not self.args.debug_skip_optimizer_step: # ShardedGradScaler keeps fp16 policy grads from underflowing @@ -461,8 +485,16 @@ def _forward_train_pair_batch( device: torch.device, kl_beta: float = 0.0, pad_to_len: int | None = None, - ) -> torch.Tensor: - """One DiT forward + PPO loss over ``len(batch)`` train pairs. Returns sum of per-pair losses.""" + 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 train_pipeline_config = self.train_pipeline_config bsz = len(batch) @@ -589,8 +621,13 @@ def _compute_noise_pred(disable_adapter: bool = False) -> torch.Tensor: noise_level=noise_level, ) + 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_old_microbatch # (bsz,) + 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) diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index c123cf51..e38fbd8f 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -130,6 +130,16 @@ def add_train_arguments(parser): default=5.0, help="Max absolute value for advantage clipping in diffusion training.", ) + parser.add_argument( + "--diffusion-recompute-old-log-prob", + action="store_true", + help=( + "Recompute old log-probs with the trainer forward (pre-update weights) " + "instead of using rollout-stored values, making the PPO ratio " + "implementation-consistent. The first optimizer window skips the extra " + "pass and reuses its training forward's log-prob (ratio == 1)." + ), + ) parser.add_argument( "--diffusion-kl-beta", type=float,