[Diffusion] Return scheduler sigmas snapshot in rollout dit_trajectory (#32683)
This commit is contained in:
@@ -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)
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user