[diffusion] model: move tp_rmsnorm check to WanTransformerBlock (#17792)
This commit is contained in:
@@ -10,10 +10,10 @@ import torch.nn as nn
|
||||
|
||||
from sglang.multimodal_gen.configs.models.dits import WanVideoConfig
|
||||
from sglang.multimodal_gen.configs.sample.wan import WanTeaCacheParams
|
||||
from sglang.multimodal_gen.runtime.distributed import divide
|
||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||
from sglang.multimodal_gen.runtime.distributed import (
|
||||
divide,
|
||||
get_sp_world_size,
|
||||
get_tensor_model_parallel_world_size,
|
||||
get_tp_world_size,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.attention import (
|
||||
MinimalA2AAttnOp,
|
||||
@@ -138,7 +138,7 @@ class WanSelfAttention(nn.Module):
|
||||
self.qk_norm = qk_norm
|
||||
self.eps = eps
|
||||
self.parallel_attention = parallel_attention
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
tp_size = get_tp_world_size()
|
||||
|
||||
# layers
|
||||
self.to_q = ColumnParallelLinear(dim, dim, gather_output=False)
|
||||
@@ -147,11 +147,12 @@ class WanSelfAttention(nn.Module):
|
||||
self.to_out = RowParallelLinear(dim, dim, input_is_parallel=True)
|
||||
self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.tp_rmsnorm = self.tp_size > 1 and qk_norm
|
||||
self.tp_rmsnorm = tp_size > 1 and qk_norm
|
||||
self.local_num_heads = divide(num_heads, tp_size)
|
||||
|
||||
# Scaled dot product attention
|
||||
self.attn = USPAttention(
|
||||
num_heads=num_heads // self.tp_size,
|
||||
num_heads=self.local_num_heads,
|
||||
head_size=self.head_dim,
|
||||
dropout_rate=0,
|
||||
softmax_scale=None,
|
||||
@@ -179,25 +180,22 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
context(Tensor): Shape [B, L2, C]
|
||||
context_lens(Tensor): Shape [B]
|
||||
"""
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
num_heads_per_rank = n // self.tp_size
|
||||
|
||||
q, _ = self.to_q(x)
|
||||
if self.tp_rmsnorm:
|
||||
q = tensor_parallel_rms_norm(q, self.norm_q)
|
||||
else:
|
||||
q = self.norm_q(q)
|
||||
q = q.view(b, -1, num_heads_per_rank, d)
|
||||
q = q.unflatten(2, (self.local_num_heads, self.head_dim))
|
||||
|
||||
k, _ = self.to_k(context)
|
||||
if self.tp_rmsnorm:
|
||||
k = tensor_parallel_rms_norm(k, self.norm_k)
|
||||
else:
|
||||
k = self.norm_k(k)
|
||||
k = k.view(b, -1, num_heads_per_rank, d)
|
||||
k = k.unflatten(2, (self.local_num_heads, self.head_dim))
|
||||
|
||||
v, _ = self.to_v(context)
|
||||
v = v.view(b, -1, num_heads_per_rank, d)
|
||||
v = v.unflatten(2, (self.local_num_heads, self.head_dim))
|
||||
|
||||
# compute attention
|
||||
x = self.attn(q, k, v)
|
||||
@@ -242,35 +240,33 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
"""
|
||||
context_img = context[:, :257]
|
||||
context = context[:, 257:]
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
num_heads_per_rank = n // self.tp_size
|
||||
|
||||
q, _ = self.to_q(x)
|
||||
if self.tp_rmsnorm:
|
||||
q = tensor_parallel_rms_norm(q, self.norm_q)
|
||||
else:
|
||||
q = self.norm_q(q)
|
||||
q = q.view(b, -1, num_heads_per_rank, d)
|
||||
q = q.unflatten(2, (self.local_num_heads, self.head_dim))
|
||||
|
||||
k, _ = self.to_k(context)
|
||||
if self.tp_rmsnorm:
|
||||
k = tensor_parallel_rms_norm(k, self.norm_k)
|
||||
else:
|
||||
k = self.norm_k(k)
|
||||
k = k.view(b, -1, num_heads_per_rank, d)
|
||||
k = k.unflatten(2, (self.local_num_heads, self.head_dim))
|
||||
|
||||
v, _ = self.to_v(context)
|
||||
v = v.view(b, -1, num_heads_per_rank, d)
|
||||
v = v.unflatten(2, (self.local_num_heads, self.head_dim))
|
||||
|
||||
k_img, _ = self.add_k_proj(context_img)
|
||||
if self.tp_rmsnorm:
|
||||
k_img = tensor_parallel_rms_norm(k_img, self.norm_added_k)
|
||||
else:
|
||||
k_img = self.norm_added_k(k_img)
|
||||
k_img = k_img.view(b, -1, num_heads_per_rank, d)
|
||||
k_img = k_img.unflatten(2, (self.local_num_heads, self.head_dim))
|
||||
|
||||
v_img, _ = self.add_v_proj(context_img)
|
||||
v_img = v_img.view(b, -1, num_heads_per_rank, d)
|
||||
v_img = v_img.unflatten(2, (self.local_num_heads, self.head_dim))
|
||||
|
||||
img_x = self.attn(q, k_img, v_img)
|
||||
x = self.attn(q, k, v)
|
||||
@@ -308,9 +304,11 @@ class WanTransformerBlock(nn.Module):
|
||||
self.to_v = ColumnParallelLinear(dim, dim, bias=True, gather_output=False)
|
||||
|
||||
self.to_out = RowParallelLinear(dim, dim, bias=True, reduce_results=True)
|
||||
tp_size = get_tp_world_size()
|
||||
self.local_num_heads = divide(num_heads, tp_size)
|
||||
if attention_type in ("sla", "sagesla"):
|
||||
self.attn1 = MinimalA2AAttnOp(
|
||||
num_heads=divide(num_heads, get_tensor_model_parallel_world_size()),
|
||||
num_heads=self.local_num_heads,
|
||||
head_size=dim // num_heads,
|
||||
attention_type=attention_type,
|
||||
topk=sla_topk,
|
||||
@@ -321,7 +319,7 @@ class WanTransformerBlock(nn.Module):
|
||||
)
|
||||
else:
|
||||
self.attn1 = USPAttention(
|
||||
num_heads=divide(num_heads, get_tensor_model_parallel_world_size()),
|
||||
num_heads=self.local_num_heads,
|
||||
head_size=dim // num_heads,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
@@ -343,6 +341,7 @@ class WanTransformerBlock(nn.Module):
|
||||
raise Exception
|
||||
assert cross_attn_norm is True
|
||||
self.qk_norm = qk_norm
|
||||
self.tp_rmsnorm = qk_norm == "rms_norm_across_heads" and tp_size > 1
|
||||
self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift(
|
||||
dim,
|
||||
norm_type="layer",
|
||||
@@ -424,24 +423,20 @@ class WanTransformerBlock(nn.Module):
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
tp_rmsnorm = (
|
||||
self.qk_norm == "rms_norm_across_heads"
|
||||
and get_tensor_model_parallel_world_size() > 1
|
||||
)
|
||||
|
||||
if self.norm_q is not None:
|
||||
if tp_rmsnorm:
|
||||
if self.tp_rmsnorm:
|
||||
query = tensor_parallel_rms_norm(query, self.norm_q)
|
||||
else:
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
if tp_rmsnorm:
|
||||
if self.tp_rmsnorm:
|
||||
key = tensor_parallel_rms_norm(key, self.norm_k)
|
||||
else:
|
||||
key = self.norm_k(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (-1, self.dim_head))
|
||||
key = key.squeeze(1).unflatten(2, (-1, self.dim_head))
|
||||
value = value.squeeze(1).unflatten(2, (-1, self.dim_head))
|
||||
query = query.squeeze(1).unflatten(2, (self.local_num_heads, self.dim_head))
|
||||
key = key.squeeze(1).unflatten(2, (self.local_num_heads, self.dim_head))
|
||||
value = value.squeeze(1).unflatten(2, (self.local_num_heads, self.dim_head))
|
||||
|
||||
# Apply rotary embeddings
|
||||
cos, sin = freqs_cis
|
||||
|
||||
Reference in New Issue
Block a user