[diffusion] improve: tiny speedup qwen-image-edit-2511 by avoiding unnecessary calculation (#15896)

This commit is contained in:
Mick
2025-12-30 10:10:32 +08:00
committed by GitHub
parent 3de23274ee
commit 26e17f9076
6 changed files with 53 additions and 15 deletions
@@ -596,7 +596,7 @@ def calculate_metrics(outputs: List[RequestFuncOutput], total_duration: float):
return metrics 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}...") print(f"Waiting for service at {base_url}...")
start_time = time.time() start_time = time.time()
while True: while True:
@@ -38,8 +38,22 @@ class QwenImageArchConfig(DiTArchConfig):
self.num_channels_latents = self.out_channels self.num_channels_latents = self.out_channels
@dataclass
class QwenImageEditPlus_2511_ArchConfig(DiTArchConfig):
zero_cond_t: bool = True
@dataclass @dataclass
class QwenImageDitConfig(DiTConfig): class QwenImageDitConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=QwenImageArchConfig) arch_config: DiTArchConfig = field(default_factory=QwenImageArchConfig)
prefix: str = "qwenimage" prefix: str = "qwenimage"
@dataclass
class QwenImageEditPlus_2511_DitConfig(DiTConfig):
arch_config: DiTArchConfig = field(
default_factory=QwenImageEditPlus_2511_ArchConfig
)
prefix: str = "qwenimageedit"
@@ -6,7 +6,10 @@ from typing import Callable
import torch import torch
from sglang.multimodal_gen.configs.models import DiTConfig, EncoderConfig, VAEConfig 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.encoders.qwen_image import Qwen2_5VLConfig
from sglang.multimodal_gen.configs.models.vaes.qwenimage import QwenImageVAEConfig from sglang.multimodal_gen.configs.models.vaes.qwenimage import QwenImageVAEConfig
from sglang.multimodal_gen.configs.pipeline_configs.base import ( from sglang.multimodal_gen.configs.pipeline_configs.base import (
@@ -415,7 +418,6 @@ class QwenImageEditPlusPipelineConfig(QwenImageEditPipelineConfig):
assert batch_size == 1 assert batch_size == 1
height = batch.height height = batch.height
width = batch.width width = batch.width
image_size = batch.original_condition_image_size
vae_scale_factor = self.get_vae_scale_factor() 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 @dataclass
class QwenImageLayeredPipelineConfig(QwenImageEditPipelineConfig): class QwenImageLayeredPipelineConfig(QwenImageEditPipelineConfig):
resolution: int = 640 # TODO: allow user to set resolution resolution: int = 640 # TODO: allow user to set resolution
+2 -1
View File
@@ -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.flux import Flux2PipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import ( from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImageEditPipelineConfig, QwenImageEditPipelineConfig,
QwenImageEditPlus_2511_PipelineConfig,
QwenImageEditPlusPipelineConfig, QwenImageEditPlusPipelineConfig,
QwenImageLayeredPipelineConfig, QwenImageLayeredPipelineConfig,
QwenImagePipelineConfig, QwenImagePipelineConfig,
@@ -429,7 +430,7 @@ def _register_configs():
register_configs( register_configs(
sampling_param_cls=QwenImageEditPlusSamplingParams, sampling_param_cls=QwenImageEditPlusSamplingParams,
pipeline_config_cls=QwenImageEditPlusPipelineConfig, pipeline_config_cls=QwenImageEditPlus_2511_PipelineConfig,
hf_model_paths=["Qwen/Qwen-Image-Edit-2511"], hf_model_paths=["Qwen/Qwen-Image-Edit-2511"],
) )
@@ -3,7 +3,6 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
import functools import functools
from math import prod
from typing import Any, Dict, List, Optional, Tuple, Union from typing import Any, Dict, List, Optional, Tuple, Union
import numpy as np import numpy as np
@@ -16,6 +15,7 @@ from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.normalization import AdaLayerNormContinuous from diffusers.models.normalization import AdaLayerNormContinuous
from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig 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.attention import USPAttention
from sglang.multimodal_gen.runtime.layers.layernorm import ( from sglang.multimodal_gen.runtime.layers.layernorm import (
LayerNorm, LayerNorm,
@@ -792,6 +792,12 @@ class QwenImageTransformerBlock(nn.Module):
return encoder_hidden_states, hidden_states 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): class QwenImageTransformer2DModel(CachableDiT):
""" """
The Transformer model introduced in Qwen. 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.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( def forward(
self, self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
@@ -923,16 +945,9 @@ class QwenImageTransformer2DModel(CachableDiT):
timestep = (timestep / 1000).to(hidden_states.dtype) timestep = (timestep / 1000).to(hidden_states.dtype)
if self.zero_cond_t: if self.zero_cond_t:
timestep = torch.cat([timestep, timestep * 0], dim=0) timestep = torch.cat([timestep, self.timestep_zero], dim=0)
# Use torch operations for GPU efficiency device = timestep.device
modulate_index = torch.tensor( modulate_index = self.build_modulate_index(to_hashable(img_shapes), device)
[
[0] * prod(sample[0]) + [1] * sum([prod(s) for s in sample[1:]])
for sample in img_shapes
],
device=timestep.device,
dtype=torch.int,
)
else: else:
modulate_index = None modulate_index = None
@@ -487,6 +487,7 @@ class DenoisingStage(PipelineStage):
Returns: Returns:
A dictionary containing all the prepared variables for the denoising loop. 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 pipeline = self.pipeline() if self.pipeline else None
if not server_args.model_loaded["transformer"]: if not server_args.model_loaded["transformer"]:
loader = TransformerLoader() loader = TransformerLoader()