Skip to content

vllm.model_executor.layers.layernorm

Custom normalization layers.

Classes:

GemmaRMSNorm

Bases: CustomOp

RMS normalization for Gemma.

Two differences from the above RMSNorm
  1. x * (1 + w) instead of x * w.
  2. (x * w).to(orig_dtype) instead of x.to(orig_dtype) * w.

Methods:

  • forward_native –

    PyTorch-native implementation equivalent to forward().

Source code in vllm/model_executor/layers/layernorm.py
@CustomOp.register("gemma_rms_norm")
class GemmaRMSNorm(CustomOp):
    """RMS normalization for Gemma.

    Two differences from the above RMSNorm:
        1. x * (1 + w) instead of x * w.
        2. (x * w).to(orig_dtype) instead of x.to(orig_dtype) * w.
    """

    # --8<-- [end:gemma_rms_norm]

    def __init__(
        self,
        hidden_size: int,
        eps: float = 1e-6,
    ) -> None:
        super().__init__()
        self.weight = nn.Parameter(torch.zeros(hidden_size))
        self.variance_epsilon = eps

    def forward_native(
        self,
        x: torch.Tensor,
        residual: torch.Tensor | None = None,
    ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
        """PyTorch-native implementation equivalent to forward()."""
        weight = self.weight.float() + 1.0
        if residual is None:
            return ir.ops.rms_norm(x, weight, self.variance_epsilon)
        return ir.ops.fused_add_rms_norm(x, residual, weight, self.variance_epsilon)

    def forward_cuda(
        self,
        x: torch.Tensor,
        residual: torch.Tensor | None = None,
    ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
        return self.forward_native(x, residual)

    def forward_xpu(
        self,
        x: torch.Tensor,
        residual: torch.Tensor | None = None,
    ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
        if envs.VLLM_BATCH_INVARIANT:
            if residual is None:
                return self.forward_cuda(x, residual)
            return ir.ops.fused_add_rms_norm.impls["native"].impl_fn(
                x,
                residual,
                self.weight.float() + 1.0,
                self.variance_epsilon,
            )
        import vllm._xpu_ops  # noqa: F401 registers torch.ops.vllm.xpu_gemma_rms_norm

        # Fall back to the native path if the fused gemma kernels are not
        # available in the installed vllm-xpu-kernels package.
        if not hasattr(torch.ops._C, "gemma_rms_norm"):
            return self.forward_native(x, residual)

        # Pass the raw (bf16/fp16) weight; the +1 offset and the fp32 multiply
        # are folded into the kernel (matches forward_native numerics).
        if residual is not None:
            torch.ops.vllm.xpu_fused_add_gemma_rms_norm(
                x, residual, self.weight.data, self.variance_epsilon
            )
            return x, residual
        # empty_like preserves x's strides, but the kernel requires a
        # contiguous out (unlike x, which it can handle non-contiguous).
        out = torch.empty(x.shape, device=x.device, dtype=x.dtype)
        torch.ops.vllm.xpu_gemma_rms_norm(
            out, x, self.weight.data, self.variance_epsilon
        )
        return out

forward_native(x, residual=None)

PyTorch-native implementation equivalent to forward().

Source code in vllm/model_executor/layers/layernorm.py
def forward_native(
    self,
    x: torch.Tensor,
    residual: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
    """PyTorch-native implementation equivalent to forward()."""
    weight = self.weight.float() + 1.0
    if residual is None:
        return ir.ops.rms_norm(x, weight, self.variance_epsilon)
    return ir.ops.fused_add_rms_norm(x, residual, weight, self.variance_epsilon)

LayerNorm

Bases: CustomOp

Standard (mean-centered) LayerNorm.

Drop-in for a bare nn.LayerNorm that dispatches to a fused XPU kernel when one is available, and to the native implementation otherwise.

Normalization runs in the wider of the input and parameter dtypes and the result is cast back to the input dtype, so dtype=torch.float32 gives an fp32 reduction for a lower-precision input.

Source code in vllm/model_executor/layers/layernorm.py
@CustomOp.register("layer_norm")
class LayerNorm(CustomOp):
    """Standard (mean-centered) LayerNorm.

    Drop-in for a bare `nn.LayerNorm` that dispatches to a fused XPU kernel
    when one is available, and to the native implementation otherwise.

    Normalization runs in the wider of the input and parameter dtypes and the
    result is cast back to the input dtype, so `dtype=torch.float32` gives an
    fp32 reduction for a lower-precision input.
    """

    def __init__(
        self,
        hidden_size: int,
        eps: float = 1e-5,
        elementwise_affine: bool = True,
        bias: bool = True,
        dtype: torch.dtype | None = None,
    ) -> None:
        super().__init__()
        self.normalized_shape = (hidden_size,)
        self.eps = eps
        weight_dtype = dtype or torch.get_default_dtype()
        self.weight: nn.Parameter | None = None
        self.bias: nn.Parameter | None = None
        if elementwise_affine:
            self.weight = nn.Parameter(torch.ones(hidden_size, dtype=weight_dtype))
            if bias:
                self.bias = nn.Parameter(torch.zeros(hidden_size, dtype=weight_dtype))

    def _compute_dtype(self, x: torch.Tensor) -> torch.dtype:
        if self.weight is None:
            return x.dtype
        return torch.promote_types(x.dtype, self.weight.dtype)

    def forward_native(self, x: torch.Tensor) -> torch.Tensor:
        compute_dtype = self._compute_dtype(x)
        if compute_dtype != x.dtype:
            return F.layer_norm(
                x.to(compute_dtype),
                self.normalized_shape,
                self.weight,
                self.bias,
                self.eps,
            ).type_as(x)
        return F.layer_norm(x, self.normalized_shape, self.weight, self.bias, self.eps)

    def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
        return self.forward_native(x)

    def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
        import vllm._xpu_ops  # noqa: F401 registers torch.ops.vllm.xpu_layer_norm

        if (
            self.weight is None
            or self.bias is None
            or not hasattr(torch.ops._C, "layer_norm")
            # The kernel reduces in the input dtype, so a widened reduction
            # takes the native path.
            or self._compute_dtype(x) != x.dtype
        ):
            return self.forward_native(x)
        # empty_like preserves x's strides, but the kernel requires a
        # contiguous out (unlike x, which it can handle non-contiguous).
        out = torch.empty(x.shape, device=x.device, dtype=x.dtype)
        torch.ops.vllm.xpu_layer_norm(out, x, self.weight, self.bias, self.eps)
        return out

    def extra_repr(self) -> str:
        return f"hidden_size={self.normalized_shape[0]}, eps={self.eps}"

RMSNorm

Bases: CustomOp

Root mean square normalization.

Computes x -> w * x / sqrt(E[x^2] + eps) where w is the learned weight. Refer to https://arxiv.org/abs/1910.07467

Methods:

  • forward_native –

    PyTorch-native implementation equivalent to forward().

Source code in vllm/model_executor/layers/layernorm.py
@CustomOp.register("rms_norm")
class RMSNorm(CustomOp):
    """Root mean square normalization.

    Computes x -> w * x / sqrt(E[x^2] + eps) where w is the learned weight.
    Refer to https://arxiv.org/abs/1910.07467
    """

    # --8<-- [end:rms_norm]

    def __init__(
        self,
        hidden_size: int,
        eps: float = 1e-6,
        var_hidden_size: int | None = None,
        has_weight: bool = True,
        dtype: torch.dtype | None = None,
    ) -> None:
        super().__init__()

        self.hidden_size = hidden_size
        self.variance_epsilon = eps
        self.variance_size_override = (
            None if var_hidden_size == hidden_size else var_hidden_size
        )
        weight_dtype = dtype or torch.get_default_dtype()
        self.has_weight = has_weight
        self.weight = torch.ones(hidden_size, dtype=weight_dtype)
        if self.has_weight:
            self.weight = nn.Parameter(self.weight)

        # When has_weight=False, pass weight=None so implementations that
        # support a weightless path can skip the per-channel multiply.
        # Implementations that require weight (e.g. oink) fall back via IR
        # op priority when weight=None is unsupported.
        self.pass_weight = self.has_weight
        self.pass_weight_add = self.has_weight

    def forward_native(
        self,
        x: torch.Tensor,
        residual: torch.Tensor | None = None,
    ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
        """PyTorch-native implementation equivalent to forward()."""
        if residual is None:
            return ir.ops.rms_norm(
                x,
                self.weight.data if self.pass_weight else None,
                self.variance_epsilon,
                self.variance_size_override,
            )
        else:
            return ir.ops.fused_add_rms_norm.maybe_inplace(
                x,
                residual,
                self.weight.data if self.pass_weight_add else None,
                self.variance_epsilon,
                self.variance_size_override,
            )

    def forward_cuda(
        self,
        x: torch.Tensor,
        residual: torch.Tensor | None = None,
    ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
        if envs.VLLM_BATCH_INVARIANT:
            assert self.variance_size_override is None, (
                "Batch invariance is not supported for variance_size_override"
            )
            pass_weight = (
                self.pass_weight_add if residual is not None else self.pass_weight
            )
            return rms_norm_batch_invariant(
                x,
                self.weight.data if pass_weight else None,
                self.variance_epsilon,
                residual=residual,
            )

        return self.forward_native(x, residual)

    def forward_xpu(
        self,
        x: torch.Tensor,
        residual: torch.Tensor | None = None,
    ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
        if envs.VLLM_BATCH_INVARIANT and residual is not None:
            assert self.variance_size_override is None, (
                "Batch invariance is not supported for variance_size_override"
            )
            weight = self.weight.data if self.pass_weight_add else None
            return ir.ops.fused_add_rms_norm.impls["native"].impl_fn(
                x, residual, weight, self.variance_epsilon
            )
        return self.forward_cuda(x, residual)

    def extra_repr(self) -> str:
        s = f"hidden_size={self.weight.data.size(0)}"
        s += f", eps={self.variance_epsilon}"
        return s

forward_native(x, residual=None)

PyTorch-native implementation equivalent to forward().

Source code in vllm/model_executor/layers/layernorm.py
def forward_native(
    self,
    x: torch.Tensor,
    residual: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
    """PyTorch-native implementation equivalent to forward()."""
    if residual is None:
        return ir.ops.rms_norm(
            x,
            self.weight.data if self.pass_weight else None,
            self.variance_epsilon,
            self.variance_size_override,
        )
    else:
        return ir.ops.fused_add_rms_norm.maybe_inplace(
            x,
            residual,
            self.weight.data if self.pass_weight_add else None,
            self.variance_epsilon,
            self.variance_size_override,
        )

RMSNormGated

Bases: CustomOp

RMS Normalization with optional gating.

This is a native PyTorch implementation that supports: - Standard RMS normalization - Group RMS normalization - Optional gating with SiLU activation

Methods:

  • __init__ –

    Initialize RMSNormGated.

  • forward_native –

    PyTorch-native implementation equivalent to forward().

  • forward_static –

    Pure-PyTorch RMS normalization with optional gating.

Source code in vllm/model_executor/layers/layernorm.py
@CustomOp.register("rms_norm_gated")
class RMSNormGated(CustomOp):
    """RMS Normalization with optional gating.

    This is a native PyTorch implementation that supports:
    - Standard RMS normalization
    - Group RMS normalization
    - Optional gating with SiLU activation
    """

    # --8<-- [end:rms_norm_gated]

    def __init__(
        self,
        hidden_size: int,
        eps: float = 1e-5,
        group_size: int | None = None,
        norm_before_gate: bool = False,
        device: torch.device | None = None,
        dtype: torch.dtype | None = None,
        activation: str = "swish",
    ):
        """Initialize RMSNormGated.

        Args:
            hidden_size: Size of the hidden dimension
            eps: Epsilon for numerical stability
            group_size: If not None, do GroupNorm with each group
                        having group_size elements.
                        group_size=None is equivalent to group_size=hidden_size
                        (i.e. there's only 1 group).
            norm_before_gate: If True and z is provided: out = norm(x) * silu(z)
                              If False and z is provided: out = norm(x * silu(z))
            device: Device to create parameters on
            dtype: Data type for parameters
            activation: Activation function name for gating

        """
        factory_kwargs = {"device": device, "dtype": dtype}
        super().__init__()
        self.eps = eps
        self.activation = activation
        self.weight = nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
        self.register_parameter("bias", None)
        self.group_size = group_size
        self.norm_before_gate = norm_before_gate
        self.reset_parameters()

    def reset_parameters(self):
        torch.nn.init.ones_(self.weight)

    @staticmethod
    def forward_static(
        x: torch.Tensor,
        z: torch.Tensor | None,
        weight: torch.Tensor,
        epsilon: float,
        orig_dtype: torch.dtype,
        group_size: int | None = None,
        norm_before_gate: bool = True,
        activation: str = "swish",
    ) -> torch.Tensor:
        """Pure-PyTorch RMS normalization with optional gating.

        This static method contains the full native logic so that both
        ``forward_native`` and ``MatcherRMSNormGated`` (used by the
        compilation pattern matcher) can share the same implementation.

        If *z* is not None and *norm_before_gate* is True:
            ``out = rms_norm(x) * act(z)``
        If *z* is not None and *norm_before_gate* is False:
            ``out = rms_norm(x * act(z))``
        """
        x = x.float()
        weight = weight.float()
        if z is not None:
            z = z.float()

        assert activation in ["silu", "sigmoid", "swish"]
        act_fn = F.sigmoid if activation == "sigmoid" else F.silu

        if z is not None and not norm_before_gate:
            x = x * act_fn(z)

        if group_size is None:
            variance = x.pow(2).mean(dim=-1, keepdim=True)
            x_normed = x * torch.rsqrt(variance + epsilon)
            out = x_normed * weight
        else:
            from einops import rearrange

            x_group = rearrange(x, "... (g d) -> ... g d", d=group_size)
            variance = x_group.pow(2).mean(dim=-1, keepdim=True)
            x_normed = x_group * torch.rsqrt(variance + epsilon)
            out = rearrange(x_normed, "... g d -> ... (g d)") * weight

        if z is not None and norm_before_gate:
            out = out * act_fn(z)

        return out.to(orig_dtype)

    def forward_native(
        self, x: torch.Tensor, z: torch.Tensor | None = None
    ) -> torch.Tensor:
        """PyTorch-native implementation equivalent to forward()."""
        return self.forward_static(
            x,
            z,
            self.weight,
            self.eps,
            x.dtype,
            group_size=self.group_size,
            norm_before_gate=self.norm_before_gate,
            activation=self.activation,
        )

    def forward_cuda(
        self, x: torch.Tensor, z: torch.Tensor | None = None
    ) -> torch.Tensor:
        from vllm.third_party.flash_linear_attention.ops.layernorm_guard import (
            rmsnorm_fn,
        )

        return rmsnorm_fn(
            x,
            self.weight,
            self.bias,
            z=z,
            eps=self.eps,
            group_size=self.group_size,
            norm_before_gate=self.norm_before_gate,
            activation=self.activation,
        )

    def forward_xpu(
        self, x: torch.Tensor, z: torch.Tensor | None = None
    ) -> torch.Tensor:
        return self.forward_cuda(x, z)

__init__(hidden_size, eps=1e-05, group_size=None, norm_before_gate=False, device=None, dtype=None, activation='swish')

Initialize RMSNormGated.

Parameters:

  • hidden_size

    (int) –

    Size of the hidden dimension

  • eps

    (float, default: 1e-05 ) –

    Epsilon for numerical stability

  • group_size

    (int | None, default: None ) –

    If not None, do GroupNorm with each group having group_size elements. group_size=None is equivalent to group_size=hidden_size (i.e. there's only 1 group).

  • norm_before_gate

    (bool, default: False ) –

    If True and z is provided: out = norm(x) * silu(z) If False and z is provided: out = norm(x * silu(z))

  • device

    (device | None, default: None ) –

    Device to create parameters on

  • dtype

    (dtype | None, default: None ) –

    Data type for parameters

  • activation

    (str, default: 'swish' ) –

    Activation function name for gating

Source code in vllm/model_executor/layers/layernorm.py
def __init__(
    self,
    hidden_size: int,
    eps: float = 1e-5,
    group_size: int | None = None,
    norm_before_gate: bool = False,
    device: torch.device | None = None,
    dtype: torch.dtype | None = None,
    activation: str = "swish",
):
    """Initialize RMSNormGated.

    Args:
        hidden_size: Size of the hidden dimension
        eps: Epsilon for numerical stability
        group_size: If not None, do GroupNorm with each group
                    having group_size elements.
                    group_size=None is equivalent to group_size=hidden_size
                    (i.e. there's only 1 group).
        norm_before_gate: If True and z is provided: out = norm(x) * silu(z)
                          If False and z is provided: out = norm(x * silu(z))
        device: Device to create parameters on
        dtype: Data type for parameters
        activation: Activation function name for gating

    """
    factory_kwargs = {"device": device, "dtype": dtype}
    super().__init__()
    self.eps = eps
    self.activation = activation
    self.weight = nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
    self.register_parameter("bias", None)
    self.group_size = group_size
    self.norm_before_gate = norm_before_gate
    self.reset_parameters()

forward_native(x, z=None)

PyTorch-native implementation equivalent to forward().

Source code in vllm/model_executor/layers/layernorm.py
def forward_native(
    self, x: torch.Tensor, z: torch.Tensor | None = None
) -> torch.Tensor:
    """PyTorch-native implementation equivalent to forward()."""
    return self.forward_static(
        x,
        z,
        self.weight,
        self.eps,
        x.dtype,
        group_size=self.group_size,
        norm_before_gate=self.norm_before_gate,
        activation=self.activation,
    )

forward_static(x, z, weight, epsilon, orig_dtype, group_size=None, norm_before_gate=True, activation='swish') staticmethod

Pure-PyTorch RMS normalization with optional gating.

This static method contains the full native logic so that both forward_native and MatcherRMSNormGated (used by the compilation pattern matcher) can share the same implementation.

If z is not None and norm_before_gate is True: out = rms_norm(x) * act(z) If z is not None and norm_before_gate is False: out = rms_norm(x * act(z))

Source code in vllm/model_executor/layers/layernorm.py
@staticmethod
def forward_static(
    x: torch.Tensor,
    z: torch.Tensor | None,
    weight: torch.Tensor,
    epsilon: float,
    orig_dtype: torch.dtype,
    group_size: int | None = None,
    norm_before_gate: bool = True,
    activation: str = "swish",
) -> torch.Tensor:
    """Pure-PyTorch RMS normalization with optional gating.

    This static method contains the full native logic so that both
    ``forward_native`` and ``MatcherRMSNormGated`` (used by the
    compilation pattern matcher) can share the same implementation.

    If *z* is not None and *norm_before_gate* is True:
        ``out = rms_norm(x) * act(z)``
    If *z* is not None and *norm_before_gate* is False:
        ``out = rms_norm(x * act(z))``
    """
    x = x.float()
    weight = weight.float()
    if z is not None:
        z = z.float()

    assert activation in ["silu", "sigmoid", "swish"]
    act_fn = F.sigmoid if activation == "sigmoid" else F.silu

    if z is not None and not norm_before_gate:
        x = x * act_fn(z)

    if group_size is None:
        variance = x.pow(2).mean(dim=-1, keepdim=True)
        x_normed = x * torch.rsqrt(variance + epsilon)
        out = x_normed * weight
    else:
        from einops import rearrange

        x_group = rearrange(x, "... (g d) -> ... g d", d=group_size)
        variance = x_group.pow(2).mean(dim=-1, keepdim=True)
        x_normed = x_group * torch.rsqrt(variance + epsilon)
        out = rearrange(x_normed, "... g d -> ... (g d)") * weight

    if z is not None and norm_before_gate:
        out = out * act_fn(z)

    return out.to(orig_dtype)