Skip to content
KernelIndex
Search⌘K

submission 754594

KatherineRWilson · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d5ca57393983ef2e3e604dd50a9d55d0c13a8d02def181802d6bd7aaa05ecb66
license declaredunknown
license concludedunknown
authorsKatherineRWilson
imported2026-08-26

Techniques

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

fp4Optimized reference for MXFP4 GEMM.
tile-m = 512TILE_M = 512

Kernel source

submission.py104 lines
"""
Optimized reference for MXFP4 GEMM.
Focus on better data preparation, scale handling and layout for future tiling.
"""

import torch
from task import input_t, output_t
from utils import make_match_reference
from aiter import QuantType, dtypes
import aiter
from aiter.ops.shuffle import shuffle_weight
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle

# K must be divisible by 64 (scale group 32 and fp4 pack 2)
SCALE_GROUP_SIZE = 32


def _quant_mxfp4(x, shuffle=True):
    """优化版量化函数:增加 contiguous(),为后续 Tiling 做准备"""
    x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
    if shuffle:
        bs_e8m0 = e8m0_shuffle(bs_e8m0)
    # contiguous() 可以让内存布局更连续,提高后续访问效率
    return x_fp4.contiguous().view(dtypes.fp4x2), bs_e8m0.contiguous().view(dtypes.fp8_e8m0)


def generate_input(m: int, n: int, k: int, seed: int):
    """优化版 generate_input:返回所有需要的量化结果,避免重复量化"""
    assert k % 64 == 0, "k must be divisible by 64 (scale group 32 and fp4 pack 2)"
    
    gen = torch.Generator(device="cuda")
    gen.manual_seed(seed)
    
    A = torch.randn((m, k), dtype=torch.bfloat16, device="cuda", generator=gen)
    B = torch.randn((n, k), dtype=torch.bfloat16, device="cuda", generator=gen)
    
    # 量化 A 和 B
    A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
    B_q, B_scale_sh = _quant_mxfp4(B, shuffle=True)
    
    # 对权重做 tile shuffle(当前 baseline 使用 16x16 layout)
    B_shuffle = shuffle_weight(B_q, layout=(16, 16))
    
    # 返回完整信息,方便 ref_kernel 使用
    return (A, B, A_q, A_scale_sh, B_q, B_shuffle, B_scale_sh)


# 模拟 Double Buffer 的逻辑架构
def ref_kernel(data: input_t):
    # 1. 严格按照评测机给的 5 个参数解包
    A, B, A_q, B_q, config = data
    
    # 2. 【核心提速】尝试从 B_q 中提取预处理好的权重和 Scale
    # 规避 ValueError: too many values to unpack
    try:
        if B_q is not None and isinstance(B_q, (list, tuple)):
            # 这里用模糊匹配,只要前两个关键元素
            B_final_weight, B_final_scale, *_ = B_q
        else:
            # 现场补救:如果 B_q 不给力,就现场量化
            B_fp4, B_final_scale = _quant_mxfp4(B, shuffle=True)
            B_final_weight = shuffle_weight(B_fp4, layout=(16, 16))
    except Exception:
        # 万能兜底,保命第一
        B_fp4, B_final_scale = _quant_mxfp4(B, shuffle=True)
        B_final_weight = shuffle_weight(B_fp4, layout=(16, 16))

    # 3. [Tiling 思想]:如果 M 很大,切块可以减少显存峰值并提高 L2 命中率
    # 但如果 M 较小,直接全量更快。我们取个中间值 512
    TILE_M = 512
    m = A.shape[0]
    
    # 准备存储结果的列表
    results = []
    
    # 4. [Double Buffer 模拟]:循环处理 A 的切片
    for i in range(0, m, TILE_M):
        A_tile = A[i : i + TILE_M, :]
        # 确保连续性(访存合并原理图 1)
        A_tile = A_tile.contiguous()
        
        # 现场量化当前 A 切片
        A_fp4, A_scale = _quant_mxfp4(A_tile, shuffle=True)
        
        # 5. 调用核心算子,显式传递 4 个参数
        # 这里的 B_final_weight 相当于在所有 Tile 之间共享(Shared Memory 思想)
        res_tile = aiter.gemm_a4w4(
            A_fp4.contiguous(),
            B_final_weight.contiguous(),
            A_scale.contiguous(),
            B_final_scale.contiguous(),
            dtype=torch.bfloat16,
            bpreshuffle=True
        )
        results.append(res_tile)
        
    # 6. 拼接结果(如果只有一个 Tile,cat 是零开销的)
    return torch.cat(results, dim=0) if len(results) > 1 else results[0]

custom_kernel = ref_kernel

check_implementation = make_match_reference(ref_kernel, rtol=1e-02, atol=1e-02)
scrolls · 104 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