[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 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 (
get_ring_parallel_world_size,
get_sp_parallel_rank,
@@ -37,8 +40,10 @@ from sglang.multimodal_gen.runtime.layers.layernorm import (
apply_qk_norm_with_optional_rope,
)
from sglang.multimodal_gen.runtime.layers.linear import (
ColumnParallelLinear,
MergedColumnParallelLinear,
ReplicatedLinear,
RowParallelLinear,
)
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
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_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:
# Use fused QKV projection for nunchaku quantization
self.to_qkv = MergedColumnParallelLinear(
@@ -627,24 +638,27 @@ class QwenImageCrossAttention(nn.Module):
prefix=f"{prefix}.to_qkv",
)
else:
self.to_q = ReplicatedLinear(
self.to_q = ColumnParallelLinear(
dim,
self.inner_dim,
bias=True,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.to_q",
)
self.to_k = ReplicatedLinear(
self.to_k = ColumnParallelLinear(
dim,
self.inner_dim,
bias=True,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.to_k",
)
self.to_v = ReplicatedLinear(
self.to_v = ColumnParallelLinear(
dim,
self.inner_dim,
bias=True,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.to_v",
)
@@ -664,33 +678,37 @@ class QwenImageCrossAttention(nn.Module):
prefix=f"{prefix}.to_added_qkv",
)
else:
self.add_q_proj = ReplicatedLinear(
self.add_q_proj = ColumnParallelLinear(
added_kv_proj_dim,
self.inner_dim,
bias=True,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.add_q_proj",
)
self.add_k_proj = ReplicatedLinear(
self.add_k_proj = ColumnParallelLinear(
added_kv_proj_dim,
self.inner_dim,
bias=True,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.add_k_proj",
)
self.add_v_proj = ReplicatedLinear(
self.add_v_proj = ColumnParallelLinear(
added_kv_proj_dim,
self.inner_dim,
bias=True,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.add_v_proj",
)
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.dim,
bias=out_bias,
input_is_parallel=True,
quant_config=quant_config,
prefix=f"{prefix}.to_add_out",
)
@@ -700,10 +718,11 @@ class QwenImageCrossAttention(nn.Module):
if not pre_only:
self.to_out = nn.ModuleList([])
self.to_out.append(
ReplicatedLinear(
RowParallelLinear(
self.inner_dim,
self.dim,
bias=out_bias,
input_is_parallel=True,
quant_config=quant_config,
prefix=f"{prefix}.to_out.0",
)
@@ -716,7 +735,7 @@ class QwenImageCrossAttention(nn.Module):
# Scaled dot product attention
self.attn = USPAttention(
num_heads=num_heads,
num_heads=self.local_num_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
@@ -768,13 +787,13 @@ class QwenImageCrossAttention(nn.Module):
) = _get_qkv_projections(self, hidden_states, encoder_hidden_states)
# Reshape for multi-head attention
img_query = img_query.unflatten(-1, (self.num_heads, -1))
img_key = img_key.unflatten(-1, (self.num_heads, -1))
img_value = img_value.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.local_num_heads, self.head_dim))
img_value = img_value.unflatten(-1, (self.local_num_heads, self.head_dim))
txt_query = txt_query.unflatten(-1, (self.num_heads, -1))
txt_key = txt_key.unflatten(-1, (self.num_heads, -1))
txt_value = txt_value.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.local_num_heads, self.head_dim))
txt_value = txt_value.unflatten(-1, (self.local_num_heads, self.head_dim))
img_cache = txt_cache = None
if image_rotary_emb is not None:
@@ -866,15 +885,26 @@ class QwenImageGELU(nn.Module):
inner_dim: int,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
replicated: bool = False,
) -> None:
super().__init__()
self.proj = ReplicatedLinear(
dim,
inner_dim,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.proj",
)
if replicated:
self.proj = ReplicatedLinear(
dim,
inner_dim,
bias=True,
quant_config=quant_config,
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:
hidden_states, _ = self.proj(hidden_states)
@@ -889,9 +919,30 @@ class QwenImageFeedForward(nn.Module):
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
mult: int = 4,
replicated: bool = False,
) -> None:
super().__init__()
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(
[
QwenImageGELU(
@@ -899,15 +950,10 @@ class QwenImageFeedForward(nn.Module):
inner_dim,
quant_config=quant_config,
prefix=f"{prefix}.net.0",
replicated=replicated,
),
nn.Dropout(0.0),
ReplicatedLinear(
inner_dim,
dim_out,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.net.2",
),
down,
]
)
@@ -1022,6 +1068,13 @@ class QwenImageTransformerBlock(nn.Module):
dim_out=dim,
quant_config=quant_config,
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:
@@ -1308,17 +1361,19 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
self.txt_norm = RMSNorm(joint_attention_dim, eps=1e-6)
self.img_in = ReplicatedLinear(
self.img_in = ColumnParallelLinear(
in_channels,
self.inner_dim,
bias=True,
gather_output=True,
quant_config=quant_config,
prefix="img_in",
)
self.txt_in = ReplicatedLinear(
self.txt_in = ColumnParallelLinear(
joint_attention_dim,
self.inner_dim,
bias=True,
gather_output=True,
quant_config=quant_config,
prefix="txt_in",
)
@@ -1340,10 +1395,11 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
self.norm_out = AdaLayerNormContinuous(
self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6
)
self.proj_out = ReplicatedLinear(
self.proj_out = ColumnParallelLinear(
self.inner_dim,
patch_size * patch_size * self.out_channels,
bias=True,
gather_output=True,
quant_config=quant_config,
prefix="proj_out",
)