[diffusion] rl: revamp rollout Log-Prob support with SDE/CPS for RL post-training (#21204)

Co-authored-by: MikukuOvO <mikukuovo@gmail.com>
This commit is contained in:
Kangrui Du
2026-04-09 09:00:00 +08:00
committed by GitHub
co-authored by MikukuOvO
parent 2c4e113dd7
commit 1b7c33a5b7
17 changed files with 944 additions and 11 deletions
@@ -0,0 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
"""Configuration for post-training and RL-related features (e.g. rollout)."""
from sglang.multimodal_gen.configs.post_training.rl_rollout import RLRolloutArgs
__all__ = ["RLRolloutArgs"]
@@ -0,0 +1,121 @@
# SPDX-License-Identifier: Apache-2.0
"""CLI- and API-facing configuration for diffusion post-training / rollout paths."""
from __future__ import annotations
import argparse
import math
from dataclasses import dataclass
from typing import Any, Callable
from sglang.multimodal_gen.utils import StoreBoolean
_VALID_ROLLOUT_SDE_TYPES = ("sde", "cps", "ode")
@dataclass
class RLRolloutArgs:
"""Rollout (log-prob trajectory) options used by SamplingParams and APIs."""
rollout: bool = False
rollout_sde_type: str = "sde"
rollout_noise_level: float = 0.7
rollout_log_prob_no_const: bool = False
rollout_debug_mode: bool = False
def validate(self) -> None:
noise = self.rollout_noise_level
if isinstance(noise, bool) or not isinstance(noise, (int, float)):
raise ValueError(f"rollout_noise_level must be a number, got {noise!r}")
if not math.isfinite(float(noise)):
raise ValueError(f"rollout_noise_level must be finite, got {noise!r}")
if float(noise) < 0.0:
raise ValueError(f"rollout_noise_level must be non-negative, got {noise!r}")
if self.rollout_sde_type not in _VALID_ROLLOUT_SDE_TYPES:
raise ValueError(
f"rollout_sde_type must be one of {_VALID_ROLLOUT_SDE_TYPES}, "
f"got {self.rollout_sde_type!r}"
)
@classmethod
def validate_sampling_params(cls, params: Any) -> None:
"""Validate rollout fields on a duck-typed object (e.g. ``SamplingParams``).
Mirrors how ``ServerArgs`` runs ``NunchakuSVDQuantArgs.validate()`` from
``_adjust_quant_config`` instead of inlining checks in a large validator.
"""
cls(
rollout=params.rollout,
rollout_sde_type=params.rollout_sde_type,
rollout_noise_level=params.rollout_noise_level,
rollout_log_prob_no_const=params.rollout_log_prob_no_const,
rollout_debug_mode=params.rollout_debug_mode,
).validate()
@staticmethod
def add_cli_args(
parser: Any,
add_argument: Callable[..., Any] | None = None,
) -> None:
"""Register rollout-related CLI flags on ``parser``.
If ``add_argument`` is provided (e.g. SamplingParams' wrapper with
``default=argparse.SUPPRESS``), it is used; otherwise a local wrapper
is applied.
"""
if add_argument is None:
def _add(*name_or_flags: Any, **kwargs: Any):
kwargs.setdefault("default", argparse.SUPPRESS)
return parser.add_argument(*name_or_flags, **kwargs)
add_argument = _add
add_argument(
"--rollout",
action="store_true",
help="Enable rollout mode and return per-step log_prob trajectory",
)
add_argument(
"--rollout-sde-type",
type=str,
choices=list(_VALID_ROLLOUT_SDE_TYPES),
help="Rollout step objective type used in log-prob computation.",
)
add_argument(
"--rollout-noise-level",
type=float,
help="Noise level used by rollout SDE/CPS step objective.",
)
add_argument(
"--rollout-log-prob-no-const",
action=StoreBoolean,
help="If true, return rollout log-prob without constant terms.",
)
add_argument(
"--rollout-debug-mode",
action=StoreBoolean,
help=(
"If true, return rollout debug tensors "
"(variance noise, mean, std, model output)."
),
)
@classmethod
def from_dict(cls, kwargs: dict[str, Any]) -> RLRolloutArgs:
return cls(
rollout=bool(kwargs.get("rollout", cls.rollout)),
rollout_sde_type=str(kwargs.get("rollout_sde_type", cls.rollout_sde_type)),
rollout_noise_level=float(
kwargs.get("rollout_noise_level", cls.rollout_noise_level)
),
rollout_log_prob_no_const=bool(
kwargs.get("rollout_log_prob_no_const", cls.rollout_log_prob_no_const)
),
rollout_debug_mode=bool(
kwargs.get("rollout_debug_mode", cls.rollout_debug_mode)
),
)
@@ -16,6 +16,7 @@ from dataclasses import dataclass
from enum import Enum, auto
from typing import TYPE_CHECKING, Any, ClassVar
from sglang.multimodal_gen.configs.post_training import RLRolloutArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.utils import StoreBoolean, expand_path_fields
@@ -177,6 +178,13 @@ class SamplingParams:
# Misc
save_output: bool = True
return_frames: bool = False
rollout: bool = False
rollout_sde_type: str = "sde"
rollout_noise_level: float = 0.7
rollout_log_prob_no_const: bool = False # exclude constants in rollout logprob
rollout_debug_mode: bool = (
False # return rollout debug tensors (intermediate states)
)
return_trajectory_latents: bool = False # returns all latents for each timestep
return_trajectory_decoded: bool = False # returns decoded latents for each timestep
# if True, disallow user params to override subclass-defined protected fields
@@ -358,6 +366,8 @@ class SamplingParams:
f"boundary_ratio must be within [0, 1], got {self.boundary_ratio!r}"
)
RLRolloutArgs.validate_sampling_params(self)
def check_sampling_param(self):
# Keep backward-compatibility for old call sites.
self._validate()
@@ -820,6 +830,10 @@ class SamplingParams:
action="store_true",
help="Whether to return the trajectory",
)
# Rollout arguments
RLRolloutArgs.add_cli_args(parser, add_argument=add_argument)
add_argument(
"--return-trajectory-decoded",
action="store_true",
@@ -44,6 +44,11 @@ def sequence_model_parallel_all_gather(
return get_sp_group().all_gather(input_, dim)
def sequence_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
"""All-reduce the input tensor across model parallel group."""
return get_sp_group().all_reduce(input_)
def cfg_model_parallel_all_gather(
input_: torch.Tensor, dim: int = -1, separate_tensors: bool = False
) -> torch.Tensor:
@@ -256,6 +256,7 @@ class DiffGenerator:
),
trajectory_latents=output_batch.trajectory_latents,
trajectory_timesteps=output_batch.trajectory_timesteps,
rollout_trajectory_data=output_batch.rollout_trajectory_data,
trajectory_decoded=output_batch.trajectory_decoded,
)
@@ -108,6 +108,7 @@ class GenerationResult:
metrics: dict = field(default_factory=dict)
trajectory_latents: Any = None
trajectory_timesteps: Any = None
rollout_trajectory_data: Any = None
trajectory_decoded: Any = None
prompt_index: int = 0
output_file_path: str | None = None
@@ -237,6 +237,9 @@ class GPUWorker:
metrics=result.metrics,
trajectory_timesteps=getattr(result, "trajectory_timesteps", None),
trajectory_latents=getattr(result, "trajectory_latents", None),
rollout_trajectory_data=getattr(
result, "rollout_trajectory_data", None
),
noise_pred=getattr(result, "noise_pred", None),
trajectory_decoded=getattr(result, "trajectory_decoded", None),
)
@@ -943,7 +943,8 @@ class ZImageTransformer2DModel(CachableDiT, OffloadableDiTMixin):
x = list(unified.unbind(dim=0))
x = self.unpatchify(x, x_size, patch_size, f_patch_size)
return -x[0]
# Keep batch dim so output shape matches input (e.g. rollout/scheduler expect same ndim).
return -torch.stack(x)
EntryClass = ZImageTransformer2DModel
@@ -32,6 +32,9 @@ from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput
from sglang.multimodal_gen.runtime.models.schedulers.base import BaseScheduler
from sglang.multimodal_gen.runtime.post_training.scheduler_rl_mixin import (
SchedulerRLMixin,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
@@ -51,7 +54,9 @@ class FlowMatchEulerDiscreteSchedulerOutput(BaseOutput):
prev_sample: torch.FloatTensor
class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
class FlowMatchEulerDiscreteScheduler(
SchedulerMixin, ConfigMixin, BaseScheduler, SchedulerRLMixin
):
"""
Euler scheduler.
@@ -447,6 +452,7 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler
s_noise: float = 1.0,
generator: torch.Generator | None = None,
per_token_timesteps: torch.Tensor | None = None,
batch=None,
return_dict: bool = True,
) -> FlowMatchEulerDiscreteSchedulerOutput | tuple[torch.FloatTensor, ...]:
"""
@@ -516,12 +522,19 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler
next_sigma = sigma_next
dt = sigma_next - sigma
if self.config.stochastic_sampling:
x0 = sample - current_sigma * model_output
noise = torch.randn_like(sample)
prev_sample = (1.0 - next_sigma) * x0 + next_sigma * noise
if batch is not None and batch.rollout:
if not self.already_prepared_rollout(batch):
raise RuntimeError("Rollout not prepared before step")
prev_sample = self.flow_sde_sampling(
batch, model_output, sample, current_sigma, next_sigma, generator
)
else:
prev_sample = sample + dt * model_output
if self.config.stochastic_sampling:
x0 = sample - current_sigma * model_output
noise = torch.randn_like(sample)
prev_sample = (1.0 - next_sigma) * x0 + next_sigma * noise
else:
prev_sample = sample + dt * model_output
# upon completion increase step index by one
assert self._step_index is not None, "_step_index should not be None"
@@ -22,6 +22,9 @@ import PIL.Image
import torch
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
from sglang.multimodal_gen.runtime.post_training.rl_dataclasses import (
RolloutTrajectoryData,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import (
_sanitize_for_logging,
@@ -131,8 +134,9 @@ class Req:
# Component modules (populated by the pipeline)
modules: dict[str, Any] = field(default_factory=dict)
trajectory_timesteps: list[torch.Tensor] | None = None
trajectory_timesteps: torch.Tensor | None = None
trajectory_latents: torch.Tensor | None = None
rollout_trajectory_data: RolloutTrajectoryData | None = None
trajectory_audio_latents: torch.Tensor | None = None
# Extra parameters that might be needed by specific pipeline implementations (e.g., LTX2.3 DenoisingAVStage)
@@ -333,8 +337,9 @@ class OutputBatch:
output: torch.Tensor | None = None
audio: torch.Tensor | None = None
audio_sample_rate: int | None = None
trajectory_timesteps: list[torch.Tensor] | None = None
trajectory_timesteps: torch.Tensor | None = None
trajectory_latents: torch.Tensor | None = None
rollout_trajectory_data: RolloutTrajectoryData | None = None
trajectory_decoded: list[torch.Tensor] | None = None
error: str | None = None
output_file_paths: list[str] | None = None
@@ -236,6 +236,7 @@ class DecodingStage(PipelineStage):
output=frames,
trajectory_timesteps=batch.trajectory_timesteps,
trajectory_latents=batch.trajectory_latents,
rollout_trajectory_data=batch.rollout_trajectory_data,
trajectory_decoded=trajectory_decoded,
metrics=batch.metrics,
noise_pred=None,
@@ -78,6 +78,12 @@ from sglang.multimodal_gen.runtime.platforms import (
AttentionBackendEnum,
current_platform,
)
from sglang.multimodal_gen.runtime.post_training.rl_dataclasses import (
RolloutTrajectoryData,
)
from sglang.multimodal_gen.runtime.post_training.scheduler_rl_mixin import (
SchedulerRLMixin,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -134,6 +140,43 @@ class DenoisingStage(PipelineStage):
self._cached_num_steps = None
self._is_warmed_up = False
def _maybe_prepare_rollout(self, batch: Req):
"""Prepare denoising loop for rollout."""
if not isinstance(self.scheduler, SchedulerRLMixin):
if batch.rollout:
raise ValueError(
f"Scheduler {type(self.scheduler)} does not support rollout"
)
return
self.scheduler.release_rollout_resources(batch)
if batch.rollout:
self.scheduler.prepare_rollout(
batch=batch,
pipeline_config=self.server_args.pipeline_config,
)
def _maybe_collect_rollout_log_probs(self, batch: Req):
"""Get rollout log probs and store into batch for reward calculation."""
if not isinstance(self.scheduler, SchedulerRLMixin):
if batch.rollout:
raise ValueError(
f"Scheduler {type(self.scheduler)} does not support rollout"
)
return
if batch.rollout:
if batch.rollout_trajectory_data is None:
batch.rollout_trajectory_data = RolloutTrajectoryData()
batch.rollout_trajectory_data.rollout_log_probs = (
self.scheduler.collect_rollout_log_probs(batch)
)
if getattr(batch, "rollout_debug_mode", False):
batch.rollout_trajectory_data.rollout_debug_tensors = (
self.scheduler.collect_rollout_debug_tensors(batch)
)
self.scheduler.release_rollout_resources(batch)
def _maybe_enable_torch_compile(self, module: object) -> None:
"""
Compile a module with torch.compile, and enable inductor overlap tweak if available.
@@ -563,10 +606,13 @@ class DenoisingStage(PipelineStage):
else:
self._maybe_enable_cache_dit(cache_dit_num_inference_steps, batch)
if batch.rollout:
self._maybe_prepare_rollout(batch)
# Prepare extra step kwargs for scheduler
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step,
{"generator": batch.generator, "eta": batch.eta},
{"generator": batch.generator, "eta": batch.eta, "batch": batch},
)
# Setup precision and autocast settings
@@ -726,6 +772,10 @@ class DenoisingStage(PipelineStage):
trajectory_tensor = None
trajectory_timesteps_tensor = None
# Gather log probs for rollout
if batch.rollout:
self._maybe_collect_rollout_log_probs(batch)
# Gather results if using sequence parallelism
latents, trajectory_tensor = self._postprocess_sp_latents(
batch, latents, trajectory_tensor
@@ -1093,7 +1143,6 @@ class DenoisingStage(PipelineStage):
guidance=guidance,
latents=latents,
)
# Save noise_pred to batch for external access (e.g., ComfyUI)
if server_args.comfyui_mode:
batch.noise_pred = noise_pred
@@ -0,0 +1,46 @@
# SPDX-License-Identifier: Apache-2.0
"""RL-specific dataclasses used by post-training and rollout paths."""
from dataclasses import dataclass, field
from typing import Any
import torch
@dataclass
class RolloutSessionData:
"""Per-batch rollout state created by prepare_rollout(), lives on the batch object.
Cleared by setting ``batch._rollout_session_data = None``.
"""
pipeline_config: Any = None
sigma_max: float = 0.0
latents_shape: tuple | None = None
noise_buffer: torch.Tensor | None = None
local_log_prob_sum: list[torch.Tensor] = field(default_factory=list)
local_log_prob_count: list[torch.Tensor] = field(default_factory=list)
local_variance_noises: list[torch.Tensor] = field(default_factory=list)
local_prev_sample_means: list[torch.Tensor] = field(default_factory=list)
local_noise_std_devs: list[torch.Tensor] = field(default_factory=list)
local_model_outputs: list[torch.Tensor] = field(default_factory=list)
@dataclass
class RolloutDebugTensors:
"""Container for rollout debug tensors collected during denoising."""
rollout_variance_noises: torch.Tensor | None = None
rollout_prev_sample_means: torch.Tensor | None = None
rollout_noise_std_devs: torch.Tensor | None = None
rollout_model_outputs: torch.Tensor | None = None
@dataclass
class RolloutTrajectoryData:
"""Container for rollout-specific trajectory outputs."""
rollout_log_probs: torch.Tensor | None = None
rollout_debug_tensors: RolloutDebugTensors | None = None
@@ -0,0 +1,115 @@
# SPDX-License-Identifier: Apache-2.0
"""Debug tensor helpers for rollout-enabled schedulers."""
import torch
from sglang.multimodal_gen.runtime.distributed import (
get_local_torch_device,
get_sp_world_size,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.post_training.rl_dataclasses import (
RolloutDebugTensors,
RolloutSessionData,
)
class SchedulerRLDebugMixin:
@staticmethod
def _reset_rollout_debug_tensors(rollout_session_data: RolloutSessionData) -> None:
rollout_session_data.local_variance_noises = []
rollout_session_data.local_prev_sample_means = []
rollout_session_data.local_noise_std_devs = []
rollout_session_data.local_model_outputs = []
def append_local_rollout_debug_tensors(
self,
batch,
*,
variance_noise: torch.Tensor,
prev_sample_mean: torch.Tensor,
noise_std_dev: torch.Tensor,
model_output: torch.Tensor,
) -> None:
rollout_session_data = batch._rollout_session_data
batch_size = variance_noise.shape[0]
rollout_session_data.local_variance_noises.append(variance_noise)
rollout_session_data.local_prev_sample_means.append(prev_sample_mean)
rollout_session_data.local_noise_std_devs.append(
noise_std_dev.expand((batch_size, 1))
)
rollout_session_data.local_model_outputs.append(model_output)
def consume_local_rollout_debug_tensors(
self,
batch,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
rollout_session_data = batch._rollout_session_data
variance_noises = torch.stack(rollout_session_data.local_variance_noises, dim=1)
prev_sample_means = torch.stack(
rollout_session_data.local_prev_sample_means, dim=1
)
noise_std_devs = torch.stack(rollout_session_data.local_noise_std_devs, dim=1)
model_outputs = torch.stack(rollout_session_data.local_model_outputs, dim=1)
self._reset_rollout_debug_tensors(rollout_session_data)
return variance_noises, prev_sample_means, noise_std_devs, model_outputs
def collect_rollout_debug_tensors(self, batch: Req) -> RolloutDebugTensors:
"""
Consume rollout debug tensors and merge for all SP ranks.
Returns rollout debug tensors with shape [B, T, ...].
"""
rollout_session_data = batch._rollout_session_data
variance_noises, prev_sample_means, noise_std_devs, model_outputs = (
self.consume_local_rollout_debug_tensors(batch)
)
if get_sp_world_size() > 1 and getattr(batch, "did_sp_shard_latents", False):
variance_noises = variance_noises.to(get_local_torch_device())
prev_sample_means = prev_sample_means.to(get_local_torch_device())
noise_std_devs = noise_std_devs.to(get_local_torch_device())
model_outputs = model_outputs.to(get_local_torch_device())
pipeline_config = rollout_session_data.pipeline_config
bsz, num_steps = variance_noises.shape[0], variance_noises.shape[1]
# [B, T, ...] -> [B*T, ...]
variance_noises_packed = variance_noises.contiguous().reshape(
bsz * num_steps, *variance_noises.shape[2:]
)
prev_sample_means_packed = prev_sample_means.contiguous().reshape(
bsz * num_steps, *prev_sample_means.shape[2:]
)
model_outputs_packed = model_outputs.contiguous().reshape(
bsz * num_steps, *model_outputs.shape[2:]
)
# Gather on packed tensors first.
variance_noises_packed = pipeline_config.gather_latents_for_sp(
variance_noises_packed
)
prev_sample_means_packed = pipeline_config.gather_latents_for_sp(
prev_sample_means_packed
)
model_outputs_packed = pipeline_config.gather_latents_for_sp(
model_outputs_packed
)
# Unpack back to [B, T, ...].
variance_noises = variance_noises_packed.reshape(
bsz, num_steps, *variance_noises_packed.shape[1:]
)
prev_sample_means = prev_sample_means_packed.reshape(
bsz, num_steps, *prev_sample_means_packed.shape[1:]
)
model_outputs = model_outputs_packed.reshape(
bsz, num_steps, *model_outputs_packed.shape[1:]
)
# noise_std_devs is same on every device, not a sharded latent tensor.
return RolloutDebugTensors(
rollout_variance_noises=variance_noises.cpu(),
rollout_prev_sample_means=prev_sample_means.cpu(),
rollout_noise_std_devs=noise_std_devs.cpu(),
rollout_model_outputs=model_outputs.cpu(),
)
@@ -0,0 +1,269 @@
# SPDX-License-Identifier: Apache-2.0
"""Flow-matching rollout step utilities for log-prob computation."""
import math
from typing import Any, Union
import torch
from sglang.multimodal_gen.runtime.distributed import (
get_local_torch_device,
get_sp_world_size,
)
from sglang.multimodal_gen.runtime.distributed.communication_op import (
sequence_model_parallel_all_reduce,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.post_training.rl_dataclasses import (
RolloutSessionData,
)
from sglang.multimodal_gen.runtime.post_training.scheduler_rl_debug_mixin import (
SchedulerRLDebugMixin,
)
_LOG_SQRT_2PI = math.log(math.sqrt(2 * math.pi))
class SchedulerRLMixin(SchedulerRLDebugMixin):
@staticmethod
def _get_rollout_session_data(batch) -> RolloutSessionData:
"""Return the RolloutSessionData attached to *batch*, or raise if not prepared."""
rollout_session_data = getattr(batch, "_rollout_session_data", None)
if rollout_session_data is None:
raise RuntimeError("prepare_rollout() not called before rollout")
return rollout_session_data
def release_rollout_resources(self, batch) -> None:
"""Release rollout-owned resources. Call when denoising ends or before a new rollout."""
batch._rollout_session_data = None
def prepare_rollout(self, batch: Req, pipeline_config: Any = None) -> None:
"""Enable rollout and set SDE/CPS params. Call once before the denoising loop."""
if get_sp_world_size() > 1 and pipeline_config is None:
raise RuntimeError(
"SP rollout requires pipeline_config to be passed to prepare_rollout()."
)
batch._rollout_session_data = RolloutSessionData(
pipeline_config=pipeline_config,
sigma_max=self.sigmas[min(1, len(self.sigmas) - 1)].item(),
latents_shape=(
tuple(batch.latents.shape) if batch.latents is not None else None
),
)
def already_prepared_rollout(self, batch) -> bool:
return getattr(batch, "_rollout_session_data", None) is not None
def _get_or_create_rollout_noise_buffer(
self,
rollout_session_data: RolloutSessionData,
full_shape: tuple,
device: torch.device,
dtype: torch.dtype,
) -> torch.Tensor:
"""Get or create the reusable noise buffer (local or full shape) for rollout."""
buffer = rollout_session_data.noise_buffer
if (
buffer is None
or buffer.shape != full_shape
or buffer.dtype != dtype
or buffer.device != device
):
buffer = torch.empty(full_shape, device=device, dtype=dtype)
rollout_session_data.noise_buffer = buffer
return buffer
def _rollout_variance_noise(
self,
batch,
model_output: torch.FloatTensor,
generator: Union[torch.Generator, list[torch.Generator]],
) -> torch.FloatTensor:
"""Generate variance noise for rollout. If generator is a list, use generator[i] for the i-th batch item."""
assert generator is not None, "Generator must be provided"
rollout_session_data = self._get_rollout_session_data(batch)
device = model_output.device
dtype = model_output.dtype
local_shape = tuple(model_output.shape)
B = local_shape[0]
if isinstance(generator, torch.Generator):
assert B == 1, "Generator must be a list if batch size is not 1"
generator = [generator]
else:
assert (
len(generator) == B
), "Generator list must have the same length as batch size"
buffer = self._get_or_create_rollout_noise_buffer(
rollout_session_data, rollout_session_data.latents_shape, device, dtype
)
for i in range(B):
torch.randn(
rollout_session_data.latents_shape,
out=buffer[i : i + 1],
generator=generator[i],
)
sharded_noise, _ = rollout_session_data.pipeline_config.shard_latents_for_sp(
batch, buffer
)
if tuple(sharded_noise.shape) != local_shape:
raise ValueError(
"Rollout SP noise shape mismatch after shard. "
f"Expected local_shape={local_shape}, got {tuple(sharded_noise.shape)}."
)
return sharded_noise
def flow_sde_sampling(
self,
batch,
model_output: torch.FloatTensor,
sample: torch.FloatTensor,
current_sigma: torch.FloatTensor,
next_sigma: torch.FloatTensor,
generator: torch.Generator,
) -> torch.Tensor:
"""Flow rollout step for log-prob / sampling (see FlowGRPO-style references).
``rollout_sde_type`` (from batch SamplingParams):
1. ``"sde"``: Standard stochastic differential equation transition (Gaussian).
2. ``"cps"``: Coupled Particle Sampling.
3. ``"ode"``: Deterministic ODE step (no diffusion noise).
"""
rollout_session_data = self._get_rollout_session_data(batch)
sde_type = batch.rollout_sde_type
noise_level = float(batch.rollout_noise_level)
log_prob_no_const = batch.rollout_log_prob_no_const
debug_mode = bool(getattr(batch, "rollout_debug_mode", False))
if not log_prob_no_const and sde_type != "ode":
assert (
noise_level > 0
), "True log-probability computation requires a non-zero noise level."
dt = next_sigma - current_sigma
if sde_type == "sde":
variance_noise = self._rollout_variance_noise(
batch, model_output, generator
)
std_dev_t = (
torch.sqrt(
current_sigma
/ (
1
- torch.where(
torch.isclose(current_sigma, current_sigma.new_tensor(1.0)),
rollout_session_data.sigma_max,
current_sigma,
)
)
)
* noise_level
)
noise_std_dev = std_dev_t * torch.sqrt(-1 * dt)
prev_sample_mean = (
sample * (1 + std_dev_t**2 / (2 * current_sigma) * dt)
+ model_output
* (1 + std_dev_t**2 * (1 - current_sigma) / (2 * current_sigma))
* dt
)
weighted_variance_noise = variance_noise * noise_std_dev
prev_sample = prev_sample_mean + weighted_variance_noise
log_prob_no_const_val = -(weighted_variance_noise**2)
elif sde_type == "cps":
variance_noise = self._rollout_variance_noise(
batch, model_output, generator
)
std_dev_t = next_sigma * math.sin(noise_level * math.pi / 2)
noise_std_dev = std_dev_t
pred_original_sample = sample - current_sigma * model_output
noise_estimate = sample + model_output * (1 - current_sigma)
prev_sample_mean = pred_original_sample * (
1 - next_sigma
) + noise_estimate * torch.sqrt(next_sigma**2 - std_dev_t**2)
weighted_variance_noise = variance_noise * noise_std_dev
prev_sample = prev_sample_mean + weighted_variance_noise
log_prob_no_const_val = -(weighted_variance_noise**2)
elif sde_type == "ode":
prev_sample = sample + dt * model_output
prev_sample_mean = prev_sample
variance_noise = torch.zeros_like(model_output)
noise_std_dev = torch.zeros(
(), device=model_output.device, dtype=model_output.dtype
)
log_prob_no_const_val = torch.zeros_like(model_output)
assert (
log_prob_no_const
), "p_ode is always 0, true log_prob is meaningless, set rollout_log_prob_no_const to True to enable log_prob computation"
else:
raise ValueError(f"Unsupported sde_type: {sde_type}")
reduce_dims = list(range(1, len(log_prob_no_const_val.shape)))
local_elem_count = log_prob_no_const_val.new_full(
(log_prob_no_const_val.shape[0],),
float(math.prod(log_prob_no_const_val.shape[1:])),
)
if log_prob_no_const:
log_prob_local_sum = log_prob_no_const_val.sum(dim=reduce_dims)
else:
log_prob_local_sum = (
log_prob_no_const_val / (2 * (noise_std_dev**2))
- torch.log(noise_std_dev)
- _LOG_SQRT_2PI
).sum(dim=list(range(1, len(log_prob_no_const_val.shape))))
if debug_mode:
self.append_local_rollout_debug_tensors(
batch,
variance_noise=variance_noise,
prev_sample_mean=prev_sample_mean,
noise_std_dev=noise_std_dev,
model_output=model_output,
)
self.append_local_rollout_log_probs(batch, log_prob_local_sum, local_elem_count)
return prev_sample
def append_local_rollout_log_probs(
self, batch, log_prob_sum: torch.Tensor, log_prob_count: torch.Tensor
) -> None:
rollout_session_data = self._get_rollout_session_data(batch)
rollout_session_data.local_log_prob_sum.append(log_prob_sum)
rollout_session_data.local_log_prob_count.append(log_prob_count)
def consume_local_rollout_log_probs(
self, batch
) -> tuple[torch.Tensor, torch.Tensor]:
rollout_session_data = self._get_rollout_session_data(batch)
values_sum = torch.stack(rollout_session_data.local_log_prob_sum, dim=-1)
values_count = torch.stack(rollout_session_data.local_log_prob_count, dim=-1)
rollout_session_data.local_log_prob_sum = []
rollout_session_data.local_log_prob_count = []
return values_sum, values_count
def collect_rollout_log_probs(self, batch: Req) -> torch.Tensor | None:
"""Consume local rollout log probs and merge for all SP ranks."""
trajectory_log_prob_sum, trajectory_log_prob_count = (
self.consume_local_rollout_log_probs(batch)
)
if get_sp_world_size() > 1 and getattr(batch, "did_sp_shard_latents", False):
packed = torch.stack(
[trajectory_log_prob_sum, trajectory_log_prob_count], dim=0
).to(get_local_torch_device())
sequence_model_parallel_all_reduce(packed)
trajectory_log_prob_sum = packed[0]
trajectory_log_prob_count = packed[1]
rollout_log_probs_tensor = trajectory_log_prob_sum / trajectory_log_prob_count
return rollout_log_probs_tensor.cpu()
@@ -0,0 +1,282 @@
import math
import types
import unittest
import torch
import sglang.multimodal_gen.runtime.post_training.scheduler_rl_mixin as rl_mixin_module
from sglang.multimodal_gen.runtime.post_training.scheduler_rl_mixin import (
SchedulerRLMixin,
)
class _DummyScheduler(SchedulerRLMixin):
def __init__(self):
self.sigmas = torch.tensor([1.0, 0.8, 0.6, 0.4, 0.2, 0.0], dtype=torch.float32)
class TestSchedulerRolloutOdeUnit(unittest.TestCase):
def setUp(self):
self._orig_get_sp_world_size = rl_mixin_module.get_sp_world_size
rl_mixin_module.get_sp_world_size = lambda: 1
def tearDown(self):
rl_mixin_module.get_sp_world_size = self._orig_get_sp_world_size
def _build_batch(self, *, debug_mode: bool) -> types.SimpleNamespace:
return types.SimpleNamespace(
rollout_log_prob_no_const=True,
rollout_noise_level=0.5,
rollout_sde_type="ode",
rollout_debug_mode=debug_mode,
latents=torch.zeros(2, 4, 8, 8, dtype=torch.float32),
_rollout_session_data=None,
)
def test_ode_step_does_not_call_variance_noise_sampler(self):
scheduler = _DummyScheduler()
batch = self._build_batch(debug_mode=False)
scheduler.prepare_rollout(batch)
def _raise_if_called(*args, **kwargs):
raise AssertionError("ODE path should not call _rollout_variance_noise")
scheduler._rollout_variance_noise = _raise_if_called # type: ignore[method-assign]
sample = torch.randn(2, 4, 8, 8, dtype=torch.float32)
model_output = torch.randn_like(sample)
current_sigma = torch.tensor(0.6, dtype=torch.float32)
next_sigma = torch.tensor(0.4, dtype=torch.float32)
prev_sample = scheduler.flow_sde_sampling(
batch,
model_output=model_output,
sample=sample,
current_sigma=current_sigma,
next_sigma=next_sigma,
generator=torch.Generator(device=sample.device).manual_seed(1),
)
log_prob_local_sum, local_elem_count = (
scheduler.consume_local_rollout_log_probs(batch)
)
log_prob_local_sum = log_prob_local_sum.squeeze(-1)
local_elem_count = local_elem_count.squeeze(-1)
expected_prev = sample + (next_sigma - current_sigma) * model_output
self.assertTrue(torch.allclose(prev_sample, expected_prev, atol=1e-6, rtol=0.0))
self.assertTrue(
torch.allclose(log_prob_local_sum, torch.zeros_like(log_prob_local_sum))
)
self.assertEqual(tuple(log_prob_local_sum.shape), (sample.shape[0],))
self.assertEqual(tuple(local_elem_count.shape), (sample.shape[0],))
self.assertTrue(torch.all(local_elem_count == float(sample[0].numel())))
def test_ode_debug_tensors_have_shape_safe_noise_std(self):
scheduler = _DummyScheduler()
batch = self._build_batch(debug_mode=True)
scheduler.prepare_rollout(batch)
sample = torch.randn(2, 4, 8, 8, dtype=torch.float32)
model_output = torch.randn_like(sample)
current_sigma = torch.tensor(0.6, dtype=torch.float32)
next_sigma = torch.tensor(0.4, dtype=torch.float32)
scheduler.flow_sde_sampling(
batch,
model_output=model_output,
sample=sample,
current_sigma=current_sigma,
next_sigma=next_sigma,
generator=torch.Generator(device=sample.device).manual_seed(2),
)
(
variance_noises,
prev_sample_means,
noise_std_devs,
model_outputs,
) = scheduler.consume_local_rollout_debug_tensors(batch)
# [B, T, ...] with one step in this test.
self.assertEqual(tuple(variance_noises.shape), (2, 1, 4, 8, 8))
self.assertEqual(tuple(prev_sample_means.shape), (2, 1, 4, 8, 8))
self.assertEqual(tuple(model_outputs.shape), (2, 1, 4, 8, 8))
self.assertEqual(tuple(noise_std_devs.shape), (2, 1, 1))
self.assertTrue(
torch.allclose(noise_std_devs, torch.zeros_like(noise_std_devs))
)
self.assertTrue(
torch.allclose(variance_noises, torch.zeros_like(variance_noises))
)
def _flowgrpo_sde_step_with_logprob(
*,
model_output: torch.Tensor,
sample: torch.Tensor,
variance_noise: torch.Tensor,
sigma: torch.Tensor,
sigma_prev: torch.Tensor,
sigma_max: float,
noise_level: float,
sde_type: str,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
"""Verbatim from FlowGRPO sd3_sde_with_logprob.py ``sde_step_with_logprob``.
Returns (prev_sample, log_prob, prev_sample_mean, noise_std_dev).
``sigma`` / ``sigma_prev`` follow FlowGRPO convention (current / next).
"""
model_output = model_output.float()
sample = sample.float()
dt = sigma_prev - sigma
if sde_type == "sde":
std_dev_t = (
torch.sqrt(sigma / (1 - torch.where(sigma == 1, sigma_max, sigma)))
* noise_level
)
prev_sample_mean = (
sample * (1 + std_dev_t**2 / (2 * sigma) * dt)
+ model_output * (1 + std_dev_t**2 * (1 - sigma) / (2 * sigma)) * dt
)
noise_std_dev = std_dev_t * torch.sqrt(-1 * dt)
prev_sample = prev_sample_mean + noise_std_dev * variance_noise
log_prob = (
-((prev_sample.detach() - prev_sample_mean) ** 2)
/ (2 * ((std_dev_t * torch.sqrt(-1 * dt)) ** 2))
- torch.log(std_dev_t * torch.sqrt(-1 * dt))
- torch.log(torch.sqrt(2 * torch.as_tensor(math.pi)))
)
elif sde_type == "cps":
std_dev_t = sigma_prev * math.sin(noise_level * math.pi / 2)
noise_std_dev = std_dev_t
pred_original_sample = sample - sigma * model_output
noise_estimate = sample + model_output * (1 - sigma)
prev_sample_mean = pred_original_sample * (
1 - sigma_prev
) + noise_estimate * torch.sqrt(sigma_prev**2 - std_dev_t**2)
prev_sample = prev_sample_mean + std_dev_t * variance_noise
log_prob = -((prev_sample.detach() - prev_sample_mean) ** 2)
else:
raise ValueError(f"Unsupported sde_type: {sde_type}")
log_prob = log_prob.mean(dim=tuple(range(1, log_prob.ndim)))
return prev_sample, log_prob, prev_sample_mean, noise_std_dev
# FlowGRPO convention: SDE uses full Gaussian log-prob, CPS uses no_const.
_FLOWGRPO_LOG_PROB_NO_CONST = {"sde": False, "cps": True}
class TestSchedulerFlowGRPOStepAlignmentUnit(unittest.TestCase):
def setUp(self):
self._orig_get_sp_world_size = rl_mixin_module.get_sp_world_size
rl_mixin_module.get_sp_world_size = lambda: 1
def tearDown(self):
rl_mixin_module.get_sp_world_size = self._orig_get_sp_world_size
def _build_batch(
self, *, sde_type: str, shape: tuple[int, ...]
) -> types.SimpleNamespace:
return types.SimpleNamespace(
rollout_log_prob_no_const=_FLOWGRPO_LOG_PROB_NO_CONST[sde_type],
rollout_noise_level=0.5,
rollout_sde_type=sde_type,
rollout_debug_mode=True,
latents=torch.empty(shape, dtype=torch.float32),
_rollout_session_data=None,
)
def test_single_step_matches_flowgrpo_reference(self):
"""Verify prev_sample, prev_sample_mean, noise_std_dev, and log_prob
all match FlowGRPO's ``sde_step_with_logprob`` for SDE and CPS."""
scheduler = _DummyScheduler()
current_sigma = torch.tensor(0.5, dtype=torch.float32)
next_sigma = torch.tensor(0.3, dtype=torch.float32)
shape = (1, 16, 1, 32, 32)
atol = 1e-6
pipeline_config = types.SimpleNamespace(
shard_latents_for_sp=lambda _batch, latents: (latents, False)
)
for sde_type in ("sde", "cps"):
for seed in (0, 1, 2, 3):
batch = self._build_batch(sde_type=sde_type, shape=shape)
scheduler.release_rollout_resources(batch)
scheduler.prepare_rollout(batch=batch, pipeline_config=pipeline_config)
g = torch.Generator(device="cpu").manual_seed(seed)
model_output = torch.randn(shape, generator=g, dtype=torch.float32)
sample = torch.randn(shape, generator=g, dtype=torch.float32)
variance_noise = torch.randn(shape, generator=g, dtype=torch.float32)
scheduler._rollout_variance_noise = ( # type: ignore[method-assign]
lambda _batch, *_args, **_kwargs: variance_noise
)
prev_sgl = scheduler.flow_sde_sampling(
batch,
model_output=model_output,
sample=sample,
current_sigma=current_sigma,
next_sigma=next_sigma,
generator=g,
)
log_prob_sum, elem_count = scheduler.consume_local_rollout_log_probs(
batch
)
log_prob_sum = log_prob_sum.squeeze(-1)
elem_count = elem_count.squeeze(-1)
(
_variance_noises,
prev_sample_means,
noise_std_devs,
_model_outputs,
) = scheduler.consume_local_rollout_debug_tensors(batch)
prev_ref, log_prob_ref, prev_mean_ref, noise_std_ref = (
_flowgrpo_sde_step_with_logprob(
model_output=model_output,
sample=sample,
variance_noise=variance_noise,
sigma=current_sigma,
sigma_prev=next_sigma,
sigma_max=scheduler.sigmas[1].item(),
noise_level=0.5,
sde_type=sde_type,
)
)
log_prob_mean = log_prob_sum / elem_count
errs = {
"prev_sample": float((prev_sgl - prev_ref).abs().max().item()),
"prev_sample_mean": float(
(prev_sample_means[:, 0] - prev_mean_ref).abs().max().item()
),
"noise_std": float(
(noise_std_devs[:, 0, 0] - noise_std_ref.reshape(-1))
.abs()
.max()
.item()
),
"log_prob": float(
(log_prob_mean - log_prob_ref).abs().max().item()
),
}
for name, err in errs.items():
self.assertLessEqual(
err,
atol,
msg=f"{sde_type} seed={seed} {name} max_abs={err:.9f}",
)
if __name__ == "__main__":
unittest.main()