[Perf] Precompute gemma_weight to avoid redundant add on every forward (#22673)
This commit is contained in:
@@ -476,6 +476,16 @@ class GemmaRMSNorm(MultiPlatformOp):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.weight = nn.Parameter(torch.zeros(hidden_size))
|
self.weight = nn.Parameter(torch.zeros(hidden_size))
|
||||||
self.variance_epsilon = eps
|
self.variance_epsilon = eps
|
||||||
|
self.register_buffer("gemma_weight", self.weight.data + 1.0, persistent=False)
|
||||||
|
# (Chen-0210) Gemma weight = standard_weight + 1. Precompute once.
|
||||||
|
# If TRTLLM allreduce fusion ever provides gemma-style norm
|
||||||
|
# natively, this can be removed.
|
||||||
|
self.weight.weight_loader = self._weight_loader
|
||||||
|
|
||||||
|
def _weight_loader(self, param: torch.Tensor, loaded_weight: torch.Tensor) -> None:
|
||||||
|
assert param.size() == loaded_weight.size()
|
||||||
|
param.data.copy_(loaded_weight)
|
||||||
|
self.gemma_weight = param.data + 1.0
|
||||||
|
|
||||||
def _forward_impl(
|
def _forward_impl(
|
||||||
self,
|
self,
|
||||||
@@ -536,7 +546,7 @@ class GemmaRMSNorm(MultiPlatformOp):
|
|||||||
if not _has_vllm_rms_norm:
|
if not _has_vllm_rms_norm:
|
||||||
return self.forward_native(x, residual, post_residual_addition)
|
return self.forward_native(x, residual, post_residual_addition)
|
||||||
|
|
||||||
w = self.weight.data + 1.0
|
w = self.gemma_weight
|
||||||
if _use_aiter:
|
if _use_aiter:
|
||||||
# aiter API: rms_norm(input, weight, eps) -> output
|
# aiter API: rms_norm(input, weight, eps) -> output
|
||||||
# fused_add_rms_norm(output, input, residual, residual_out, weight, eps)
|
# fused_add_rms_norm(output, input, residual, residual_out, weight, eps)
|
||||||
@@ -620,13 +630,12 @@ class GemmaRMSNorm(MultiPlatformOp):
|
|||||||
use_attn_tp_group: bool = True,
|
use_attn_tp_group: bool = True,
|
||||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
"""Forward with allreduce fusion; uses 1 + weight for fused kernels."""
|
"""Forward with allreduce fusion; uses 1 + weight for fused kernels."""
|
||||||
# TODO(brayden): we can see if TRTLLM allreduce fusion can provide gemma-style norm
|
|
||||||
return _forward_with_allreduce_fusion(
|
return _forward_with_allreduce_fusion(
|
||||||
self,
|
self,
|
||||||
x,
|
x,
|
||||||
residual,
|
residual,
|
||||||
post_residual_addition,
|
post_residual_addition,
|
||||||
self.weight + 1.0,
|
self.gemma_weight,
|
||||||
use_attn_tp_group=True,
|
use_attn_tp_group=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user