[diffusion] ZImage support Tensor Parallel (#15849)
This commit is contained in:
@@ -8,7 +8,13 @@ from sglang.multimodal_gen.configs.models.dits.zimage import ZImageDitConfig
|
||||
from sglang.multimodal_gen.runtime.layers.activation import SiluAndMul
|
||||
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
|
||||
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm
|
||||
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
||||
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||
ColumnParallelLinear,
|
||||
MergedColumnParallelLinear,
|
||||
QKVParallelLinear,
|
||||
ReplicatedLinear,
|
||||
RowParallelLinear,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import _apply_rotary_emb
|
||||
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
@@ -36,9 +42,13 @@ class TimestepEmbedder(nn.Module):
|
||||
|
||||
self.mlp = nn.ModuleList(
|
||||
[
|
||||
ReplicatedLinear(frequency_embedding_size, mid_size, bias=True),
|
||||
ColumnParallelLinear(
|
||||
frequency_embedding_size, mid_size, bias=True, gather_output=False
|
||||
),
|
||||
nn.SiLU(),
|
||||
ReplicatedLinear(mid_size, out_size, bias=True),
|
||||
RowParallelLinear(
|
||||
mid_size, out_size, bias=True, input_is_parallel=True
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -74,9 +84,11 @@ class TimestepEmbedder(nn.Module):
|
||||
class FeedForward(nn.Module):
|
||||
def __init__(self, dim: int, hidden_dim: int):
|
||||
super().__init__()
|
||||
# Use ReplicatedLinear for gate and up projection (fused)
|
||||
self.w13 = ReplicatedLinear(dim, hidden_dim * 2, bias=False)
|
||||
self.w2 = ReplicatedLinear(hidden_dim, dim, bias=False)
|
||||
# Use MergedColumnParallelLinear for gate and up projection (fused)
|
||||
self.w13 = MergedColumnParallelLinear(
|
||||
dim, [hidden_dim, hidden_dim], bias=False, gather_output=False
|
||||
)
|
||||
self.w2 = RowParallelLinear(hidden_dim, dim, bias=False, input_is_parallel=True)
|
||||
self.act = SiluAndMul()
|
||||
|
||||
def forward(self, x):
|
||||
@@ -97,14 +109,17 @@ class ZImageAttention(nn.Module):
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.num_kv_heads = num_kv_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.qk_norm = qk_norm
|
||||
|
||||
# Use ReplicatedLinear for QKV projection (fused)
|
||||
qkv_dim = dim + 2 * (num_kv_heads * self.head_dim)
|
||||
self.to_qkv = ReplicatedLinear(dim, qkv_dim, bias=False)
|
||||
# Use QKVParallelLinear for QKV projection (fused)
|
||||
self.to_qkv = QKVParallelLinear(
|
||||
hidden_size=dim,
|
||||
head_size=self.head_dim,
|
||||
total_num_heads=num_heads,
|
||||
total_num_kv_heads=num_kv_heads,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
if self.qk_norm:
|
||||
self.norm_q = RMSNorm(self.head_dim, eps=eps)
|
||||
@@ -113,7 +128,9 @@ class ZImageAttention(nn.Module):
|
||||
self.norm_q = None
|
||||
self.norm_k = None
|
||||
|
||||
self.to_out = nn.ModuleList([ReplicatedLinear(dim, dim, bias=False)])
|
||||
self.to_out = nn.ModuleList(
|
||||
[RowParallelLinear(dim, dim, bias=False, input_is_parallel=True)]
|
||||
)
|
||||
|
||||
self.attn = USPAttention(
|
||||
num_heads=num_heads,
|
||||
@@ -130,12 +147,13 @@ class ZImageAttention(nn.Module):
|
||||
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
):
|
||||
qkv, _ = self.to_qkv(hidden_states)
|
||||
kv_dim = self.head_dim * self.num_kv_heads
|
||||
q, k, v = torch.split(qkv, [self.dim, kv_dim, kv_dim], dim=-1)
|
||||
q_dim = self.to_qkv.num_heads * self.head_dim
|
||||
kv_dim = self.to_qkv.num_kv_heads * self.head_dim
|
||||
q, k, v = torch.split(qkv, [q_dim, kv_dim, kv_dim], dim=-1)
|
||||
|
||||
q = q.view(*q.shape[:-1], self.num_heads, self.head_dim)
|
||||
k = k.view(*k.shape[:-1], self.num_kv_heads, self.head_dim)
|
||||
v = v.view(*v.shape[:-1], self.num_kv_heads, self.head_dim)
|
||||
q = q.view(*q.shape[:-1], self.to_qkv.num_heads, self.head_dim)
|
||||
k = k.view(*k.shape[:-1], self.to_qkv.num_kv_heads, self.head_dim)
|
||||
v = v.view(*v.shape[:-1], self.to_qkv.num_kv_heads, self.head_dim)
|
||||
|
||||
if self.norm_q is not None:
|
||||
q = self.norm_q(q)
|
||||
@@ -251,7 +269,9 @@ class FinalLayer(nn.Module):
|
||||
def __init__(self, hidden_size, out_channels):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.linear = ReplicatedLinear(hidden_size, out_channels, bias=True)
|
||||
self.linear = ColumnParallelLinear(
|
||||
hidden_size, out_channels, bias=True, gather_output=True
|
||||
)
|
||||
|
||||
self.act = nn.SiLU()
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
@@ -373,10 +393,11 @@ class ZImageTransformer2DModel(CachableDiT):
|
||||
for patch_idx, (patch_size, f_patch_size) in enumerate(
|
||||
zip(self.all_patch_size, self.all_f_patch_size)
|
||||
):
|
||||
x_embedder = ReplicatedLinear(
|
||||
x_embedder = ColumnParallelLinear(
|
||||
f_patch_size * patch_size * patch_size * self.in_channels,
|
||||
self.dim,
|
||||
bias=True,
|
||||
gather_output=True,
|
||||
)
|
||||
all_x_embedder[f"{patch_size}-{f_patch_size}"] = x_embedder
|
||||
|
||||
|
||||
Reference in New Issue
Block a user