Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
45 changes: 41 additions & 4 deletions miles/backends/fsdp_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down
10 changes: 10 additions & 0 deletions miles/utils/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
Loading