From 166f4e2dc3a6ed1dfc3641611c9a07d8f0e87f9c Mon Sep 17 00:00:00 2001 From: Anuj Bolewar <143988950+anujbolewar@users.noreply.github.com> Date: Tue, 4 Aug 2026 21:33:09 +0530 Subject: [PATCH] fix: fuse weightless RMSNorms at their declared width (#50867) Signed-off-by: Anuj Bolewar Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com> Co-authored-by: Anuj Bolewar Co-authored-by: Harry Mellor <19981378+hmellor@users.noreply.github.com> --- .../transformers/fusers/test_rms_norm.py | 43 +++++++++++-------- .../models/transformers/fusers/rms_norm.py | 9 ++-- 2 files changed, 29 insertions(+), 23 deletions(-) diff --git a/tests/models/transformers/fusers/test_rms_norm.py b/tests/models/transformers/fusers/test_rms_norm.py index 131897b1924..711f9b8b221 100644 --- a/tests/models/transformers/fusers/test_rms_norm.py +++ b/tests/models/transformers/fusers/test_rms_norm.py @@ -146,26 +146,39 @@ def test_rms_norm_builds_vllm_class(cls, expected, zero_centered, default_vllm_c from vllm.model_executor.layers.layernorm import GemmaRMSNorm as VLLMGemmaRMSNorm from vllm.model_executor.layers.layernorm import RMSNorm as VLLMRMSNorm - # `default_vllm_config` supplies the config context the CustomOp needs; the - # weightless path reads hidden size from the model config, so stub it. - vllm_config = SimpleNamespace( - model_config=SimpleNamespace(get_hidden_size=lambda: 16) - ) with torch.device("meta"): module = cls() fuser = get_fuser(module) - built = fuser.fuse(module, "norm", vllm_config) + built = fuser.fuse(module, "norm", default_vllm_config) from vllm.model_executor.models.transformers.fusers.rms_norm import ( TPAwareNormMixin, ) types_by_name = {"RMSNorm": VLLMRMSNorm, "GemmaRMSNorm": VLLMGemmaRMSNorm} assert isinstance(built, types_by_name[expected]) - assert isinstance(built, TPAwareNormMixin) # fused norms self-correct under TP + assert isinstance(built, TPAwareNormMixin) assert built.variance_epsilon == module.variance_epsilon - assert isinstance(built.weight, nn.Parameter) == ( - getattr(module, "weight", None) is not None - ) + # The weight states the hidden size (and is what the TP check reads); a norm + # without one has none to state. + weight = getattr(module, "weight", None) + assert isinstance(built.weight, nn.Parameter) == (weight is not None) + assert built.weight.shape[0] == (weight.size(0) if weight is not None else 0) + + +def test_weightless_norm_has_no_hidden_size(default_vllm_config): + """A weightless norm states no hidden size, and the model's (LM) hidden size + would be wrong for one on a sub-dimension: Llama 4's `qk_norm` normalizes + head_dim, so sizing it at hidden_size made the TP gather reject its input. + It fuses with none, and so normalizes the width it is given at any TP size, + matching the unfused norm and vLLM's native `RMSNorm(hidden_size=head_dim)`. + """ + module = WeightlessRMSNorm(72) + built = get_fuser(module).fuse(module, "norm", default_vllm_config) + assert built.hidden_size == 0 + + built.tp_size = 2 # emulate TP=2 without a real process group + x = torch.randn(4, 72) + torch.testing.assert_close(built(x), module(x)) def test_fused_rms_norm_op_default_eps(default_vllm_config): @@ -178,8 +191,7 @@ def test_fused_rms_norm_op_default_eps(default_vllm_config): fuser = get_fuser(module) assert isinstance(fuser, RMSNormFuser) assert not fuser.zero_centered - mc = SimpleNamespace(get_hidden_size=lambda: 16, dtype=torch.float32) - vllm_config = SimpleNamespace(model_config=mc) + vllm_config = SimpleNamespace(model_config=SimpleNamespace(dtype=torch.float32)) built = fuser.fuse(module, "norm", vllm_config) assert isinstance(built, VLLMRMSNorm) assert built.variance_epsilon == torch.finfo(torch.float32).eps @@ -188,18 +200,15 @@ def test_fused_rms_norm_op_default_eps(default_vllm_config): def test_eps_is_derived_per_instance(default_vllm_config): """Two instances of the same norm class with different eps must fuse to their own eps: the type-cached fuser holds only structure, not this value.""" - vllm_config = SimpleNamespace( - model_config=SimpleNamespace(get_hidden_size=lambda: 16) - ) with torch.device("meta"): for eps in (1e-5, 1e-6): module = RMSNorm(16, eps=eps) - built = get_fuser(module).fuse(module, "norm", vllm_config) + built = get_fuser(module).fuse(module, "norm", default_vllm_config) assert built.variance_epsilon == eps def test_fused_norm_is_gather_capable(default_vllm_config): - """Every fused norm is emitted gather-capable, so a norm on a head-sharded + """Every weighted fused norm is emitted gather-capable, so a norm on a head-sharded projection (OLMoE-style) self-corrects at runtime with no QKV-specific plumbing. A full-width input skips the gather and equals a plain norm.""" from vllm.model_executor.layers.layernorm import GemmaRMSNorm, RMSNorm diff --git a/vllm/model_executor/models/transformers/fusers/rms_norm.py b/vllm/model_executor/models/transformers/fusers/rms_norm.py index 781b417177a..396ae4c5b2d 100644 --- a/vllm/model_executor/models/transformers/fusers/rms_norm.py +++ b/vllm/model_executor/models/transformers/fusers/rms_norm.py @@ -190,20 +190,17 @@ class RMSNormFuser(BaseFuser): self, module: nn.Module, prefix: str, vllm_config: "VllmConfig" ) -> nn.Module: """Fuse the matched RMSNorm pattern into a vLLM fused RMSNorm CustomOp.""" - model_config = vllm_config.model_config weight = getattr(module, "weight", None) - hidden_size = ( - weight.size(0) if weight is not None else model_config.get_hidden_size() - ) + has_weight = weight is not None + hidden_size = weight.size(0) if has_weight else 0 graph = trace(module) eps = self._eps_from_graph(graph) if graph is not None else None if eps is None: # If eps not in graph, match torch behaviour. - dtype = weight.dtype if weight is not None else model_config.dtype + dtype = weight.dtype if has_weight else vllm_config.model_config.dtype eps = torch.finfo(dtype).eps if self.zero_centered: return TPAwareGemmaRMSNorm(hidden_size=hidden_size, eps=eps) - has_weight = weight is not None return TPAwareRMSNorm( hidden_size=hidden_size, eps=eps,