From c879f3da5ceaaef3cb197c4e59ce683d420ce96c Mon Sep 17 00:00:00 2001 From: Andy Ye Date: Wed, 15 Jul 2026 07:30:51 -0700 Subject: [PATCH] [diffusion] rl: support rl rollout for the wan pipeline via a per-request scheduler switch (#30036) Co-authored-by: Claude Fable 5 --- .../runtime/disaggregation/scheduler_mixin.py | 32 +++++++++++---- .../stages/timestep_preparation.py | 9 +++- .../post_training/rollout_denoising_mixin.py | 20 +++++---- .../post_training/rollout_scheduler.py | 41 +++++++++++++++++++ 4 files changed, 85 insertions(+), 17 deletions(-) create mode 100644 python/sglang/multimodal_gen/runtime/post_training/rollout_scheduler.py diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py index 24dc62806..a50316629 100644 --- a/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py +++ b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py @@ -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]: diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/timestep_preparation.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/timestep_preparation.py index ea581f9b8..09f8a9cb8 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/timestep_preparation.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/timestep_preparation.py @@ -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 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 7b6e222ca..c161973a1 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 @@ -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, diff --git a/python/sglang/multimodal_gen/runtime/post_training/rollout_scheduler.py b/python/sglang/multimodal_gen/runtime/post_training/rollout_scheduler.py new file mode 100644 index 000000000..da2d3b7bc --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/post_training/rollout_scheduler.py @@ -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, + )