Skip to content
KernelIndex
Search⌘K

submission 531035

josusanmartin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

mxfp4_v119_xcd_only.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-531035?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
11.2µs
#315 of 1143
2026-03-11

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

split-kand (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)

Kernel source

mxfp4_v119_xcd_only.py295 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
Version 119: v103 base + ONLY remap_xcd optimization.
Isolate the effect of XCD-aware scheduling without .wt store or acc pattern.
"""
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.utils._triton.pid_preprocessing import pid_grid, remap_xcd
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)


_mxfp4_quant_op = _mxfp4_quant_op_asm_exact

import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module
_kernel_module._mxfp4_quant_op = _mxfp4_quant_op_asm_exact


# Kernel with ONLY remap_xcd (no .wt store, no acc pattern change)
@triton.heuristics(
    {
        "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
        and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
        and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
        "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
        * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
    }
)
@triton.jit
def _gemm_a16wfp4_preshuffle_kernel_xcd(
    a_ptr, b_ptr, c_ptr, b_scales_ptr,
    M, N, K,
    stride_am, stride_ak,
    stride_bn, stride_bk,
    stride_ck, stride_cm, stride_cn,
    stride_bsn, stride_bsk,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr,
    NUM_KSPLIT: tl.constexpr,
    SPLITK_BLOCK_SIZE: tl.constexpr,
    EVEN_K: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr,
    matrix_instr_nonkdim: tl.constexpr,
    GRID_MN: tl.constexpr,
    PREQUANT: tl.constexpr,
    cache_modifier: tl.constexpr,
):
    tl.assume(stride_am > 0)
    tl.assume(stride_ak > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)
    tl.assume(stride_bsk > 0)
    tl.assume(stride_bsn > 0)

    pid_unified = tl.program_id(axis=0)
    # ONLY change: XCD-aware remapping
    pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)

    pid_k = pid_unified % NUM_KSPLIT
    pid = pid_unified // NUM_KSPLIT
    num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
    num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)

    if NUM_KSPLIT == 1:
        pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
    else:
        pid_m = pid // num_pid_n
        pid_n = pid % num_pid_n

    tl.assume(pid_m >= 0)
    tl.assume(pid_n >= 0)
    tl.assume(pid_k >= 0)

    SCALE_GROUP_SIZE: tl.constexpr = 32

    if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
        num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)

        offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
        offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
        offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        a_ptrs = a_ptr + (
            offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
        )

        offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
        offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
        offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
        b_ptrs = b_ptr + (
            offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
        )

        offs_bsn = (
            pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))
        ) % N
        offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(
            0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32
        )
        b_scale_ptrs = (
            b_scales_ptr
            + offs_bsn[:, None] * stride_bsn
            + offs_ks[None, :] * stride_bsk
        )

        accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)

        for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
            b_scales = (
                tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
                .reshape(
                    BLOCK_SIZE_N // 32,
                    BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8,
                    4, 16, 2, 2, 1,
                )
                .permute(0, 5, 3, 1, 4, 2, 6)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
            )

            if EVEN_K:
                a_bf16 = tl.load(a_ptrs)
                b = tl.load(b_ptrs, cache_modifier=cache_modifier)

            b = (
                b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
                .permute(0, 1, 4, 2, 3, 5)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
                .trans(1, 0)
            )

            if PREQUANT:
                a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)

            # Keep original += pattern
            accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")

            a_ptrs += BLOCK_SIZE_K * stride_ak
            b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
            b_scale_ptrs += BLOCK_SIZE_K * stride_bsk

        c = accumulator.to(c_ptr.type.element_ty)

        offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
        offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
        c_ptrs = (
            c_ptr
            + stride_cm * offs_cm[:, None]
            + stride_cn * offs_cn[None, :]
            + pid_k * stride_ck
        )
        c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
        # Keep original store (no .wt)
        tl.store(c_ptrs, c, mask=c_mask)


import aiter.ops.triton.gemm.basic.gemm_a16wfp4 as _wrapper_module
_wrapper_module._gemm_a16wfp4_preshuffle_kernel = _gemm_a16wfp4_preshuffle_kernel_xcd

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"

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

_FUSED_CONFIGS = {
    (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,
    },
}

_QUANT_BLOCK = 32
_QUANT_TILE = 128
_bufs = {}


def _get_asm_bufs(m, k, n, 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)
    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


@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)

    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]

        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,
            BLOCK_SIZE=_QUANT_TILE,
            MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK,
            SCALING_MODE=0, SHUFFLE=True,
        )

        splitK = _ASM_SPLITK[key]
        gemm_a4w4_asm(
            x_fp4.view(_fp4x2), B_shuffle, scale.view(_fp8_e8m0), B_scale_sh,
            out, _kernel_32x128,
            bpreshuffle=True, log2_k_split=splitK,
        )
        return out[:m]
    else:
        if key not in _bufs:
            _bufs[key] = ('fused', torch.empty((m, n), dtype=torch.bfloat16, 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)

        config = _FUSED_CONFIGS.get(key)
        return gemm_a16wfp4_preshuffle(A, w, w_scales, prequant=True, y=out, config=config)
scrolls · 295 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 530427.

⋯ 1 unchanged lines
#!POPCORN gpu MI355X
"""
- Version 87: Best hybrid combining ALL discoveries:
- - M=4,K=512: Fused GEMM with BLOCK_SIZE_M=8 → 7.38µs (v80 finding)
- - M=16,K=7168: Fused GEMM with KSPLIT=14 → 16.8µs (v79 finding)
- - M=32,K=512: Fused GEMM with DEFAULT config → 9.47/9.50µs (v77 finding)
- - M=64,K=2048: ASM splitK=2 → 13.7µs (v62 finding)
- - M=256,K=1536: ASM splitK=1 → 12.5µs (v62 finding)
-
- Expected geomean: ~11.1µs (vs 12.7µs baseline)
+ Version 119: v103 base + ONLY remap_xcd optimization.
+ Isolate the effect of XCD-aware scheduling without .wt store or acc pattern.
"""
import torch
import triton
⋯ 2 unchanged lines
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.utils._triton.pid_preprocessing import pid_grid, remap_xcd
from task import input_t, output_t
⋯ 35 unchanged lines
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
+ _mxfp4_quant_op = _mxfp4_quant_op_asm_exact
+
import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module
_kernel_module._mxfp4_quant_op = _mxfp4_quant_op_asm_exact
+
+ # Kernel with ONLY remap_xcd (no .wt store, no acc pattern change)
+ @triton.heuristics(
+ {
+ "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
+ and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
+ and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
+ "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
+ * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
+ }
+ )
+ @triton.jit
+ def _gemm_a16wfp4_preshuffle_kernel_xcd(
+ a_ptr, b_ptr, c_ptr, b_scales_ptr,
+ M, N, K,
+ stride_am, stride_ak,
+ stride_bn, stride_bk,
+ stride_ck, stride_cm, stride_cn,
+ stride_bsn, stride_bsk,
+ BLOCK_SIZE_M: tl.constexpr,
+ BLOCK_SIZE_N: tl.constexpr,
+ BLOCK_SIZE_K: tl.constexpr,
+ GROUP_SIZE_M: tl.constexpr,
+ NUM_KSPLIT: tl.constexpr,
+ SPLITK_BLOCK_SIZE: tl.constexpr,
+ EVEN_K: tl.constexpr,
+ num_warps: tl.constexpr,
+ num_stages: tl.constexpr,
+ waves_per_eu: tl.constexpr,
+ matrix_instr_nonkdim: tl.constexpr,
+ GRID_MN: tl.constexpr,
+ PREQUANT: tl.constexpr,
+ cache_modifier: tl.constexpr,
+ ):
+ tl.assume(stride_am > 0)
+ tl.assume(stride_ak > 0)
+ tl.assume(stride_bk > 0)
+ tl.assume(stride_bn > 0)
+ tl.assume(stride_cm > 0)
+ tl.assume(stride_cn > 0)
+ tl.assume(stride_bsk > 0)
+ tl.assume(stride_bsn > 0)
+
+ pid_unified = tl.program_id(axis=0)
+ # ONLY change: XCD-aware remapping
+ pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)
+
+ pid_k = pid_unified % NUM_KSPLIT
+ pid = pid_unified // NUM_KSPLIT
+ num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
+ num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
+
+ if NUM_KSPLIT == 1:
+ pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
+ else:
+ pid_m = pid // num_pid_n
+ pid_n = pid % num_pid_n
+
+ tl.assume(pid_m >= 0)
+ tl.assume(pid_n >= 0)
+ tl.assume(pid_k >= 0)
+
+ SCALE_GROUP_SIZE: tl.constexpr = 32
+
+ if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
+ num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
+
+ offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
+ offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
+ offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
+ a_ptrs = a_ptr + (
+ offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
+ )
+
+ offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
+ offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
+ offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
+ b_ptrs = b_ptr + (
+ offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
+ )
+
+ offs_bsn = (
+ pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))
+ ) % N
+ offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(
+ 0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32
+ )
+ b_scale_ptrs = (
+ b_scales_ptr
+ + offs_bsn[:, None] * stride_bsn
+ + offs_ks[None, :] * stride_bsk
+ )
+
+ accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
+
+ for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
+ b_scales = (
+ tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
+ .reshape(
+ BLOCK_SIZE_N // 32,
+ BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8,
+ 4, 16, 2, 2, 1,
+ )
+ .permute(0, 5, 3, 1, 4, 2, 6)
+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
+ )
+
+ if EVEN_K:
+ a_bf16 = tl.load(a_ptrs)
+ b = tl.load(b_ptrs, cache_modifier=cache_modifier)
+
+ b = (
+ b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
+ .permute(0, 1, 4, 2, 3, 5)
+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
+ .trans(1, 0)
+ )
+
+ if PREQUANT:
+ a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
+
+ # Keep original += pattern
+ accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
+
+ a_ptrs += BLOCK_SIZE_K * stride_ak
+ b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
+ b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
+
+ c = accumulator.to(c_ptr.type.element_ty)
+
+ offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
+ offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
+ c_ptrs = (
+ c_ptr
+ + stride_cm * offs_cm[:, None]
+ + stride_cn * offs_cn[None, :]
+ + pid_k * stride_ck
+ )
+ c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
+ # Keep original store (no .wt)
+ tl.store(c_ptrs, c, mask=c_mask)
+
+
+ import aiter.ops.triton.gemm.basic.gemm_a16wfp4 as _wrapper_module
+ _wrapper_module._gemm_a16wfp4_preshuffle_kernel = _gemm_a16wfp4_preshuffle_kernel_xcd
+
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
_bf16 = dtypes.bf16
⋯ 2 unchanged lines
_kernel_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
- # ASM splitK for large-M shapes
_ASM_SPLITK = {
- (64, 7168, 2048): 2, # Proven best: 4 K-splits
- (256, 3072, 1536): 1, # Proven best: 2 K-splits
+ (64, 7168, 2048): 2,
+ (256, 3072, 1536): 1,
}
- # Fused configs - only override where needed
_FUSED_CONFIGS = {
- # M=4: BLOCK_SIZE_M=8 reduces padding waste (4→8 vs 4→32)
(4, 2880, 512): {
- "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 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,
},
- # M=16: Split-K=14 for better CU utilization
(16, 2112, 7168): {
- "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,
+ "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": 14,
+ "cache_modifier": ".cg", "NUM_KSPLIT": 7,
},
- # M=32: Use default config (BLOCK_SIZE_N=64 gives more N-tiles → better CU util)
- # No config override → uses default
}
_QUANT_BLOCK = 32
_QUANT_TILE = 128
-
_bufs = {}
⋯ 16 unchanged lines
key = (m, n, k)
if key in _ASM_SPLITK:
- # ASM 2-kernel path for M>=64 shapes
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]
⋯ 19 unchanged lines
)
return out[:m]
else:
- # Fused preshuffle GEMM for M<=32 shapes
if key not in _bufs:
_bufs[key] = ('fused', torch.empty((m, n), dtype=torch.bfloat16, device=A.device))
_, out = _bufs[key]
scrolls · 242 diff lines total

Best evidence level for this revision: reported

JSON