From 3a0324fb9bba280af291c3e27ea5d54b96987315 Mon Sep 17 00:00:00 2001 From: Mick Date: Mon, 21 Sep 2026 08:42:02 +0800 Subject: [PATCH] [diffusion] optimization: reduce Qwen-Image 2.1 vae and graph warmup memory (#40481) Co-authored-by: Mick Qian --- .../models/vaes/autoencoder_kl_qwenimage21.py | 3 ++ .../model_specific_stages/qwen_image21.py | 7 ++++ .../test/unit/test_qwen_image21.py | 23 +++++++++++++ .../test/unit/test_qwen_image21_cuda.py | 32 +++++++++++++++++-- 4 files changed, 62 insertions(+), 3 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage21.py b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage21.py index 2272670df..a227f8a84 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage21.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/autoencoder_kl_qwenimage21.py @@ -197,6 +197,9 @@ class QwenImage21RMS_norm(nn.Module): class QwenImage21Upsample(nn.Upsample): def forward(self, x): + # Nearest interpolation copies values; no FP32 arithmetic is needed. + if self.mode == "nearest-exact" and x.dtype in (torch.float16, torch.bfloat16): + return super().forward(x) return super().forward(x.float()).type_as(x) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image21.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image21.py index f83c0e865..f1efc17f6 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image21.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image21.py @@ -234,6 +234,13 @@ def prepare_qwen21_mu(batch, server_args): class QwenImage21DenoisingStage(DenoisingStage): + def _bcg_pad_prompt_kwargs( + self, call_kwargs, current_model=None, force_bucket=None + ): + # Prefill runs eagerly. Later steps use exact-length prefix KV, so text + # padding only creates duplicate graphs without enabling more replay. + return call_kwargs + def _predict_noise( self, current_model, diff --git a/python/sglang/multimodal_gen/test/unit/test_qwen_image21.py b/python/sglang/multimodal_gen/test/unit/test_qwen_image21.py index d38737612..198a15c0f 100644 --- a/python/sglang/multimodal_gen/test/unit/test_qwen_image21.py +++ b/python/sglang/multimodal_gen/test/unit/test_qwen_image21.py @@ -36,6 +36,7 @@ from sglang.multimodal_gen.runtime.models.encoders.qwen3vl_vision import ( from sglang.multimodal_gen.runtime.models.vaes.autoencoder_kl_qwenimage21 import ( AutoencoderKLQwenImage21, QwenImage21RMS_norm, + QwenImage21Upsample, _patchify, _unpatchify, ) @@ -50,6 +51,28 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.q ) +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +@pytest.mark.parametrize("layout", ["contiguous", "channels_last", "transposed"]) +@pytest.mark.parametrize("device", ["cpu", "cuda"]) +def test_nearest_upsample_preserves_every_finite_low_precision_value( + dtype, layout, device +): + if device == "cuda" and not torch.cuda.is_available(): + pytest.skip("CUDA required") + values = torch.arange(65536, dtype=torch.int32).to(torch.int16).view(dtype) + values = values[torch.isfinite(values)].reshape(1, 2, -1, 128).to(device) + if layout == "channels_last": + values = values.contiguous(memory_format=torch.channels_last) + elif layout == "transposed": + values = values.transpose(2, 3) + upsample = QwenImage21Upsample(scale_factor=2, mode="nearest-exact") + expected = torch.nn.functional.interpolate( + values.float(), scale_factor=2, mode="nearest-exact" + ).to(dtype) + actual = upsample(values) + assert torch.equal(actual.view(torch.int16), expected.view(torch.int16)) + + @pytest.mark.parametrize("prompt", ["edit", ""]) @pytest.mark.parametrize("image_count", [0, 1, 2]) def test_prompt_conditioning_uses_training_template_and_pre_norm(prompt, image_count): diff --git a/python/sglang/multimodal_gen/test/unit/test_qwen_image21_cuda.py b/python/sglang/multimodal_gen/test/unit/test_qwen_image21_cuda.py index 4b6e23476..3b672c50f 100644 --- a/python/sglang/multimodal_gen/test/unit/test_qwen_image21_cuda.py +++ b/python/sglang/multimodal_gen/test/unit/test_qwen_image21_cuda.py @@ -16,6 +16,7 @@ from sglang.multimodal_gen.configs.models.dits.qwenimage21 import ( from sglang.multimodal_gen.configs.pipeline_configs.qwen_image21 import ( QwenImage21PipelineConfig, ) +from sglang.multimodal_gen.configs.sample.qwenimage21 import QwenImage21SamplingParams from sglang.multimodal_gen.runtime.breakable_cuda_graph.runner import ( DiffusionBreakableCudaGraphRunner, ) @@ -33,6 +34,10 @@ from sglang.multimodal_gen.runtime.pipelines.qwen_image21 import QwenImage21Pipe from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import ( ComposedPipelineBase, ) +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.qwen_image21 import ( + QwenImage21DenoisingStage, +) from sglang.multimodal_gen.runtime.server_args import ( ServerArgs, set_global_server_args, @@ -261,14 +266,35 @@ def test_cached_prefix_matches_full_recomputation(model, edit): def test_graph_replay_uses_new_request_prefix(model, edit, sample_count): first = batched_inputs([inputs(5 + i, edit) for i in range(sample_count)]) second = batched_inputs([inputs(9 + i, edit) for i in range(sample_count)]) + for kwargs in (first, second): + kwargs["encoder_hidden_states_mask"] = torch.ones( + kwargs["encoder_hidden_states"].shape[:2], device="cuda", dtype=torch.bool + ) + stage = object.__new__(QwenImage21DenoisingStage) runner = DiffusionBreakableCudaGraphRunner(model, torch.device("cuda")) try: - with torch.no_grad(), set_forward_context(None, None): + with ( + torch.no_grad(), + set_forward_context( + None, + None, + Req(sampling_params=QwenImage21SamplingParams(), is_warmup=True), + ), + ): model(**first) - assert runner.capture(**first) + stage._bcg_run(runner, first, model) + assert len(runner.entries) == 1 + with ( + torch.no_grad(), + set_forward_context( + None, + None, + Req(sampling_params=QwenImage21SamplingParams()), + ), + ): model(**second) expected = model(**second) - actual = runner(**second) + actual = stage._bcg_run(runner, second, model) assert len(runner.entries) == 1 torch.testing.assert_close(actual, expected, atol=1e-6, rtol=1e-6) finally: