Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -168,21 +168,35 @@ 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,
video_target: torch.Tensor,
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())

Expand All @@ -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
Expand All @@ -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
),
),
)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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)
Expand All @@ -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:
Expand Down
Loading