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