[codex] Optimize Helios fused norm modulation (#24059)
This commit is contained in:
@@ -28,7 +28,8 @@ from sglang.multimodal_gen.runtime.distributed.communication_op import (
|
|||||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import get_sp_group
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import get_sp_group
|
||||||
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 (
|
from sglang.multimodal_gen.runtime.layers.layernorm import (
|
||||||
FP32LayerNorm,
|
LayerNorm,
|
||||||
|
LayerNormScaleShift,
|
||||||
RMSNorm,
|
RMSNorm,
|
||||||
tensor_parallel_rms_norm,
|
tensor_parallel_rms_norm,
|
||||||
)
|
)
|
||||||
@@ -92,7 +93,9 @@ class HeliosOutputNorm(nn.Module):
|
|||||||
def __init__(self, dim: int, eps: float = 1e-6):
|
def __init__(self, dim: int, eps: float = 1e-6):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5)
|
self.scale_shift_table = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5)
|
||||||
self.norm = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
self.norm = LayerNormScaleShift(
|
||||||
|
dim, eps=eps, elementwise_affine=False, dtype=torch.float32
|
||||||
|
)
|
||||||
|
|
||||||
def forward(self, hidden_states, temb, original_context_length):
|
def forward(self, hidden_states, temb, original_context_length):
|
||||||
temb = temb[:, -original_context_length:, :]
|
temb = temb[:, -original_context_length:, :]
|
||||||
@@ -102,9 +105,7 @@ class HeliosOutputNorm(nn.Module):
|
|||||||
shift = shift.squeeze(2).to(hidden_states.device)
|
shift = shift.squeeze(2).to(hidden_states.device)
|
||||||
scale = scale.squeeze(2).to(hidden_states.device)
|
scale = scale.squeeze(2).to(hidden_states.device)
|
||||||
hidden_states = hidden_states[:, -original_context_length:, :]
|
hidden_states = hidden_states[:, -original_context_length:, :]
|
||||||
hidden_states = (
|
hidden_states = self.norm(hidden_states, shift, scale)
|
||||||
self.norm(hidden_states.float()) * (1 + scale) + shift
|
|
||||||
).type_as(hidden_states)
|
|
||||||
return hidden_states
|
return hidden_states
|
||||||
|
|
||||||
|
|
||||||
@@ -418,7 +419,9 @@ class HeliosTransformerBlock(nn.Module):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
# 1. Self-attention
|
# 1. Self-attention
|
||||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
self.norm1 = LayerNormScaleShift(
|
||||||
|
dim, eps=eps, elementwise_affine=False, dtype=torch.float32
|
||||||
|
)
|
||||||
self.attn1 = HeliosSelfAttention(
|
self.attn1 = HeliosSelfAttention(
|
||||||
dim=dim,
|
dim=dim,
|
||||||
num_heads=num_heads,
|
num_heads=num_heads,
|
||||||
@@ -436,7 +439,7 @@ class HeliosTransformerBlock(nn.Module):
|
|||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
)
|
)
|
||||||
self.self_attn_residual_norm = (
|
self.self_attn_residual_norm = (
|
||||||
FP32LayerNorm(dim, eps, elementwise_affine=True)
|
LayerNorm(dim, eps=eps, elementwise_affine=True, dtype=torch.float32)
|
||||||
if cross_attn_norm
|
if cross_attn_norm
|
||||||
else nn.Identity()
|
else nn.Identity()
|
||||||
)
|
)
|
||||||
@@ -445,7 +448,9 @@ class HeliosTransformerBlock(nn.Module):
|
|||||||
self.ffn = MLP(
|
self.ffn = MLP(
|
||||||
dim, ffn_dim, act_type="gelu_pytorch_tanh", quant_config=quant_config
|
dim, ffn_dim, act_type="gelu_pytorch_tanh", quant_config=quant_config
|
||||||
)
|
)
|
||||||
self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
self.norm3 = LayerNormScaleShift(
|
||||||
|
dim, eps=eps, elementwise_affine=False, dtype=torch.float32
|
||||||
|
)
|
||||||
|
|
||||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||||
|
|
||||||
@@ -476,9 +481,7 @@ class HeliosTransformerBlock(nn.Module):
|
|||||||
).chunk(6, dim=1)
|
).chunk(6, dim=1)
|
||||||
|
|
||||||
# 1. Self-attention
|
# 1. Self-attention
|
||||||
norm_hidden_states = (
|
norm_hidden_states = self.norm1(hidden_states, shift_msa, scale_msa)
|
||||||
self.norm1(hidden_states.float()) * (1 + scale_msa) + shift_msa
|
|
||||||
).type_as(hidden_states)
|
|
||||||
attn_output = self.attn1(
|
attn_output = self.attn1(
|
||||||
norm_hidden_states, rotary_emb, original_context_length
|
norm_hidden_states, rotary_emb, original_context_length
|
||||||
)
|
)
|
||||||
@@ -508,9 +511,7 @@ class HeliosTransformerBlock(nn.Module):
|
|||||||
hidden_states = hidden_states + attn_output
|
hidden_states = hidden_states + attn_output
|
||||||
|
|
||||||
# 3. Feed-forward
|
# 3. Feed-forward
|
||||||
norm_hidden_states = (
|
norm_hidden_states = self.norm3(hidden_states, c_shift_msa, c_scale_msa)
|
||||||
self.norm3(hidden_states.float()) * (1 + c_scale_msa) + c_shift_msa
|
|
||||||
).type_as(hidden_states)
|
|
||||||
ff_output = self.ffn(norm_hidden_states)
|
ff_output = self.ffn(norm_hidden_states)
|
||||||
hidden_states = (
|
hidden_states = (
|
||||||
hidden_states.float() + ff_output.float() * c_gate_msa
|
hidden_states.float() + ff_output.float() * c_gate_msa
|
||||||
|
|||||||
Reference in New Issue
Block a user