Skip to content
KernelIndex
Search⌘K

submission 754826

tao_yafan · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754826?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
12.6µs
#388 of 1143
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f1c00a6c4ac46f66e74a4cdb46d90d645c6257ae674c026808ad44e2b3d65d1a
license declaredunknown
license concludedunknown
authorstao_yafan
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4Fully-fused single Triton kernel: BF16 A ___ on-the-fly MXFP4 quant ___ FP4 GEMM.
split-k- Split-K for small M improves GPU utilisation.
tile-k = 128BLOCK_SIZE_K = 128
tile-m = 32BLOCK_SIZE_M = 32

Kernel source

submission_v3.py578 lines
"""
submission_v3.py

Fully-fused single Triton kernel: BF16 A ___ on-the-fly MXFP4 quant ___ FP4 GEMM.

Pipeline comparison:
  Original: [dynamic_mxfp4_quant] ___ [e8m0_shuffle] ___ [gemm_a4w4]   (3 kernels)
  v2:       [fused_quant_shuffle]                   ___ [gemm_a4w4]   (2 kernels)
  v3:       [fused_bf16_fp4gemm_kernel]                               (1 kernel)

HBM savings over original:
  - A_q    : never stored  (M__K/2  bytes, e.g. 128 KB for M=64, K=4096)
  - A_scale: never stored  (M__K/32 bytes, e.g.   8 KB)
  - A is read once (unchanged)

Design notes:
  - Uses Triton gemm_afp4wfp4 inner loop pattern with tl.dot_scaled (AMD MFMA).
  - B_q [N, K//2] and B_scale_sh (shuffled e8m0) are loaded directly from
    the original input tensors ___ no transpose or unshuffle preprocessing.
  - B_q is loaded as [BN, BK/2] (coalesced) then tl.trans'd to [BK/2, BN].
  - Shuffled scale indices are computed in-kernel via integer arithmetic,
    matching the e8m0_shuffle permutation pattern.
  - Split-K for small M improves GPU utilisation.
"""

import torch
import triton
import triton.language as tl
from task import input_t, output_t
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op


# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________
# Case 1: M=4, N=2880, K=512, BM=32, BN=64, BK=128, sk=4
#   k_tiles=4, k_tiles_per_split=1 (no loop), num_pid_n=45, grid=(45, 4)
#   M<BM ___ M-mask only; N,K tile-aligned ___ no other masks
#   stride_am=512, stride_bn=256, stride_pp_m=2880, scale_n_pad=16
# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________

@triton.jit
def _splitk_case1(a_ptr, b_ptr, y_pp_ptr, stride_pp_s, bs_ptr):
    pid_n = tl.program_id(0)
    tile_idx = tl.program_id(1)  # = pid_split, each split handles 1 tile

    offs_m = tl.arange(0, 32)
    offs_n = pid_n * 64 + tl.arange(0, 64)
    offs_k = tl.arange(0, 128)
    offs_k2 = tl.arange(0, 64)
    offs_ks = tl.arange(0, 4)

    a_ptrs = a_ptr + offs_m[:, None] * 512 + offs_k[None, :] + tile_idx * 128
    b_ptrs = b_ptr + offs_n[:, None] * 256 + offs_k2[None, :] + tile_idx * 64

    a_bf16 = tl.load(a_ptrs, mask=offs_m[:, None] < 4, other=0.0, cache_modifier=".cg")
    a_fp4, a_scales = _mxfp4_quant_op(a_bf16.to(tl.float32), 128, 32, 32)
    b = tl.trans(tl.load(b_ptrs, cache_modifier=".cg"))

    s_d0 = offs_n // 32
    s_d1 = (offs_n % 32) // 16
    s_d2 = offs_n % 16
    bs_n_part = s_d0 * 512 + s_d2 * 4 + s_d1  # 32*16=512
    bs_k_part = (tile_idx >> 1) * 256 + offs_ks * 64 + (tile_idx & 1) * 2
    b_scales = tl.load(bs_ptr + bs_n_part[:, None] + bs_k_part[None, :])

    acc = tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1")

    c_ptrs = y_pp_ptr + tile_idx * stride_pp_s + offs_m[:, None] * 2880 + offs_n[None, :]
    tl.store(c_ptrs, acc, mask=offs_m[:, None] < 4)


