[diffusion] feat: shard qwen-image dit across tp ranks (#29774)
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
@@ -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",
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user