[diffusion] warmup: default to model sampling resolution (declare Z-Image default) (#29519)

This commit is contained in:
Mick
2026-06-30 10:32:11 +08:00
committed by GitHub
parent 3e16be2122
commit 25b6051c70
4 changed files with 40 additions and 17 deletions
@@ -13,8 +13,9 @@ class ZImageTurboSamplingParams(SamplingParams):
num_frames: int = 1
negative_prompt: str = None
# height: int = 720
# width: int = 1280
# Z-Image officially recommends starting at 1024x1024
height: int = 1024
width: int = 1024
# fps: int = 24
guidance_scale: float = 0.0
@@ -65,15 +65,28 @@ def _resolve_default_warmup_resolution(
*,
server_based_warmup: bool,
) -> tuple[int, int]:
"""returns a default resolution to warmup"""
if server_based_warmup:
return _resolve_representative_warmup_resolution(server_args, sampling_defaults)
"""Return the default warmup resolution.
Prefer the model's sampling-default resolution — the most likely real
request shape — so warmup specializes kernels for it. Server-based image
warmup used to shrink this to an area cap (``SERVER_WARMUP_IMAGE_MAX_AREA``,
768x768) to bound startup, but that left a residual first-request
cold-start when the real request is larger (e.g. 1024x1024 paid ~0.1s of
first-shape kernel autotuning, measured on H100).
"""
width = sampling_defaults.width
height = sampling_defaults.height
if width is not None and height is not None:
is_image_gen = server_args.pipeline_config.task_type.is_image_gen()
if (
width is not None
and height is not None
and (not server_based_warmup or is_image_gen)
):
return width, height
if server_based_warmup:
return _resolve_representative_warmup_resolution(server_args, sampling_defaults)
supported_resolutions = sampling_defaults.supported_resolutions
if supported_resolutions:
return min(supported_resolutions, key=lambda size: size[0] * size[1])
@@ -374,10 +374,10 @@
"ImageVAEEncodingStage": 0.01
},
"denoise_step_ms": {
"0": 16.21,
"1": 13.42,
"2": 17.64,
"3": 63.6
"0": 28.54,
"1": 12.61,
"2": 57.24,
"3": 64.49
},
"expected_e2e_ms": 434.57,
"expected_avg_denoise_ms": 39.98,
@@ -2101,7 +2101,7 @@
"3": 650.0,
"4": 271.46,
"5": 266.71,
"6": 232.55,
"6": 650.0,
"7": 270.08,
"8": 259.96
},
@@ -317,7 +317,11 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
self.assertEqual(req.negative_prompt, "model default negative")
self.assertIs(req.do_classifier_free_guidance, True)
def test_server_based_warmup_uses_supported_resolution_within_budget(self):
def test_server_based_image_warmup_uses_model_default_over_supported(self):
"""Server-based image warmup uses the model's default resolution so it
warms up at the real inference shape (avoiding a residual
cudagraph/compile gap), rather than shrinking to the smallest supported
resolution within an area budget."""
server_args = MagicMock()
server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False
@@ -344,9 +348,12 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_based_warmup=True,
)
self.assertEqual((reqs[0].width, reqs[0].height), (512, 512))
self.assertEqual((reqs[0].width, reqs[0].height), (1024, 1024))
def test_server_based_warmup_scales_large_image_default(self):
def test_server_based_image_warmup_uses_full_model_default(self):
"""Server-based image warmup keeps the model's full default resolution
instead of scaling down to a server-warmup area budget, so warmup hits
the real inference shape."""
server_args = MagicMock()
server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False
@@ -370,9 +377,11 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_based_warmup=True,
)
self.assertEqual((reqs[0].width, reqs[0].height), (768, 768))
self.assertEqual((reqs[0].width, reqs[0].height), (1024, 1024))
def test_server_based_warmup_uses_diffusers_image_budget(self):
def test_server_based_image_warmup_diffusers_uses_model_default(self):
"""Even on the diffusers backend, server-based image warmup uses the
model default resolution rather than the diffusers image area budget."""
server_args = MagicMock()
server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False
@@ -396,7 +405,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_based_warmup=True,
)
self.assertEqual((reqs[0].width, reqs[0].height), (512, 512))
self.assertEqual((reqs[0].width, reqs[0].height), (1024, 1024))
def test_server_based_warmup_keeps_video_warmup_lightweight(self):
server_args = MagicMock()