@triton.jit
def _reduce_case1(y_pp_ptr, stride_pp_s, c_ptr):
    pid_n = tl.program_id(0)
    offs_m = tl.arange(0, 32)
    offs_n = pid_n * 64 + tl.arange(0, 64)
    mask = offs_m[:, None] < 4

    acc = tl.zeros((32, 64), dtype=tl.float32)
    for k in range(4):
        acc += tl.load(y_pp_ptr + k * stride_pp_s + offs_m[:, None] * 2880 + offs_n[None, :],
                       mask=mask, other=0.0)

    tl.store(c_ptr + offs_m[:, None] * 2880 + offs_n[None, :], acc.to(tl.bfloat16), mask=mask)


# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________
# Case 2: M=16, N=2112, K=7168, BM=32, BN=128, BK=128, sk=14
#   k_tiles=56, k_tiles_per_split=4, num_pid_n=17, grid=(17, 14)
#   M<BM ___ M-mask; N%BN=64 ___ last-N-tile needs N-mask; K tile-aligned
#   stride_am=7168, stride_bn=3584, stride_pp_m=2112, scale_n_pad=224
# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________

@triton.jit
def _splitk_case2(a_ptr, b_ptr, y_pp_ptr, stride_pp_s, bs_ptr):
    K_TILES_PER_SPLIT: tl.constexpr = 4

    pid_n = tl.program_id(0)
    pid_split = tl.program_id(1)

    offs_m = tl.arange(0, 32)
    offs_n = pid_n * 128 + tl.arange(0, 128)
    offs_k = tl.arange(0, 128)
    offs_k2 = tl.arange(0, 64)
    offs_ks = tl.arange(0, 4)
    m_mask = offs_m[:, None] < 16
    n_mask = offs_n[:, None] < 2112  # only last tile (pid_n=16) partial

    start_tile = pid_split * K_TILES_PER_SPLIT

    a_ptrs = a_ptr + offs_m[:, None] * 7168 + offs_k[None, :] + start_tile * 128
    b_ptrs = b_ptr + offs_n[:, None] * 3584 + offs_k2[None, :] + start_tile * 64

    s_d0 = offs_n // 32
    s_d1 = (offs_n % 32) // 16
    s_d2 = offs_n % 16
    bs_n_part = s_d0 * 7168 + s_d2 * 4 + s_d1  # 32*224=7168

    acc = tl.zeros((32, 128), dtype=tl.float32)

    for i in range(K_TILES_PER_SPLIT):
        tile_idx = start_tile + i

        a_bf16 = tl.load(a_ptrs, mask=m_mask, other=0.0, cache_modifier=".cg")
        a_fp4, a_scales = _mxfp4_quant_op(a_bf16.to(tl.float32), 128, 32, 32)

        b = tl.trans(tl.load(b_ptrs, mask=n_mask, other=0, cache_modifier=".cg"))

        bs_k_part = (tile_idx >> 1) * 256 + offs_ks * 64 + (tile_idx & 1) * 2
        b_scales = tl.load(bs_ptr + bs_n_part[:, None] + bs_k_part[None, :],
                           mask=offs_n[:, None] < 2112, other=0)

        acc = tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1", acc)

        a_ptrs += 128
        b_ptrs += 64

    c_ptrs = y_pp_ptr + pid_split * stride_pp_s + offs_m[:, None] * 2112 + offs_n[None, :]
    tl.store(c_ptrs, acc, mask=(offs_m[:, None] < 16) & (offs_n[None, :] < 2112))


