From f4eac503897bf4b8965d0e774a459c0a0a871b95 Mon Sep 17 00:00:00 2001 From: Jiajun Li <48857426+guapisolo@users.noreply.github.com> Date: Thu, 28 May 2026 03:42:46 -0700 Subject: [PATCH] Fix GemmaRMSNorm gemma_weight buffer storage for Qwen3.5 (#26430) --- python/sglang/srt/layers/layernorm.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index 7db57684a..951d704f7 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -613,7 +613,9 @@ 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) + self.register_buffer( + "gemma_weight", torch.ones_like(self.weight), 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. @@ -622,7 +624,8 @@ class GemmaRMSNorm(MultiPlatformOp): 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 + # Keep storage stable for CUDA graphs or fused paths that capture this buffer. + torch.add(param.data, 1.0, out=self.gemma_weight) def _forward_impl( self,