Fix GemmaRMSNorm gemma_weight buffer storage for Qwen3.5 (#26430)
This commit is contained in:
@@ -613,7 +613,9 @@ 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)
|
self.register_buffer(
|
||||||
|
"gemma_weight", torch.ones_like(self.weight), persistent=False
|
||||||
|
)
|
||||||
# (Chen-0210) Gemma weight = standard_weight + 1. Precompute once.
|
# (Chen-0210) Gemma weight = standard_weight + 1. Precompute once.
|
||||||
# If TRTLLM allreduce fusion ever provides gemma-style norm
|
# If TRTLLM allreduce fusion ever provides gemma-style norm
|
||||||
# natively, this can be removed.
|
# natively, this can be removed.
|
||||||
@@ -622,7 +624,8 @@ class GemmaRMSNorm(MultiPlatformOp):
|
|||||||
def _weight_loader(self, param: torch.Tensor, loaded_weight: torch.Tensor) -> None:
|
def _weight_loader(self, param: torch.Tensor, loaded_weight: torch.Tensor) -> None:
|
||||||
assert param.size() == loaded_weight.size()
|
assert param.size() == loaded_weight.size()
|
||||||
param.data.copy_(loaded_weight)
|
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(
|
def _forward_impl(
|
||||||
self,
|
self,
|
||||||
|
|||||||
Reference in New Issue
Block a user