@triton.jit
def _reduce_case2(y_pp_ptr, stride_pp_s, c_ptr):
    pid_n = tl.program_id(0)
    offs_m = tl.arange(0, 32)
    offs_n = pid_n * 128 + tl.arange(0, 128)
    mask = (offs_m[:, None] < 16) & (offs_n[None, :] < 2112)

    acc = tl.zeros((32, 128), dtype=tl.float32)
    for k in range(14):
        acc += tl.load(y_pp_ptr + k * stride_pp_s + offs_m[:, None] * 2112 + offs_n[None, :],
                       mask=mask, other=0.0)

    tl.store(c_ptr + offs_m[:, None] * 2112 + offs_n[None, :], acc.to(tl.bfloat16), mask=mask)


# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________
# Case 3: M=32, N=4096, K=512, BM=32, BN=64, BK=128, sk=4
#   k_tiles=4, k_tiles_per_split=1 (no loop), num_pid_n=64, grid=(64, 4)
#   All dims tile-aligned ___ NO masks
#   stride_am=512, stride_bn=256, stride_pp_m=4096, scale_n_pad=16
# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________

@triton.jit
def _splitk_case3(a_ptr, b_ptr, y_pp_ptr, stride_pp_s, bs_ptr):
    pid_n = tl.program_id(0)
    tile_idx = tl.program_id(1)

    offs_m = tl.arange(0, 32)
    offs_n = pid_n * 64 + tl.arange(0, 64)
    offs_k = tl.arange(0, 128)
    offs_k2 = tl.arange(0, 64)
    offs_ks = tl.arange(0, 4)

    a_ptrs = a_ptr + offs_m[:, None] * 512 + offs_k[None, :] + tile_idx * 128
    b_ptrs = b_ptr + offs_n[:, None] * 256 + offs_k2[None, :] + tile_idx * 64

    a_bf16 = tl.load(a_ptrs, cache_modifier=".cg")
    a_fp4, a_scales = _mxfp4_quant_op(a_bf16.to(tl.float32), 128, 32, 32)
    b = tl.trans(tl.load(b_ptrs, cache_modifier=".cg"))

    s_d0 = offs_n // 32
    s_d1 = (offs_n % 32) // 16
    s_d2 = offs_n % 16
    bs_n_part = s_d0 * 512 + s_d2 * 4 + s_d1
    bs_k_part = (tile_idx >> 1) * 256 + offs_ks * 64 + (tile_idx & 1) * 2
    b_scales = tl.load(bs_ptr + bs_n_part[:, None] + bs_k_part[None, :])

    acc = tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1")

    tl.store(y_pp_ptr + tile_idx * stride_pp_s + offs_m[:, None] * 4096 + offs_n[None, :], acc)


@triton.jit
def _reduce_case3(y_pp_ptr, stride_pp_s, c_ptr):
    pid_n = tl.program_id(0)
    offs_m = tl.arange(0, 32)
    offs_n = pid_n * 64 + tl.arange(0, 64)

    acc = tl.zeros((32, 64), dtype=tl.float32)
    for k in range(4):
        acc += tl.load(y_pp_ptr + k * stride_pp_s + offs_m[:, None] * 4096 + offs_n[None, :])

    tl.store(c_ptr + offs_m[:, None] * 4096 + offs_n[None, :], acc.to(tl.bfloat16))


# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________
# Case 4: M=32, N=2880, K=512, BM=32, BN=64, BK=128, sk=4
#   k_tiles=4, k_tiles_per_split=1 (no loop), num_pid_n=45, grid=(45, 4)
#   All dims tile-aligned ___ NO masks
#   stride_am=512, stride_bn=256, stride_pp_m=2880, scale_n_pad=16
# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________

