From 26e17f9076809b34828d61d0eabe354066f74e35 Mon Sep 17 00:00:00 2001 From: Mick Date: Tue, 30 Dec 2025 10:10:32 +0800 Subject: [PATCH] [diffusion] improve: tiny speedup qwen-image-edit-2511 by avoiding unnecessary calculation (#15896) --- .../benchmarks/bench_serving.py | 2 +- .../configs/models/dits/qwenimage.py | 14 +++++++ .../configs/pipeline_configs/qwen_image.py | 11 +++++- python/sglang/multimodal_gen/registry.py | 3 +- .../runtime/models/dits/qwen_image.py | 37 +++++++++++++------ .../pipelines_core/stages/denoising.py | 1 + 6 files changed, 53 insertions(+), 15 deletions(-) diff --git a/python/sglang/multimodal_gen/benchmarks/bench_serving.py b/python/sglang/multimodal_gen/benchmarks/bench_serving.py index 698c43353..86a0ef081 100644 --- a/python/sglang/multimodal_gen/benchmarks/bench_serving.py +++ b/python/sglang/multimodal_gen/benchmarks/bench_serving.py @@ -596,7 +596,7 @@ def calculate_metrics(outputs: List[RequestFuncOutput], total_duration: float): return metrics -def wait_for_service(base_url: str, timeout: int = 120) -> None: +def wait_for_service(base_url: str, timeout: int = 1200) -> None: print(f"Waiting for service at {base_url}...") start_time = time.time() while True: diff --git a/python/sglang/multimodal_gen/configs/models/dits/qwenimage.py b/python/sglang/multimodal_gen/configs/models/dits/qwenimage.py index 0738b63b0..53bf3aaba 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/qwenimage.py +++ b/python/sglang/multimodal_gen/configs/models/dits/qwenimage.py @@ -38,8 +38,22 @@ class QwenImageArchConfig(DiTArchConfig): self.num_channels_latents = self.out_channels +@dataclass +class QwenImageEditPlus_2511_ArchConfig(DiTArchConfig): + zero_cond_t: bool = True + + @dataclass class QwenImageDitConfig(DiTConfig): arch_config: DiTArchConfig = field(default_factory=QwenImageArchConfig) prefix: str = "qwenimage" + + +@dataclass +class QwenImageEditPlus_2511_DitConfig(DiTConfig): + arch_config: DiTArchConfig = field( + default_factory=QwenImageEditPlus_2511_ArchConfig + ) + + prefix: str = "qwenimageedit" diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py b/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py index f3fe94188..e29df2f9f 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py @@ -6,7 +6,10 @@ from typing import Callable import torch from sglang.multimodal_gen.configs.models import DiTConfig, EncoderConfig, VAEConfig -from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig +from sglang.multimodal_gen.configs.models.dits.qwenimage import ( + QwenImageDitConfig, + QwenImageEditPlus_2511_DitConfig, +) from sglang.multimodal_gen.configs.models.encoders.qwen_image import Qwen2_5VLConfig from sglang.multimodal_gen.configs.models.vaes.qwenimage import QwenImageVAEConfig from sglang.multimodal_gen.configs.pipeline_configs.base import ( @@ -415,7 +418,6 @@ class QwenImageEditPlusPipelineConfig(QwenImageEditPipelineConfig): assert batch_size == 1 height = batch.height width = batch.width - image_size = batch.original_condition_image_size vae_scale_factor = self.get_vae_scale_factor() @@ -474,6 +476,11 @@ class QwenImageEditPlusPipelineConfig(QwenImageEditPipelineConfig): } +@dataclass +class QwenImageEditPlus_2511_PipelineConfig(QwenImageEditPlusPipelineConfig): + dit_config: DiTConfig = field(default_factory=QwenImageEditPlus_2511_DitConfig) + + @dataclass class QwenImageLayeredPipelineConfig(QwenImageEditPipelineConfig): resolution: int = 640 # TODO: allow user to set resolution diff --git a/python/sglang/multimodal_gen/registry.py b/python/sglang/multimodal_gen/registry.py index 42cb20259..b81391170 100644 --- a/python/sglang/multimodal_gen/registry.py +++ b/python/sglang/multimodal_gen/registry.py @@ -28,6 +28,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.flux import Flux2PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import ( QwenImageEditPipelineConfig, + QwenImageEditPlus_2511_PipelineConfig, QwenImageEditPlusPipelineConfig, QwenImageLayeredPipelineConfig, QwenImagePipelineConfig, @@ -429,7 +430,7 @@ def _register_configs(): register_configs( sampling_param_cls=QwenImageEditPlusSamplingParams, - pipeline_config_cls=QwenImageEditPlusPipelineConfig, + pipeline_config_cls=QwenImageEditPlus_2511_PipelineConfig, hf_model_paths=["Qwen/Qwen-Image-Edit-2511"], ) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py index 53e138641..c3a2f2406 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py @@ -3,7 +3,6 @@ # SPDX-License-Identifier: Apache-2.0 import functools -from math import prod from typing import Any, Dict, List, Optional, Tuple, Union import numpy as np @@ -16,6 +15,7 @@ from diffusers.models.modeling_outputs import Transformer2DModelOutput from diffusers.models.normalization import AdaLayerNormContinuous from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig +from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.layers.attention import USPAttention from sglang.multimodal_gen.runtime.layers.layernorm import ( LayerNorm, @@ -792,6 +792,12 @@ class QwenImageTransformerBlock(nn.Module): return encoder_hidden_states, hidden_states +def to_hashable(obj): + if isinstance(obj, list): + return tuple(to_hashable(x) for x in obj) + return obj + + class QwenImageTransformer2DModel(CachableDiT): """ The Transformer model introduced in Qwen. @@ -868,6 +874,22 @@ class QwenImageTransformer2DModel(CachableDiT): self.inner_dim, patch_size * patch_size * self.out_channels, bias=True ) + self.timestep_zero = torch.zeros( + (1,), dtype=torch.int, device=get_local_torch_device() + ) + + @functools.lru_cache(maxsize=50) + def build_modulate_index(self, img_shapes: tuple[int, int, int], device): + modulate_index_list = [] + for sample in img_shapes: + first_size = sample[0][0] * sample[0][1] * sample[0][2] + total_size = sum(s[0] * s[1] * s[2] for s in sample) + idx = (torch.arange(total_size, device=device) >= first_size).int() + modulate_index_list.append(idx) + + modulate_index = torch.stack(modulate_index_list) + return modulate_index + def forward( self, hidden_states: torch.Tensor, @@ -923,16 +945,9 @@ class QwenImageTransformer2DModel(CachableDiT): timestep = (timestep / 1000).to(hidden_states.dtype) if self.zero_cond_t: - timestep = torch.cat([timestep, timestep * 0], dim=0) - # Use torch operations for GPU efficiency - modulate_index = torch.tensor( - [ - [0] * prod(sample[0]) + [1] * sum([prod(s) for s in sample[1:]]) - for sample in img_shapes - ], - device=timestep.device, - dtype=torch.int, - ) + timestep = torch.cat([timestep, self.timestep_zero], dim=0) + device = timestep.device + modulate_index = self.build_modulate_index(to_hashable(img_shapes), device) else: modulate_index = None diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index 4c4dcdc7e..8f71617aa 100755 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -487,6 +487,7 @@ class DenoisingStage(PipelineStage): Returns: A dictionary containing all the prepared variables for the denoising loop. """ + assert self.transformer is not None pipeline = self.pipeline() if self.pipeline else None if not server_args.model_loaded["transformer"]: loader = TransformerLoader()