Skip to content

vllm.model_executor.kernels.linear.mxfp4.aiter

Classes:

AiterMxfp4LinearKernel

Bases: MxFp4LinearKernel

AITER-based native MXFP4 GEMM kernel for ROCm.

Source code in vllm/model_executor/kernels/linear/mxfp4/aiter.py
class AiterMxfp4LinearKernel(MxFp4LinearKernel):
    """AITER-based native MXFP4 GEMM kernel for ROCm."""

    def __init__(self, config: MxFp4LinearLayerConfig) -> None:
        super().__init__(config)
        self.use_asm_gemm = rocm_aiter_ops.is_asm_fp4_gemm_dynamic_quant_enabled()
        self.out_dtype = torch.get_default_dtype()

    @classmethod
    def is_supported(
        cls, compute_capability: int | None = None
    ) -> tuple[bool, str | None]:
        if not current_platform.supports_mx():
            return False, "current platform does not support native MXFP4 computation"
        if is_aiter_found_and_supported():
            return True, None
        return False, "AITER not found or not supported on the current platform"

    @classmethod
    def can_implement(cls, c: MxFp4LinearLayerConfig) -> tuple[bool, str | None]:
        return True, None

    def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
        if self.use_asm_gemm:
            from aiter.ops.shuffle import shuffle_weight

            weight_scale = layer.weight_scale.data
            sm, sn = weight_scale.shape
            weight_scale = weight_scale.view(sm // 32, 2, 16, sn // 8, 2, 4, 1)
            weight_scale = weight_scale.permute(0, 3, 5, 2, 4, 1, 6).contiguous()
            weight_scale = weight_scale.view(sm, sn)
            layer.weight_scale = Parameter(weight_scale, requires_grad=False)

            layer.weight = Parameter(
                shuffle_weight(layer.weight.data, layout=(16, 16)),
                requires_grad=False,
            )
        else:
            layer.weight_scale = Parameter(
                layer.weight_scale.data.T.contiguous(), requires_grad=False
            )

    def apply_weights(
        self,
        layer: torch.nn.Module,
        x: torch.Tensor,
        bias: torch.Tensor | None = None,
    ) -> torch.Tensor:
        y = torch.ops.vllm.gemm_with_dynamic_quant(
            x,
            layer.weight,
            layer.weight_scale,
            self.use_asm_gemm,
            self.out_dtype,
        )
        if bias is not None:
            y = y + bias
        return y