[diffusion] feat: shard qwen-image dit across tp ranks (#29774)

Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
Yiqi Yang
2026-07-03 13:17:53 +08:00
committed by GitHub
co-authored by Mick
parent f011d8c2e5
commit e878c6ebdd
@@ -15,7 +15,10 @@ 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.distributed import (
get_local_torch_device,
get_tp_world_size,
)
from sglang.multimodal_gen.runtime.distributed.parallel_state import ( from sglang.multimodal_gen.runtime.distributed.parallel_state import (
get_ring_parallel_world_size, get_ring_parallel_world_size,
get_sp_parallel_rank, get_sp_parallel_rank,
@@ -37,8 +40,10 @@ from sglang.multimodal_gen.runtime.layers.layernorm import (
apply_qk_norm_with_optional_rope, apply_qk_norm_with_optional_rope,
) )
from sglang.multimodal_gen.runtime.layers.linear import ( from sglang.multimodal_gen.runtime.layers.linear import (
ColumnParallelLinear,
MergedColumnParallelLinear, MergedColumnParallelLinear,
ReplicatedLinear, ReplicatedLinear,
RowParallelLinear,
) )
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import ( from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
QuantizationConfig, QuantizationConfig,
@@ -617,6 +622,12 @@ class QwenImageCrossAttention(nn.Module):
self.inner_dim = out_dim if out_dim is not None else head_dim * num_heads self.inner_dim = out_dim if out_dim is not None else head_dim * num_heads
self.inner_kv_dim = self.inner_dim self.inner_kv_dim = self.inner_dim
tp_size = get_tp_world_size()
assert (
self.num_heads % tp_size == 0
), f"num_heads ({self.num_heads}) must be divisible by tp_size ({tp_size})"
self.local_num_heads = self.num_heads // tp_size
if self.use_fused_qkv: if self.use_fused_qkv:
# Use fused QKV projection for nunchaku quantization # Use fused QKV projection for nunchaku quantization
self.to_qkv = MergedColumnParallelLinear( self.to_qkv = MergedColumnParallelLinear(
@@ -627,24 +638,27 @@ class QwenImageCrossAttention(nn.Module):
prefix=f"{prefix}.to_qkv", prefix=f"{prefix}.to_qkv",
) )
else: else:
self.to_q = ReplicatedLinear( self.to_q = ColumnParallelLinear(
dim, dim,
self.inner_dim, self.inner_dim,
bias=True, bias=True,
gather_output=False,
quant_config=quant_config, quant_config=quant_config,
prefix=f"{prefix}.to_q", prefix=f"{prefix}.to_q",
) )
self.to_k = ReplicatedLinear( self.to_k = ColumnParallelLinear(
dim, dim,
self.inner_dim, self.inner_dim,
bias=True, bias=True,
gather_output=False,
quant_config=quant_config, quant_config=quant_config,
prefix=f"{prefix}.to_k", prefix=f"{prefix}.to_k",
) )
self.to_v = ReplicatedLinear( self.to_v = ColumnParallelLinear(
dim, dim,
self.inner_dim, self.inner_dim,
bias=True, bias=True,
gather_output=False,
quant_config=quant_config, quant_config=quant_config,
prefix=f"{prefix}.to_v", prefix=f"{prefix}.to_v",
) )
@@ -664,33 +678,37 @@ class QwenImageCrossAttention(nn.Module):
prefix=f"{prefix}.to_added_qkv", prefix=f"{prefix}.to_added_qkv",
) )
else: else:
self.add_q_proj = ReplicatedLinear( self.add_q_proj = ColumnParallelLinear(
added_kv_proj_dim, added_kv_proj_dim,
self.inner_dim, self.inner_dim,
bias=True, bias=True,
gather_output=False,
quant_config=quant_config, quant_config=quant_config,
prefix=f"{prefix}.add_q_proj", prefix=f"{prefix}.add_q_proj",
) )
self.add_k_proj = ReplicatedLinear( self.add_k_proj = ColumnParallelLinear(
added_kv_proj_dim, added_kv_proj_dim,
self.inner_dim, self.inner_dim,
bias=True, bias=True,
gather_output=False,
quant_config=quant_config, quant_config=quant_config,
prefix=f"{prefix}.add_k_proj", prefix=f"{prefix}.add_k_proj",
) )
self.add_v_proj = ReplicatedLinear( self.add_v_proj = ColumnParallelLinear(
added_kv_proj_dim, added_kv_proj_dim,
self.inner_dim, self.inner_dim,
bias=True, bias=True,
gather_output=False,
quant_config=quant_config, quant_config=quant_config,
prefix=f"{prefix}.add_v_proj", prefix=f"{prefix}.add_v_proj",
) )
if context_pre_only is not None and not context_pre_only: if context_pre_only is not None and not context_pre_only:
self.to_add_out = ReplicatedLinear( self.to_add_out = RowParallelLinear(
self.inner_dim, self.inner_dim,
self.dim, self.dim,
bias=out_bias, bias=out_bias,
input_is_parallel=True,
quant_config=quant_config, quant_config=quant_config,
prefix=f"{prefix}.to_add_out", prefix=f"{prefix}.to_add_out",
) )
@@ -700,10 +718,11 @@ class QwenImageCrossAttention(nn.Module):
if not pre_only: if not pre_only:
self.to_out = nn.ModuleList([]) self.to_out = nn.ModuleList([])
self.to_out.append( self.to_out.append(
ReplicatedLinear( RowParallelLinear(
self.inner_dim, self.inner_dim,
self.dim, self.dim,
bias=out_bias, bias=out_bias,
input_is_parallel=True,
quant_config=quant_config, quant_config=quant_config,
prefix=f"{prefix}.to_out.0", prefix=f"{prefix}.to_out.0",
) )
@@ -716,7 +735,7 @@ class QwenImageCrossAttention(nn.Module):
# Scaled dot product attention # Scaled dot product attention
self.attn = USPAttention( self.attn = USPAttention(
num_heads=num_heads, num_heads=self.local_num_heads,
head_size=self.head_dim, head_size=self.head_dim,
dropout_rate=0, dropout_rate=0,
softmax_scale=None, softmax_scale=None,
@@ -768,13 +787,13 @@ class QwenImageCrossAttention(nn.Module):
) = _get_qkv_projections(self, hidden_states, encoder_hidden_states) ) = _get_qkv_projections(self, hidden_states, encoder_hidden_states)
# Reshape for multi-head attention # Reshape for multi-head attention
img_query = img_query.unflatten(-1, (self.num_heads, -1)) img_query = img_query.unflatten(-1, (self.local_num_heads, self.head_dim))
img_key = img_key.unflatten(-1, (self.num_heads, -1)) img_key = img_key.unflatten(-1, (self.local_num_heads, self.head_dim))
img_value = img_value.unflatten(-1, (self.num_heads, -1)) img_value = img_value.unflatten(-1, (self.local_num_heads, self.head_dim))
txt_query = txt_query.unflatten(-1, (self.num_heads, -1)) txt_query = txt_query.unflatten(-1, (self.local_num_heads, self.head_dim))
txt_key = txt_key.unflatten(-1, (self.num_heads, -1)) txt_key = txt_key.unflatten(-1, (self.local_num_heads, self.head_dim))
txt_value = txt_value.unflatten(-1, (self.num_heads, -1)) txt_value = txt_value.unflatten(-1, (self.local_num_heads, self.head_dim))
img_cache = txt_cache = None img_cache = txt_cache = None
if image_rotary_emb is not None: if image_rotary_emb is not None:
@@ -866,8 +885,10 @@ class QwenImageGELU(nn.Module):
inner_dim: int, inner_dim: int,
quant_config: Optional[QuantizationConfig] = None, quant_config: Optional[QuantizationConfig] = None,
prefix: str = "", prefix: str = "",
replicated: bool = False,
) -> None: ) -> None:
super().__init__() super().__init__()
if replicated:
self.proj = ReplicatedLinear( self.proj = ReplicatedLinear(
dim, dim,
inner_dim, inner_dim,
@@ -875,6 +896,15 @@ class QwenImageGELU(nn.Module):
quant_config=quant_config, quant_config=quant_config,
prefix=f"{prefix}.proj", prefix=f"{prefix}.proj",
) )
else:
self.proj = ColumnParallelLinear(
dim,
inner_dim,
bias=True,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.proj",
)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states, _ = self.proj(hidden_states) hidden_states, _ = self.proj(hidden_states)
@@ -889,9 +919,30 @@ class QwenImageFeedForward(nn.Module):
quant_config: Optional[QuantizationConfig] = None, quant_config: Optional[QuantizationConfig] = None,
prefix: str = "", prefix: str = "",
mult: int = 4, mult: int = 4,
replicated: bool = False,
) -> None: ) -> None:
super().__init__() super().__init__()
inner_dim = dim * mult inner_dim = dim * mult
if replicated:
# Keep the whole FFN resident on every rank: no per-block
# all-reduce. Only worth it when the branch's token count is small
# enough that the duplicated GEMM is cheaper than the all-reduce.
down = ReplicatedLinear(
inner_dim,
dim_out,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.net.2",
)
else:
down = RowParallelLinear(
inner_dim,
dim_out,
bias=True,
input_is_parallel=True,
quant_config=quant_config,
prefix=f"{prefix}.net.2",
)
self.net = nn.ModuleList( self.net = nn.ModuleList(
[ [
QwenImageGELU( QwenImageGELU(
@@ -899,15 +950,10 @@ class QwenImageFeedForward(nn.Module):
inner_dim, inner_dim,
quant_config=quant_config, quant_config=quant_config,
prefix=f"{prefix}.net.0", prefix=f"{prefix}.net.0",
replicated=replicated,
), ),
nn.Dropout(0.0), nn.Dropout(0.0),
ReplicatedLinear( down,
inner_dim,
dim_out,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.net.2",
),
] ]
) )
@@ -1022,6 +1068,13 @@ class QwenImageTransformerBlock(nn.Module):
dim_out=dim, dim_out=dim,
quant_config=quant_config, quant_config=quant_config,
prefix=f"{prefix}.txt_mlp", prefix=f"{prefix}.txt_mlp",
# The text branch is ~1K tokens regardless of image size, so
# sharding its FFN saves less GEMM time than the per-block
# all-reduce it adds. Measured e2e crossover (H100, 1024x1024):
# tp=2 replication wins (5.58s -> 5.12s, -8%) but at tp=4 the
# duplicated GEMM outgrows the all-reduce saved (~1% loss), so
# gate on the TP degree.
replicated=get_tp_world_size() <= 2,
) )
if nunchaku_enabled: if nunchaku_enabled:
@@ -1308,17 +1361,19 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
self.txt_norm = RMSNorm(joint_attention_dim, eps=1e-6) self.txt_norm = RMSNorm(joint_attention_dim, eps=1e-6)
self.img_in = ReplicatedLinear( self.img_in = ColumnParallelLinear(
in_channels, in_channels,
self.inner_dim, self.inner_dim,
bias=True, bias=True,
gather_output=True,
quant_config=quant_config, quant_config=quant_config,
prefix="img_in", prefix="img_in",
) )
self.txt_in = ReplicatedLinear( self.txt_in = ColumnParallelLinear(
joint_attention_dim, joint_attention_dim,
self.inner_dim, self.inner_dim,
bias=True, bias=True,
gather_output=True,
quant_config=quant_config, quant_config=quant_config,
prefix="txt_in", prefix="txt_in",
) )
@@ -1340,10 +1395,11 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
self.norm_out = AdaLayerNormContinuous( self.norm_out = AdaLayerNormContinuous(
self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6 self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6
) )
self.proj_out = ReplicatedLinear( self.proj_out = ColumnParallelLinear(
self.inner_dim, self.inner_dim,
patch_size * patch_size * self.out_channels, patch_size * patch_size * self.out_channels,
bias=True, bias=True,
gather_output=True,
quant_config=quant_config, quant_config=quant_config,
prefix="proj_out", prefix="proj_out",
) )