diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/rollout_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/rollout_api.py index ce2a77a9c..3b600bbba 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/rollout_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/rollout_api.py @@ -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( diff --git a/python/sglang/multimodal_gen/runtime/post_training/rl_dataclasses.py b/python/sglang/multimodal_gen/runtime/post_training/rl_dataclasses.py index 351c17736..1882776d5 100644 --- a/python/sglang/multimodal_gen/runtime/post_training/rl_dataclasses.py +++ b/python/sglang/multimodal_gen/runtime/post_training/rl_dataclasses.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/post_training/rollout_denoising_mixin.py b/python/sglang/multimodal_gen/runtime/post_training/rollout_denoising_mixin.py index c161973a1..8b7545198 100644 --- a/python/sglang/multimodal_gen/runtime/post_training/rollout_denoising_mixin.py +++ b/python/sglang/multimodal_gen/runtime/post_training/rollout_denoising_mixin.py @@ -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: diff --git a/python/sglang/multimodal_gen/test/unit/test_rollout_api.py b/python/sglang/multimodal_gen/test/unit/test_rollout_api.py index 791e61ad5..8f31e6531 100644 --- a/python/sglang/multimodal_gen/test/unit/test_rollout_api.py +++ b/python/sglang/multimodal_gen/test/unit/test_rollout_api.py @@ -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) )