[diffusion] rl: support rl rollout for the wan pipeline via a per-request scheduler switch (#30036)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Andy Ye
2026-07-15 22:30:51 +08:00
committed by GitHub
co-authored by Claude Fable 5
parent d2b1243be0
commit c879f3da5c
4 changed files with 85 additions and 17 deletions
@@ -171,10 +171,7 @@ def _extract_extra_fields(extra: dict, scalar_fields: dict) -> None:
pass
def _init_request_scheduler_from_template(
scheduler_template: Any, req: Req, device: torch.device
) -> None:
scheduler = clone_scheduler_runtime(scheduler_template)
def _init_request_scheduler(scheduler: Any, req: Req, device: torch.device) -> None:
extra_kwargs = {}
mu = req.extra.get("mu") if hasattr(req, "extra") else None
if mu is not None:
@@ -202,12 +199,33 @@ def _init_request_scheduler_from_template(
req.timesteps = scheduler.timesteps
def _init_request_scheduler_from_template(
scheduler_template: Any, req: Req, device: torch.device
) -> None:
scheduler = clone_scheduler_runtime(scheduler_template)
_init_request_scheduler(scheduler, req, device)
def _init_disagg_request_scheduler(self: Scheduler, req: Req) -> None:
scheduler_template = self.worker.pipeline.get_module("scheduler")
if scheduler_template is None:
serving_scheduler = self.worker.pipeline.get_module("scheduler")
if serving_scheduler is None:
return
device = torch.device(f"{current_platform.device_type}:{self.worker.local_rank}")
_init_request_scheduler_from_template(scheduler_template, req, device)
if not req.rollout:
_init_request_scheduler_from_template(serving_scheduler, req, device)
return
from sglang.multimodal_gen.runtime.post_training.rollout_scheduler import (
get_or_create_rollout_request_scheduler,
)
scheduler = get_or_create_rollout_request_scheduler(
req,
serving_scheduler,
isolate=True,
)
_init_request_scheduler(scheduler, req, device)
def extract_transfer_fields(req) -> tuple[dict, dict]:
@@ -90,7 +90,14 @@ class TimestepPreparationStage(PipelineStage):
if batch.scheduler is not None and batch.timesteps is not None:
return batch
scheduler = get_or_create_request_scheduler(batch, self.scheduler)
if batch.rollout:
from sglang.multimodal_gen.runtime.post_training.rollout_scheduler import (
get_or_create_rollout_request_scheduler,
)
scheduler = get_or_create_rollout_request_scheduler(batch, self.scheduler)
else:
scheduler = get_or_create_request_scheduler(batch, self.scheduler)
device = get_local_torch_device()
num_inference_steps = batch.num_inference_steps
timesteps = batch.timesteps
@@ -40,25 +40,27 @@ class RolloutDenoisingMixin:
def _maybe_prepare_rollout(self, batch: Req):
"""Prepare denoising loop for rollout."""
if not isinstance(self.scheduler, SchedulerRLMixin):
scheduler = batch.scheduler
if not isinstance(scheduler, SchedulerRLMixin):
if batch.rollout:
raise ValueError(
f"Scheduler {type(self.scheduler)} does not support rollout"
f"Scheduler {type(scheduler)} does not support rollout"
)
return
self.scheduler.release_rollout_resources(batch)
scheduler.release_rollout_resources(batch)
if batch.rollout:
self.scheduler.prepare_rollout(
scheduler.prepare_rollout(
batch=batch,
pipeline_config=self.server_args.pipeline_config,
)
def _maybe_collect_rollout_log_probs(self, batch: Req):
if not isinstance(self.scheduler, SchedulerRLMixin):
scheduler = batch.scheduler
if not isinstance(scheduler, SchedulerRLMixin):
if batch.rollout:
raise ValueError(
f"Scheduler {type(self.scheduler)} does not support rollout"
f"Scheduler {type(scheduler)} does not support rollout"
)
return
@@ -66,13 +68,13 @@ class RolloutDenoisingMixin:
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)
scheduler.collect_rollout_log_probs(batch)
)
if batch.rollout_debug_mode:
batch.rollout_trajectory_data.rollout_debug_tensors = (
self.scheduler.collect_rollout_debug_tensors(batch)
scheduler.collect_rollout_debug_tensors(batch)
)
self.scheduler.release_rollout_resources(batch)
scheduler.release_rollout_resources(batch)
def _postprocess_rollout_outputs(
self,
@@ -0,0 +1,41 @@
from typing import Any
from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler,
)
from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler,
)
from sglang.multimodal_gen.runtime.pipelines_core.diffusion_scheduler_utils import (
get_or_create_request_scheduler,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
def rollout_scheduler_for(serving):
"""Some serving schedulers cannot be used for rollout; map them to one
that can. Schedulers without a mapping pass through unchanged.
"""
if isinstance(serving, FlowUniPCMultistepScheduler):
return FlowMatchEulerDiscreteScheduler(shift=serving.config.shift)
return serving
def get_or_create_rollout_request_scheduler(
batch: Req,
serving_scheduler: Any,
*,
isolate: bool = False,
) -> Any:
"""Return the scheduler runtime for a rollout request."""
if batch.scheduler is not None:
return batch.scheduler
scheduler = rollout_scheduler_for(serving_scheduler)
scheduler_is_shared = scheduler is serving_scheduler
needs_clone = isolate and scheduler_is_shared
return get_or_create_request_scheduler(
batch,
scheduler,
isolate=needs_clone,
)