[diffusion] RL: add per-step rollout options for SDE and trajectory capture (#23151)
This commit is contained in:
@@ -508,7 +508,7 @@ class PipelineConfig:
|
|||||||
def _unpad_and_unpack_latents(self, latents, audio_latents, batch, vae, audio_vae):
|
def _unpad_and_unpack_latents(self, latents, audio_latents, batch, vae, audio_vae):
|
||||||
raise NotImplementedError("not yet implemented")
|
raise NotImplementedError("not yet implemented")
|
||||||
|
|
||||||
def gather_dit_env_static_for_sp(self, batch, cond_kwargs: dict | None):
|
def gather_denoising_env_static_for_sp(self, batch, cond_kwargs: dict | None):
|
||||||
return cond_kwargs
|
return cond_kwargs
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
+1
-1
@@ -13,7 +13,7 @@ from sglang.multimodal_gen.runtime.post_training.sp_utils import (
|
|||||||
|
|
||||||
class QwenImageRolloutPipelineMixin:
|
class QwenImageRolloutPipelineMixin:
|
||||||
|
|
||||||
def gather_dit_env_static_for_sp(self, batch, cond_kwargs: dict | None):
|
def gather_denoising_env_static_for_sp(self, batch, cond_kwargs: dict | None):
|
||||||
if cond_kwargs is None:
|
if cond_kwargs is None:
|
||||||
return None
|
return None
|
||||||
out = dict(cond_kwargs)
|
out = dict(cond_kwargs)
|
||||||
|
|||||||
+1
-1
@@ -13,7 +13,7 @@ from sglang.multimodal_gen.runtime.post_training.sp_utils import (
|
|||||||
|
|
||||||
class ZImageRolloutPipelineMixin:
|
class ZImageRolloutPipelineMixin:
|
||||||
|
|
||||||
def gather_dit_env_static_for_sp(self, batch, cond_kwargs: dict | None):
|
def gather_denoising_env_static_for_sp(self, batch, cond_kwargs: dict | None):
|
||||||
if cond_kwargs is None:
|
if cond_kwargs is None:
|
||||||
return None
|
return None
|
||||||
out = dict(cond_kwargs)
|
out = dict(cond_kwargs)
|
||||||
|
|||||||
@@ -193,6 +193,9 @@ class SamplingParams:
|
|||||||
rollout_return_dit_trajectory: bool = (
|
rollout_return_dit_trajectory: bool = (
|
||||||
False # per-step noisy latents + final latent + timesteps (RolloutDitTrajectory)
|
False # per-step noisy latents + final latent + timesteps (RolloutDitTrajectory)
|
||||||
)
|
)
|
||||||
|
# 0-indexed denoising-loop step filters; None = all steps.
|
||||||
|
rollout_sde_step_indices: list[int] | None = None
|
||||||
|
rollout_return_step_indices: list[int] | None = None
|
||||||
# if True, disallow user params to override subclass-defined protected fields
|
# if True, disallow user params to override subclass-defined protected fields
|
||||||
no_override_protected_fields: bool = False
|
no_override_protected_fields: bool = False
|
||||||
# whether to adjust num_frames for multi-GPU friendly splitting (default: True)
|
# whether to adjust num_frames for multi-GPU friendly splitting (default: True)
|
||||||
|
|||||||
@@ -51,6 +51,10 @@ class RolloutRequest(BaseModel):
|
|||||||
rollout_return_denoising_env: bool = False
|
rollout_return_denoising_env: bool = False
|
||||||
rollout_return_dit_trajectory: bool = False
|
rollout_return_dit_trajectory: bool = False
|
||||||
|
|
||||||
|
# 0-indexed denoising-loop step filters. None = all steps.
|
||||||
|
rollout_sde_step_indices: Optional[list[int]] = None
|
||||||
|
rollout_return_step_indices: Optional[list[int]] = None
|
||||||
|
|
||||||
image_path: Optional[list[str]] = None
|
image_path: Optional[list[str]] = None
|
||||||
|
|
||||||
# suppress verbose per-request logging (also gates peak_memory_mb collection)
|
# suppress verbose per-request logging (also gates peak_memory_mb collection)
|
||||||
|
|||||||
@@ -249,10 +249,7 @@ def _build_response(
|
|||||||
return responses
|
return responses
|
||||||
|
|
||||||
|
|
||||||
@router.post("/generate", response_model=list[RolloutResponse])
|
def _build_sampling_kwargs(request: RolloutRequest) -> dict:
|
||||||
async def rollout_generate(request: RolloutRequest):
|
|
||||||
request_id = generate_request_id()
|
|
||||||
server_args = get_global_server_args()
|
|
||||||
sampling_kwargs: dict = dict(
|
sampling_kwargs: dict = dict(
|
||||||
prompt=request.prompt,
|
prompt=request.prompt,
|
||||||
negative_prompt=request.negative_prompt,
|
negative_prompt=request.negative_prompt,
|
||||||
@@ -274,6 +271,8 @@ async def rollout_generate(request: RolloutRequest):
|
|||||||
rollout_debug_mode=request.rollout_debug_mode,
|
rollout_debug_mode=request.rollout_debug_mode,
|
||||||
rollout_return_denoising_env=request.rollout_return_denoising_env,
|
rollout_return_denoising_env=request.rollout_return_denoising_env,
|
||||||
rollout_return_dit_trajectory=request.rollout_return_dit_trajectory,
|
rollout_return_dit_trajectory=request.rollout_return_dit_trajectory,
|
||||||
|
rollout_sde_step_indices=request.rollout_sde_step_indices,
|
||||||
|
rollout_return_step_indices=request.rollout_return_step_indices,
|
||||||
suppress_logs=request.suppress_logs,
|
suppress_logs=request.suppress_logs,
|
||||||
save_output=False,
|
save_output=False,
|
||||||
return_trajectory_latents=False,
|
return_trajectory_latents=False,
|
||||||
@@ -282,7 +281,14 @@ async def rollout_generate(request: RolloutRequest):
|
|||||||
if request.extra_sampling_params:
|
if request.extra_sampling_params:
|
||||||
sampling_kwargs.update(request.extra_sampling_params)
|
sampling_kwargs.update(request.extra_sampling_params)
|
||||||
sampling_kwargs["rollout"] = request.rollout
|
sampling_kwargs["rollout"] = request.rollout
|
||||||
sampling_kwargs = {k: v for k, v in sampling_kwargs.items() if v is not None}
|
return {k: v for k, v in sampling_kwargs.items() if v is not None}
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/generate", response_model=list[RolloutResponse])
|
||||||
|
async def rollout_generate(request: RolloutRequest):
|
||||||
|
request_id = generate_request_id()
|
||||||
|
server_args = get_global_server_args()
|
||||||
|
sampling_kwargs = _build_sampling_kwargs(request)
|
||||||
try:
|
try:
|
||||||
sampling_params = build_sampling_params(request_id, **sampling_kwargs)
|
sampling_params = build_sampling_params(request_id, **sampling_kwargs)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
|
|||||||
@@ -1140,10 +1140,12 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
# pre-step value. Gated on batch.rollout to keep the
|
# pre-step value. Gated on batch.rollout to keep the
|
||||||
# non-rollout path strictly untouched.
|
# non-rollout path strictly untouched.
|
||||||
if batch.rollout:
|
if batch.rollout:
|
||||||
|
batch._rollout_loop_step_index = step_index
|
||||||
self._maybe_append_dit_trajectory_step(
|
self._maybe_append_dit_trajectory_step(
|
||||||
batch=batch,
|
batch=batch,
|
||||||
latents=ctx.latents,
|
latents=ctx.latents,
|
||||||
timestep_value=step.t_host,
|
timestep_value=step.t_host,
|
||||||
|
step_index=step_index,
|
||||||
)
|
)
|
||||||
self._run_denoising_step(ctx, step, batch, server_args)
|
self._run_denoising_step(ctx, step, batch, server_args)
|
||||||
self._record_trajectory(ctx, step, batch, server_args)
|
self._record_trajectory(ctx, step, batch, server_args)
|
||||||
@@ -1174,6 +1176,8 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
self._postprocess_rollout_outputs(
|
self._postprocess_rollout_outputs(
|
||||||
batch=batch,
|
batch=batch,
|
||||||
latents=ctx.latents,
|
latents=ctx.latents,
|
||||||
|
num_inference_steps=num_timesteps,
|
||||||
|
final_timestep=timesteps_cpu.new_zeros(()),
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
)
|
)
|
||||||
self._finalize_denoising_loop(ctx, batch, server_args)
|
self._finalize_denoising_loop(ctx, batch, server_args)
|
||||||
|
|||||||
@@ -78,6 +78,8 @@ class RolloutDenoisingMixin:
|
|||||||
self,
|
self,
|
||||||
batch: Req,
|
batch: Req,
|
||||||
latents: torch.Tensor,
|
latents: torch.Tensor,
|
||||||
|
num_inference_steps: int,
|
||||||
|
final_timestep: torch.Tensor,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Finalize rollout-only outputs.
|
"""Finalize rollout-only outputs.
|
||||||
@@ -87,12 +89,16 @@ class RolloutDenoisingMixin:
|
|||||||
uniformly with the per-step trajectory latents.
|
uniformly with the per-step trajectory latents.
|
||||||
"""
|
"""
|
||||||
self._maybe_collect_rollout_log_probs(batch)
|
self._maybe_collect_rollout_log_probs(batch)
|
||||||
# Append the final denoised latent as the (T+1)-th entry of the
|
# Append final denoised latent as the (T+1)-th entry (step_index=T),
|
||||||
# dit-trajectory latents list.
|
# routed through the same filter so rollout_return_step_indices can
|
||||||
state = getattr(batch, "_rollout_dit_env_state", None)
|
# include/exclude it.
|
||||||
if state is not None and batch.rollout and batch.rollout_return_dit_trajectory:
|
self._maybe_append_dit_trajectory_step(
|
||||||
state["step_latents"].append(latents.detach())
|
batch=batch,
|
||||||
self._maybe_finalize_dit_env_collection(
|
latents=latents,
|
||||||
|
timestep_value=final_timestep,
|
||||||
|
step_index=num_inference_steps,
|
||||||
|
)
|
||||||
|
self._maybe_finalize_denoising_env_collection(
|
||||||
batch=batch,
|
batch=batch,
|
||||||
pipeline_config=server_args.pipeline_config,
|
pipeline_config=server_args.pipeline_config,
|
||||||
)
|
)
|
||||||
@@ -109,18 +115,15 @@ class RolloutDenoisingMixin:
|
|||||||
collect_env = batch.rollout_return_denoising_env
|
collect_env = batch.rollout_return_denoising_env
|
||||||
collect_traj = batch.rollout_return_dit_trajectory
|
collect_traj = batch.rollout_return_dit_trajectory
|
||||||
if not (collect_env or collect_traj):
|
if not (collect_env or collect_traj):
|
||||||
batch._rollout_dit_env_state = None
|
batch._rollout_denoising_env_state = None
|
||||||
return
|
return
|
||||||
|
|
||||||
sanitize = getattr(pipeline_config, "sanitize_dit_env_kwargs", lambda x: x)
|
|
||||||
if collect_env:
|
if collect_env:
|
||||||
env = RolloutDenoisingEnv(
|
env = RolloutDenoisingEnv(
|
||||||
image_kwargs=_kwargs_to_cpu(sanitize(image_kwargs)),
|
image_kwargs=_kwargs_to_cpu(image_kwargs),
|
||||||
pos_cond_kwargs=_kwargs_to_cpu(sanitize(pos_cond_kwargs)),
|
pos_cond_kwargs=_kwargs_to_cpu(pos_cond_kwargs),
|
||||||
neg_cond_kwargs=(
|
neg_cond_kwargs=(
|
||||||
_kwargs_to_cpu(sanitize(neg_cond_kwargs))
|
_kwargs_to_cpu(neg_cond_kwargs) if neg_cond_kwargs else None
|
||||||
if neg_cond_kwargs
|
|
||||||
else None
|
|
||||||
),
|
),
|
||||||
guidance=guidance.detach().cpu() if guidance is not None else None,
|
guidance=guidance.detach().cpu() if guidance is not None else None,
|
||||||
)
|
)
|
||||||
@@ -131,7 +134,7 @@ class RolloutDenoisingMixin:
|
|||||||
pos_src = None
|
pos_src = None
|
||||||
neg_src = None
|
neg_src = None
|
||||||
|
|
||||||
batch._rollout_dit_env_state = {
|
batch._rollout_denoising_env_state = {
|
||||||
"env": env,
|
"env": env,
|
||||||
"step_latents": [],
|
"step_latents": [],
|
||||||
"step_timesteps": [],
|
"step_timesteps": [],
|
||||||
@@ -144,18 +147,23 @@ class RolloutDenoisingMixin:
|
|||||||
batch,
|
batch,
|
||||||
latents: torch.Tensor,
|
latents: torch.Tensor,
|
||||||
timestep_value: torch.Tensor,
|
timestep_value: torch.Tensor,
|
||||||
|
step_index: int,
|
||||||
) -> None:
|
) -> None:
|
||||||
if not batch.rollout or not batch.rollout_return_dit_trajectory:
|
if not batch.rollout or not batch.rollout_return_dit_trajectory:
|
||||||
return
|
return
|
||||||
state = getattr(batch, "_rollout_dit_env_state", None)
|
state = getattr(batch, "_rollout_denoising_env_state", None)
|
||||||
if state is None:
|
if state is None:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
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_latents"].append(latents.detach())
|
||||||
state["step_timesteps"].append(timestep_value.detach().cpu())
|
state["step_timesteps"].append(timestep_value.detach().cpu())
|
||||||
|
|
||||||
def _maybe_finalize_dit_env_collection(self, batch, pipeline_config) -> None:
|
def _maybe_finalize_denoising_env_collection(self, batch, pipeline_config) -> None:
|
||||||
state = getattr(batch, "_rollout_dit_env_state", None)
|
state = getattr(batch, "_rollout_denoising_env_state", None)
|
||||||
if state is None:
|
if state is None:
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -179,19 +187,20 @@ class RolloutDenoisingMixin:
|
|||||||
)
|
)
|
||||||
|
|
||||||
if env is not None and batch.rollout_return_denoising_env:
|
if env is not None and batch.rollout_return_denoising_env:
|
||||||
sanitize = getattr(pipeline_config, "sanitize_dit_env_kwargs", lambda x: x)
|
gather_fn = getattr(
|
||||||
gather_fn = getattr(pipeline_config, "gather_dit_env_static_for_sp", None)
|
pipeline_config, "gather_denoising_env_static_for_sp", None
|
||||||
|
)
|
||||||
|
|
||||||
pos_src = state.get("pos_cond_kwargs_src")
|
pos_src = state.get("pos_cond_kwargs_src")
|
||||||
if pos_src is not None and env.pos_cond_kwargs is not None:
|
if pos_src is not None and env.pos_cond_kwargs is not None:
|
||||||
gathered_pos = gather_fn(batch, pos_src) if gather_fn else pos_src
|
gathered_pos = gather_fn(batch, pos_src) if gather_fn else pos_src
|
||||||
env.pos_cond_kwargs = _kwargs_to_cpu(sanitize(gathered_pos))
|
env.pos_cond_kwargs = _kwargs_to_cpu(gathered_pos)
|
||||||
|
|
||||||
neg_src = state.get("neg_cond_kwargs_src")
|
neg_src = state.get("neg_cond_kwargs_src")
|
||||||
if neg_src is not None and env.neg_cond_kwargs is not None:
|
if neg_src is not None and env.neg_cond_kwargs is not None:
|
||||||
gathered_neg = gather_fn(batch, neg_src) if gather_fn else neg_src
|
gathered_neg = gather_fn(batch, neg_src) if gather_fn else neg_src
|
||||||
env.neg_cond_kwargs = _kwargs_to_cpu(sanitize(gathered_neg))
|
env.neg_cond_kwargs = _kwargs_to_cpu(gathered_neg)
|
||||||
|
|
||||||
batch.rollout_trajectory_data.denoising_env = env
|
batch.rollout_trajectory_data.denoising_env = env
|
||||||
|
|
||||||
batch._rollout_dit_env_state = None
|
batch._rollout_denoising_env_state = None
|
||||||
|
|||||||
@@ -141,13 +141,29 @@ class SchedulerRLMixin(SchedulerRLDebugMixin):
|
|||||||
), "True log-probability computation requires a non-zero noise level."
|
), "True log-probability computation requires a non-zero noise level."
|
||||||
|
|
||||||
dt = next_sigma - current_sigma
|
dt = next_sigma - current_sigma
|
||||||
|
|
||||||
|
# step_index comes from the denoising-loop counter stashed by
|
||||||
|
# DenoisingStage — scheduler._step_index would differ when
|
||||||
|
# _begin_index != 0 (e.g. partial denoising).
|
||||||
|
sde_step_indices = getattr(batch, "rollout_sde_step_indices", None)
|
||||||
|
loop_step_index = getattr(batch, "_rollout_loop_step_index", None)
|
||||||
|
if (
|
||||||
|
sde_type != "ode"
|
||||||
|
and sde_step_indices is not None
|
||||||
|
and loop_step_index is not None
|
||||||
|
and loop_step_index not in sde_step_indices
|
||||||
|
):
|
||||||
|
effective_sde_type = "ode"
|
||||||
|
else:
|
||||||
|
effective_sde_type = sde_type
|
||||||
|
|
||||||
# sde/cps: cast to fp32 to match flowGRPO semantics and avoid the
|
# sde/cps: cast to fp32 to match flowGRPO semantics and avoid the
|
||||||
# 0-dim-fp32 wrapped-scalar promotion demoting log-prob to bf16.
|
# 0-dim-fp32 wrapped-scalar promotion demoting log-prob to bf16.
|
||||||
# ode: keep dtypes unchanged so rollout(ode) stays bit-exact with
|
# ode: keep dtypes unchanged so rollout(ode) stays bit-exact with
|
||||||
# rollout=False (scheduling_flow_match_euler_discrete.step()).
|
# rollout=False (scheduling_flow_match_euler_discrete.step()).
|
||||||
# log_prob is computed on the full pre-shard noise buffer so SP ranks
|
# log_prob is computed on the full pre-shard noise buffer so SP ranks
|
||||||
# produce identical sums — see collect_rollout_log_probs().
|
# produce identical sums — see collect_rollout_log_probs().
|
||||||
if sde_type == "sde":
|
if effective_sde_type == "sde":
|
||||||
model_output = model_output.float()
|
model_output = model_output.float()
|
||||||
sample = sample.float()
|
sample = sample.float()
|
||||||
variance_noise = self._rollout_variance_noise(
|
variance_noise = self._rollout_variance_noise(
|
||||||
@@ -180,7 +196,7 @@ class SchedulerRLMixin(SchedulerRLDebugMixin):
|
|||||||
prev_sample = prev_sample_mean + weighted_variance_noise
|
prev_sample = prev_sample_mean + weighted_variance_noise
|
||||||
log_prob_no_const_val = -((full_variance_noise * noise_std_dev) ** 2)
|
log_prob_no_const_val = -((full_variance_noise * noise_std_dev) ** 2)
|
||||||
|
|
||||||
elif sde_type == "cps":
|
elif effective_sde_type == "cps":
|
||||||
model_output = model_output.float()
|
model_output = model_output.float()
|
||||||
sample = sample.float()
|
sample = sample.float()
|
||||||
variance_noise = self._rollout_variance_noise(
|
variance_noise = self._rollout_variance_noise(
|
||||||
@@ -199,7 +215,7 @@ class SchedulerRLMixin(SchedulerRLDebugMixin):
|
|||||||
prev_sample = prev_sample_mean + weighted_variance_noise
|
prev_sample = prev_sample_mean + weighted_variance_noise
|
||||||
log_prob_no_const_val = -((full_variance_noise * noise_std_dev) ** 2)
|
log_prob_no_const_val = -((full_variance_noise * noise_std_dev) ** 2)
|
||||||
|
|
||||||
elif sde_type == "ode":
|
elif effective_sde_type == "ode":
|
||||||
prev_sample = sample + dt * model_output
|
prev_sample = sample + dt * model_output
|
||||||
prev_sample_mean = prev_sample
|
prev_sample_mean = prev_sample
|
||||||
variance_noise = torch.zeros_like(model_output)
|
variance_noise = torch.zeros_like(model_output)
|
||||||
@@ -211,9 +227,12 @@ class SchedulerRLMixin(SchedulerRLDebugMixin):
|
|||||||
device=model_output.device,
|
device=model_output.device,
|
||||||
dtype=torch.float32,
|
dtype=torch.float32,
|
||||||
)
|
)
|
||||||
assert (
|
# Only enforce the "no full log-prob with ODE" constraint when the
|
||||||
log_prob_no_const
|
# user explicitly chose ODE globally.
|
||||||
), "p_ode is always 0, true log_prob is meaningless, set rollout_log_prob_no_const to True to enable log_prob computation"
|
if sde_type == "ode":
|
||||||
|
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:
|
else:
|
||||||
raise ValueError(f"Unsupported sde_type: {sde_type}")
|
raise ValueError(f"Unsupported sde_type: {sde_type}")
|
||||||
@@ -224,7 +243,7 @@ class SchedulerRLMixin(SchedulerRLDebugMixin):
|
|||||||
float(math.prod(log_prob_no_const_val.shape[1:])),
|
float(math.prod(log_prob_no_const_val.shape[1:])),
|
||||||
)
|
)
|
||||||
|
|
||||||
if log_prob_no_const:
|
if log_prob_no_const or effective_sde_type == "ode":
|
||||||
log_prob_local_sum = log_prob_no_const_val.sum(dim=reduce_dims)
|
log_prob_local_sum = log_prob_no_const_val.sum(dim=reduce_dims)
|
||||||
else:
|
else:
|
||||||
log_prob_local_sum = (
|
log_prob_local_sum = (
|
||||||
|
|||||||
@@ -356,5 +356,57 @@ class TestBuildResponse(unittest.TestCase):
|
|||||||
self.assertIsNotNone(resps[0].generated_output)
|
self.assertIsNotNone(resps[0].generated_output)
|
||||||
|
|
||||||
|
|
||||||
|
class TestBuildSamplingKwargs(unittest.TestCase):
|
||||||
|
def _make_request(self, **overrides):
|
||||||
|
from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
|
||||||
|
RolloutRequest,
|
||||||
|
)
|
||||||
|
|
||||||
|
base = dict(prompt="x", num_inference_steps=4, rollout=True)
|
||||||
|
base.update(overrides)
|
||||||
|
return RolloutRequest(**base)
|
||||||
|
|
||||||
|
def test_step_index_filters_forwarded(self):
|
||||||
|
from sglang.multimodal_gen.runtime.entrypoints.post_training.rollout_api import (
|
||||||
|
_build_sampling_kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
kwargs = _build_sampling_kwargs(
|
||||||
|
self._make_request(
|
||||||
|
rollout_sde_step_indices=[0, 2],
|
||||||
|
rollout_return_step_indices=[1, 3],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.assertEqual(kwargs["rollout_sde_step_indices"], [0, 2])
|
||||||
|
self.assertEqual(kwargs["rollout_return_step_indices"], [1, 3])
|
||||||
|
|
||||||
|
def test_step_index_filters_default_dropped_as_none(self):
|
||||||
|
from sglang.multimodal_gen.runtime.entrypoints.post_training.rollout_api import (
|
||||||
|
_build_sampling_kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
kwargs = _build_sampling_kwargs(self._make_request())
|
||||||
|
# None values are stripped; absence here is the correct default-path behavior.
|
||||||
|
self.assertNotIn("rollout_sde_step_indices", kwargs)
|
||||||
|
self.assertNotIn("rollout_return_step_indices", kwargs)
|
||||||
|
|
||||||
|
def test_sampling_params_exposes_filters_via_req_getattr(self):
|
||||||
|
from sglang.multimodal_gen.configs.sample.sampling_params import (
|
||||||
|
SamplingParams,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||||
|
|
||||||
|
sp = SamplingParams(
|
||||||
|
prompt="x",
|
||||||
|
num_inference_steps=4,
|
||||||
|
rollout=True,
|
||||||
|
rollout_sde_step_indices=[0, 2],
|
||||||
|
rollout_return_step_indices=[1, 3],
|
||||||
|
)
|
||||||
|
req = Req(sampling_params=sp)
|
||||||
|
self.assertEqual(req.rollout_sde_step_indices, [0, 2])
|
||||||
|
self.assertEqual(req.rollout_return_step_indices, [1, 3])
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -102,16 +102,25 @@ class TestSchedulerRolloutOdeUnit(unittest.TestCase):
|
|||||||
# dt * model_output` (after the shared ``sample.to(fp32)`` cast).
|
# dt * model_output` (after the shared ``sample.to(fp32)`` cast).
|
||||||
non_rollout_prev = sample + dt * model_output
|
non_rollout_prev = sample + dt * model_output
|
||||||
|
|
||||||
self.assertEqual(rollout_prev.dtype, non_rollout_prev.dtype)
|
pre_cast_max_abs_diff = (rollout_prev - non_rollout_prev).abs().max().item()
|
||||||
self.assertTrue(torch.equal(rollout_prev, non_rollout_prev))
|
post_cast_max_abs_diff = (
|
||||||
# Also verify the post-cast to model_output.dtype (what scheduler.step
|
(
|
||||||
# returns downstream) is bit-exact.
|
rollout_prev.to(model_output.dtype)
|
||||||
self.assertTrue(
|
- non_rollout_prev.to(model_output.dtype)
|
||||||
torch.equal(
|
|
||||||
rollout_prev.to(model_output.dtype),
|
|
||||||
non_rollout_prev.to(model_output.dtype),
|
|
||||||
)
|
)
|
||||||
|
.abs()
|
||||||
|
.max()
|
||||||
|
.item()
|
||||||
)
|
)
|
||||||
|
print(
|
||||||
|
f"\n[ODE rollout vs non-rollout, bf16 model_output] "
|
||||||
|
f"max |diff| pre-cast={pre_cast_max_abs_diff}, "
|
||||||
|
f"post-cast={post_cast_max_abs_diff}"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(rollout_prev.dtype, non_rollout_prev.dtype)
|
||||||
|
self.assertEqual(pre_cast_max_abs_diff, 0.0)
|
||||||
|
self.assertEqual(post_cast_max_abs_diff, 0.0)
|
||||||
|
|
||||||
def test_ode_debug_tensors_have_shape_safe_noise_std(self):
|
def test_ode_debug_tensors_have_shape_safe_noise_std(self):
|
||||||
scheduler = _DummyScheduler()
|
scheduler = _DummyScheduler()
|
||||||
@@ -381,6 +390,193 @@ class TestSchedulerFlowGRPOStepAlignmentUnit(unittest.TestCase):
|
|||||||
msg=f"{sde_type}: noise_buffer must be fp32 with bf16 model_output",
|
msg=f"{sde_type}: noise_buffer must be fp32 with bf16 model_output",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_timestep_filters_gate_sde_and_trajectory(self):
|
||||||
|
"""Per-step index filters: rollout_sde_step_indices gates variance-noise
|
||||||
|
injection (excluded steps = ODE transition + zero log-prob); independently,
|
||||||
|
rollout_return_step_indices gates the dit_trajectory append. Both features
|
||||||
|
are exercised here because they share the same step_index predicate."""
|
||||||
|
from sglang.multimodal_gen.runtime.post_training.rollout_denoising_mixin import (
|
||||||
|
RolloutDenoisingMixin,
|
||||||
|
)
|
||||||
|
|
||||||
|
# --- Part 1: rollout_sde_step_indices gates SDE noise injection ---
|
||||||
|
scheduler = _DummyScheduler()
|
||||||
|
shape = (1, 4, 8, 8)
|
||||||
|
pipeline_config = types.SimpleNamespace(
|
||||||
|
shard_latents_for_sp=lambda _batch, latents: (latents, False)
|
||||||
|
)
|
||||||
|
batch = types.SimpleNamespace(
|
||||||
|
rollout_log_prob_no_const=False,
|
||||||
|
rollout_noise_level=0.5,
|
||||||
|
rollout_sde_type="sde",
|
||||||
|
rollout_debug_mode=False,
|
||||||
|
rollout_sde_step_indices=[1], # only step 1 is stochastic
|
||||||
|
latents=torch.empty(shape, dtype=torch.float32),
|
||||||
|
_rollout_session_data=None,
|
||||||
|
)
|
||||||
|
scheduler.prepare_rollout(batch=batch, pipeline_config=pipeline_config)
|
||||||
|
|
||||||
|
g = torch.Generator(device="cpu").manual_seed(0)
|
||||||
|
sample = torch.randn(shape, generator=g, dtype=torch.float32)
|
||||||
|
model_output = torch.randn(shape, generator=g, dtype=torch.float32)
|
||||||
|
current_sigma = torch.tensor(0.6, dtype=torch.float32)
|
||||||
|
next_sigma = torch.tensor(0.4, dtype=torch.float32)
|
||||||
|
|
||||||
|
variance_noise_ref = torch.randn(shape, generator=g, dtype=torch.float32)
|
||||||
|
variance_noise_call_count = {"n": 0}
|
||||||
|
|
||||||
|
def _mock_variance_noise(_batch, *_args, **_kwargs):
|
||||||
|
variance_noise_call_count["n"] += 1
|
||||||
|
scheduler._get_rollout_session_data(_batch).noise_buffer = (
|
||||||
|
variance_noise_ref
|
||||||
|
)
|
||||||
|
return variance_noise_ref
|
||||||
|
|
||||||
|
scheduler._rollout_variance_noise = ( # type: ignore[method-assign]
|
||||||
|
_mock_variance_noise
|
||||||
|
)
|
||||||
|
|
||||||
|
# Step 0: not in filter → deterministic ODE transition, no noise draw.
|
||||||
|
batch._rollout_loop_step_index = 0
|
||||||
|
prev_0 = scheduler.flow_sde_sampling(
|
||||||
|
batch,
|
||||||
|
model_output=model_output,
|
||||||
|
sample=sample,
|
||||||
|
current_sigma=current_sigma,
|
||||||
|
next_sigma=next_sigma,
|
||||||
|
generator=g,
|
||||||
|
)
|
||||||
|
self.assertEqual(variance_noise_call_count["n"], 0)
|
||||||
|
expected_ode = sample + (next_sigma - current_sigma) * model_output
|
||||||
|
self.assertTrue(torch.allclose(prev_0, expected_ode, atol=1e-6))
|
||||||
|
|
||||||
|
# Step 1: in filter → real SDE, noise drawn, prev differs from ODE form.
|
||||||
|
batch._rollout_loop_step_index = 1
|
||||||
|
prev_1 = scheduler.flow_sde_sampling(
|
||||||
|
batch,
|
||||||
|
model_output=model_output,
|
||||||
|
sample=sample,
|
||||||
|
current_sigma=current_sigma,
|
||||||
|
next_sigma=next_sigma,
|
||||||
|
generator=g,
|
||||||
|
)
|
||||||
|
self.assertEqual(variance_noise_call_count["n"], 1)
|
||||||
|
self.assertFalse(torch.allclose(prev_1, expected_ode, atol=1e-3))
|
||||||
|
|
||||||
|
log_prob_sum, elem_count = scheduler.consume_local_rollout_log_probs(batch)
|
||||||
|
self.assertEqual(tuple(log_prob_sum.shape), (shape[0], 2))
|
||||||
|
# Filtered step contributes zero log-prob; real SDE step does not.
|
||||||
|
self.assertTrue(
|
||||||
|
torch.allclose(log_prob_sum[:, 0], torch.zeros_like(log_prob_sum[:, 0]))
|
||||||
|
)
|
||||||
|
self.assertFalse(
|
||||||
|
torch.allclose(log_prob_sum[:, 1], torch.zeros_like(log_prob_sum[:, 1]))
|
||||||
|
)
|
||||||
|
# elem_count dimension must be preserved for both steps so downstream
|
||||||
|
# consume_local_rollout_log_probs stacking stays consistent.
|
||||||
|
self.assertTrue(torch.all(elem_count > 0))
|
||||||
|
|
||||||
|
# --- Part 2: rollout_return_step_indices gates dit trajectory append ---
|
||||||
|
class _DummyDit(RolloutDenoisingMixin):
|
||||||
|
pass
|
||||||
|
|
||||||
|
dit = _DummyDit()
|
||||||
|
lat = torch.zeros(1, 4, 8, 8)
|
||||||
|
ts = torch.tensor(0.5)
|
||||||
|
|
||||||
|
# Filter [0, 2] over steps 0,1,2 → steps 0 and 2 appended, step 1 skipped.
|
||||||
|
traj_filtered = types.SimpleNamespace(
|
||||||
|
rollout=True,
|
||||||
|
rollout_return_dit_trajectory=True,
|
||||||
|
rollout_return_step_indices=[0, 2],
|
||||||
|
_rollout_denoising_env_state={"step_latents": [], "step_timesteps": []},
|
||||||
|
)
|
||||||
|
for i in range(3):
|
||||||
|
dit._maybe_append_dit_trajectory_step(
|
||||||
|
batch=traj_filtered,
|
||||||
|
latents=lat,
|
||||||
|
timestep_value=ts,
|
||||||
|
step_index=i,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
len(traj_filtered._rollout_denoising_env_state["step_latents"]), 2
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
len(traj_filtered._rollout_denoising_env_state["step_timesteps"]), 2
|
||||||
|
)
|
||||||
|
|
||||||
|
# None (default) → all steps appended (back-compat).
|
||||||
|
traj_all = types.SimpleNamespace(
|
||||||
|
rollout=True,
|
||||||
|
rollout_return_dit_trajectory=True,
|
||||||
|
rollout_return_step_indices=None,
|
||||||
|
_rollout_denoising_env_state={"step_latents": [], "step_timesteps": []},
|
||||||
|
)
|
||||||
|
for i in range(3):
|
||||||
|
dit._maybe_append_dit_trajectory_step(
|
||||||
|
batch=traj_all,
|
||||||
|
latents=lat,
|
||||||
|
timestep_value=ts,
|
||||||
|
step_index=i,
|
||||||
|
)
|
||||||
|
self.assertEqual(len(traj_all._rollout_denoising_env_state["step_latents"]), 3)
|
||||||
|
|
||||||
|
# Filter excludes step_index=T (the final/(T+1)-th latent appended by
|
||||||
|
# _postprocess_rollout_outputs). Simulate T=3 loop steps + final append.
|
||||||
|
traj_exclude_final = types.SimpleNamespace(
|
||||||
|
rollout=True,
|
||||||
|
rollout_return_dit_trajectory=True,
|
||||||
|
rollout_return_step_indices=[0, 1, 2], # excludes T=3
|
||||||
|
_rollout_denoising_env_state={"step_latents": [], "step_timesteps": []},
|
||||||
|
)
|
||||||
|
for i in range(3):
|
||||||
|
dit._maybe_append_dit_trajectory_step(
|
||||||
|
batch=traj_exclude_final,
|
||||||
|
latents=lat,
|
||||||
|
timestep_value=ts,
|
||||||
|
step_index=i,
|
||||||
|
)
|
||||||
|
# Mimic the final append routed through the same filter.
|
||||||
|
dit._maybe_append_dit_trajectory_step(
|
||||||
|
batch=traj_exclude_final,
|
||||||
|
latents=lat,
|
||||||
|
timestep_value=torch.zeros(()),
|
||||||
|
step_index=3,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
len(traj_exclude_final._rollout_denoising_env_state["step_latents"]), 3
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
len(traj_exclude_final._rollout_denoising_env_state["step_timesteps"]), 3
|
||||||
|
)
|
||||||
|
|
||||||
|
# Filter includes only step_index=T → only the final latent survives.
|
||||||
|
traj_only_final = types.SimpleNamespace(
|
||||||
|
rollout=True,
|
||||||
|
rollout_return_dit_trajectory=True,
|
||||||
|
rollout_return_step_indices=[3],
|
||||||
|
_rollout_denoising_env_state={"step_latents": [], "step_timesteps": []},
|
||||||
|
)
|
||||||
|
for i in range(3):
|
||||||
|
dit._maybe_append_dit_trajectory_step(
|
||||||
|
batch=traj_only_final,
|
||||||
|
latents=lat,
|
||||||
|
timestep_value=ts,
|
||||||
|
step_index=i,
|
||||||
|
)
|
||||||
|
dit._maybe_append_dit_trajectory_step(
|
||||||
|
batch=traj_only_final,
|
||||||
|
latents=lat,
|
||||||
|
timestep_value=torch.zeros(()),
|
||||||
|
step_index=3,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
len(traj_only_final._rollout_denoising_env_state["step_latents"]), 1
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
len(traj_only_final._rollout_denoising_env_state["step_timesteps"]), 1
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user