Fix GemmaRMSNorm gemma_weight buffer storage for Qwen3.5 (#26430)

This commit is contained in:
Jiajun Li
2026-05-28 18:42:46 +08:00
committed by GitHub
parent 8e0ed75f2d
commit f4eac50389
+5 -2
View File
@@ -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,