[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:
co-authored by
Claude Fable 5
parent
d2b1243be0
commit
c879f3da5c
@@ -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,
|
||||
)
|
||||
Reference in New Issue
Block a user