[diffusion] fix: fix pack qkv opt break tensor parallel (#15225)

This commit is contained in:
Xiaoyu Zhang
2025-12-16 14:33:49 +08:00
committed by GitHub
parent c843419562
commit 6292d97135
4 changed files with 25 additions and 72 deletions
@@ -36,10 +36,7 @@ from sglang.multimodal_gen.runtime.layers.attention import USPAttention
# from sglang.multimodal_gen.runtime.layers.layernorm import LayerNorm as LayerNorm
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm
from sglang.multimodal_gen.runtime.layers.linear import (
QKVParallelLinear,
ReplicatedLinear,
)
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
from sglang.multimodal_gen.runtime.layers.mlp import MLP
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
NDRotaryEmbedding,
@@ -102,13 +99,8 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
self.norm_q = RMSNorm(dim_head, eps=eps)
self.norm_k = RMSNorm(dim_head, eps=eps)
# Use QKVParallelLinear for fused QKV projections
self.to_qkv = QKVParallelLinear(
hidden_size=query_dim,
head_size=dim_head,
total_num_heads=num_heads,
bias=bias,
)
# Use ReplicatedLinear for fused QKV projections
self.to_qkv = ReplicatedLinear(query_dim, self.inner_dim * 3, bias=bias)
if not self.pre_only:
self.to_out = torch.nn.ModuleList([])
@@ -121,12 +113,9 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
if added_kv_proj_dim is not None:
self.norm_added_q = RMSNorm(dim_head, eps=eps)
self.norm_added_k = RMSNorm(dim_head, eps=eps)
# Use QKVParallelLinear for added (encoder) QKV projections
self.to_added_qkv = QKVParallelLinear(
hidden_size=added_kv_proj_dim,
head_size=dim_head,
total_num_heads=num_heads,
bias=added_proj_bias,
# Use ReplicatedLinear for added (encoder) QKV projections
self.to_added_qkv = ReplicatedLinear(
added_kv_proj_dim, self.inner_dim * 3, bias=added_proj_bias
)
self.to_add_out = ReplicatedLinear(self.inner_dim, query_dim, bias=out_bias)