[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

View File

@@ -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:

View File

@@ -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"

View File

@@ -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

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.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"],
)

View File

@@ -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

View File

@@ -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()