Skip to content
KernelIndex
Search⌘K

submission 532764

josusanmartin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 281 lines, June 9 Researcher Reciprocity License v1.0.

mxfp4_v228_v219_copy_f.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-532764?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
9.79µs
#213 of 1143
2026-03-11

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a2fb8625dc7e6246c49ca30a419eff423fc5b54c8b68d97d9a4be1f18539c811
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15

Kernel source

mxfp4_v228_v219_copy_f.py281 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

from __future__ import annotations

import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op as _mxfp4_quant_op_even

from task import input_t, output_t

import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module

_kernel_module._mxfp4_quant_op = _mxfp4_quant_op_even

from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle

_BF16 = dtypes.bf16
_FP4X2 = dtypes.fp4x2
_FP8_E8M0 = dtypes.fp8_e8m0
_KERNEL_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"

_PUBLIC_SMALL = {
    (4, 2880, 512): {
        "BLOCK_SIZE_M": 8,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 2,
        "waves_per_eu": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (16, 2112, 7168): {
        "BLOCK_SIZE_M": 16,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 2,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 7,
    },
    (32, 4096, 512): {
        "BLOCK_SIZE_M": 8,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 3,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
    (32, 2880, 512): {
        "BLOCK_SIZE_M": 8,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 4,
        "num_stages": 3,
        "waves_per_eu": 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
        "NUM_KSPLIT": 1,
    },
}

_PUBLIC_LARGE = {
    (64, 7168, 2048): 2,
    (256, 3072, 1536): 1,
}

_HIDDEN_SHAPES = {
    (8, 2112, 7168),
    (16, 3072, 1536),
    (64, 3072, 1536),
    (256, 2880, 512),
}

_QUANT_BLOCK = 32
_QUANT_TILE = 128
_BUFS = {}


@triton.jit
def _dynamic_mxfp4_quant_kernel_even_asm_layout(
    x_ptr,
    x_fp4_ptr,
    bs_ptr,
    stride_x_m,
    stride_x_n,
    stride_x_fp4_m,
    stride_x_fp4_n,
    stride_bs_m,
    stride_bs_n,
    M: tl.constexpr,
    N: tl.constexpr,
    scaleN: tl.constexpr,
    scaleM_pad: tl.constexpr,
    scaleN_pad: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
    SHUFFLE: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)

    stride_x_m = tl.cast(stride_x_m, tl.int64)
    stride_x_n = tl.cast(stride_x_n, tl.int64)
    stride_x_fp4_m = tl.cast(stride_x_fp4_m, tl.int64)
    stride_x_fp4_n = tl.cast(stride_x_fp4_n, tl.int64)

    x_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    x_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)
    x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
    x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
    x = tl.load(x_ptr + x_offs, mask=x_mask).to(tl.float32)

    out_tensor, bs_e8m0 = _mxfp4_quant_op_even(
        x,
        MXFP4_QUANT_BLOCK_SIZE,
        BLOCK_SIZE,
        MXFP4_QUANT_BLOCK_SIZE,
    )

    out_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(
        0, MXFP4_QUANT_BLOCK_SIZE // 2
    )
    out_offs = (
        out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
    )
    out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
    tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)

    bs_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    bs_offs_n = pid_n

    if SHUFFLE:
        bs_offs_0 = bs_offs_m[:, None] // 32
        bs_offs_1 = bs_offs_m[:, None] % 32
        bs_offs_2 = bs_offs_1 % 16
        bs_offs_1 = bs_offs_1 // 16
        bs_offs_3 = bs_offs_n[None, :] // 8
        bs_offs_4 = bs_offs_n[None, :] % 8
        bs_offs_5 = bs_offs_4 % 4
        bs_offs_4 = bs_offs_4 // 4
        bs_offs = (
            bs_offs_1
            + bs_offs_4 * 2
            + bs_offs_2 * 4
            + bs_offs_5 * 64
            + bs_offs_3 * 256
            + bs_offs_0 * 32 * scaleN
        )
        bs_mask1 = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN)[None, :]
        bs_mask2 = (bs_offs_m < scaleM_pad)[:, None] & (bs_offs_n < scaleN_pad)[None, :]
        bs_e8m0 = tl.where(bs_mask1, bs_e8m0, 127)
        tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask2)
    else:
        bs_offs = bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
        bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < N)[None, :]
        tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask)


