[diffusion] model: move tp_rmsnorm check to WanTransformerBlock (#17792)

This commit is contained in:
triple-mu
2026-01-29 16:39:00 +08:00
committed by GitHub
parent cdedbf1486
commit 319f6886fe
@@ -10,10 +10,10 @@ import torch.nn as nn
from sglang.multimodal_gen.configs.models.dits import WanVideoConfig from sglang.multimodal_gen.configs.models.dits import WanVideoConfig
from sglang.multimodal_gen.configs.sample.wan import WanTeaCacheParams from sglang.multimodal_gen.configs.sample.wan import WanTeaCacheParams
from sglang.multimodal_gen.runtime.distributed import divide from sglang.multimodal_gen.runtime.distributed import (
from sglang.multimodal_gen.runtime.distributed.parallel_state import ( divide,
get_sp_world_size, get_sp_world_size,
get_tensor_model_parallel_world_size, get_tp_world_size,
) )
from sglang.multimodal_gen.runtime.layers.attention import ( from sglang.multimodal_gen.runtime.layers.attention import (
MinimalA2AAttnOp, MinimalA2AAttnOp,
@@ -138,7 +138,7 @@ class WanSelfAttention(nn.Module):
self.qk_norm = qk_norm self.qk_norm = qk_norm
self.eps = eps self.eps = eps
self.parallel_attention = parallel_attention self.parallel_attention = parallel_attention
self.tp_size = get_tensor_model_parallel_world_size() tp_size = get_tp_world_size()
# layers # layers
self.to_q = ColumnParallelLinear(dim, dim, gather_output=False) 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.to_out = RowParallelLinear(dim, dim, input_is_parallel=True)
self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity() 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.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 # Scaled dot product attention
self.attn = USPAttention( self.attn = USPAttention(
num_heads=num_heads // self.tp_size, num_heads=self.local_num_heads,
head_size=self.head_dim, head_size=self.head_dim,
dropout_rate=0, dropout_rate=0,
softmax_scale=None, softmax_scale=None,
@@ -179,25 +180,22 @@ class WanT2VCrossAttention(WanSelfAttention):
context(Tensor): Shape [B, L2, C] context(Tensor): Shape [B, L2, C]
context_lens(Tensor): Shape [B] 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) q, _ = self.to_q(x)
if self.tp_rmsnorm: if self.tp_rmsnorm:
q = tensor_parallel_rms_norm(q, self.norm_q) q = tensor_parallel_rms_norm(q, self.norm_q)
else: else:
q = self.norm_q(q) 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) k, _ = self.to_k(context)
if self.tp_rmsnorm: if self.tp_rmsnorm:
k = tensor_parallel_rms_norm(k, self.norm_k) k = tensor_parallel_rms_norm(k, self.norm_k)
else: else:
k = self.norm_k(k) 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, _ = 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 # compute attention
x = self.attn(q, k, v) x = self.attn(q, k, v)
@@ -242,35 +240,33 @@ class WanI2VCrossAttention(WanSelfAttention):
""" """
context_img = context[:, :257] context_img = context[:, :257]
context = 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) q, _ = self.to_q(x)
if self.tp_rmsnorm: if self.tp_rmsnorm:
q = tensor_parallel_rms_norm(q, self.norm_q) q = tensor_parallel_rms_norm(q, self.norm_q)
else: else:
q = self.norm_q(q) 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) k, _ = self.to_k(context)
if self.tp_rmsnorm: if self.tp_rmsnorm:
k = tensor_parallel_rms_norm(k, self.norm_k) k = tensor_parallel_rms_norm(k, self.norm_k)
else: else:
k = self.norm_k(k) 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, _ = 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) k_img, _ = self.add_k_proj(context_img)
if self.tp_rmsnorm: if self.tp_rmsnorm:
k_img = tensor_parallel_rms_norm(k_img, self.norm_added_k) k_img = tensor_parallel_rms_norm(k_img, self.norm_added_k)
else: else:
k_img = self.norm_added_k(k_img) 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, _ = 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) img_x = self.attn(q, k_img, v_img)
x = self.attn(q, k, v) 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_v = ColumnParallelLinear(dim, dim, bias=True, gather_output=False)
self.to_out = RowParallelLinear(dim, dim, bias=True, reduce_results=True) 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"): if attention_type in ("sla", "sagesla"):
self.attn1 = MinimalA2AAttnOp( 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, head_size=dim // num_heads,
attention_type=attention_type, attention_type=attention_type,
topk=sla_topk, topk=sla_topk,
@@ -321,7 +319,7 @@ class WanTransformerBlock(nn.Module):
) )
else: else:
self.attn1 = USPAttention( 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, head_size=dim // num_heads,
causal=False, causal=False,
supported_attention_backends=supported_attention_backends, supported_attention_backends=supported_attention_backends,
@@ -343,6 +341,7 @@ class WanTransformerBlock(nn.Module):
raise Exception raise Exception
assert cross_attn_norm is True assert cross_attn_norm is True
self.qk_norm = qk_norm self.qk_norm = qk_norm
self.tp_rmsnorm = qk_norm == "rms_norm_across_heads" and tp_size > 1
self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift( self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim, dim,
norm_type="layer", norm_type="layer",
@@ -424,24 +423,20 @@ class WanTransformerBlock(nn.Module):
query, _ = self.to_q(norm_hidden_states) query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states) key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(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 self.norm_q is not None:
if tp_rmsnorm: if self.tp_rmsnorm:
query = tensor_parallel_rms_norm(query, self.norm_q) query = tensor_parallel_rms_norm(query, self.norm_q)
else: else:
query = self.norm_q(query) query = self.norm_q(query)
if self.norm_k is not None: if self.norm_k is not None:
if tp_rmsnorm: if self.tp_rmsnorm:
key = tensor_parallel_rms_norm(key, self.norm_k) key = tensor_parallel_rms_norm(key, self.norm_k)
else: else:
key = self.norm_k(key) key = self.norm_k(key)
query = query.squeeze(1).unflatten(2, (self.local_num_heads, self.dim_head))
query = query.squeeze(1).unflatten(2, (-1, self.dim_head)) key = key.squeeze(1).unflatten(2, (self.local_num_heads, self.dim_head))
key = key.squeeze(1).unflatten(2, (-1, self.dim_head)) value = value.squeeze(1).unflatten(2, (self.local_num_heads, self.dim_head))
value = value.squeeze(1).unflatten(2, (-1, self.dim_head))
# Apply rotary embeddings # Apply rotary embeddings
cos, sin = freqs_cis cos, sin = freqs_cis