Skip to content

vllm.compilation.passes.fusion.rms_quant_fusion

Classes:

FusedAddRMSNormNvfp4QuantPattern

Fuse add-RMSNorm with NVFP4 quantization for either scale layout.

Source code in vllm/compilation/passes/fusion/rms_quant_fusion.py
class FusedAddRMSNormNvfp4QuantPattern:
    """Fuse add-RMSNorm with NVFP4 quantization for either scale layout."""

    def __init__(self, epsilon: float, is_sf_swizzled_layout: bool) -> None:
        assert _FLASHINFER_NVFP4_RMS_QUANT_OP is not None
        self.epsilon = epsilon
        self.is_sf_swizzled_layout = is_sf_swizzled_layout
        self.FUSED_OP = _FLASHINFER_NVFP4_RMS_QUANT_OP

    def register(self, pm_pass: PatternMatcherPass) -> None:
        def pattern(
            result: torch.Tensor,
            result_block_scale: torch.Tensor,
            input: torch.Tensor,
            weight: torch.Tensor,
            residual: torch.Tensor,
            input_global_scale: torch.Tensor,
        ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
            result_rms, updated_residual = vllm.ir.ops.fused_add_rms_norm(
                input, residual, weight, self.epsilon
            )
            at = auto_functionalized(
                torch.ops._C.scaled_fp4_quant.out,
                input=result_rms,
                input_scale=input_global_scale,
                is_sf_swizzled_layout=self.is_sf_swizzled_layout,
                output=result,
                output_scale=result_block_scale,
            )
            return at[1], updated_residual, at[2]

        def replacement(
            result: torch.Tensor,
            result_block_scale: torch.Tensor,
            input: torch.Tensor,
            weight: torch.Tensor,
            residual: torch.Tensor,
            input_global_scale: torch.Tensor,
        ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
            hidden_size = input.shape[-1]
            num_tokens = input.numel() // hidden_size
            # This full-size dummy is required by FlashInfer's TVM-FFI tensor
            # validation even though output_both_sf_layouts=False leaves it untouched.
            block_scale_unswizzled = torch.empty(
                (num_tokens, hidden_size // 16),
                dtype=torch.float8_e4m3fn,
                device=input.device,
            )
            at = auto_functionalized(
                self.FUSED_OP,
                result=result,
                result_block_scale=result_block_scale,
                residual=residual,
                input=input,
                weight=weight,
                input_global_scale=input_global_scale,
                block_scale_unswizzled=block_scale_unswizzled,
                is_sf_swizzled_layout=self.is_sf_swizzled_layout,
                epsilon=self.epsilon,
            )
            # result, updated residual, block scale in the requested layout
            return at[1], at[3], at[2]

        inputs = [
            torch.empty(
                (5, 32), dtype=torch.uint8, device=current_platform.device_type
            ),
            (
                empty_i32(128, 4)
                if self.is_sf_swizzled_layout
                else torch.empty(
                    (5, 4),
                    dtype=torch.uint8,
                    device=current_platform.device_type,
                )
            ),
            empty_bf16(5, 64),
            empty_bf16(64),
            empty_bf16(5, 64),
            empty_fp32(1),
        ]
        pm.register_replacement(
            pattern,
            replacement,
            inputs,
            pm.fwd_only,
            pm_pass,
            extra_check=_rms_input_weight_dtype_match,
        )

FusedRMSQuantKey

Bases: NamedTuple

Named tuple for identifying the type of RMSNorm + quant fusion. quant: type of quantization fused_add: does the op also perform the residual add

Source code in vllm/compilation/passes/fusion/rms_quant_fusion.py
class FusedRMSQuantKey(NamedTuple):
    """
    Named tuple for identifying the type of RMSNorm + quant fusion.
    quant: type of quantization
    fused_add: does the op also perform the residual add
    """

    quant: QuantKey
    fused_add: bool

    def __str__(self) -> str:
        return (
            f"FusedQuantKey({self.quant}, with"
            f"{'' if self.fused_add else 'out'} residual)"
        )

RMSNormQuantFusionPass

Bases: VllmPatternMatcherPass

This pass fuses rms_norm & quant custom ops into a fused rms_norm_quant op. It also supports fused_add_rms_norm.

Source code in vllm/compilation/passes/fusion/rms_quant_fusion.py
class RMSNormQuantFusionPass(VllmPatternMatcherPass):
    """
    This pass fuses rms_norm & quant custom ops into a fused rms_norm_quant op.
    It also supports fused_add_rms_norm.
    """

    @enable_fake_mode
    def __init__(self, config: VllmConfig) -> None:
        super().__init__(config)

        self.patterns: PatternMatcherPass = PatternMatcherPass(
            pass_name="rmsnorm_quant_fusion_pass"
        )

        # Make sure fused add patterns are before simple rms norm,
        # as the latter is a subset of the former in torch ops
        for epsilon in [1e-5, 1e-6]:
            if _FLASHINFER_NVFP4_RMS_QUANT_OP is not None and (
                current_platform.has_device_capability(100)
            ):
                for is_sf_swizzled_layout in (True, False):
                    FusedAddRMSNormNvfp4QuantPattern(
                        epsilon, is_sf_swizzled_layout
                    ).register(self.patterns)

            # Fuse fused_add_rms_norm + static fp8 quant
            FusedAddRMSNormStaticQuantPattern(epsilon, FP8_DTYPE).register(
                self.patterns
            )

            # Fuse rms_norm + static fp8 quant
            RMSNormStaticQuantPattern(epsilon, FP8_DTYPE).register(self.patterns)

            # Fuse fused_add_rms_norm + dynamic per-token fp8 quant
            FusedAddRMSNormDynamicQuantPattern(epsilon, FP8_DTYPE).register(
                self.patterns
            )

            # Fuse rms_norm + dynamic per-token fp8 quant
            RMSNormDynamicQuantPattern(epsilon, FP8_DTYPE).register(self.patterns)

            # Only register group quant patterns on CUDA/ROCm where the C++ op exists
            for group_shape in [GroupShape(1, 128), GroupShape(1, 64)]:
                for has_col_major_scales in [True, False]:
                    for is_e8m0 in [True, False]:
                        for is_tma_aligned in [False, True]:
                            # Fuse fused_add_rms_norm + fp8 group quant
                            FusedAddRMSNormGroupQuantPattern(
                                epsilon,
                                FP8_DTYPE,
                                group_shape=group_shape,
                                is_e8m0=is_e8m0,
                                has_col_major_scales=has_col_major_scales,
                                is_tma_aligned=is_tma_aligned,
                            ).register(self.patterns)

                            # Fuse rms_norm + fp8 group quant
                            RMSNormGroupQuantPattern(
                                epsilon,
                                FP8_DTYPE,
                                group_shape=group_shape,
                                is_e8m0=is_e8m0,
                                has_col_major_scales=has_col_major_scales,
                                is_tma_aligned=is_tma_aligned,
                            ).register(self.patterns)

        self.dump_patterns(config, self.patterns)

    @VllmInductorPass.time_and_log
    def __call__(self, graph: fx.Graph) -> None:
        self.matched_count = self.patterns.apply(graph)
        logger.debug("Replaced %s patterns", self.matched_count)

    def uuid(self) -> str:
        return self.hash_source(
            self,
            RMSNormGroupQuantPattern,
            RMSNormQuantPattern,
            RMSNormStaticQuantPattern,
            RMSNormDynamicQuantPattern,
            FusedAddRMSNormStaticQuantPattern,
            FusedAddRMSNormDynamicQuantPattern,
            FusedAddRMSNormGroupQuantPattern,
            FusedAddRMSNormNvfp4QuantPattern,
        )

_flashinfer_fused_add_rms_norm_nvfp4_quant(result, result_block_scale, residual, input, weight, input_global_scale, block_scale_unswizzled, is_sf_swizzled_layout, epsilon)

FlashInfer fused add + RMSNorm + NVFP4 quantization.

Source code in vllm/compilation/passes/fusion/rms_quant_fusion.py
def _flashinfer_fused_add_rms_norm_nvfp4_quant(
    result: torch.Tensor,
    result_block_scale: torch.Tensor,
    residual: torch.Tensor,
    input: torch.Tensor,
    weight: torch.Tensor,
    input_global_scale: torch.Tensor,
    block_scale_unswizzled: torch.Tensor,
    is_sf_swizzled_layout: bool,
    epsilon: float,
) -> None:
    """FlashInfer fused add + RMSNorm + NVFP4 quantization."""
    assert _FLASHINFER_ADD_RMSNORM_FP4QUANT is not None
    _FLASHINFER_ADD_RMSNORM_FP4QUANT(
        input,
        residual,
        weight,
        y_fp4=result.view(torch.float4_e2m1fn_x2),
        block_scale=result_block_scale.view(torch.float8_e4m3fn),
        global_scale=input_global_scale.reshape(1),
        eps=epsilon,
        block_size=16,
        scale_format="e4m3",
        is_sf_swizzled_layout=is_sf_swizzled_layout,
        output_both_sf_layouts=False,
        block_scale_unswizzled=block_scale_unswizzled,
    )

_rms_input_weight_dtype_match(match)

Prevent fusion when rms_norm input and weight dtypes differ.

Source code in vllm/compilation/passes/fusion/rms_quant_fusion.py
def _rms_input_weight_dtype_match(match: pm.Match) -> bool:
    """Prevent fusion when rms_norm input and weight dtypes differ."""
    for node in match.nodes:
        if node.target == _RMS_NORM_OP:
            # rms_norm(x, weight, epsilon, variance_size)
            x, weight = node.args[0], node.args[1]
        elif node.target == _FUSED_ADD_RMS_NORM_OP:
            # fused_add_rms_norm(x, residual, weight, epsilon, variance_size)
            x, weight = node.args[0], node.args[2]
        else:
            continue
        if isinstance(x, fx.Node) and isinstance(weight, fx.Node):
            return x.meta["val"].dtype == weight.meta["val"].dtype
    return True