Skip to content

vllm.model_executor.layers.quantization.compressed_tensors.transform.schemes.linear_qutlass_nvfp4

Classes:

QutlassNvFP4LinearMethod

Bases: CompressedTensorsLinearTransformMethod

Source code in vllm/model_executor/layers/quantization/compressed_tensors/transform/schemes/linear_qutlass_nvfp4.py
class QutlassNvFP4LinearMethod(CompressedTensorsLinearTransformMethod):
    def create_weights(
        self,
        layer,
        input_size_per_partition,
        output_partition_sizes,
        input_size,
        output_size,
        params_dtype,
        **extra_weight_attrs,
    ):
        # initializes fp4 qparams
        assert isinstance(layer.scheme, (CompressedTensorsW4A4Fp4,))
        ret = super().create_weights(
            layer,
            input_size_per_partition,
            output_partition_sizes,
            input_size,
            output_size,
            params_dtype,
            **extra_weight_attrs,
        )

        assert self.input_transform is not None
        assert len(self.input_transform.weight.partitions) >= 1

        return ret

    @staticmethod
    def _get_flashinfer_gemm_backend(kernel: NvFp4LinearKernel) -> str:
        """
        Given a kernel, find the string that is needed to be passed into
        `flashinfer_scaled_fp4_mm`, using
        vllm.model_executor.kernels.linear._LINEAR_BACKEND_KERNEL_MAP as source of truth
        """
        kernel_type = type(kernel)
        for key, kernels in _LINEAR_BACKEND_KERNEL_MAP.items():
            if not key.startswith("flashinfer_") or kernel_type not in kernels:
                continue
            backend = key.removeprefix("flashinfer_")
            # flashinfer GEMM backend uses "cute-dsl", not "cutedsl"
            return backend.replace("cutedsl", "cute-dsl")
        raise ValueError(
            f"QutlassNvFP4 transform requires a FlashInfer kernel, "
            f"got {kernel_type.__name__}"
        )

    def process_weights_after_loading(self, layer):
        super().process_weights_after_loading(layer)

        assert self.input_transform is not None
        layer.hadamard_matrix = self.input_transform.weight.partitions[0].data

        # fusedQuantizeNv stores raw absmax as block scales (sf = absmax),
        # while CT weights use sf = absmax * SFScaleVal / 6.0. The GEMM
        # computes alpha * sum(fp4_a * sf_a * fp4_w * sf_w), so alpha must
        # compensate: alpha = weight_global_scale / 6.0
        layer.fused_alpha = Parameter(
            layer.weight_global_scale / NVFP4_MAX, requires_grad=False
        )

        layer.fused_global_scale = Parameter(
            torch.tensor(
                [NVFP4_MAX],
                dtype=torch.float32,
                device=layer.weight_global_scale.device,
            ),
            requires_grad=False,
        )

        layer.flashinfer_gemm_backend = self._get_flashinfer_gemm_backend(
            layer.scheme.kernel
        )

    def apply(
        self,
        layer: torch.nn.Module,
        x: torch.Tensor,
        bias: torch.Tensor | None = None,
    ) -> torch.Tensor:
        assert bias is None
        output_size = layer.output_size_per_partition
        output_shape = [*x.shape[:-1], output_size]

        x_flat = x.contiguous().flatten(end_dim=-2)

        x_fp4, x_scales = fusedQuantizeNv(
            x_flat, layer.hadamard_matrix, layer.fused_global_scale
        )

        x_scales_blocked = to_blocked(x_scales, backend="triton").view(x_scales.shape)

        out = flashinfer_scaled_fp4_mm(
            x_fp4,
            layer.weight,
            x_scales_blocked,
            layer.weight_scale,
            layer.fused_alpha,
            x.dtype,
            backend=layer.flashinfer_gemm_backend,
        )

        out = slice_nvfp4_output(out, output_size)

        if self.output_transform is not None:
            for part_id, (start, length) in enumerate(self.partition_ranges):
                out[:, start : start + length] = self.output_transform(
                    out[:, start : start + length].clone(), part_id=part_id
                )

        return out.view(*output_shape)

_get_flashinfer_gemm_backend(kernel) staticmethod

Given a kernel, find the string that is needed to be passed into flashinfer_scaled_fp4_mm, using vllm.model_executor.kernels.linear._LINEAR_BACKEND_KERNEL_MAP as source of truth

Source code in vllm/model_executor/layers/quantization/compressed_tensors/transform/schemes/linear_qutlass_nvfp4.py
@staticmethod
def _get_flashinfer_gemm_backend(kernel: NvFp4LinearKernel) -> str:
    """
    Given a kernel, find the string that is needed to be passed into
    `flashinfer_scaled_fp4_mm`, using
    vllm.model_executor.kernels.linear._LINEAR_BACKEND_KERNEL_MAP as source of truth
    """
    kernel_type = type(kernel)
    for key, kernels in _LINEAR_BACKEND_KERNEL_MAP.items():
        if not key.startswith("flashinfer_") or kernel_type not in kernels:
            continue
        backend = key.removeprefix("flashinfer_")
        # flashinfer GEMM backend uses "cute-dsl", not "cutedsl"
        return backend.replace("cutedsl", "cute-dsl")
    raise ValueError(
        f"QutlassNvFP4 transform requires a FlashInfer kernel, "
        f"got {kernel_type.__name__}"
    )