@triton.jit
def _splitk_case4(a_ptr, b_ptr, y_pp_ptr, stride_pp_s, bs_ptr):
    pid_n = tl.program_id(0)
    tile_idx = tl.program_id(1)

    offs_m = tl.arange(0, 32)
    offs_n = pid_n * 64 + tl.arange(0, 64)
    offs_k = tl.arange(0, 128)
    offs_k2 = tl.arange(0, 64)
    offs_ks = tl.arange(0, 4)

    a_ptrs = a_ptr + offs_m[:, None] * 512 + offs_k[None, :] + tile_idx * 128
    b_ptrs = b_ptr + offs_n[:, None] * 256 + offs_k2[None, :] + tile_idx * 64

    a_bf16 = tl.load(a_ptrs, cache_modifier=".cg")
    a_fp4, a_scales = _mxfp4_quant_op(a_bf16.to(tl.float32), 128, 32, 32)
    b = tl.trans(tl.load(b_ptrs, cache_modifier=".cg"))

    s_d0 = offs_n // 32
    s_d1 = (offs_n % 32) // 16
    s_d2 = offs_n % 16
    bs_n_part = s_d0 * 512 + s_d2 * 4 + s_d1
    bs_k_part = (tile_idx >> 1) * 256 + offs_ks * 64 + (tile_idx & 1) * 2
    b_scales = tl.load(bs_ptr + bs_n_part[:, None] + bs_k_part[None, :])

    acc = tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1")

    tl.store(y_pp_ptr + tile_idx * stride_pp_s + offs_m[:, None] * 2880 + offs_n[None, :], acc)


@triton.jit
def _reduce_case4(y_pp_ptr, stride_pp_s, c_ptr):
    pid_n = tl.program_id(0)
    offs_m = tl.arange(0, 32)
    offs_n = pid_n * 64 + tl.arange(0, 64)

    acc = tl.zeros((32, 64), dtype=tl.float32)
    for k in range(4):
        acc += tl.load(y_pp_ptr + k * stride_pp_s + offs_m[:, None] * 2880 + offs_n[None, :])

    tl.store(c_ptr + offs_m[:, None] * 2880 + offs_n[None, :], acc.to(tl.bfloat16))


# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________
# Case 5: M=64, N=7168, K=2048, BM=32, BN=128, BK=128, sk=2
#   k_tiles=16, k_tiles_per_split=8, num_pid_m=2, num_pid_n=56, grid=(112, 2)
#   All dims tile-aligned ___ NO masks
#   stride_am=2048, stride_bn=1024, stride_pp_m=7168, scale_n_pad=64
# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________

@triton.jit
def _splitk_case5(a_ptr, b_ptr, y_pp_ptr, stride_pp_s, bs_ptr):
    K_TILES_PER_SPLIT: tl.constexpr = 8
    NUM_PID_N: tl.constexpr = 56

    pid = tl.program_id(0)
    pid_split = tl.program_id(1)
    pid_m = pid // NUM_PID_N
    pid_n = pid % NUM_PID_N

    offs_m = pid_m * 32 + tl.arange(0, 32)
    offs_n = pid_n * 128 + tl.arange(0, 128)
    offs_k = tl.arange(0, 128)
    offs_k2 = tl.arange(0, 64)
    offs_ks = tl.arange(0, 4)

    start_tile = pid_split * K_TILES_PER_SPLIT
    a_ptrs = a_ptr + offs_m[:, None] * 2048 + offs_k[None, :] + start_tile * 128
    b_ptrs = b_ptr + offs_n[:, None] * 1024 + offs_k2[None, :] + start_tile * 64

    s_d0 = offs_n // 32
    s_d1 = (offs_n % 32) // 16
    s_d2 = offs_n % 16
    bs_n_part = s_d0 * 2048 + s_d2 * 4 + s_d1  # 32*64=2048

    acc = tl.zeros((32, 128), dtype=tl.float32)

    for i in range(K_TILES_PER_SPLIT):
        tile_idx = start_tile + i

        a_bf16 = tl.load(a_ptrs, cache_modifier=".cg")
        a_fp4, a_scales = _mxfp4_quant_op(a_bf16.to(tl.float32), 128, 32, 32)
        b = tl.trans(tl.load(b_ptrs, cache_modifier=".cg"))

        bs_k_part = (tile_idx >> 1) * 256 + offs_ks * 64 + (tile_idx & 1) * 2
        b_scales = tl.load(bs_ptr + bs_n_part[:, None] + bs_k_part[None, :])

        acc = tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1", acc)

        a_ptrs += 128
        b_ptrs += 64

    tl.store(y_pp_ptr + pid_split * stride_pp_s + offs_m[:, None] * 7168 + offs_n[None, :], acc)