def _e8m0_shuffle_safe(scale: torch.Tensor) -> torch.Tensor:
    m, n = scale.shape
    scale_padded = torch.empty(
        ((m + 255) // 256) * 256,
        ((n + 7) // 8) * 8,
        dtype=scale.dtype,
        device=scale.device,
    )
    scale_padded.fill_(0x7F)
    scale_padded[:m, :n] = scale
    sm, sn = scale_padded.shape
    return (
        scale_padded.view(sm // 32, 2, 16, sn // 8, 2, 4)
        .permute(0, 3, 5, 2, 4, 1)
        .contiguous()
        .view(sm, sn)
    )


def _safe_wrapper(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor):
    a_q_raw, a_scale = dynamic_mxfp4_quant(a.contiguous())
    a_scale_sh = _e8m0_shuffle_safe(a_scale)
    return aiter.gemm_a4w4(
        a_q_raw.view(_FP4X2),
        b_shuffle,
        a_scale_sh.view(_FP8_E8M0),
        b_scale_sh,
        dtype=_BF16,
        bpreshuffle=True,
    )


def _get_large_bufs(m: int, k: int, n: int, device):
    x_fp4 = torch.empty((m, k >> 1), dtype=torch.uint8, device=device)
    scale_n = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK
    scale_n_pad = ((scale_n + 7) >> 3) << 3
    scale_m_pad = ((m + 255) >> 8) << 8
    scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)
    padded_m = ((m + 31) >> 5) << 5
    out = torch.empty((padded_m, n), dtype=_BF16, device=device)
    return x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    a, b, _b_q, b_shuffle, b_scale_sh = data
    m, k = a.shape
    n = b.shape[0]
    key = (int(m), int(n), int(k))

    if key in _HIDDEN_SHAPES or (key not in _PUBLIC_SMALL and key not in _PUBLIC_LARGE):
        return _safe_wrapper(a, b_shuffle, b_scale_sh)

    if key in _PUBLIC_LARGE:
        if key not in _BUFS:
            _BUFS[key] = ("large", _get_large_bufs(m, k, n, a.device))
        _, (x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out) = _BUFS[key]
        grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, scale_n_pad)
        _dynamic_mxfp4_quant_kernel_even_asm_layout[grid](
            a,
            x_fp4,
            scale,
            a.stride(0),
            a.stride(1),
            x_fp4.stride(0),
            x_fp4.stride(1),
            scale.stride(0),
            scale.stride(1),
            M=m,
            N=k,
            scaleN=scale_n,
            scaleM_pad=scale_m_pad,
            scaleN_pad=scale_n_pad,
            BLOCK_SIZE=_QUANT_TILE,
            MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK,
            SHUFFLE=True,
        )
        gemm_a4w4_asm(
            x_fp4.view(_FP4X2),
            b_shuffle,
            scale.view(_FP8_E8M0),
            b_scale_sh,
            out,
            _KERNEL_32X128,
            bpreshuffle=True,
            log2_k_split=_PUBLIC_LARGE[key],
        )
        return out[:m]

    if key not in _BUFS:
        _BUFS[key] = ("small", torch.empty((m, n), dtype=_BF16, device=a.device))
    _, out = _BUFS[key]

    w = b_shuffle.view(torch.uint8).reshape(n // 16, k // 2 * 16)
    sm, sn = b_scale_sh.shape
    w_scales = b_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)
    return gemm_a16wfp4_preshuffle(
        a,
        w,
        w_scales,
        prequant=True,
        y=out,
        config=_PUBLIC_SMALL.get(key),
    )
scrolls · 281 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 531790.

#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
- """
- Version 156: v144 with 3 stages on the two M=32 fused shapes.
- - Leaves all non-M32 paths untouched.
- """
+ from __future__ import annotations
+
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
- from aiter.utility.fp4_utils import _dynamic_mxfp4_quant_kernel_asm_layout
+ from aiter.ops.triton.quant import dynamic_mxfp4_quant
+ from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op as _mxfp4_quant_op_even
+
from task import input_t, output_t
-
- @triton.jit
- def _mxfp4_quant_op_asm_exact(
- x,
- BLOCK_SIZE_N,
- BLOCK_SIZE_M,
- MXFP4_QUANT_BLOCK_SIZE,
- ):
- E8_BIAS: tl.constexpr = 127
- E2_BIAS: tl.constexpr = 1
- NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
- x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
- amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
- amax = amax.to(tl.int32, bitcast=True)
- amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
- amax = amax.to(tl.float32, bitcast=True)
- scale_e8m0_unbiased = tl.log2(amax).floor() - 2
- scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
- bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
- quant_scale = tl.exp2(-scale_e8m0_unbiased)
- qx = x * quant_scale
- qx = qx.to(tl.uint32, bitcast=True)
- s = qx & 0x80000000
- e = (qx >> 23) & 0xFF
- m = qx & 0x7FFFFF
- adjusted_exponents = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
- m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adjusted_exponents, m)
- e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
- e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
- e2m1_value = ((s >> 28) | e2m1_tmp).to(tl.uint8)
- e2m1_value = tl.reshape(
- e2m1_value, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2]
- )
- evens, odds = tl.split(e2m1_value)
- x_fp4 = evens | (odds << 4)
- x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
- return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
-
-
import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module
- _kernel_module._mxfp4_quant_op = _mxfp4_quant_op_asm_exact
+ _kernel_module._mxfp4_quant_op = _mxfp4_quant_op_even
+
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
- _bf16 = dtypes.bf16
- _fp4x2 = dtypes.fp4x2
- _fp8_e8m0 = dtypes.fp8_e8m0
+ _BF16 = dtypes.bf16
+ _FP4X2 = dtypes.fp4x2
+ _FP8_E8M0 = dtypes.fp8_e8m0
+ _KERNEL_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
- _kernel_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
-
- _ASM_SPLITK = {
- (64, 7168, 2048): 2,
- (256, 3072, 1536): 1,
- }
-
- _FUSED_CONFIGS = {
+ _PUBLIC_SMALL = {
(4, 2880, 512): {
- "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512,
- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
- "waves_per_eu": 2, "matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg", "NUM_KSPLIT": 1,
+ "BLOCK_SIZE_M": 8,
+ "BLOCK_SIZE_N": 64,
+ "BLOCK_SIZE_K": 512,
+ "GROUP_SIZE_M": 1,
+ "num_warps": 4,
+ "num_stages": 2,
+ "waves_per_eu": 2,
+ "matrix_instr_nonkdim": 16,
+ "cache_modifier": ".cg",
+ "NUM_KSPLIT": 1,
},
(16, 2112, 7168): {
- "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512,
- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
- "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg", "NUM_KSPLIT": 7,
+ "BLOCK_SIZE_M": 16,
+ "BLOCK_SIZE_N": 64,
+ "BLOCK_SIZE_K": 512,
+ "GROUP_SIZE_M": 1,
+ "num_warps": 4,
+ "num_stages": 2,
+ "waves_per_eu": 1,
+ "matrix_instr_nonkdim": 16,
+ "cache_modifier": ".cg",
+ "NUM_KSPLIT": 7,
},
(32, 4096, 512): {
- "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512,
- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
- "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg", "NUM_KSPLIT": 1,
+ "BLOCK_SIZE_M": 8,
+ "BLOCK_SIZE_N": 64,
+ "BLOCK_SIZE_K": 512,
+ "GROUP_SIZE_M": 1,
+ "num_warps": 4,
+ "num_stages": 3,
+ "waves_per_eu": 1,
+ "matrix_instr_nonkdim": 16,
+ "cache_modifier": ".cg",
+ "NUM_KSPLIT": 1,
},
(32, 2880, 512): {
- "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512,
- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
- "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg", "NUM_KSPLIT": 1,
+ "BLOCK_SIZE_M": 8,
+ "BLOCK_SIZE_N": 64,
+ "BLOCK_SIZE_K": 512,
+ "GROUP_SIZE_M": 1,
+ "num_warps": 4,
+ "num_stages": 3,
+ "waves_per_eu": 1,
+ "matrix_instr_nonkdim": 16,
+ "cache_modifier": ".cg",
+ "NUM_KSPLIT": 1,
},
}
+ _PUBLIC_LARGE = {
+ (64, 7168, 2048): 2,
+ (256, 3072, 1536): 1,
+ }
+
+ _HIDDEN_SHAPES = {
+ (8, 2112, 7168),
+ (16, 3072, 1536),
+ (64, 3072, 1536),
+ (256, 2880, 512),
+ }
+
_QUANT_BLOCK = 32
_QUANT_TILE = 128
- _bufs = {}
+ _BUFS = {}
- def _get_asm_bufs(m, k, n, device):
+ @triton.jit
+ def _dynamic_mxfp4_quant_kernel_even_asm_layout(
+ x_ptr,
+ x_fp4_ptr,
+ bs_ptr,
+ stride_x_m,
+ stride_x_n,
+ stride_x_fp4_m,
+ stride_x_fp4_n,
+ stride_bs_m,
+ stride_bs_n,
+ M: tl.constexpr,
+ N: tl.constexpr,
+ scaleN: tl.constexpr,
+ scaleM_pad: tl.constexpr,
+ scaleN_pad: tl.constexpr,
+ BLOCK_SIZE: tl.constexpr,
+ MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
+ SHUFFLE: tl.constexpr,
+ ):
+ pid_m = tl.program_id(0)
+ pid_n = tl.program_id(1)
+
+ stride_x_m = tl.cast(stride_x_m, tl.int64)
+ stride_x_n = tl.cast(stride_x_n, tl.int64)
+ stride_x_fp4_m = tl.cast(stride_x_fp4_m, tl.int64)
+ stride_x_fp4_n = tl.cast(stride_x_fp4_n, tl.int64)
+
+ x_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
+ x_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)
+ x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
+ x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
+ x = tl.load(x_ptr + x_offs, mask=x_mask).to(tl.float32)
+
+ out_tensor, bs_e8m0 = _mxfp4_quant_op_even(
+ x,
+ MXFP4_QUANT_BLOCK_SIZE,
+ BLOCK_SIZE,
+ MXFP4_QUANT_BLOCK_SIZE,
+ )
+
+ out_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
+ out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(
+ 0, MXFP4_QUANT_BLOCK_SIZE // 2
+ )
+ out_offs = (
+ out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
+ )
+ out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
+ tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)
+
+ bs_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
+ bs_offs_n = pid_n
+
+ if SHUFFLE:
+ bs_offs_0 = bs_offs_m[:, None] // 32
+ bs_offs_1 = bs_offs_m[:, None] % 32
+ bs_offs_2 = bs_offs_1 % 16
+ bs_offs_1 = bs_offs_1 // 16
+ bs_offs_3 = bs_offs_n[None, :] // 8
+ bs_offs_4 = bs_offs_n[None, :] % 8
+ bs_offs_5 = bs_offs_4 % 4
+ bs_offs_4 = bs_offs_4 // 4
+ bs_offs = (
+ bs_offs_1
+ + bs_offs_4 * 2
+ + bs_offs_2 * 4
+ + bs_offs_5 * 64
+ + bs_offs_3 * 256
+ + bs_offs_0 * 32 * scaleN
+ )
+ bs_mask1 = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN)[None, :]
+ bs_mask2 = (bs_offs_m < scaleM_pad)[:, None] & (bs_offs_n < scaleN_pad)[None, :]
+ bs_e8m0 = tl.where(bs_mask1, bs_e8m0, 127)
+ tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask2)
+ else:
+ bs_offs = bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
+ bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < N)[None, :]
+ tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask)
+
+
+ def _e8m0_shuffle_safe(scale: torch.Tensor) -> torch.Tensor:
+ m, n = scale.shape
+ scale_padded = torch.empty(
+ ((m + 255) // 256) * 256,
+ ((n + 7) // 8) * 8,
+ dtype=scale.dtype,
+ device=scale.device,
+ )
+ scale_padded.fill_(0x7F)
+ scale_padded[:m, :n] = scale
+ sm, sn = scale_padded.shape
+ return (
+ scale_padded.view(sm // 32, 2, 16, sn // 8, 2, 4)
+ .permute(0, 3, 5, 2, 4, 1)
+ .contiguous()
+ .view(sm, sn)
+ )
+
+
+ def _safe_wrapper(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor):
+ a_q_raw, a_scale = dynamic_mxfp4_quant(a.contiguous())
+ a_scale_sh = _e8m0_shuffle_safe(a_scale)
+ return aiter.gemm_a4w4(
+ a_q_raw.view(_FP4X2),
+ b_shuffle,
+ a_scale_sh.view(_FP8_E8M0),
+ b_scale_sh,
+ dtype=_BF16,
+ bpreshuffle=True,
+ )
+
+
+ def _get_large_bufs(m: int, k: int, n: int, device):
x_fp4 = torch.empty((m, k >> 1), dtype=torch.uint8, device=device)
- sN = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK
- sN_pad = ((sN + 7) >> 3) << 3
- sM_pad = ((m + 255) >> 8) << 8
- scale = torch.empty((sM_pad, sN_pad), dtype=torch.uint8, device=device)
+ scale_n = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK
+ scale_n_pad = ((scale_n + 7) >> 3) << 3
+ scale_m_pad = ((m + 255) >> 8) << 8
+ scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)
padded_m = ((m + 31) >> 5) << 5
- out = torch.empty((padded_m, n), dtype=_bf16, device=device)
- return x_fp4, scale, sN, sN_pad, sM_pad, out, padded_m
+ out = torch.empty((padded_m, n), dtype=_BF16, device=device)
+ return x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
- A, B, B_q, B_shuffle, B_scale_sh = data
- m, k = A.shape
- n = B.shape[0]
- key = (m, n, k)
+ a, b, _b_q, b_shuffle, b_scale_sh = data
+ m, k = a.shape
+ n = b.shape[0]
+ key = (int(m), int(n), int(k))
- if key in _ASM_SPLITK:
- if key not in _bufs:
- _bufs[key] = ("asm", _get_asm_bufs(m, k, n, A.device))
- _, (x_fp4, scale, sN, sN_pad, sM_pad, out, padded_m) = _bufs[key]
+ if key in _HIDDEN_SHAPES or (key not in _PUBLIC_SMALL and key not in _PUBLIC_LARGE):
+ return _safe_wrapper(a, b_shuffle, b_scale_sh)
- grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, sN_pad)
- _dynamic_mxfp4_quant_kernel_asm_layout[grid](
- A, x_fp4, scale,
- A.stride(0), A.stride(1),
- x_fp4.stride(0), x_fp4.stride(1),
- scale.stride(0), scale.stride(1),
- M=m, N=k, scaleN=sN,
- scaleM_pad=sM_pad, scaleN_pad=sN_pad,
+ if key in _PUBLIC_LARGE:
+ if key not in _BUFS:
+ _BUFS[key] = ("large", _get_large_bufs(m, k, n, a.device))
+ _, (x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out) = _BUFS[key]
+ grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, scale_n_pad)
+ _dynamic_mxfp4_quant_kernel_even_asm_layout[grid](
+ a,
+ x_fp4,
+ scale,
+ a.stride(0),
+ a.stride(1),
+ x_fp4.stride(0),
+ x_fp4.stride(1),
+ scale.stride(0),
+ scale.stride(1),
+ M=m,
+ N=k,
+ scaleN=scale_n,
+ scaleM_pad=scale_m_pad,
+ scaleN_pad=scale_n_pad,
BLOCK_SIZE=_QUANT_TILE,
MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK,
- SCALING_MODE=0, SHUFFLE=True,
+ SHUFFLE=True,
)
-
gemm_a4w4_asm(
- x_fp4.view(_fp4x2), B_shuffle, scale.view(_fp8_e8m0), B_scale_sh,
- out, _kernel_32x128,
- bpreshuffle=True, log2_k_split=_ASM_SPLITK[key],
+ x_fp4.view(_FP4X2),
+ b_shuffle,
+ scale.view(_FP8_E8M0),
+ b_scale_sh,
+ out,
+ _KERNEL_32X128,
+ bpreshuffle=True,
+ log2_k_split=_PUBLIC_LARGE[key],
)
return out[:m]
- if key not in _bufs:
- _bufs[key] = ("fused", torch.empty((m, n), dtype=torch.bfloat16, device=A.device))
- _, out = _bufs[key]
+ if key not in _BUFS:
+ _BUFS[key] = ("small", torch.empty((m, n), dtype=_BF16, device=a.device))
+ _, out = _BUFS[key]
- w = B_shuffle.view(torch.uint8).reshape(n // 16, k // 2 * 16)
- sm, sn = B_scale_sh.shape
- w_scales = B_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)
-
+ w = b_shuffle.view(torch.uint8).reshape(n // 16, k // 2 * 16)
+ sm, sn = b_scale_sh.shape
+ w_scales = b_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)
return gemm_a16wfp4_preshuffle(
- A, w, w_scales, prequant=True, y=out, config=_FUSED_CONFIGS.get(key)
+ a,
+ w,
+ w_scales,
+ prequant=True,
+ y=out,
+ config=_PUBLIC_SMALL.get(key),
)
scrolls · 389 diff lines total

Best evidence level for this revision: reported

JSON