diff --git a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py index b767b1e49..73f3061fc 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py @@ -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