Skip to content
KernelIndex
Search⌘K

submission 670608

CaymanYang · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-670608?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
13.3µs
#417 of 1143
2026-03-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e5617be8ab35e8d790f2fd1971e0794500e055d5583cc1d37705ec8198defc06
license declaredunknown
license concludedunknown
authorsCaymanYang
imported2026-08-26

Techniques

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

fp4Two-stage FP4 path:
num-warps = 2num_warps = 2 if m <= 8 else (4 if m <= 64 else 8)
stages = 3num_stages = 3 if k >= 2048 else 2

Kernel source

submission.py820 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
Two-stage FP4 path:
1) Quantize bf16 A -> (MXFP4 A_q + E8M0 A_scales).
2) GEMM with tl.dot_scaled using pre-quantized A and shuffled B/B-scales.
"""
from __future__ import annotations

import torch
import triton
import triton.language as tl

from task import input_t, output_t

@triton.jit
def _mxfp4_quant_op(
    x,
    BLOCK_SIZE_N,
    BLOCK_SIZE_M,
    MXFP4_QUANT_BLOCK_SIZE,
):
    EXP_BIAS_FP32: tl.constexpr = 127
    EXP_BIAS_FP4: tl.constexpr = 1
    EBITS_F32: tl.constexpr = 8
    EBITS_FP4: tl.constexpr = 2
    MBITS_F32: tl.constexpr = 23
    MBITS_FP4: tl.constexpr = 1
    max_normal: tl.constexpr = 6
    min_normal: 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)
    '''
    amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
    '''#有精度问题
    # 位操作: 向上舍入到 2 的幂
    amax_bits = amax.to(tl.uint32, bitcast=True) #有精度问题
    # 简化: 直接提取指数 + 1 (向上取整效果)
    amax_exp = ((amax_bits >> 23) & 0xFF).to(tl.int32)
    # 如果尾数非零,指数+1 (向上舍入)
    has_mant = (amax_bits & 0x7FFFFF) != 0
    amax_exp = amax_exp + has_mant.to(tl.int32)
    # E8M0 指数: exp - 2 (让 max 映射到 4.0)
    scale_exp = amax_exp - 2 - 127  # 无偏指数
    '''
    amax = amax.to(tl.int32, bitcast=True)
    amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    amax = amax.to(tl.float32, bitcast=True)
    #scale_exp = tl.log2(amax).floor() - 2
    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
    #scale_exp = scale_exp.to(tl.float32)
    #scale_exp = tl.clamp(scale_exp, -127, 127)
    #bs_e8m0 = (scale_exp + 127).to(tl.uint8)
    # 扩展 scale 到每个元素
    scale = tl.exp2(scale_e8m0_unbiased.to(tl.float32))
    scale = tl.broadcast_to(
        scale, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE]
    )
    
    # ===== 步骤3: 按原始位级流程量化,保证与基准语义一致 =====
    quant_scale = tl.exp2(-scale_e8m0_unbiased)
    qx = x * quant_scale
    qx_bits = qx.to(tl.uint32, bitcast=True)
    s = qx_bits & 0x80000000
    qx_mag = qx_bits ^ s
    qx_fp32 = qx_mag.to(tl.float32, bitcast=True)
    saturate_mask = qx_fp32 >= max_normal
    denormal_mask = (~saturate_mask) & (qx_fp32 < min_normal)
    normal_mask = ~(saturate_mask | denormal_mask)
    denorm_exp: tl.constexpr = (
        (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
    )
    denorm_mask_int: tl.constexpr = denorm_exp << MBITS_F32
    denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)
    denormal_x = qx_fp32 + denorm_mask_float
    denormal_x = denormal_x.to(tl.uint32, bitcast=True)
    denormal_x -= denorm_mask_int
    denormal_x = denormal_x.to(tl.uint8)
    normal_x = qx_mag
    mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1
    val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1
    normal_x += val_to_add
    normal_x += mant_odd
    normal_x = normal_x >> (MBITS_F32 - MBITS_FP4)
    normal_x = normal_x.to(tl.uint8)
    e2m1_value = tl.full(
        [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE], 0x7, dtype=tl.uint8
    )
    e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
    e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)
    sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
    sign_lp = sign_lp.to(tl.uint8)
    e2m1_value = e2m1_value | sign_lp
    
    '''
    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
    qx = qx ^ s
    qx_fp32 = qx.to(tl.float32, bitcast=True)
    saturate_mask = qx_fp32 >= max_normal
    denormal_mask = (not saturate_mask) & (qx_fp32 < min_normal)
    normal_mask = not (saturate_mask | denormal_mask)
    denorm_exp: tl.constexpr = (
        (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
    )
    denorm_mask_int: tl.constexpr = denorm_exp << MBITS_F32
    denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)
    denormal_x = qx_fp32 + denorm_mask_float
    denormal_x = denormal_x.to(tl.uint32, bitcast=True)
    denormal_x -= denorm_mask_int
    denormal_x = denormal_x.to(tl.uint8)
    normal_x = qx
    mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1
    val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1
    normal_x += val_to_add
    normal_x += mant_odd
    normal_x = normal_x >> (MBITS_F32 - MBITS_FP4)
    normal_x = normal_x.to(tl.uint8)
    e2m1_value = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
    e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
    e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)
    sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
    sign_lp = sign_lp.to(tl.uint8)
    e2m1_value = e2m1_value | sign_lp
    '''
    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)


def _quant_config(m: int, k: int) -> tuple[int, int, int, int, int]:
    """
    Return (BLOCK_M, BLOCK_K, NUM_WARPS, NUM_STAGES, K_CHUNK) for
    _amd_mxfp4_quant_a_kernel.

    Quant kernel parallelizes M and K-chunk in 2D grid, and each CTA walks its
    K-chunk with inner tl.range. num_stages participates in software pipelining.
    """
    bench: dict[tuple[int, int], tuple[int, int, int, int, int]] = {
        # Short-K: larger K tile with moderate pipeline depth.
        (4, 512): (4, 512, 2, 2, 512),
        (32, 512): (32, 512, 4, 2, 512),
        # Long-K: reduce BLOCK_K for register pressure; increase stages to hide latency.
        (16, 7168): (16, 128, 4, 3, 256),
        (64, 2048): (32, 128, 4, 3, 256),
        (256, 1536): (32, 64, 4, 3, 256),
    }
    if (m, k) in bench:
        return bench[(m, k)]

    block_m = 16 if m <= 16 else (32 if m <= 64 else 64)
    block_k = 512 if k <= 1024 else 256
    num_warps = 2 if m <= 8 else (4 if m <= 64 else 8)
    num_stages = 3 if k >= 2048 else 2
    # Keep each K-chunk large enough for overlap but bounded for occupancy.
    k_chunk = block_k * (4 if k >= 2048 else 2)
    return (block_m, block_k, num_warps, num_stages, k_chunk)


def _gemm_config(m: int, n: int, k: int) -> tuple[int, int, int, int, int]:
    """
    Return (BLOCK_M, BLOCK_N, BLOCK_K, NUM_WARPS, NUM_STAGES).
    """
    bench: dict[tuple[int, int, int], tuple[int, int, int, int, int]] = {
        #(4, 2880, 512): (16, 128, 64, 4, 2),
        (4, 2880, 512): (16, 16, 128, 4, 3),#fix config
        # Long-K: deeper pipeline to hide global-memory latency.
        #(16, 2112, 7168): (16, 128, 128, 4, 3),
        (16, 2112, 7168): (16, 16, 512, 4, 3),
        #(8, 2112, 7168): (16, 16, 128, 4, 2),
        (32, 4096, 512): (16, 16, 128, 4, 3), #fix config
        (32, 2880, 512): (16, 16, 128, 4, 3), #fix config
        # Case5: recover opt_3 behavior (more CTAs, lower per-CTA pressure).
        (64, 7168, 2048): (64, 16, 256, 4, 3),
        (256, 3072, 1536): (32, 16, 256, 4, 3),
    }
    if (m, n, k) in bench:
        return bench[(m, n, k)]
    # Unknown benchmark shapes: avoid KeyError on CI / extra tests.
    if m <= 16:
        if k >= 4096:
            return (16, 128, 256, 8, 2)
        return (16, 128, 128, 4, 2)
    if m <= 64:
        if m <= 32:
            return (32, 128, 256, 8, 2)
        return (64, 128, 256, 8, 2)
    if k >= 2048:
        return (64, 128, 256, 8, 2)
    return (64, 128, 256, 8, 2)


@triton.jit
def _amd_mxfp4_quant_a_kernel(
    a_bf16_ptr,
    a_q_ptr,
    a_sc_ptr,
    stride_a_m,
    stride_a_k,
    stride_aq_m,
    stride_aq_kh,
    stride_asc_m,
    stride_asc_kg,
    M,
    K,
    BLOCK_M: tl.constexpr,
    BLOCK_K: tl.constexpr,
    NUM_STAGES: tl.constexpr,
    K_CHUNK: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_k_chunk = tl.program_id(1)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    m_mask = offs_m < M
    quant_block_size: tl.constexpr = 32
    k_half = K // 2
    k_scale = K // 32
    k_begin = pid_k_chunk * K_CHUNK

    for k_iter in tl.range(0, K_CHUNK, BLOCK_K, num_stages=NUM_STAGES):
        k0 = k_begin + k_iter
        offs_kh = (k0 // 2) + tl.arange(0, BLOCK_K // 2)
        offs_kg = (k0 // 32) + tl.arange(0, BLOCK_K // 32)

        a_block_ptr = tl.make_block_ptr(
            base=a_bf16_ptr,
            shape=(M, K),
            strides=(stride_a_m, stride_a_k),
            offsets=(pid_m * BLOCK_M, k0),
            block_shape=(BLOCK_M, BLOCK_K),
            order=(1, 0),
        )
        a_f32 = tl.load(a_block_ptr, boundary_check=(0, 1), padding_option="zero").to(
            tl.float32
        )

        a_q, a_scales = _mxfp4_quant_op(
            a_f32,
            BLOCK_SIZE_N=BLOCK_K,
            BLOCK_SIZE_M=BLOCK_M,
            MXFP4_QUANT_BLOCK_SIZE=quant_block_size,
        )

        tl.store(
            a_q_ptr + offs_m[:, None] * stride_aq_m + offs_kh[None, :] * stride_aq_kh,
            a_q,
            mask=m_mask[:, None] & (offs_kh[None, :] < k_half),
        )
        tl.store(
            a_sc_ptr + offs_m[:, None] * stride_asc_m + offs_kg[None, :] * stride_asc_kg,
            a_scales,
            mask=m_mask[:, None] & (offs_kg[None, :] < k_scale),
        )


@triton.jit
def _amd_mxfp4_qs_gemm_kernel(
    a_q_ptr,
    a_sc_ptr,
    b_sh_ptr,
    b_sc_sh_ptr,
    c_ptr,
    stride_aq_m,
    stride_aq_kh,
    stride_asc_m,
    stride_asc_kg,
    stride_b_m,
    stride_b_kh,
    stride_bs_m,
    stride_bs_n,
    stride_c_m,
    stride_c_n,
    M,
    N,
    K,
    BS_SN,
    N_TILES: tl.constexpr,
    K_TILES: tl.constexpr,
    USE_B_SHUFFLE_TILE_LOAD: tl.constexpr,
    USE_BS_SHUFFLE_TILE_LOAD: tl.constexpr,
    BS_BLOCK_K_TILES: tl.constexpr,
    BS_BLOCK4_TILES: tl.constexpr,
    BLOCK_K_TILES: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
    NUM_STAGES: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    offs_k = tl.arange(0, BLOCK_K)
    offs_kh = tl.arange(0, BLOCK_K // 2)
    offs_kg = tl.arange(0, BLOCK_K // 32)

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    k_half = K // 2
    k_blocks_half = k_half // 32
    k_scale_valid = K // 32
    for k0 in tl.range(0, K, BLOCK_K, num_stages=NUM_STAGES):
        kh0 = k0 // 2
        kg0 = k0 // 32

        a_kh = kh0 + offs_kh
        a_q = tl.load(
            a_q_ptr + offs_m[:, None] * stride_aq_m + a_kh[None, :] * stride_aq_kh,
            mask=(offs_m[:, None] < M) & (a_kh[None, :] < k_half),
            other=0,
            cache_modifier=".ca",
        )
        a_kg = kg0 + offs_kg
        a_scales = tl.load(
            a_sc_ptr + offs_m[:, None] * stride_asc_m + a_kg[None, :] * stride_asc_kg,
            mask=(offs_m[:, None] < M) & (a_kg[None, :] < k_scale_valid),
            other=127,
            cache_modifier=".ca",
        )

        #b_kh = kh0 + offs_kh
        #bn = offs_n[None, :] // 16
        #bi = offs_n[None, :] % 16
        #kb = b_kh[:, None] // 32
        #kk = b_kh[:, None] % 32
        #k4 = kk // 16
        #k5 = kk % 16

        # BLOCK_N % 16 == 0, BLOCK_K multiple of 64: load BLOCK_N//16 shuffle n-tiles along dim 1.
        # Raw (Nn,Kt,2,16,16) -> permute (1,2,4,0,3) -> (k_tile,sub,k5,n_tile,bi) -> row-major (kh,n) for dot_scaled.
        '''
        if (
            tl.constexpr(USE_B_SHUFFLE_TILE_LOAD)
            and tl.constexpr(BLOCK_N % 16 == 0)
            and tl.constexpr(BLOCK_K_TILES >= 1)
        ):
        '''
        kb0 = kh0 // 32
        b_block_ptr = tl.make_block_ptr(
            base=b_sh_ptr,
            shape=(1, N_TILES, K_TILES, 2, 16, 16),
            strides=(
                N_TILES * K_TILES * 512,
                K_TILES * 512,
                512,
                256,
                16,
                1,
            ),
            offsets=(0, pid_n * (BLOCK_N // 16), kb0, 0, 0, 0),
            block_shape=(1, BLOCK_N // 16, BLOCK_K_TILES, 2, 16, 16),
            order=(5, 4, 3, 2, 1, 0),
        )
        b_raw = tl.load(b_block_ptr)
        b4 = tl.reshape(b_raw, (BLOCK_N // 16, BLOCK_K_TILES, 2, 16, 16))
        b4 = tl.permute(b4, (1, 2, 4, 0, 3))
        b = tl.reshape(b4, (BLOCK_K // 2, BLOCK_N))
        '''
        else:
            b_lin = (((((bn * k_blocks_half + kb) * 2 + k4) * 16 + bi) * 16) + k5)
            b_row = b_lin // k_half
            b_col = b_lin % k_half
            b = tl.load(
                b_sh_ptr + b_row * stride_b_m + b_col * stride_b_kh,
                mask=(offs_n[None, :] < N) & (b_kh[:, None] < k_half),
                other=0,
                cache_modifier=".cg",
            )
        '''
        b_kg = kg0 + offs_kg
        b_sc_mask = (offs_n[:, None] < N) & (b_kg[None, :] < k_scale_valid)
        if tl.constexpr(USE_BS_SHUFFLE_TILE_LOAD):
            # Fast path for B_scale_sh shuffle layout:
            # load (b3, b4, b5, b2) tile and remap to (n, kg).
            n_tile16 = pid_n
            n_block32 = n_tile16 // 2
            n_half16 = n_tile16 % 2
            kg_block8 = kg0 // 8
            kg_half4 = (kg0 % 8) // 4
            bs_base = (
                b_sc_sh_ptr
                + n_block32 * 32 * stride_bs_m
                + (kg_block8 * 256 + kg_half4 * 2 + n_half16) * stride_bs_n
            )
            bs_block_ptr = tl.make_block_ptr(
                base=bs_base,
                shape=(BS_BLOCK_K_TILES, BS_BLOCK4_TILES, 4, 16),
                strides=(256, 2, 64, 4),
                offsets=(0, 0, 0, 0),
                block_shape=(BS_BLOCK_K_TILES, BS_BLOCK4_TILES, 4, 16),
                order=(3, 2, 1, 0),
            )
            bs_raw = tl.load(bs_block_ptr)
            bs4 = tl.permute(bs_raw, (3, 0, 1, 2))
            b_scales = tl.reshape(bs4, (BLOCK_N, BLOCK_K // 32))
            b_scales = tl.where(b_sc_mask, b_scales, 127)
        else:
            b0 = offs_n[:, None] // 32
            b1 = offs_n[:, None] % 32
            b2 = b1 % 16
            b1 = b1 // 16
            b3 = b_kg[None, :] // 8
            b4 = b_kg[None, :] % 8
            b5 = b4 % 4
            b4 = b4 // 4
            b_lin_sc = b1 + b4 * 2 + b2 * 4 + b5 * 64 + b3 * 256 + b0 * 32 * BS_SN
            b_sm = b_lin_sc // BS_SN
            b_sn = b_lin_sc % BS_SN
            b_scales = tl.load(
                b_sc_sh_ptr + b_sm * stride_bs_m + b_sn * stride_bs_n,
                mask=b_sc_mask,
                other=127,
                cache_modifier=".cg",
            )

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

    c = acc.to(tl.bfloat16)
    c_ptrs = c_ptr + offs_m[:, None] * stride_c_m + offs_n[None, :] * stride_c_n
    c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
    tl.store(c_ptrs, c, mask=c_mask)


@triton.jit
def _amd_fused_mxfp4_qs_gemm_kernel(
    a_bf16_ptr,
    b_sh_ptr,
    b_sc_sh_ptr,
    c_ptr,
    stride_a_m,
    stride_a_k,
    stride_b_m,
    stride_b_kh,
    stride_bs_m,
    stride_bs_n,
    stride_c_m,
    stride_c_n,
    M,
    N,
    K,
    BS_SN,
    N_TILES: tl.constexpr,
    K_TILES: tl.constexpr,
    USE_B_SHUFFLE_TILE_LOAD: tl.constexpr,
    USE_BS_SHUFFLE_TILE_LOAD: tl.constexpr,
    BS_BLOCK_K_TILES: tl.constexpr,
    BS_BLOCK4_TILES: tl.constexpr,
    BLOCK_K_TILES: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
    NUM_STAGES: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    offs_k = tl.arange(0, BLOCK_K)
    offs_kh = tl.arange(0, BLOCK_K // 2)
    offs_kg = tl.arange(0, BLOCK_K // 32)

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    k_half = K // 2
    k_blocks_half = k_half // 32
    k_scale_valid = K // 32
    quant_block_size: tl.constexpr = 32

    for k0 in tl.range(0, K, BLOCK_K, num_stages=NUM_STAGES):
        kh0 = k0 // 2
        kg0 = k0 // 32

        a_block_ptr = tl.make_block_ptr(
            base=a_bf16_ptr,
            shape=(M, K),
            strides=(stride_a_m, stride_a_k),
            offsets=(pid_m * BLOCK_M, k0),
            block_shape=(BLOCK_M, BLOCK_K),
            order=(1, 0),
        )
        a_f32 = tl.load(
            a_block_ptr, boundary_check=(0, 1), padding_option="zero"
        ).to(tl.float32)

        a_q, a_scales = _mxfp4_quant_op(
            a_f32,
            BLOCK_SIZE_N=BLOCK_K,
            BLOCK_SIZE_M=BLOCK_M,
            MXFP4_QUANT_BLOCK_SIZE=quant_block_size,
        )

        #b_kh = kh0 + offs_kh
        #bn = offs_n[None, :] // 16
        #bi = offs_n[None, :] % 16
        #kb = b_kh[:, None] // 32
        #kk = b_kh[:, None] % 32
        #k4 = kk // 16
        #k5 = kk % 16

        # Same B layout as _amd_mxfp4_qs_gemm_kernel; BLOCK_K_TILES = BLOCK_K//64.
        '''
        if (
            tl.constexpr(USE_B_SHUFFLE_TILE_LOAD)
            and tl.constexpr(BLOCK_N % 16 == 0)
            and tl.constexpr(BLOCK_K_TILES >= 1)
        ):
        '''
        kb0 = kh0 // 32
        b_block_ptr = tl.make_block_ptr(
            base=b_sh_ptr,
            shape=(1, N_TILES, K_TILES, 2, 16, 16),
            strides=(
                N_TILES * K_TILES * 512,
                K_TILES * 512,
                512,
                256,
                16,
                1,
            ),
            offsets=(0, pid_n * (BLOCK_N // 16), kb0, 0, 0, 0),
            block_shape=(1, BLOCK_N // 16, BLOCK_K_TILES, 2, 16, 16),
            order=(5, 4, 3, 2, 1, 0),
        )
        b_raw = tl.load(b_block_ptr)
        b4 = tl.reshape(b_raw, (BLOCK_N // 16, BLOCK_K_TILES, 2, 16, 16))
        b4 = tl.permute(b4, (1, 2, 4, 0, 3))
        b = tl.reshape(b4, (BLOCK_K // 2, BLOCK_N))
        '''
        else:
            b_lin = (((((bn * k_blocks_half + kb) * 2 + k4) * 16 + bi) * 16) + k5)
            b_row = b_lin // k_half
            b_col = b_lin % k_half
            b = tl.load(
                b_sh_ptr + b_row * stride_b_m + b_col * stride_b_kh,
                mask=(offs_n[None, :] < N) & (b_kh[:, None] < k_half),
                other=0,
                cache_modifier=".cg",
            )
        '''
        b_kg = kg0 + offs_kg
        b_sc_mask = (offs_n[:, None] < N) & (b_kg[None, :] < k_scale_valid)
        if tl.constexpr(USE_BS_SHUFFLE_TILE_LOAD):
            n_tile16 = pid_n
            n_block32 = n_tile16 // 2
            n_half16 = n_tile16 % 2
            kg_block8 = kg0 // 8
            kg_half4 = (kg0 % 8) // 4
            bs_base = (
                b_sc_sh_ptr
                + n_block32 * 32 * stride_bs_m
                + (kg_block8 * 256 + kg_half4 * 2 + n_half16) * stride_bs_n
            )
            bs_block_ptr = tl.make_block_ptr(
                base=bs_base,
                shape=(BS_BLOCK_K_TILES, BS_BLOCK4_TILES, 4, 16),
                strides=(256, 2, 64, 4),
                offsets=(0, 0, 0, 0),
                block_shape=(BS_BLOCK_K_TILES, BS_BLOCK4_TILES, 4, 16),
                order=(3, 2, 1, 0),
            )
            bs_raw = tl.load(bs_block_ptr)
            bs4 = tl.permute(bs_raw, (3, 0, 1, 2))
            b_scales = tl.reshape(bs4, (BLOCK_N, BLOCK_K // 32))
            b_scales = tl.where(b_sc_mask, b_scales, 127)
        else:
            b0 = offs_n[:, None] // 32
            b1 = offs_n[:, None] % 32
            b2 = b1 % 16
            b1 = b1 // 16
            b3 = b_kg[None, :] // 8
            b4 = b_kg[None, :] % 8
            b5 = b4 % 4
            b4 = b4 // 4
            b_lin_sc = b1 + b4 * 2 + b2 * 4 + b5 * 64 + b3 * 256 + b0 * 32 * BS_SN
            b_sm = b_lin_sc // BS_SN
            b_sn = b_lin_sc % BS_SN
            b_scales = tl.load(
                b_sc_sh_ptr + b_sm * stride_bs_m + b_sn * stride_bs_n,
                mask=b_sc_mask,
                other=127,
                cache_modifier=".cg",
            )

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

    c = acc.to(tl.bfloat16)
    c_ptrs = c_ptr + offs_m[:, None] * stride_c_m + offs_n[None, :] * stride_c_n
    c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
    tl.store(c_ptrs, c, mask=c_mask)


#mxfp4_gemm_shuffle_direct = _amd_fused_mxfp4_qs_gemm_kernel


def _b_shuffle_tile_load_ok(B_shuffle: torch.Tensor, block_n: int, block_k: int) -> bool:
    """
    Block contiguous load matches shuffle_weight [N, K/2] layout when:
    - BLOCK_N is a multiple of 16 (covers BLOCK_N//16 shuffle n-tiles per program),
    - K-tile count BLOCK_K//64 is integer (BLOCK_K divisible by 64),
    - tensor is contiguous row-major (stride (kh, 1)).
    """
    if block_n % 16 != 0:
        return False
    if block_k % 64 != 0:
        return False
    if not B_shuffle.is_contiguous():
        return False
    n, kh = B_shuffle.shape
    return B_shuffle.stride(0) == kh and B_shuffle.stride(1) == 1


def _b_scale_shuffle_tile_load_ok(
    B_scale_sh: torch.Tensor, block_n: int, block_k: int
) -> bool:
    """
    Fast path for B_scale_sh shuffle layout.

    Conditions match current kernel mapping:
    - BLOCK_N == 16 (single half-32 n tile per program),
    - BLOCK_K in {128, 256, 512} (and 128-aligned),
    - B_scale_sh is contiguous row-major with stride (BS_SN, 1).
    """
    if block_n != 16:
        return False
    if block_k % 128 != 0:
        return False
    if block_k > 512:
        return False
    if not B_scale_sh.is_contiguous():
        return False
    bs_m, bs_n = B_scale_sh.shape
    if bs_n % 8 != 0:
        return False
    return B_scale_sh.stride(0) == bs_n and B_scale_sh.stride(1) == 1


def amd_fused_mxfp4_qs_gemm(
    A: torch.Tensor,
    B_shuffle: torch.Tensor,
    B_scale_sh: torch.Tensor,
    *,
    dtype: torch.dtype,
) -> torch.Tensor:
    m, k = A.shape
    n, kh_b = B_shuffle.shape
    assert kh_b * 2 == k, "A and B_shuffle K mismatch"
    assert k % 64 == 0, "K must be divisible by 64"

    if B_shuffle.dtype != torch.uint8:
        B_shuffle = B_shuffle.view(torch.uint8)
    if B_scale_sh.dtype != torch.uint8:
        B_scale_sh = B_scale_sh.view(torch.uint8)

    out = torch.empty((m, n), device=A.device, dtype=dtype)
    block_m, block_n, block_k, num_warps, num_stages = _gemm_config(m, n, k)
    grid = (triton.cdiv(m, block_m), triton.cdiv(n, block_n))
    n_tiles = n // 16
    k_tiles = k // 64
    use_b_tile = _b_shuffle_tile_load_ok(B_shuffle, block_n, block_k)
    use_bs_tile = _b_scale_shuffle_tile_load_ok(B_scale_sh, block_n, block_k)
    block_k_tiles = block_k // 64
    bs_block_k_tiles = max(1, block_k // 256)
    bs_block4_tiles = 1 if block_k == 128 else 2

    _amd_fused_mxfp4_qs_gemm_kernel[grid](
        A,
        B_shuffle,
        B_scale_sh,
        out,
        *A.stride(),
        *B_shuffle.stride(),
        *B_scale_sh.stride(),
        *out.stride(),
        m,
        n,
        k,
        B_scale_sh.shape[1],
        N_TILES=n_tiles,
        K_TILES=k_tiles,
        USE_B_SHUFFLE_TILE_LOAD=use_b_tile,
        USE_BS_SHUFFLE_TILE_LOAD=use_bs_tile,
        BS_BLOCK_K_TILES=bs_block_k_tiles,
        BS_BLOCK4_TILES=bs_block4_tiles,
        BLOCK_K_TILES=block_k_tiles,
        BLOCK_M=block_m,
        BLOCK_N=block_n,
        BLOCK_K=block_k,
        NUM_STAGES=num_stages,
        num_warps=num_warps,
        num_stages=num_stages,
    )
    return out


def amd_mxfp4_qs_gemm_two_stage(
    A: torch.Tensor,
    B_shuffle: torch.Tensor,
    B_scale_sh: torch.Tensor,
    *,
    dtype: torch.dtype,
) -> torch.Tensor:
    m, k = A.shape
    n, kh_b = B_shuffle.shape
    assert kh_b * 2 == k, "A and B_shuffle K mismatch"
    assert k % 64 == 0, "K must be divisible by 64"

    if B_shuffle.dtype != torch.uint8:
        B_shuffle = B_shuffle.view(torch.uint8)
    if B_scale_sh.dtype != torch.uint8:
        B_scale_sh = B_scale_sh.view(torch.uint8)

    out = torch.empty((m, n), device=A.device, dtype=dtype)
    q_block_m, q_block_k, q_num_warps, q_num_stages, q_k_chunk = _quant_config(m, k)
    g_block_m, g_block_n, g_block_k, g_num_warps, g_num_stages = _gemm_config(m, n, k)

    a_q = torch.empty((m, k // 2), device=A.device, dtype=torch.uint8)
    a_scales = torch.empty((m, k // 32), device=A.device, dtype=torch.uint8)

    grid_quant = (triton.cdiv(m, q_block_m), triton.cdiv(k, q_k_chunk))
    _amd_mxfp4_quant_a_kernel[grid_quant](
        A,
        a_q,
        a_scales,
        *A.stride(),
        *a_q.stride(),
        *a_scales.stride(),
        m,
        k,
        BLOCK_M=q_block_m,
        BLOCK_K=q_block_k,
        NUM_STAGES=q_num_stages,
        K_CHUNK=q_k_chunk,
        num_warps=q_num_warps,
        num_stages=q_num_stages,
    )

    grid = (triton.cdiv(m, g_block_m), triton.cdiv(n, g_block_n))
    n_tiles = n // 16
    k_tiles = k // 64
    use_b_tile = _b_shuffle_tile_load_ok(B_shuffle, g_block_n, g_block_k)
    use_bs_tile = _b_scale_shuffle_tile_load_ok(B_scale_sh, g_block_n, g_block_k)
    block_k_tiles = g_block_k // 64
    bs_block_k_tiles = max(1, g_block_k // 256)
    bs_block4_tiles = 1 if g_block_k == 128 else 2

    _amd_mxfp4_qs_gemm_kernel[grid](
        a_q,
        a_scales,
        B_shuffle,
        B_scale_sh,
        out,
        *a_q.stride(),
        *a_scales.stride(),
        *B_shuffle.stride(),
        *B_scale_sh.stride(),
        *out.stride(),
        m,
        n,
        k,
        B_scale_sh.shape[1],
        N_TILES=n_tiles,
        K_TILES=k_tiles,
        USE_B_SHUFFLE_TILE_LOAD=use_b_tile,
        USE_BS_SHUFFLE_TILE_LOAD=use_bs_tile,
        BS_BLOCK_K_TILES=bs_block_k_tiles,
        BS_BLOCK4_TILES=bs_block4_tiles,
        BLOCK_K_TILES=block_k_tiles,
        BLOCK_M=g_block_m,
        BLOCK_N=g_block_n,
        BLOCK_K=g_block_k,
        NUM_STAGES=g_num_stages,
        num_warps=g_num_warps,
        num_stages=g_num_stages,
    )
    return out


def custom_kernel(data: input_t) -> output_t:
    A, _B, _B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    m, k = A.shape

    # First threshold pass: small-M/small-K tends to favor fused path.
    use_fused = (k <1536 )
    if use_fused:
        return amd_fused_mxfp4_qs_gemm(
            A,
            B_shuffle,
            B_scale_sh,
            dtype=torch.bfloat16,
        )

    return amd_mxfp4_qs_gemm_two_stage(
    #return amd_fused_mxfp4_qs_gemm(
        A,
        B_shuffle,
        B_scale_sh,
        dtype=torch.bfloat16,
    )
scrolls · 820 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