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