[diffusion] warmup: default to model sampling resolution (declare Z-Image default) (#29519)
This commit is contained in:
@@ -13,8 +13,9 @@ class ZImageTurboSamplingParams(SamplingParams):
|
|||||||
|
|
||||||
num_frames: int = 1
|
num_frames: int = 1
|
||||||
negative_prompt: str = None
|
negative_prompt: str = None
|
||||||
# height: int = 720
|
# Z-Image officially recommends starting at 1024x1024
|
||||||
# width: int = 1280
|
height: int = 1024
|
||||||
|
width: int = 1024
|
||||||
# fps: int = 24
|
# fps: int = 24
|
||||||
|
|
||||||
guidance_scale: float = 0.0
|
guidance_scale: float = 0.0
|
||||||
|
|||||||
@@ -65,15 +65,28 @@ def _resolve_default_warmup_resolution(
|
|||||||
*,
|
*,
|
||||||
server_based_warmup: bool,
|
server_based_warmup: bool,
|
||||||
) -> tuple[int, int]:
|
) -> tuple[int, int]:
|
||||||
"""returns a default resolution to warmup"""
|
"""Return the default warmup resolution.
|
||||||
if server_based_warmup:
|
|
||||||
return _resolve_representative_warmup_resolution(server_args, sampling_defaults)
|
|
||||||
|
|
||||||
|
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
|
width = sampling_defaults.width
|
||||||
height = sampling_defaults.height
|
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
|
return width, height
|
||||||
|
|
||||||
|
if server_based_warmup:
|
||||||
|
return _resolve_representative_warmup_resolution(server_args, sampling_defaults)
|
||||||
|
|
||||||
supported_resolutions = sampling_defaults.supported_resolutions
|
supported_resolutions = sampling_defaults.supported_resolutions
|
||||||
if supported_resolutions:
|
if supported_resolutions:
|
||||||
return min(supported_resolutions, key=lambda size: size[0] * size[1])
|
return min(supported_resolutions, key=lambda size: size[0] * size[1])
|
||||||
|
|||||||
@@ -374,10 +374,10 @@
|
|||||||
"ImageVAEEncodingStage": 0.01
|
"ImageVAEEncodingStage": 0.01
|
||||||
},
|
},
|
||||||
"denoise_step_ms": {
|
"denoise_step_ms": {
|
||||||
"0": 16.21,
|
"0": 28.54,
|
||||||
"1": 13.42,
|
"1": 12.61,
|
||||||
"2": 17.64,
|
"2": 57.24,
|
||||||
"3": 63.6
|
"3": 64.49
|
||||||
},
|
},
|
||||||
"expected_e2e_ms": 434.57,
|
"expected_e2e_ms": 434.57,
|
||||||
"expected_avg_denoise_ms": 39.98,
|
"expected_avg_denoise_ms": 39.98,
|
||||||
@@ -2101,7 +2101,7 @@
|
|||||||
"3": 650.0,
|
"3": 650.0,
|
||||||
"4": 271.46,
|
"4": 271.46,
|
||||||
"5": 266.71,
|
"5": 266.71,
|
||||||
"6": 232.55,
|
"6": 650.0,
|
||||||
"7": 270.08,
|
"7": 270.08,
|
||||||
"8": 259.96
|
"8": 259.96
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -317,7 +317,11 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
|
|||||||
self.assertEqual(req.negative_prompt, "model default negative")
|
self.assertEqual(req.negative_prompt, "model default negative")
|
||||||
self.assertIs(req.do_classifier_free_guidance, True)
|
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 = MagicMock()
|
||||||
server_args.warmup_steps = 1
|
server_args.warmup_steps = 1
|
||||||
server_args.enable_cfg_parallel = False
|
server_args.enable_cfg_parallel = False
|
||||||
@@ -344,9 +348,12 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
|
|||||||
server_based_warmup=True,
|
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 = MagicMock()
|
||||||
server_args.warmup_steps = 1
|
server_args.warmup_steps = 1
|
||||||
server_args.enable_cfg_parallel = False
|
server_args.enable_cfg_parallel = False
|
||||||
@@ -370,9 +377,11 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
|
|||||||
server_based_warmup=True,
|
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 = MagicMock()
|
||||||
server_args.warmup_steps = 1
|
server_args.warmup_steps = 1
|
||||||
server_args.enable_cfg_parallel = False
|
server_args.enable_cfg_parallel = False
|
||||||
@@ -396,7 +405,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
|
|||||||
server_based_warmup=True,
|
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):
|
def test_server_based_warmup_keeps_video_warmup_lightweight(self):
|
||||||
server_args = MagicMock()
|
server_args = MagicMock()
|
||||||
|
|||||||
Reference in New Issue
Block a user