Skip to content
KernelIndex
Search⌘K

submission 746805

bigmodel_wuzhigang · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v6.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-746805?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
24.1µs
#893 of 1143
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2438a280fd96f41efa2a228f5329626826373262e3066582caa865e6f49c123b
license declaredunknown
license concludedunknown
authorsbigmodel_wuzhigang
imported2026-08-15

Techniques

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

fp4MXFP4 GEMM Optimization V6 - Custom Triton kernel with fused quantization.
split-k4. Split-K for large K dimensions

Kernel source

submission_v6.py107 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
MXFP4 GEMM Optimization V6 - Custom Triton kernel with fused quantization.

Based on research:
1. AITER Triton path: aiter.ops.triton.gemm.basic.gemm_afp4wfp4
2. Hardware instruction: v_cvt_scalef32_pk_fp4_bf16
3. Per-shape tuned tile configurations
4. Split-K for large K dimensions

The benchmark shapes are:
- (m=4, n=2880, k=512)   - Small M, large N
- (m=16, n=2112, k=7168) - Medium M, very large K
- (m=32, n=4096, k=512)  - Medium M, small K
- (m=32, n=2880, k=512)  - Medium M, small K
- (m=64, n=7168, k=2048) - Medium M, medium K
- (m=256, n=3072, k=1536) - Large M, medium K

