[diffusion] fix: fix pack qkv opt break tensor parallel (#15225)
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user