@triton.jit
def _reduce_case5(y_pp_ptr, stride_pp_s, c_ptr):
    NUM_PID_N: tl.constexpr = 56

    pid = tl.program_id(0)
    pid_m = pid // NUM_PID_N
    pid_n = pid % NUM_PID_N

    offs_m = pid_m * 32 + tl.arange(0, 32)
    offs_n = pid_n * 128 + tl.arange(0, 128)

    acc = tl.zeros((32, 128), dtype=tl.float32)
    for k in range(2):
        acc += tl.load(y_pp_ptr + k * stride_pp_s + offs_m[:, None] * 7168 + offs_n[None, :])

    tl.store(c_ptr + offs_m[:, None] * 7168 + offs_n[None, :], acc.to(tl.bfloat16))


# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________
# Case 6: M=256, N=3072, K=1536, BM=32, BN=128, BK=128, sk=1
#   k_iters=12, num_pid_m=8, num_pid_n=24, grid=192
#   All dims tile-aligned ___ NO masks
#   stride_am=1536, stride_bn=768, stride_cm=3072, scale_n_pad=48
#   XCD-aware pid mapping (8 XCDs)
# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________

@triton.jit
def _kernel_case6(a_ptr, b_ptr, c_ptr, bs_ptr):
    NUM_K_ITERS: tl.constexpr = 12
    NUM_XCDS: tl.constexpr = 8

    pid = tl.program_id(0)
    xcd = pid % NUM_XCDS
    local_id = pid // NUM_XCDS
    pid_m = (xcd // 4) * 4 + (local_id % 4)
    pid_n = (xcd % 4) * 6 + (local_id // 4)

    offs_m = pid_m * 32 + tl.arange(0, 32)
    offs_n = pid_n * 128 + tl.arange(0, 128)
    offs_k = tl.arange(0, 128)
    offs_k2 = tl.arange(0, 64)
    offs_ks = tl.arange(0, 4)

    a_ptrs = a_ptr + offs_m[:, None] * 1536 + offs_k[None, :]
    b_ptrs = b_ptr + offs_n[:, None] * 768 + offs_k2[None, :]

    s_d0 = offs_n // 32
    s_d1 = (offs_n % 32) // 16
    s_d2 = offs_n % 16
    bs_n_part = s_d0 * 1536 + s_d2 * 4 + s_d1  # 32*48=1536

    acc = tl.zeros((32, 128), dtype=tl.float32)

    for k in range(NUM_K_ITERS):
        a_bf16 = tl.load(a_ptrs, cache_modifier=".cg")
        a_fp4, a_scales = _mxfp4_quant_op(a_bf16.to(tl.float32), 128, 32, 32)
        b = tl.trans(tl.load(b_ptrs, cache_modifier=".cg"))

        bs_k_part = (k >> 1) * 256 + offs_ks * 64 + (k & 1) * 2
        b_scales = tl.load(bs_ptr + bs_n_part[:, None] + bs_k_part[None, :])

        acc = tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1", acc)

        a_ptrs += 128
        b_ptrs += 64

    tl.store(c_ptr + offs_m[:, None] * 3072 + offs_n[None, :], acc.to(tl.bfloat16))


# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________
# Fallback generic kernel (for unknown problem sizes)
# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________

@triton.jit
def _fallback_splitk_kernel(
    a_ptr, stride_am,
    b_ptr, stride_bn,
    y_pp_ptr, stride_pp_s, stride_pp_m,
    bs_ptr, scale_n_pad,
    M, N, K,
    BLOCK_SIZE_N: tl.constexpr,
    SPLIT_K: tl.constexpr,
):
    pid = tl.program_id(0)
    pid_split = tl.program_id(1)
    num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
    pid_m = pid // num_pid_n
    pid_n = pid % num_pid_n

    offs_m = pid_m * 32 + tl.arange(0, 32)
    offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
    offs_k = tl.arange(0, 128)
    offs_k2 = tl.arange(0, 64)
    offs_ks = tl.arange(0, 4)

    total_k_tiles = tl.cdiv(K, 128)
    k_tiles_per_split = tl.cdiv(total_k_tiles, SPLIT_K)
    start_tile = pid_split * k_tiles_per_split
    end_tile = tl.minimum(start_tile + k_tiles_per_split, total_k_tiles)

    a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k[None, :] + start_tile * 128
    b_ptrs = b_ptr + offs_n[:, None] * stride_bn + offs_k2[None, :] + start_tile * 64

    s_d0 = offs_n // 32
    s_d1 = (offs_n % 32) // 16
    s_d2 = offs_n % 16
    bs_n_part = s_d0 * (32 * scale_n_pad) + s_d2 * 4 + s_d1

    acc = tl.zeros((32, BLOCK_SIZE_N), dtype=tl.float32)

    for tile_idx in range(start_tile, end_tile):
        tile_k_start = tile_idx * 128

        a_mask = (offs_m[:, None] < M) & (offs_k[None, :] < K - tile_k_start)
        a_bf16 = tl.load(a_ptrs, mask=a_mask, other=0.0, cache_modifier=".cg")
        a_fp4, a_scales = _mxfp4_quant_op(a_bf16.to(tl.float32), 128, 32, 32)

        b_mask = (offs_n[:, None] < N) & (offs_k2[None, :] < (K - tile_k_start) // 2)
        b = tl.trans(tl.load(b_ptrs, mask=b_mask, other=0, cache_modifier=".cg"))

        bs_k_part = (tile_idx >> 1) * 256 + offs_ks * 64 + (tile_idx & 1) * 2
        bs_idx = bs_n_part[:, None] + bs_k_part[None, :]
        kg = tile_idx * 4 + offs_ks
        b_scale_mask = (offs_n[:, None] < N) & (kg[None, :] < K // 32)
        b_scales = tl.load(bs_ptr + bs_idx, mask=b_scale_mask, other=0)

        acc = tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1", acc)

        a_ptrs += 128
        b_ptrs += 64

    c_ptrs = y_pp_ptr + pid_split * stride_pp_s + offs_m[:, None] * stride_pp_m + offs_n[None, :]
    tl.store(c_ptrs, acc, mask=(offs_m[:, None] < M) & (offs_n[None, :] < N))


@triton.jit
def _fallback_reduce_kernel(
    y_pp_ptr, stride_pp_s, stride_pp_m,
    c_ptr, stride_cm,
    M, N,
    SPLIT_K: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
):
    pid = tl.program_id(0)
    num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
    pid_m = pid // num_pid_n
    pid_n = pid % num_pid_n

    offs_m = pid_m * 32 + tl.arange(0, 32)
    offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
    mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)

    acc = tl.zeros((32, BLOCK_SIZE_N), dtype=tl.float32)
    for k in range(SPLIT_K):
        acc += tl.load(y_pp_ptr + k * stride_pp_s + offs_m[:, None] * stride_pp_m + offs_n[None, :],
                       mask=mask, other=0.0)

    tl.store(c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :], acc.to(tl.bfloat16), mask=mask)


# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________
# Dispatch
# _____________________________________________________________________________________________________________________________________________________________________________________________________________________________________________

def _fused_bf16_fp4gemm(
    A: torch.Tensor,          # BF16 [M, K]
    B_q: torch.Tensor,        # uint8 [N, K//2]  (FP4 packed, original layout)
    B_scale_sh: torch.Tensor, # uint8 [scaleM_pad, scaleN_pad] (shuffled e8m0)
) -> torch.Tensor:
    M, K = A.shape
    N    = B_q.shape[0]
    scale_n_pad = B_scale_sh.shape[1]
    assert K % 32 == 0, "K must be divisible by 32 for MXFP4"

    C = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
    key = (M, N, K)
    # -- Case 1: M=4, N=2880, K=512, sk=4 --
    if key == (4, 2880, 512):
        y_pp = torch.empty((4, M, N), dtype=torch.float32, device=A.device)
        _splitk_case1[(45, 4)](A, B_q, y_pp, y_pp.stride(0), B_scale_sh)
        _reduce_case1[(45,)](y_pp, y_pp.stride(0), C)
        return C

    # -- Case 2: M=16, N=2112, K=7168, sk=14 --
    if key == (16, 2112, 7168):
        y_pp = torch.empty((14, M, N), dtype=torch.float32, device=A.device)
        _splitk_case2[(17, 14)](A, B_q, y_pp, y_pp.stride(0), B_scale_sh)
        _reduce_case2[(17,)](y_pp, y_pp.stride(0), C)
        return C

    # -- Case 3: M=32, N=4096, K=512, sk=4 --
    if key == (32, 4096, 512):
        y_pp = torch.empty((4, M, N), dtype=torch.float32, device=A.device)
        _splitk_case3[(64, 4)](A, B_q, y_pp, y_pp.stride(0), B_scale_sh)
        _reduce_case3[(64,)](y_pp, y_pp.stride(0), C)
        return C

    # -- Case 4: M=32, N=2880, K=512, sk=4 --
    if key == (32, 2880, 512):
        y_pp = torch.empty((4, M, N), dtype=torch.float32, device=A.device)
        _splitk_case4[(45, 4)](A, B_q, y_pp, y_pp.stride(0), B_scale_sh)
        _reduce_case4[(45,)](y_pp, y_pp.stride(0), C)
        return C

    # -- Case 5: M=64, N=7168, K=2048, sk=2 --
    if key == (64, 7168, 2048):
        y_pp = torch.empty((2, M, N), dtype=torch.float32, device=A.device)
        _splitk_case5[(112, 2)](A, B_q, y_pp, y_pp.stride(0), B_scale_sh)
        _reduce_case5[(112,)](y_pp, y_pp.stride(0), C)
        return C

    # -- Case 6: M=256, N=3072, K=1536, sk=1 --
    if key == (256, 3072, 1536):
        _kernel_case6[(192,)](A, B_q, C, B_scale_sh)
        return C

    # -- Fallback: generic kernel for unknown sizes --
    BLOCK_SIZE_M = 32
    BLOCK_SIZE_K = 128
    NUM_CUS = 256
    k_tiles = triton.cdiv(K, BLOCK_SIZE_K)
    best_bn = 128
    best_split_k = 1
    found_max_split = False
    for bn in (128, 64, 32):
        mn = triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, bn)
        sk = max(1, min(k_tiles, NUM_CUS // max(mn, 1)))
        if sk == k_tiles:
            best_bn = bn
            best_split_k = sk
            found_max_split = True
            continue
        if found_max_split:
            break
        best_bn = bn
        best_split_k = sk
        break
    BLOCK_SIZE_N = best_bn
    split_k = best_split_k

    mn_tiles = triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)
    y_pp = torch.empty((split_k, M, N), dtype=torch.float32, device=A.device)

    _fallback_splitk_kernel[(mn_tiles, split_k)](
        A, A.stride(0),
        B_q, B_q.stride(0),
        y_pp, y_pp.stride(0), y_pp.stride(1),
        B_scale_sh, scale_n_pad,
        M, N, K,
        BLOCK_SIZE_N=BLOCK_SIZE_N,
        SPLIT_K=split_k,
    )
    _fallback_reduce_kernel[(mn_tiles,)](
        y_pp, y_pp.stride(0), y_pp.stride(1),
        C, C.stride(0),
        M, N,
        SPLIT_K=split_k,
        BLOCK_SIZE_N=BLOCK_SIZE_N,
    )
    return C

def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()

    # B_q [N, K//2] and B_scale_sh [scaleM_pad, scaleN_pad] are used directly ___
    # no transpose or unshuffle needed; shuffled scale indices computed in-kernel.
    B_q_u8 = B_q.view(torch.uint8)
    B_scale_sh_u8 = B_scale_sh.view(torch.uint8)
    return _fused_bf16_fp4gemm(A, B_q_u8, B_scale_sh_u8)
scrolls · 578 lines total

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

Best evidence level for this revision: reported

JSON