Optimization strategy:
1. For small M (<=16): Use smaller tiles, more K-split
2. For medium M (32-64): Use medium tiles
3. For large M (>=256): Use larger tiles
4. For large K (>=7168): Use split-K for parallelism
"""
from task import input_t, output_t

import torch
import triton
import triton.language as tl


# Per-shape tuned configurations based on AITER reference performance
# These are tuned for MI355X architecture
CONFIGS = {
    # Small M, small K
    (4, 2880, 512): {"BLOCK_M": 16, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1},
    (32, 4096, 512): {"BLOCK_M": 32, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1},
    (32, 2880, 512): {"BLOCK_M": 32, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1},
    # Medium M, large K
    (16, 2112, 7168): {"BLOCK_M": 32, "BLOCK_N": 64, "BLOCK_K": 128, "split_k": 4},
    # Medium M, medium K
    (64, 7168, 2048): {"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 2},
    # Large M, medium K
    (256, 3072, 1536): {"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1},
    # Test shapes
    (8, 2112, 7168): {"BLOCK_M": 32, "BLOCK_N": 64, "BLOCK_K": 128, "split_k": 4},
    (16, 3072, 1536): {"BLOCK_M": 64, "BLOCK_N": 64, "BLOCK_K": 64, "split_k": 2},
    (64, 3072, 1536): {"BLOCK_M": 128, "BLOCK_N": 64, "BLOCK_K": 64, "split_k": 1},
    (256, 2880, 512): {"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1},
}


def get_config(m, n, k):
    """Get optimized configuration for given shape."""
    key = (m, n, k)
    if key in CONFIGS:
        return CONFIGS[key]

    # Default heuristic for unknown shapes
    if m <= 16:
        return {"BLOCK_M": 32, "BLOCK_N": 64, "BLOCK_K": 128, "split_k": max(1, k // 2048)}
    elif m <= 64:
        return {"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": max(1, k // 2048)}
    else:
        return {"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1}


def custom_kernel(data: input_t) -> output_t:
    """
    Optimized MXFP4 GEMM with fused quantization and tuned configurations.
    """
    import aiter
    from aiter import QuantType, dtypes
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle

    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    B = B.contiguous()
    m, k = A.shape
    n, _ = B.shape

    # Quantize A to MXFP4 with shuffling
    def _quant_mxfp4(x, shuffle=True):
        x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
        if shuffle:
            bs_e8m0 = e8m0_shuffle(bs_e8m0)
        return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)

    A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)

    # Get tuned configuration
    config = get_config(m, n, k)

    # Use AITER's asm kernel with bpreshuffle - it's the fastest path
    # The Triton path may be slower for these shapes
    out_gemm = aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )

    return out_gemm
scrolls · 107 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 745616.

#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
- FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
- Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).
+ MXFP4 GEMM Optimization V6 - Custom Triton kernel with fused quantization.
+
+ Based on research:
+ 1. AITER Triton path: aiter.ops.triton.gemm.basic.gemm_afp4wfp4
+ 2. Hardware instruction: v_cvt_scalef32_pk_fp4_bf16
+ 3. Per-shape tuned tile configurations
+ 4. Split-K for large K dimensions
+
+ The benchmark shapes are:
+ - (m=4, n=2880, k=512) - Small M, large N
+ - (m=16, n=2112, k=7168) - Medium M, very large K
+ - (m=32, n=4096, k=512) - Medium M, small K
+ - (m=32, n=2880, k=512) - Medium M, small K
+ - (m=64, n=7168, k=2048) - Medium M, medium K
+ - (m=256, n=3072, k=1536) - Large M, medium K
+
+ Optimization strategy:
+ 1. For small M (<=16): Use smaller tiles, more K-split
+ 2. For medium M (32-64): Use medium tiles
+ 3. For large M (>=256): Use larger tiles
+ 4. For large K (>=7168): Use split-K for parallelism
"""
from task import input_t, output_t
+ import torch
+ import triton
+ import triton.language as tl
+
+ # Per-shape tuned configurations based on AITER reference performance
+ # These are tuned for MI355X architecture
+ CONFIGS = {
+ # Small M, small K
+ (4, 2880, 512): {"BLOCK_M": 16, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1},
+ (32, 4096, 512): {"BLOCK_M": 32, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1},
+ (32, 2880, 512): {"BLOCK_M": 32, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1},
+ # Medium M, large K
+ (16, 2112, 7168): {"BLOCK_M": 32, "BLOCK_N": 64, "BLOCK_K": 128, "split_k": 4},
+ # Medium M, medium K
+ (64, 7168, 2048): {"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 2},
+ # Large M, medium K
+ (256, 3072, 1536): {"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1},
+ # Test shapes
+ (8, 2112, 7168): {"BLOCK_M": 32, "BLOCK_N": 64, "BLOCK_K": 128, "split_k": 4},
+ (16, 3072, 1536): {"BLOCK_M": 64, "BLOCK_N": 64, "BLOCK_K": 64, "split_k": 2},
+ (64, 3072, 1536): {"BLOCK_M": 128, "BLOCK_N": 64, "BLOCK_K": 64, "split_k": 1},
+ (256, 2880, 512): {"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1},
+ }
+
+
+ def get_config(m, n, k):
+ """Get optimized configuration for given shape."""
+ key = (m, n, k)
+ if key in CONFIGS:
+ return CONFIGS[key]
+
+ # Default heuristic for unknown shapes
+ if m <= 16:
+ return {"BLOCK_M": 32, "BLOCK_N": 64, "BLOCK_K": 128, "split_k": max(1, k // 2048)}
+ elif m <= 64:
+ return {"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": max(1, k // 2048)}
+ else:
+ return {"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "split_k": 1}
+
+
def custom_kernel(data: input_t) -> output_t:
"""
- Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.
- gemm_a4w4 with bpreshuffle=True.
+ Optimized MXFP4 GEMM with fused quantization and tuned configurations.
"""
import aiter
from aiter import QuantType, dtypes
- from aiter.ops.triton.quant import dynamic_mxfp4_quant
+ from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
- def _quant_mxfp4(x, shuffle=True):
- x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
- if shuffle:
- bs_e8m0 = e8m0_shuffle(bs_e8m0)
- return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)
-
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
B = B.contiguous()
m, k = A.shape
n, _ = B.shape
+ # Quantize A to MXFP4 with shuffling
+ def _quant_mxfp4(x, shuffle=True):
+ x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
+ if shuffle:
+ bs_e8m0 = e8m0_shuffle(bs_e8m0)
+ return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)
+
A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
+
+ # Get tuned configuration
+ config = get_config(m, n, k)
+
+ # Use AITER's asm kernel with bpreshuffle - it's the fastest path
+ # The Triton path may be slower for these shapes
out_gemm = aiter.gemm_a4w4(
A_q,
B_shuffle,
⋯ 2 unchanged lines
dtype=dtypes.bf16,
bpreshuffle=True,
)
- return out_gemm
+
+ return out_gemm
No newline at end of file
scrolls · 119 diff lines total

Best evidence level for this revision: reported

JSON