[Diffusion] Return scheduler sigmas snapshot in rollout dit_trajectory (#32683)

This commit is contained in:
Kangrui Du
2026-07-31 00:29:07 -07:00
committed by GitHub
parent 0d6bef6b6d
commit 585a7d05e3
4 changed files with 20 additions and 0 deletions
@@ -130,6 +130,7 @@ def _slice_rollout_trajectory_for_sample(
dit_trajectory = RolloutDitTrajectory(
latents=_extract_single_sample_tensor(dit.latents, sample_idx, batch_size),
timesteps=dit.timesteps,
sigmas=dit.sigmas,
)
return RolloutTrajectoryData(
rollout_log_probs=log_probs,
@@ -143,6 +144,7 @@ def _serialize_rollout_trajectory(
rtd: RolloutTrajectoryData | None,
*,
serialized_dit_timesteps: dict | None = None,
serialized_dit_sigmas: dict | None = None,
) -> tuple[dict | None, dict | None, dict | None, dict | None]:
"""Return order: rollout_log_probs, rollout_debug_tensors, denoising_env, dit_trajectory."""
if rtd is None:
@@ -182,6 +184,7 @@ def _serialize_rollout_trajectory(
_maybe_serialize(dit.latents) if dit.latents is not None else None
),
"timesteps": serialized_dit_timesteps,
"sigmas": serialized_dit_sigmas,
}
return (
serialized_log_probs,
@@ -211,10 +214,14 @@ def _build_response(
), "rollout_trajectory_data must be present when rollout=True"
serialized_dit_timesteps = None
serialized_dit_sigmas = None
if rollout and rollout_trajectory_data and rollout_trajectory_data.dit_trajectory:
serialized_dit_timesteps = _maybe_serialize(
rollout_trajectory_data.dit_trajectory.timesteps
)
serialized_dit_sigmas = _maybe_serialize(
rollout_trajectory_data.dit_trajectory.sigmas
)
responses: list[RolloutResponse] = []
for sample_idx in range(batch_size):
@@ -245,6 +252,7 @@ def _build_response(
) = _serialize_rollout_trajectory(
per_sample_trajectory,
serialized_dit_timesteps=serialized_dit_timesteps,
serialized_dit_sigmas=serialized_dit_sigmas,
)
responses.append(
RolloutResponse(
@@ -52,6 +52,8 @@ class RolloutDitTrajectory:
# final denoised latent x_{t_T} (last scheduler.step output).
latents: torch.Tensor | None = None
timesteps: torch.Tensor | None = None # [T]
# [T+1] scheduler.sigmas snapshot (post-shift, includes terminal 0).
sigmas: torch.Tensor | None = None
@dataclass
@@ -186,6 +186,7 @@ class RolloutDenoisingMixin:
batch.rollout_trajectory_data.dit_trajectory = RolloutDitTrajectory(
latents=step_latents_tensor.cpu(),
timesteps=torch.stack(step_timesteps, dim=0).cpu(),
sigmas=batch.scheduler.sigmas.detach().cpu().clone(),
)
if env is not None and batch.rollout_return_denoising_env:
@@ -207,11 +207,13 @@ class TestSerializeRolloutTrajectory(unittest.TestCase):
dit_trajectory=RolloutDitTrajectory(
latents=torch.randn(1, 5, 4, 2, 2, 2),
timesteps=torch.tensor([1.0, 0.75, 0.5, 0.25]),
sigmas=torch.tensor([1.0, 0.75, 0.5, 0.25, 0.0]),
),
)
_, _, env, dit_traj = _serialize_rollout_trajectory(
rtd,
serialized_dit_timesteps=_maybe_serialize(rtd.dit_trajectory.timesteps),
serialized_dit_sigmas=_maybe_serialize(rtd.dit_trajectory.sigmas),
)
self.assertIsNotNone(env)
self.assertIn("pos_cond_kwargs", env)
@@ -219,8 +221,10 @@ class TestSerializeRolloutTrajectory(unittest.TestCase):
self.assertIsNotNone(dit_traj)
self.assertIn("latents", dit_traj)
self.assertIn("timesteps", dit_traj)
self.assertIn("sigmas", dit_traj)
self.assertTrue(dit_traj["latents"]["__tensor__"])
self.assertTrue(dit_traj["timesteps"]["__tensor__"])
self.assertTrue(dit_traj["sigmas"]["__tensor__"])
class TestBuildResponse(unittest.TestCase):
@@ -327,6 +331,7 @@ class TestBuildResponse(unittest.TestCase):
dit_trajectory=RolloutDitTrajectory(
latents=torch.randn(B, T + 1, D),
timesteps=torch.linspace(1.0, 0.0, T),
sigmas=torch.linspace(1.0, 0.0, T + 1),
),
),
)
@@ -339,6 +344,10 @@ class TestBuildResponse(unittest.TestCase):
ts1 = bytes_to_tensor(resps[1].dit_trajectory["timesteps"]["data"])
self.assertEqual(ts0.shape, (T,))
self.assertTrue(torch.equal(ts0, ts1))
sg0 = bytes_to_tensor(resps[0].dit_trajectory["sigmas"]["data"])
sg1 = bytes_to_tensor(resps[1].dit_trajectory["sigmas"]["data"])
self.assertEqual(sg0.shape, (T + 1,))
self.assertTrue(torch.equal(sg0, sg1))
self.assertEqual(
_maybe_deserialize(resps[1].dit_trajectory["latents"]).shape, (T + 1, D)
)