From 2bac219d0cc16c2e76972d837079347d20807177 Mon Sep 17 00:00:00 2001 From: Jincong Chen Date: Fri, 17 Apr 2026 23:37:41 +0800 Subject: [PATCH] [Perf] Precompute gemma_weight to avoid redundant add on every forward (#22673) --- python/sglang/srt/layers/layernorm.py | 15 ++++++++++++--- 1 file changed, 12 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index 3d914813a..e995d8378 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -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, )