From e350ba8729031e9104a53a3377275f7d84b86698 Mon Sep 17 00:00:00 2001 From: rockdu Date: Sun, 2 Aug 2026 16:06:36 -0700 Subject: [PATCH 1/3] refactor(diffusion): require rollout scheduler sigmas, drop timesteps-derived fallbacks --- miles/backends/fsdp_utils/actor.py | 2 -- miles/ray/data_conversion_hub/flow_grpo.py | 16 +++++----- miles/ray/data_conversion_hub/nft.py | 9 ++---- miles/utils/train_data_utils.py | 8 ++--- miles/utils/types.py | 8 ++--- .../backends/fsdp_utils/test_loss_hub_nft.py | 20 +++++++++++++ tests/fast/utils/test_grouping_parity.py | 30 +++++++++++++++---- 7 files changed, 62 insertions(+), 31 deletions(-) diff --git a/miles/backends/fsdp_utils/actor.py b/miles/backends/fsdp_utils/actor.py index 4844cbaf..d9430378 100644 --- a/miles/backends/fsdp_utils/actor.py +++ b/miles/backends/fsdp_utils/actor.py @@ -364,7 +364,6 @@ def _train_core(self, rollout_id: int, rollout_data) -> None: raise ValueError("rollout_data['train_data'] is empty") num_pairs = len(train_pairs) - num_train_timesteps = self.scheduler.config.num_train_timesteps ref_mode = self.args.ref_mode if ref_mode == "lora_base" and not all(hasattr(m, "disable_adapter") for m in self.models.values()): @@ -376,7 +375,6 @@ def _train_core(self, rollout_id: int, rollout_data) -> None: scheduler_timesteps, scheduler_sigmas = scheduler_meta_from_rollout( rollout_data, device=device, - num_train_timesteps=num_train_timesteps, ) self.scheduler.timesteps = scheduler_timesteps self.scheduler.sigmas = scheduler_sigmas diff --git a/miles/ray/data_conversion_hub/flow_grpo.py b/miles/ray/data_conversion_hub/flow_grpo.py index b7bfabe7..5f080c70 100644 --- a/miles/ray/data_conversion_hub/flow_grpo.py +++ b/miles/ray/data_conversion_hub/flow_grpo.py @@ -26,10 +26,12 @@ def _expand_samples_to_train_pairs( 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() + if first_traj.sigmas is None: + raise ValueError("sample 0 missing dit_trajectory.sigmas; rollout engine must return the sigmas snapshot") + scheduler_meta: dict[str, torch.Tensor] = { + "scheduler_timesteps": first_traj.timesteps.detach().cpu().float(), + "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) @@ -38,10 +40,8 @@ def _expand_samples_to_train_pairs( 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) + if traj.sigmas is None or not torch.equal( + traj.sigmas.detach().cpu().float(), scheduler_meta["scheduler_sigmas"] ): raise ValueError( f"sample {sample.index} has different scheduler_sigmas than sample 0; " diff --git a/miles/ray/data_conversion_hub/nft.py b/miles/ray/data_conversion_hub/nft.py index 1a6e4af8..70b95b9c 100644 --- a/miles/ray/data_conversion_hub/nft.py +++ b/miles/ray/data_conversion_hub/nft.py @@ -56,14 +56,11 @@ 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") - 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)]) + if first_traj.sigmas is None: + raise ValueError("NFT needs dit_trajectory.sigmas from rollout; no timesteps-derived fallback") scheduler_meta = { "scheduler_timesteps": first_traj.timesteps.detach().cpu().float(), - "scheduler_sigmas": scheduler_sigmas, + "scheduler_sigmas": first_traj.sigmas.detach().cpu().float(), } sigmas = resolve_nft_sigmas( scheduler_meta["scheduler_sigmas"], diff --git a/miles/utils/train_data_utils.py b/miles/utils/train_data_utils.py index d4ea726c..ceb05969 100644 --- a/miles/utils/train_data_utils.py +++ b/miles/utils/train_data_utils.py @@ -26,16 +26,14 @@ def scheduler_meta_from_rollout( rollout_data: dict, *, device: torch.device, - num_train_timesteps: int, ) -> tuple[torch.Tensor, torch.Tensor]: """Use rollout-side scheduler metadata for train/rollout alignment.""" if "scheduler_timesteps" not in rollout_data: raise ValueError("rollout_data missing scheduler_timesteps") + if "scheduler_sigmas" not in rollout_data: + raise ValueError("rollout_data missing scheduler_sigmas; rollout engine must return the sigmas snapshot") timesteps = rollout_data["scheduler_timesteps"].to(device=device, dtype=torch.float32) - if "scheduler_sigmas" in rollout_data: - sigmas = rollout_data["scheduler_sigmas"].to(device=device, dtype=torch.float32) - else: - sigmas = torch.cat([timesteps / float(num_train_timesteps), timesteps.new_zeros(1)]) + sigmas = rollout_data["scheduler_sigmas"].to(device=device, dtype=torch.float32) return timesteps, sigmas diff --git a/miles/utils/types.py b/miles/utils/types.py index d133447e..0864a1c6 100644 --- a/miles/utils/types.py +++ b/miles/utils/types.py @@ -47,11 +47,9 @@ class DenoisingEnv: class DiTTrajectory: latents: torch.Tensor | None = None timesteps: torch.Tensor | None = None - # Rollout's scheduler.sigmas snapshot [T+1] (post-shift, includes - # terminal 0). Use this on the training side instead of recomputing - # sigmas from `timesteps / num_train_timesteps` — that round-trips - # σ * 1000 / 1000 in fp32 and drifts 1-2 ULPs, amplifying to ~3e-5 - # log_prob diff. + # Rollout's scheduler.sigmas snapshot [T+1] (post-shift, includes terminal 0). + # Required for training — converters raise if missing; recomputing from + # `timesteps / num_train_timesteps` drifts 1-2 ULPs (~3e-5 log_prob diff). sigmas: torch.Tensor | None = None 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 ae8da9df..65a355e1 100644 --- a/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py +++ b/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py @@ -75,6 +75,26 @@ class _Env: assert out["train_data"][0]["advantage"] == rewards[0] assert out["train_data"][2]["advantage"] == rewards[1] + def test_convert_requires_rollout_sigmas(self): + # Sigmas come from the rollout scheduler snapshot; no timesteps/1000 fallback. + class _Traj: + def __init__(self): + self.timesteps = torch.tensor([999.0, 500.0, 0.0]) + self.sigmas = None + 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())] + try: + expand_samples_to_train_pairs(_args(), samples, [1.0], [1.0]) + except ValueError as e: + assert "sigmas" in str(e) + else: + raise AssertionError("expected ValueError for missing dit_trajectory.sigmas") + class TestEmaShadow: def _model(self): diff --git a/tests/fast/utils/test_grouping_parity.py b/tests/fast/utils/test_grouping_parity.py index 9fbdfda2..21cebf41 100644 --- a/tests/fast/utils/test_grouping_parity.py +++ b/tests/fast/utils/test_grouping_parity.py @@ -32,7 +32,7 @@ import torch from miles.ray.data_conversion_hub.flow_grpo import expand_samples_to_train_pairs -from miles.utils.train_data_utils import TrainDataDPSplitter +from miles.utils.train_data_utils import TrainDataDPSplitter, scheduler_meta_from_rollout # -------------------------------------------------------------------------------------- @@ -154,12 +154,16 @@ def test_l2_converter_pairs_match_direct_indexing(): assert torch.equal(d["rollout_step_noise_std_dev"], s.rollout_debug_tensors.rollout_noise_std_devs[idx]) -def test_l2_converter_sigmas_optional(): +def test_l2_converter_requires_sigmas(): + """A trajectory without the rollout sigmas snapshot must raise (no derived fallback).""" T, sde = 4, [0, 2] samples = [_mk_sample(i, T, sde, with_sigmas=False, with_debug=False) for i in range(2)] - 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) + try: + expand_samples_to_train_pairs(None, samples, [1.0, 2.0], [1.0, 2.0]) + except ValueError: + pass + else: + raise AssertionError("expected ValueError for missing dit_trajectory.sigmas") def test_l2_converter_rejects_mismatched_scheduler_timesteps(): @@ -186,6 +190,22 @@ def test_l2_converter_rejects_mismatched_scheduler_sigmas(): raise AssertionError("expected ValueError for mismatched scheduler_sigmas") +def test_scheduler_meta_from_rollout_requires_sigmas(): + """The train actor consumes the rollout sigmas snapshot verbatim; no timesteps/N fallback.""" + ts = torch.tensor([999.0, 500.0, 1.0]) + sig = torch.tensor([1.0, 0.5, 0.001, 0.0]) + out_ts, out_sig = scheduler_meta_from_rollout( + {"scheduler_timesteps": ts, "scheduler_sigmas": sig}, device=torch.device("cpu") + ) + assert torch.equal(out_ts, ts) and torch.equal(out_sig, sig) + try: + scheduler_meta_from_rollout({"scheduler_timesteps": ts}, device=torch.device("cpu")) + except ValueError: + pass + else: + raise AssertionError("expected ValueError for missing scheduler_sigmas") + + # -------------------------------------------------------------------------------------- # L3 — Cond window padding: per-microbatch collate(pad_to_len) == legacy window-slice # -------------------------------------------------------------------------------------- From d6479df3e2e520ad15bcefd9cccf1e0e800de79b Mon Sep 17 00:00:00 2001 From: rockdu Date: Sun, 2 Aug 2026 16:50:16 -0700 Subject: [PATCH 2/3] fix(nft): reject samples whose scheduler meta differs from sample 0 --- miles/ray/data_conversion_hub/nft.py | 17 +++++++++++++ .../backends/fsdp_utils/test_loss_hub_nft.py | 24 +++++++++++++++++++ 2 files changed, 41 insertions(+) diff --git a/miles/ray/data_conversion_hub/nft.py b/miles/ray/data_conversion_hub/nft.py index 70b95b9c..cf12ab2f 100644 --- a/miles/ray/data_conversion_hub/nft.py +++ b/miles/ray/data_conversion_hub/nft.py @@ -72,6 +72,23 @@ def expand_samples_to_train_pairs( 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") + traj = sample.dit_trajectory + if ( + traj is None + or traj.timesteps is None + or 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" + ) + if traj.sigmas is None or not torch.equal( + traj.sigmas.detach().cpu().float(), scheduler_meta["scheduler_sigmas"] + ): + raise ValueError( + f"sample {sample.index} has different scheduler_sigmas than sample 0; " + "the converter assumes one shared schedule across the batch" + ) 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(): 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 65a355e1..5b6406f8 100644 --- a/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py +++ b/tests/fast/backends/fsdp_utils/test_loss_hub_nft.py @@ -75,6 +75,30 @@ class _Env: assert out["train_data"][0]["advantage"] == rewards[0] assert out["train_data"][2]["advantage"] == rewards[1] + def test_convert_rejects_mismatched_scheduler_meta(self): + # One shared schedule per batch, same contract as the flow_grpo converter. + 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()), + ] + samples[1].dit_trajectory.sigmas = samples[1].dit_trajectory.sigmas + 1.0 # tamper + try: + expand_samples_to_train_pairs(_args(), samples, [1.0, 2.0], [1.0, 2.0]) + except ValueError as e: + assert "scheduler_sigmas" in str(e) + else: + raise AssertionError("expected ValueError for mismatched scheduler_sigmas") + def test_convert_requires_rollout_sigmas(self): # Sigmas come from the rollout scheduler snapshot; no timesteps/1000 fallback. class _Traj: From 1000f13661de5b43e1eb68ea37db750f6e6be17a Mon Sep 17 00:00:00 2001 From: rockdu Date: Sun, 2 Aug 2026 19:05:32 -0700 Subject: [PATCH 3/3] refactor(diffusion): extract shared scheduler_meta_from_samples for converters --- miles/ray/data_conversion_hub/flow_grpo.py | 21 ++------------ miles/ray/data_conversion_hub/nft.py | 30 ++------------------ miles/utils/train_data_utils.py | 32 ++++++++++++++++++++++ 3 files changed, 36 insertions(+), 47 deletions(-) diff --git a/miles/ray/data_conversion_hub/flow_grpo.py b/miles/ray/data_conversion_hub/flow_grpo.py index 5f080c70..3394e275 100644 --- a/miles/ray/data_conversion_hub/flow_grpo.py +++ b/miles/ray/data_conversion_hub/flow_grpo.py @@ -4,6 +4,7 @@ import torch +from miles.utils.train_data_utils import scheduler_meta_from_samples from miles.utils.types import RolloutDebugTensors, Sample @@ -25,28 +26,10 @@ def _expand_samples_to_train_pairs( """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 - if first_traj.sigmas is None: - raise ValueError("sample 0 missing dit_trajectory.sigmas; rollout engine must return the sigmas snapshot") - scheduler_meta: dict[str, torch.Tensor] = { - "scheduler_timesteps": first_traj.timesteps.detach().cpu().float(), - "scheduler_sigmas": first_traj.sigmas.detach().cpu().float(), - } + scheduler_meta = scheduler_meta_from_samples(samples) 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" - ) - if traj.sigmas is None or not torch.equal( - traj.sigmas.detach().cpu().float(), scheduler_meta["scheduler_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, diff --git a/miles/ray/data_conversion_hub/nft.py b/miles/ray/data_conversion_hub/nft.py index cf12ab2f..f871d9fa 100644 --- a/miles/ray/data_conversion_hub/nft.py +++ b/miles/ray/data_conversion_hub/nft.py @@ -5,6 +5,7 @@ import torch +from miles.utils.train_data_utils import scheduler_meta_from_samples from miles.utils.types import Sample @@ -51,17 +52,7 @@ def expand_samples_to_train_pairs( 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 None: - raise ValueError("NFT needs dit_trajectory.sigmas from rollout; no timesteps-derived fallback") - scheduler_meta = { - "scheduler_timesteps": first_traj.timesteps.detach().cpu().float(), - "scheduler_sigmas": first_traj.sigmas.detach().cpu().float(), - } + scheduler_meta = scheduler_meta_from_samples(samples) sigmas = resolve_nft_sigmas( scheduler_meta["scheduler_sigmas"], training_timestep_fraction=args.diffusion_nft_timestep_fraction, @@ -72,23 +63,6 @@ def expand_samples_to_train_pairs( 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") - traj = sample.dit_trajectory - if ( - traj is None - or traj.timesteps is None - or 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" - ) - if traj.sigmas is None or not torch.equal( - traj.sigmas.detach().cpu().float(), scheduler_meta["scheduler_sigmas"] - ): - raise ValueError( - f"sample {sample.index} has different scheduler_sigmas than sample 0; " - "the converter assumes one shared schedule across the batch" - ) 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(): diff --git a/miles/utils/train_data_utils.py b/miles/utils/train_data_utils.py index ceb05969..d2a3cbb8 100644 --- a/miles/utils/train_data_utils.py +++ b/miles/utils/train_data_utils.py @@ -22,6 +22,38 @@ def stack_train_pair_rollout_debug( return torch.stack([item["rollout_debug_tensors"][key] for item in batch], dim=0) +def scheduler_meta_from_samples(samples: list) -> dict[str, torch.Tensor]: + """Build the batch scheduler meta from sample 0's trajectory, enforcing one shared schedule.""" + first = samples[0].dit_trajectory + if first is None: + raise ValueError("sample 0 missing dit_trajectory") + if first.timesteps is None: + raise ValueError("sample 0 missing dit_trajectory.timesteps") + if first.sigmas is None: + raise ValueError("sample 0 missing dit_trajectory.sigmas; rollout engine must return the sigmas snapshot") + meta = { + "scheduler_timesteps": first.timesteps.detach().cpu().float(), + "scheduler_sigmas": first.sigmas.detach().cpu().float(), + } + for sample in samples[1:]: + traj = sample.dit_trajectory + if ( + traj is None + or traj.timesteps is None + or not torch.equal(traj.timesteps.detach().cpu().float(), 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" + ) + if traj.sigmas is None or not torch.equal(traj.sigmas.detach().cpu().float(), meta["scheduler_sigmas"]): + raise ValueError( + f"sample {sample.index} has different scheduler_sigmas than sample 0; " + "the converter assumes one shared schedule across the batch" + ) + return meta + + def scheduler_meta_from_rollout( rollout_data: dict, *,