[diffusion] ZImage support Tensor Parallel (#15849)

This commit is contained in:
Annis
2025-12-26 14:58:14 +08:00
committed by GitHub
parent a1e9b4edfa
commit 73c0c66ff8

View File

@@ -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