[diffusion] feat: support tp for qwen-image-edit-2511 (#18619)

Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
xiaoye
2026-02-13 13:04:29 +08:00
committed by GitHub
parent d29e331491
commit 5700b19cbf
3 changed files with 109 additions and 24 deletions

View File

@@ -218,7 +218,7 @@ class GPUWorker:
# Save output to file and return file path only if requested. Avoid the serialization
# and deserialization overhead between scheduler_client and gpu_worker.
if req.save_output and req.return_file_paths_only:
if req.save_output and req.return_file_paths_only and self.rank == 0:
output_paths = save_outputs(
output_batch.output,
req.data_type,

View File

@@ -9,7 +9,6 @@ import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from diffusers.models.attention import FeedForward
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.normalization import AdaLayerNormContinuous
@@ -25,7 +24,10 @@ from sglang.multimodal_gen.runtime.layers.layernorm import (
apply_layernorm_only,
apply_qk_norm,
)
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
from sglang.multimodal_gen.runtime.layers.linear import (
ColumnParallelLinear,
RowParallelLinear,
)
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
apply_flashinfer_rope_qk_inplace,
)
@@ -60,6 +62,81 @@ def _get_qkv_projections(
return img_query, img_key, img_value, txt_query, txt_key, txt_value
class GELU(nn.Module):
r"""
GELU activation function with tanh approximation support with `approximate="tanh"`.
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation.
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
"""
def __init__(
self, dim_in: int, dim_out: int, approximate: str = "none", bias: bool = True
):
super().__init__()
self.proj = ColumnParallelLinear(
dim_in, dim_out, bias=bias, gather_output=False
)
self.approximate = approximate
def forward(self, hidden_states):
hidden_states = self.proj(hidden_states)
return F.gelu(hidden_states[0], approximate=self.approximate)
class FeedForward(nn.Module):
r"""
A feed-forward layer.
Parameters:
dim (`int`): The number of channels in the input.
dim_out (`int`, *optional*): The number of channels in the output. If not given, defaults to `dim`.
mult (`int`, *optional*, defaults to 4): The multiplier to use for the hidden dimension.
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
final_dropout (`bool` *optional*, defaults to False): Apply a final dropout.
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
"""
def __init__(
self,
dim: int,
dim_out: Optional[int] = None,
mult: int = 4,
activation_fn: str = "geglu",
inner_dim=None,
bias: bool = True,
):
super().__init__()
if inner_dim is None:
inner_dim = int(dim * mult)
dim_out = dim_out if dim_out is not None else dim
if activation_fn == "gelu":
act_fn = GELU(dim, inner_dim, bias=bias)
if activation_fn == "gelu-approximate":
act_fn = GELU(dim, inner_dim, approximate="tanh", bias=bias)
else:
raise NotImplementedError(
f"activation_fn '{activation_fn}' is not supported."
)
self.net = nn.ModuleList([])
self.net.append(act_fn)
self.net.append(nn.Identity())
self.net.append(
RowParallelLinear(inner_dim, dim_out, bias=True, input_is_parallel=True)
)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
for module in self.net:
hidden_states = module(hidden_states)
return hidden_states
class QwenTimestepProjEmbeddings(nn.Module):
def __init__(self, embedding_dim, use_additional_t_cond=False):
super().__init__()
@@ -484,34 +561,44 @@ class QwenImageCrossAttention(nn.Module):
# Use separate Q/K/V projections
self.inner_dim = out_dim if out_dim is not None else head_dim * num_heads
self.inner_kv_dim = self.inner_dim
self.to_q = ReplicatedLinear(dim, self.inner_dim, bias=True)
self.to_k = ReplicatedLinear(dim, self.inner_dim, bias=True)
self.to_v = ReplicatedLinear(dim, self.inner_dim, bias=True)
self.to_q = ColumnParallelLinear(
dim, self.inner_dim, bias=True, gather_output=True
)
self.to_k = ColumnParallelLinear(
dim, self.inner_dim, bias=True, gather_output=True
)
self.to_v = ColumnParallelLinear(
dim, self.inner_dim, bias=True, gather_output=True
)
if self.qk_norm:
self.norm_q = RMSNorm(head_dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = RMSNorm(head_dim, eps=eps) if qk_norm else nn.Identity()
if added_kv_proj_dim is not None:
self.add_q_proj = ReplicatedLinear(
added_kv_proj_dim, self.inner_dim, bias=True
self.add_q_proj = ColumnParallelLinear(
added_kv_proj_dim, self.inner_dim, bias=True, gather_output=True
)
self.add_k_proj = ReplicatedLinear(
added_kv_proj_dim, self.inner_dim, bias=True
self.add_k_proj = ColumnParallelLinear(
added_kv_proj_dim, self.inner_dim, bias=True, gather_output=True
)
self.add_v_proj = ReplicatedLinear(
added_kv_proj_dim, self.inner_dim, bias=True
self.add_v_proj = ColumnParallelLinear(
added_kv_proj_dim, self.inner_dim, bias=True, gather_output=True
)
if context_pre_only is not None and not context_pre_only:
self.to_add_out = ReplicatedLinear(self.inner_dim, self.dim, bias=out_bias)
self.to_add_out = ColumnParallelLinear(
self.inner_dim, self.dim, bias=out_bias, gather_output=True
)
else:
self.to_add_out = None
if not pre_only:
self.to_out = nn.ModuleList([])
self.to_out.append(
ReplicatedLinear(self.inner_dim, self.dim, bias=out_bias)
ColumnParallelLinear(
self.inner_dim, self.dim, bias=out_bias, gather_output=True
)
)
else:
self.to_out = None
@@ -644,8 +731,8 @@ class QwenImageTransformerBlock(nn.Module):
# Image processing modules
self.img_mod = nn.Sequential(
nn.SiLU(),
nn.Linear(
dim, 6 * dim, bias=True
ColumnParallelLinear(
dim, 6 * dim, bias=True, gather_output=True
), # For scale, shift, gate for norm1 and norm2
)
self.img_norm1 = LayerNormScaleShift(
@@ -669,8 +756,8 @@ class QwenImageTransformerBlock(nn.Module):
# Text processing modules
self.txt_mod = nn.Sequential(
nn.SiLU(),
nn.Linear(
dim, 6 * dim, bias=True
ColumnParallelLinear(
dim, 6 * dim, bias=True, gather_output=True
), # For scale, shift, gate for norm1 and norm2
)
self.txt_norm1 = LayerNormScaleShift(
@@ -790,8 +877,8 @@ class QwenImageTransformerBlock(nn.Module):
modulate_index: Optional[List[int]] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
# Get modulation parameters for both streams
img_mod_params = self.img_mod[1](temb_img_silu) # [B, 6*dim]
txt_mod_params = self.txt_mod[1](temb_txt_silu) # [B, 6*dim]
img_mod_params = self.img_mod[1](temb_img_silu)[0] # [B, 6*dim]
txt_mod_params = self.txt_mod[1](temb_txt_silu)[0] # [B, 6*dim]
# Split modulation parameters for norm1 and norm2
img_mod1, img_mod2 = img_mod_params.chunk(2, dim=-1) # Each [B, 3*dim]
@@ -837,7 +924,7 @@ class QwenImageTransformerBlock(nn.Module):
residual_x=hidden_states,
)
img_mlp_output = self.img_mlp(img_modulated2)
hidden_states = self.fuse_mul_add(img_mlp_output, img_gate2, hidden_states)
hidden_states = self.fuse_mul_add(img_mlp_output[0], img_gate2, hidden_states)
# Process text stream - norm2 + MLP
txt_shift2, txt_scale2, txt_gate2_raw = txt_mod2.chunk(3, dim=-1)
@@ -851,7 +938,7 @@ class QwenImageTransformerBlock(nn.Module):
txt_gate2 = txt_gate2_raw.unsqueeze(1)
txt_mlp_output = self.txt_mlp(txt_modulated2)
encoder_hidden_states = self.fuse_mul_add(
txt_mlp_output, txt_gate2, encoder_hidden_states
txt_mlp_output[0], txt_gate2, encoder_hidden_states
)
# Clip to prevent overflow for fp16

View File

@@ -76,8 +76,6 @@ class ImageEncodingStage(PipelineStage):
def move_to_device(self, device):
fields = [
"image_processor",
"image_encoder",
"text_encoder",
]
for field in fields:
processor = getattr(self, field, None)