[diffusion] fix: fix --warmup-resolutions hang with --enable-cfg-parallel (#23198)

This commit is contained in:
mispa-ms
2026-04-23 13:39:20 +08:00
committed by GitHub
parent 18359aadc8
commit 3c5b1f0810
3 changed files with 237 additions and 15 deletions
@@ -48,6 +48,13 @@ logger = init_logger(__name__)
MINIMUM_PICTURE_BASE64_FOR_WARMUP = "data:image/jpg;base64,iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg=="
# Placeholder negative_prompt used in synthesized warmup Reqs when
# --enable-cfg-parallel is on. A non-empty, real word (vs "" or " ") so
# every tokenizer backend emits a predictable, non-degenerate token
# sequence — rank 1's uncond branch then produces a valid tensor for
# _combine_cfg_parallel's all-reduce.
DEFAULT_PLACEHOLDER_PROMPT = "warmup"
class Scheduler(SchedulerDisaggMixin):
"""
@@ -241,22 +248,25 @@ class Scheduler(SchedulerDisaggMixin):
for resolution in self.server_args.warmup_resolutions:
width, height = _parse_size(resolution)
# CFG-parallel splits cond/uncond across ranks, so rank 1
# needs a real uncond pass. Force do_classifier_free_guidance
# + non-empty negative_prompt when cfg-parallel is on, so the
# synthesized warmup Req exercises both ranks' denoising paths.
# When cfg-parallel is off, the Req construction is
# byte-identical to the pre-fix behavior.
req_kwargs = dict(
data_type=task_type.data_type(),
width=width,
height=height,
prompt="",
)
if requires_warmup_image:
req = Req(
data_type=task_type.data_type(),
width=width,
height=height,
prompt="",
negative_prompt="",
image_path=[warmup_input_path],
)
else:
req = Req(
data_type=task_type.data_type(),
width=width,
height=height,
prompt="",
)
req_kwargs["negative_prompt"] = ""
req_kwargs["image_path"] = [warmup_input_path]
if self.server_args.enable_cfg_parallel:
req_kwargs["negative_prompt"] = DEFAULT_PLACEHOLDER_PROMPT
req_kwargs["do_classifier_free_guidance"] = True
req = Req(**req_kwargs)
req.set_as_warmup(self.server_args.warmup_steps)
self.waiting_queue.append((None, req))
# if server is warmed-up, set this flag to avoid req-based warmup
@@ -313,6 +313,34 @@ class InputValidationStage(PipelineStage):
f"Guidance scale must be positive, but got {batch.guidance_scale}"
)
# Reject requests that do not enable CFG on a server launched with
# --enable-cfg-parallel. CFG-parallel splits cond/uncond across ranks,
# so rank 1 has no work and returns None for noise_pred, which crashes
# scheduler.step() ~30 minutes later under a gloo broadcast timeout.
# Earlier, field-specific checks above (negative_prompt missing,
# guidance_scale < 0) fire first and produce better messages for those
# cases; this is the catch-all for any combination that still leaves
# do_classifier_free_guidance=False under cfg-parallel.
if server_args.enable_cfg_parallel and not batch.do_classifier_free_guidance:
neg_prompt_state = (
"not set"
if batch.negative_prompt is None
else "empty" if batch.negative_prompt == "" else "set"
)
raise ValueError(
f"Server was launched with --enable-cfg-parallel but this "
f"request does not use classifier-free guidance "
f"(do_classifier_free_guidance={batch.do_classifier_free_guidance}, "
f"guidance_scale={batch.guidance_scale}, "
f"true_cfg_scale={batch.true_cfg_scale}, "
f"negative_prompt={neg_prompt_state}). "
f"CFG-parallel splits cond/uncond across ranks and requires "
f"both to be active. Either disable --enable-cfg-parallel or "
f"ensure the request enables CFG (set guidance_scale > 1.0 or "
f"true_cfg_scale > 1.0, with a non-empty negative_prompt or "
f"negative_prompt_embeds)."
)
# for i2v, get image from image_path
# @TODO(Wei) hard-coded for wan2.2 5b ti2v for now. Should put this in image_encoding stage
if batch.image_path is not None:
@@ -0,0 +1,184 @@
"""Unit tests for the --enable-cfg-parallel warmup fix and guard.
Covers two code paths introduced alongside this file:
- Scheduler.prepare_server_warmup_reqs synthesizes warmup Reqs that
actually enable classifier-free guidance when cfg-parallel is on.
- InputValidationStage.forward rejects non-CFG requests when the server
has cfg-parallel on.
All tests are CPU-only; no model loading, no distributed init.
"""
import unittest
from collections import deque
from unittest.mock import MagicMock, patch
from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType
from sglang.multimodal_gen.runtime.managers.scheduler import (
DEFAULT_PLACEHOLDER_PROMPT,
Scheduler,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import (
InputValidationStage,
)
# Patch path for get_global_server_args used by Stage.__init__
_GLOBAL_ARGS_PATCH = (
"sglang.multimodal_gen.runtime.pipelines_core.stages.base.get_global_server_args"
)
def _make_bare_scheduler(enable_cfg_parallel: bool) -> Scheduler:
"""
Build a minimal Scheduler without calling __init__ (which requires
distributed init, ZMQ sockets, pipeline load, etc.). Populates only
the attributes prepare_server_warmup_reqs reads/writes for a
text-only task so _prepare_shared_warmup_image_path is skipped.
"""
scheduler = object.__new__(Scheduler)
server_args = MagicMock()
server_args.warmup = True
server_args.warmup_steps = 1
server_args.warmup_resolutions = ["512x512"]
server_args.enable_cfg_parallel = enable_cfg_parallel
# Text-only task — accepts_image_input() False skips the image-path
# branch entirely, so we don't need to mock
# _prepare_shared_warmup_image_path.
task_type = MagicMock()
task_type.accepts_image_input.return_value = False
task_type.data_type.return_value = ModelTaskType.T2I.data_type()
server_args.pipeline_config.task_type = task_type
scheduler.server_args = server_args
scheduler.warmed_up = False
scheduler.waiting_queue = deque()
return scheduler
def _make_input_validation_stage() -> InputValidationStage:
"""Construct InputValidationStage with the global server-args patch
that existing tests in this suite use (see test_input_validation.py)."""
with patch(_GLOBAL_ARGS_PATCH) as m:
m.return_value = MagicMock()
return InputValidationStage()
def _make_validation_server_args(enable_cfg_parallel: bool) -> MagicMock:
sa = MagicMock()
sa.enable_cfg_parallel = enable_cfg_parallel
sa.pipeline_config.task_type = ModelTaskType.T2I
return sa
class TestWarmupReqCfgParallel(unittest.TestCase):
"""Commit 1 regression: prepare_server_warmup_reqs."""
def test_warmup_req_cfg_parallel_sets_do_cfg(self):
scheduler = _make_bare_scheduler(enable_cfg_parallel=True)
scheduler.prepare_server_warmup_reqs()
self.assertEqual(len(scheduler.waiting_queue), 1)
_, req = scheduler.waiting_queue[0]
self.assertIs(req.do_classifier_free_guidance, True)
self.assertEqual(req.negative_prompt, DEFAULT_PLACEHOLDER_PROMPT)
def test_warmup_req_no_cfg_parallel_unchanged(self):
# Regression guard: the cfg-parallel=on fix must not bleed into
# the cfg-parallel=off path. Key invariant is do_cfg stays False
# AND the synthesized Req is not using the cfg-parallel-specific
# "warmup" placeholder for negative_prompt (which would indicate
# the fix's kwargs leaked into this branch).
scheduler = _make_bare_scheduler(enable_cfg_parallel=False)
scheduler.prepare_server_warmup_reqs()
self.assertEqual(len(scheduler.waiting_queue), 1)
_, req = scheduler.waiting_queue[0]
self.assertIs(req.do_classifier_free_guidance, False)
self.assertNotEqual(req.negative_prompt, DEFAULT_PLACEHOLDER_PROMPT)
class TestInputValidationCfgParallelGuard(unittest.TestCase):
"""Commit 2: per-request cfg-parallel check.
Both tests patch _generate_seeds (the first statement of
InputValidationStage.forward, input_validation.py:274) to sidestep
its device-lookup / generator-creation code which pulls in torch
CUDA bindings — keeps the suite strictly CPU-only. We still need
num_inference_steps on the Req because the stage's
"num_inference_steps <= 0" check at L305-308 raises TypeError on
None before the new commit-2 check is reached.
"""
def test_input_validation_rejects_cfg_parallel_without_cfg(self):
# negative_prompt="" (non-None) ensures the existing
# negative_prompt-is-None check at input_validation.py:295-298
# does NOT fire first — this isolates the new commit-2 check.
# width/height/num_outputs_per_prompt pre-set so the stage's
# default-dimension block at L352-361 doesn't mutate the Req
# in a way that obscures the assertion target.
req = Req(
prompt="test",
negative_prompt="",
guidance_scale=1.0,
true_cfg_scale=None,
num_inference_steps=4,
num_outputs_per_prompt=1,
width=512,
height=512,
)
self.assertIs(
req.do_classifier_free_guidance,
False,
"Sanity: test setup must leave do_cfg=False so the "
"commit-2 check is the one that fires, not an upstream check.",
)
stage = _make_input_validation_stage()
server_args = _make_validation_server_args(enable_cfg_parallel=True)
with patch.object(InputValidationStage, "_generate_seeds"):
with self.assertRaises(ValueError) as ctx:
stage.forward(req, server_args)
msg = str(ctx.exception).lower()
self.assertIn("cfg-parallel", msg)
for field in (
"do_classifier_free_guidance",
"guidance_scale",
"true_cfg_scale",
"negative_prompt",
):
self.assertIn(field, str(ctx.exception))
def test_input_validation_passes_cfg_parallel_with_cfg(self):
req = Req(
prompt="test",
negative_prompt="bad",
guidance_scale=4.0,
true_cfg_scale=4.0,
num_inference_steps=4,
num_outputs_per_prompt=1,
width=512,
height=512,
)
self.assertIs(
req.do_classifier_free_guidance,
True,
"Sanity: req must enable CFG for this positive-case test.",
)
stage = _make_input_validation_stage()
server_args = _make_validation_server_args(enable_cfg_parallel=True)
with patch.object(InputValidationStage, "_generate_seeds"):
try:
stage.forward(req, server_args)
except ValueError as e:
self.fail(f"forward() raised ValueError on a valid CFG request: {e}")
if __name__ == "__main__":
unittest.main()