[Perf] Precompute gemma_weight to avoid redundant add on every forward (#22673)

This commit is contained in:
Jincong Chen
2026-04-17 23:37:41 +08:00
committed by GitHub
parent 83c5119d01
commit 2bac219d0c
+12 -3
View File
@@ -476,6 +476,16 @@ class GemmaRMSNorm(MultiPlatformOp):
super().__init__()
self.weight = nn.Parameter(torch.zeros(hidden_size))
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(
self,
@@ -536,7 +546,7 @@ class GemmaRMSNorm(MultiPlatformOp):
if not _has_vllm_rms_norm:
return self.forward_native(x, residual, post_residual_addition)
w = self.weight.data + 1.0
w = self.gemma_weight
if _use_aiter:
# aiter API: rms_norm(input, weight, eps) -> output
# fused_add_rms_norm(output, input, residual, residual_out, weight, eps)
@@ -620,13 +630,12 @@ class GemmaRMSNorm(MultiPlatformOp):
use_attn_tp_group: bool = True,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""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(
self,
x,
residual,
post_residual_addition,
self.weight + 1.0,
self.gemma_weight,
use_attn_tp_group=True,
)