[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
|
pass
|
||||||
|
|
||||||
|
|
||||||
def _init_request_scheduler_from_template(
|
def _init_request_scheduler(scheduler: Any, req: Req, device: torch.device) -> None:
|
||||||
scheduler_template: Any, req: Req, device: torch.device
|
|
||||||
) -> None:
|
|
||||||
scheduler = clone_scheduler_runtime(scheduler_template)
|
|
||||||
extra_kwargs = {}
|
extra_kwargs = {}
|
||||||
mu = req.extra.get("mu") if hasattr(req, "extra") else None
|
mu = req.extra.get("mu") if hasattr(req, "extra") else None
|
||||||
if mu is not None:
|
if mu is not None:
|
||||||
@@ -202,12 +199,33 @@ def _init_request_scheduler_from_template(
|
|||||||
req.timesteps = scheduler.timesteps
|
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:
|
def _init_disagg_request_scheduler(self: Scheduler, req: Req) -> None:
|
||||||
scheduler_template = self.worker.pipeline.get_module("scheduler")
|
serving_scheduler = self.worker.pipeline.get_module("scheduler")
|
||||||
if scheduler_template is None:
|
if serving_scheduler is None:
|
||||||
return
|
return
|
||||||
device = torch.device(f"{current_platform.device_type}:{self.worker.local_rank}")
|
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]:
|
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:
|
if batch.scheduler is not None and batch.timesteps is not None:
|
||||||
return batch
|
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()
|
device = get_local_torch_device()
|
||||||
num_inference_steps = batch.num_inference_steps
|
num_inference_steps = batch.num_inference_steps
|
||||||
timesteps = batch.timesteps
|
timesteps = batch.timesteps
|
||||||
|
|||||||
@@ -40,25 +40,27 @@ class RolloutDenoisingMixin:
|
|||||||
|
|
||||||
def _maybe_prepare_rollout(self, batch: Req):
|
def _maybe_prepare_rollout(self, batch: Req):
|
||||||
"""Prepare denoising loop for rollout."""
|
"""Prepare denoising loop for rollout."""
|
||||||
if not isinstance(self.scheduler, SchedulerRLMixin):
|
scheduler = batch.scheduler
|
||||||
|
if not isinstance(scheduler, SchedulerRLMixin):
|
||||||
if batch.rollout:
|
if batch.rollout:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Scheduler {type(self.scheduler)} does not support rollout"
|
f"Scheduler {type(scheduler)} does not support rollout"
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
self.scheduler.release_rollout_resources(batch)
|
scheduler.release_rollout_resources(batch)
|
||||||
if batch.rollout:
|
if batch.rollout:
|
||||||
self.scheduler.prepare_rollout(
|
scheduler.prepare_rollout(
|
||||||
batch=batch,
|
batch=batch,
|
||||||
pipeline_config=self.server_args.pipeline_config,
|
pipeline_config=self.server_args.pipeline_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _maybe_collect_rollout_log_probs(self, batch: Req):
|
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:
|
if batch.rollout:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Scheduler {type(self.scheduler)} does not support rollout"
|
f"Scheduler {type(scheduler)} does not support rollout"
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -66,13 +68,13 @@ class RolloutDenoisingMixin:
|
|||||||
if batch.rollout_trajectory_data is None:
|
if batch.rollout_trajectory_data is None:
|
||||||
batch.rollout_trajectory_data = RolloutTrajectoryData()
|
batch.rollout_trajectory_data = RolloutTrajectoryData()
|
||||||
batch.rollout_trajectory_data.rollout_log_probs = (
|
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:
|
if batch.rollout_debug_mode:
|
||||||
batch.rollout_trajectory_data.rollout_debug_tensors = (
|
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(
|
def _postprocess_rollout_outputs(
|
||||||
self,
|
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