[diffusion] improve: tiny speedup qwen-image-edit-2511 by avoiding unnecessary calculation (#15896)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"],
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user