[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
|
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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user