[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(
|
dit_trajectory = RolloutDitTrajectory(
|
||||||
latents=_extract_single_sample_tensor(dit.latents, sample_idx, batch_size),
|
latents=_extract_single_sample_tensor(dit.latents, sample_idx, batch_size),
|
||||||
timesteps=dit.timesteps,
|
timesteps=dit.timesteps,
|
||||||
|
sigmas=dit.sigmas,
|
||||||
)
|
)
|
||||||
return RolloutTrajectoryData(
|
return RolloutTrajectoryData(
|
||||||
rollout_log_probs=log_probs,
|
rollout_log_probs=log_probs,
|
||||||
@@ -143,6 +144,7 @@ def _serialize_rollout_trajectory(
|
|||||||
rtd: RolloutTrajectoryData | None,
|
rtd: RolloutTrajectoryData | None,
|
||||||
*,
|
*,
|
||||||
serialized_dit_timesteps: dict | None = None,
|
serialized_dit_timesteps: dict | None = None,
|
||||||
|
serialized_dit_sigmas: dict | None = None,
|
||||||
) -> tuple[dict | None, dict | None, dict | None, dict | None]:
|
) -> tuple[dict | None, dict | None, dict | None, dict | None]:
|
||||||
"""Return order: rollout_log_probs, rollout_debug_tensors, denoising_env, dit_trajectory."""
|
"""Return order: rollout_log_probs, rollout_debug_tensors, denoising_env, dit_trajectory."""
|
||||||
if rtd is None:
|
if rtd is None:
|
||||||
@@ -182,6 +184,7 @@ def _serialize_rollout_trajectory(
|
|||||||
_maybe_serialize(dit.latents) if dit.latents is not None else None
|
_maybe_serialize(dit.latents) if dit.latents is not None else None
|
||||||
),
|
),
|
||||||
"timesteps": serialized_dit_timesteps,
|
"timesteps": serialized_dit_timesteps,
|
||||||
|
"sigmas": serialized_dit_sigmas,
|
||||||
}
|
}
|
||||||
return (
|
return (
|
||||||
serialized_log_probs,
|
serialized_log_probs,
|
||||||
@@ -211,10 +214,14 @@ def _build_response(
|
|||||||
), "rollout_trajectory_data must be present when rollout=True"
|
), "rollout_trajectory_data must be present when rollout=True"
|
||||||
|
|
||||||
serialized_dit_timesteps = None
|
serialized_dit_timesteps = None
|
||||||
|
serialized_dit_sigmas = None
|
||||||
if rollout and rollout_trajectory_data and rollout_trajectory_data.dit_trajectory:
|
if rollout and rollout_trajectory_data and rollout_trajectory_data.dit_trajectory:
|
||||||
serialized_dit_timesteps = _maybe_serialize(
|
serialized_dit_timesteps = _maybe_serialize(
|
||||||
rollout_trajectory_data.dit_trajectory.timesteps
|
rollout_trajectory_data.dit_trajectory.timesteps
|
||||||
)
|
)
|
||||||
|
serialized_dit_sigmas = _maybe_serialize(
|
||||||
|
rollout_trajectory_data.dit_trajectory.sigmas
|
||||||
|
)
|
||||||
|
|
||||||
responses: list[RolloutResponse] = []
|
responses: list[RolloutResponse] = []
|
||||||
for sample_idx in range(batch_size):
|
for sample_idx in range(batch_size):
|
||||||
@@ -245,6 +252,7 @@ def _build_response(
|
|||||||
) = _serialize_rollout_trajectory(
|
) = _serialize_rollout_trajectory(
|
||||||
per_sample_trajectory,
|
per_sample_trajectory,
|
||||||
serialized_dit_timesteps=serialized_dit_timesteps,
|
serialized_dit_timesteps=serialized_dit_timesteps,
|
||||||
|
serialized_dit_sigmas=serialized_dit_sigmas,
|
||||||
)
|
)
|
||||||
responses.append(
|
responses.append(
|
||||||
RolloutResponse(
|
RolloutResponse(
|
||||||
|
|||||||
@@ -52,6 +52,8 @@ class RolloutDitTrajectory:
|
|||||||
# final denoised latent x_{t_T} (last scheduler.step output).
|
# final denoised latent x_{t_T} (last scheduler.step output).
|
||||||
latents: torch.Tensor | None = None
|
latents: torch.Tensor | None = None
|
||||||
timesteps: torch.Tensor | None = None # [T]
|
timesteps: torch.Tensor | None = None # [T]
|
||||||
|
# [T+1] scheduler.sigmas snapshot (post-shift, includes terminal 0).
|
||||||
|
sigmas: torch.Tensor | None = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -186,6 +186,7 @@ class RolloutDenoisingMixin:
|
|||||||
batch.rollout_trajectory_data.dit_trajectory = RolloutDitTrajectory(
|
batch.rollout_trajectory_data.dit_trajectory = RolloutDitTrajectory(
|
||||||
latents=step_latents_tensor.cpu(),
|
latents=step_latents_tensor.cpu(),
|
||||||
timesteps=torch.stack(step_timesteps, dim=0).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:
|
if env is not None and batch.rollout_return_denoising_env:
|
||||||
|
|||||||
@@ -207,11 +207,13 @@ class TestSerializeRolloutTrajectory(unittest.TestCase):
|
|||||||
dit_trajectory=RolloutDitTrajectory(
|
dit_trajectory=RolloutDitTrajectory(
|
||||||
latents=torch.randn(1, 5, 4, 2, 2, 2),
|
latents=torch.randn(1, 5, 4, 2, 2, 2),
|
||||||
timesteps=torch.tensor([1.0, 0.75, 0.5, 0.25]),
|
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(
|
_, _, env, dit_traj = _serialize_rollout_trajectory(
|
||||||
rtd,
|
rtd,
|
||||||
serialized_dit_timesteps=_maybe_serialize(rtd.dit_trajectory.timesteps),
|
serialized_dit_timesteps=_maybe_serialize(rtd.dit_trajectory.timesteps),
|
||||||
|
serialized_dit_sigmas=_maybe_serialize(rtd.dit_trajectory.sigmas),
|
||||||
)
|
)
|
||||||
self.assertIsNotNone(env)
|
self.assertIsNotNone(env)
|
||||||
self.assertIn("pos_cond_kwargs", env)
|
self.assertIn("pos_cond_kwargs", env)
|
||||||
@@ -219,8 +221,10 @@ class TestSerializeRolloutTrajectory(unittest.TestCase):
|
|||||||
self.assertIsNotNone(dit_traj)
|
self.assertIsNotNone(dit_traj)
|
||||||
self.assertIn("latents", dit_traj)
|
self.assertIn("latents", dit_traj)
|
||||||
self.assertIn("timesteps", dit_traj)
|
self.assertIn("timesteps", dit_traj)
|
||||||
|
self.assertIn("sigmas", dit_traj)
|
||||||
self.assertTrue(dit_traj["latents"]["__tensor__"])
|
self.assertTrue(dit_traj["latents"]["__tensor__"])
|
||||||
self.assertTrue(dit_traj["timesteps"]["__tensor__"])
|
self.assertTrue(dit_traj["timesteps"]["__tensor__"])
|
||||||
|
self.assertTrue(dit_traj["sigmas"]["__tensor__"])
|
||||||
|
|
||||||
|
|
||||||
class TestBuildResponse(unittest.TestCase):
|
class TestBuildResponse(unittest.TestCase):
|
||||||
@@ -327,6 +331,7 @@ class TestBuildResponse(unittest.TestCase):
|
|||||||
dit_trajectory=RolloutDitTrajectory(
|
dit_trajectory=RolloutDitTrajectory(
|
||||||
latents=torch.randn(B, T + 1, D),
|
latents=torch.randn(B, T + 1, D),
|
||||||
timesteps=torch.linspace(1.0, 0.0, T),
|
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"])
|
ts1 = bytes_to_tensor(resps[1].dit_trajectory["timesteps"]["data"])
|
||||||
self.assertEqual(ts0.shape, (T,))
|
self.assertEqual(ts0.shape, (T,))
|
||||||
self.assertTrue(torch.equal(ts0, ts1))
|
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(
|
self.assertEqual(
|
||||||
_maybe_deserialize(resps[1].dit_trajectory["latents"]).shape, (T + 1, D)
|
_maybe_deserialize(resps[1].dit_trajectory["latents"]).shape, (T + 1, D)
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user