From cbeeeb5fe42c4aaaef64d57002011477c326474a Mon Sep 17 00:00:00 2001 From: rockdu Date: Sat, 29 Aug 2026 00:50:20 -0700 Subject: [PATCH] [Diffusion] Rollout API: filter trajectory latents to the requested window with step-index provenance --- .../entrypoints/post_training/rollout_api.py | 7 ++++++ .../minimax_h3/minimax_h3_rollout.py | 23 ++++++++++++++++--- .../minimax_h3/stages/denoising.py | 12 +++++++++- .../runtime/post_training/rl_dataclasses.py | 5 +++- .../post_training/rollout_denoising_mixin.py | 7 +++++- 5 files changed, 48 insertions(+), 6 deletions(-) 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 55253ab49703..331a1f016cbd 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 @@ -140,6 +140,7 @@ def _slice_rollout_trajectory_for_sample( latents=_extract_single_sample_tensor(dit.latents, sample_idx, batch_size), timesteps=dit.timesteps, sigmas=dit.sigmas, + latent_step_indices=dit.latent_step_indices, ) return RolloutTrajectoryData( rollout_log_probs=log_probs, @@ -194,6 +195,11 @@ def _serialize_rollout_trajectory( ), "timesteps": serialized_dit_timesteps, "sigmas": serialized_dit_sigmas, + "latent_step_indices": ( + _maybe_serialize(dit.latent_step_indices) + if dit.latent_step_indices is not None + else None + ), } return ( serialized_log_probs, @@ -360,6 +366,7 @@ async def rollout_generate(request: RolloutRequest): ) from exc if output_batch.error: raise HTTPException(status_code=500, detail=output_batch.error) + def _serialize_response() -> list[bytes]: with stamps.span("build"): rollout_responses = _build_response( diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/minimax_h3_rollout.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/minimax_h3_rollout.py index e2cd6877b66b..2bccab0e0ca4 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/minimax_h3_rollout.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/minimax_h3_rollout.py @@ -168,13 +168,27 @@ class MiniMaxH3RolloutCollector: """Accumulates video-target trajectory and per-step log probs.""" sigmas_video: list[float] + return_step_indices: set[int] | None = None latent_steps: list[torch.Tensor] = field(default_factory=list) + latent_step_indices: list[int] = field(default_factory=list) log_prob_sums: list[torch.Tensor] = field(default_factory=list) log_prob_counts: list[torch.Tensor] = field(default_factory=list) pos_cond_kwargs: dict[str, Any] = field(default_factory=dict) + _next_latent_index: int = 0 + + def _record_latent(self, video_target: torch.Tensor) -> None: + index = self._next_latent_index + self._next_latent_index += 1 + if ( + self.return_step_indices is not None + and index not in self.return_step_indices + ): + return + self.latent_steps.append(video_target.detach().cpu().clone()) + self.latent_step_indices.append(index) def record_initial(self, video_target: torch.Tensor) -> None: - self.latent_steps.append(video_target.detach().cpu().clone()) + self._record_latent(video_target) def record_step( self, @@ -182,7 +196,7 @@ def record_step( log_prob_sum: torch.Tensor, log_prob_count: torch.Tensor, ) -> None: - self.latent_steps.append(video_target.detach().cpu().clone()) + self._record_latent(video_target) self.log_prob_sums.append(log_prob_sum.detach().cpu()) self.log_prob_counts.append(log_prob_count.detach().cpu()) @@ -191,7 +205,7 @@ def build_trajectory_data(self) -> RolloutTrajectoryData: stacked = torch.stack(self.latent_steps, dim=0).unsqueeze(0) divisor = 1000.0 step_sigmas = torch.tensor( - [float(s) for s in self.sigmas_video[:-1]], + [float(s) for s in self.sigmas_video], dtype=torch.float32, ) timesteps = step_sigmas * divisor @@ -214,6 +228,9 @@ def build_trajectory_data(self) -> RolloutTrajectoryData: latents=stacked, timesteps=timesteps, sigmas=sigmas, + latent_step_indices=torch.tensor( + self.latent_step_indices, dtype=torch.long + ), ), ) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py index 991a3d81ee50..91efb95ec3f3 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py @@ -679,7 +679,17 @@ def _run_full_loop(self, batch: Req, server_args: ServerArgs) -> None: if isinstance(seed, list): seed = seed[0] generator.manual_seed(int(seed)) - collector = MiniMaxH3RolloutCollector(sigmas_video=sigmas_video) + return_step_indices = getattr( + batch, "rollout_return_step_indices", None + ) + collector = MiniMaxH3RolloutCollector( + sigmas_video=sigmas_video, + return_step_indices=( + set(return_step_indices) + if return_step_indices is not None + else None + ), + ) packed_cpu = { k: (v.detach().cpu() if isinstance(v, torch.Tensor) else v) for k, v in packed.items() 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 1882776d573d..c48a4fcb3bdd 100644 --- a/python/sglang/multimodal_gen/runtime/post_training/rl_dataclasses.py +++ b/python/sglang/multimodal_gen/runtime/post_training/rl_dataclasses.py @@ -51,9 +51,12 @@ class RolloutDitTrajectory: # [B, T+1, ...]: per-step noisy latents x_{t_0..t_{T-1}} followed by the # final denoised latent x_{t_T} (last scheduler.step output). latents: torch.Tensor | None = None - timesteps: torch.Tensor | None = None # [T] + timesteps: torch.Tensor | None = None # [T+1], includes the terminal timestep # [T+1] scheduler.sigmas snapshot (post-shift, includes terminal 0). sigmas: torch.Tensor | None = None + # [K] original step index of each kept latent (0..T); None means the full + # 0..T trajectory. Set whenever rollout_return_step_indices filtered it. + latent_step_indices: 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 8b754519812a..f94084db7387 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 @@ -139,6 +139,7 @@ def _maybe_init_denoising_env_collection( batch._rollout_denoising_env_state = { "env": env, "step_latents": [], + "step_latent_indices": [], "step_timesteps": [], "pos_cond_kwargs_src": pos_src, "neg_cond_kwargs_src": neg_src, @@ -157,12 +158,13 @@ def _maybe_append_dit_trajectory_step( if state is None: return + state["step_timesteps"].append(timestep_value.detach().cpu()) return_step_indices = getattr(batch, "rollout_return_step_indices", None) if return_step_indices is not None and step_index not in return_step_indices: return state["step_latents"].append(latents.detach()) - state["step_timesteps"].append(timestep_value.detach().cpu()) + state["step_latent_indices"].append(step_index) def _maybe_finalize_denoising_env_collection(self, batch, pipeline_config) -> None: state = getattr(batch, "_rollout_denoising_env_state", None) @@ -187,6 +189,9 @@ def _maybe_finalize_denoising_env_collection(self, batch, pipeline_config) -> No latents=step_latents_tensor.cpu(), timesteps=torch.stack(step_timesteps, dim=0).cpu(), sigmas=batch.scheduler.sigmas.detach().cpu().clone(), + latent_step_indices=torch.tensor( + state["step_latent_indices"], dtype=torch.long + ), ) if env is not None and batch.rollout_return_denoising_env: