Skip to content
KernelIndex
Search⌘K

submission 733857

guojun21 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_triton_handwritten_v25.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-733857?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
9.08µs
#118 of 1143
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:29b20e50fbec03a5c3f0a8252eddac346959e09afeb3d48911ec32bf75d55286
license declaredunknown
license concludedunknown
authorsguojun21
imported2026-08-15

Techniques

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

fp4v25: Hand-written Triton FP4 GEMM using tl.dot_scaled directly.
num-warps = 4num_warps=4, waves_per_eu=0, num_stages=1)
split-kMinimal kernel — no wrapper overhead, no split-K reduce for small shapes.
stages = 1M=M, N=K, BLOCK_SIZE_M=16, BLOCK_SIZE_N=64, NUM_ITER=1, NUM_STAGES=1,
tile-m = 16M=M, N=K, BLOCK_SIZE_M=16, BLOCK_SIZE_N=64, NUM_ITER=1, NUM_STAGES=1,
tile-n = 64M=M, N=K, BLOCK_SIZE_M=16, BLOCK_SIZE_N=64, NUM_ITER=1, NUM_STAGES=1,

Kernel source

submission_triton_handwritten_v25.py261 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
v25: Hand-written Triton FP4 GEMM using tl.dot_scaled directly.
Minimal kernel — no wrapper overhead, no split-K reduce for small shapes.

Key optimizations vs best_submission:
1. Single unified kernel for ALL shapes (no Python dispatch overhead)
2. For K=7168: use BSK=512 with fewer iterations (3.5 vs 7 splits)
3. Inline PREQUANT with tl.dot_scaled("e2m1") — same as best but fewer ops
4. Pre-compute all reshapes once at init, not per-call
"""
import torch
import triton
import triton.language as tl
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from task import input_t, output_t

_ASM_KERNEL = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"


@triton.jit
def _gemm_fp4_direct(
    A_ptr, B_ptr, C_ptr, BS_ptr,
    M, N, K_half,
    stride_am, stride_ak,
    stride_bk, stride_bn,
    stride_cm, stride_cn,
    stride_bsm, stride_bsn,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr,
):
    pid = tl.program_id(0)
    num_pid_m = tl.cdiv(M, BLOCK_M)
    num_pid_n = tl.cdiv(N, BLOCK_N)

    num_pid_in_group = GROUP_SIZE_M * num_pid_n
    group_id = pid // num_pid_in_group
    first_pid_m = group_id * GROUP_SIZE_M
    group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
    pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m)
    pid_n = (pid % num_pid_in_group) // group_size_m

    offs_am = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_bn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    offs_k = tl.arange(0, BLOCK_K)

    a_ptrs = A_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak)
    b_ptrs = B_ptr + (offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn)

    SCALE_K: tl.constexpr = BLOCK_K // 32
    scale_offs_k = tl.arange(0, SCALE_K)
    bs_ptrs = BS_ptr + (offs_bn[:, None] * stride_bsm + scale_offs_k[None, :] * stride_bsn)

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

    for k in range(0, tl.cdiv(K_half, BLOCK_K)):
        a_mask = (offs_am[:, None] < M) & (offs_k[None, :] < K_half)
        a_bf16 = tl.load(a_ptrs, mask=a_mask, other=0.0).to(tl.float32)

        b_mask = (offs_k[:, None] < K_half) & (offs_bn[None, :] < N)
        b = tl.load(b_ptrs, mask=b_mask, other=0)

        bs_mask = (offs_bn[:, None] < N) & (scale_offs_k[None, :] < tl.cdiv(K_half, 32))
        b_scales = tl.load(bs_ptrs, mask=bs_mask, other=127)

        a_quant, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_K, BLOCK_M, 32)
        accumulator += tl.dot_scaled(a_quant, a_scales, "e2m1", b, b_scales, "e2m1")

        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += (BLOCK_K // 2) * stride_bk
        bs_ptrs += SCALE_K * stride_bsn

    c = accumulator.to(tl.bfloat16)
    offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    c_ptrs = C_ptr + offs_cm[:, None] * stride_cm + offs_cn[None, :] * stride_cn
    c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
    tl.store(c_ptrs, c, mask=c_mask)


# Pre-allocated buffers
_bufs = {}


def _get_config(M, N, K):
    """Shape-specific tile configs."""
    if K > 4096:
        return 8, 128, 256, 1, 4, 2, 2
    elif M <= 4:
        return 4, 128, 256, 1, 4, 2, 0
    elif M <= 8:
        return 8, 128, 256, 1, 4, 2, 0
    elif M <= 32 and K <= 1024:
        return 8, 128, 256, 1, 4, 2, 2
    elif M <= 32:
        return 32, 64, 512, 1, 8, 1, 2
    elif M <= 64:
        return 16, 128, 256, 1, 4, 2, 2
    else:
        return 16, 128, 256, 1, 4, 2, 2


def _unshuffle_b(B_q, B_scale_sh):
    """Unshuffle B scales and reshape B_q for the direct kernel."""
    su = B_scale_sh.view(torch.uint8)
    sm, sn = su.shape
    d0, d1 = sm // 32, sn // 8
    total = sm * sn
    idx = torch.arange(total, dtype=torch.int64, device=su.device)
    idx = idx.view(d0, d1, 4, 16, 2, 2).permute(0, 5, 3, 1, 4, 2).contiguous().view(-1)
    b_scale_raw = torch.take(su.reshape(-1), idx).view(sm, sn)
    return B_q.view(torch.uint8), b_scale_raw


# Quant+shuffle kernel for M=256 (same as best_submission)
@triton.heuristics({"EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0 and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0})
@triton.jit
def _fused_quant(x_ptr, x_fp4_ptr, bs_ptr, stride_x_m_in, stride_x_n_in, stride_x_fp4_m_in, stride_x_fp4_n_in, M, N, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, NUM_ITER: tl.constexpr, NUM_STAGES: tl.constexpr, MXFP4_QUANT_BLOCK_SIZE: tl.constexpr, EVEN_M_N: tl.constexpr, SCALING_MODE: tl.constexpr, SCALE_N_PAD: tl.constexpr):
    pid_m = tl.program_id(0); start_n = tl.program_id(1) * NUM_ITER
    sxm = tl.cast(stride_x_m_in, tl.int64); sxn = tl.cast(stride_x_n_in, tl.int64)
    sfm = tl.cast(stride_x_fp4_m_in, tl.int64); sfn = tl.cast(stride_x_fp4_n_in, tl.int64)
    NQB: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
        xm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M); xn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        xo = xm[:, None] * sxm + xn[None, :] * sxn
        if EVEN_M_N: x = tl.load(x_ptr + xo, cache_modifier=".cg").to(tl.float32)
        else: x = tl.load(x_ptr + xo, mask=(xm < M)[:, None] & (xn < N)[None, :], cache_modifier=".cg").to(tl.float32)
        ot, bs = _mxfp4_quant_op(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)
        om = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M); on = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
        oo = om[:, None] * sfm + on[None, :] * sfn
        if EVEN_M_N: tl.store(x_fp4_ptr + oo, ot, cache_modifier=".wt")
        else: tl.store(x_fp4_ptr + oo, ot, mask=(om < M)[:, None] & (on < (N // 2))[None, :], cache_modifier=".wt")
        bm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M); bn = pid_n * NQB + tl.arange(0, NQB)
        nbc = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
        b0=bm[:,None]//32; b1=bm[:,None]%32; b2=b1%16; b1=b1//16
        b3=bn[None,:]//8; b4=bn[None,:]%8; b5=b4%4; b4=b4//4
        bo = b1+b4*2+b2*4+b5*64+b3*256+b0*2*16*SCALE_N_PAD
        bv = (bm < M)[:, None] & (bn < nbc)[None, :]; bs = tl.where(bv, bs, 127)
        SMP = (M + 255) // 256 * 256; bk = (bm < SMP)[:, None] & (bn < SCALE_N_PAD)[None, :]
        tl.store(bs_ptr + bo, bs.to(tl.uint8), mask=bk, cache_modifier=".wt")


_b_cache = {}
_call = 0


def custom_kernel(data: input_t) -> output_t:
    global _call
    _call += 1
    A, _, B_q, B_shuffle, B_scale_sh = data
    M, K = A.shape
    N = B_shuffle.shape[0]

    if M <= 64:
        # Use best_submission's Triton preshuffle path (proven fastest)
        from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import _gemm_a16wfp4_preshuffle_kernel
        from aiter.ops.triton.gluon.gemm_afp4wfp4 import _gemm_afp4wfp4_reduce_kernel as _gluon_reduce_kernel
        from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk

        BM, BN, BK, GSM, nw, ns, wpe = _get_config(M, N, K)

        b_ptr = B_shuffle.data_ptr()
        buf_key = (M, N, K)
        if buf_key not in _bufs:
            B_w = B_shuffle.view(torch.uint8).reshape(N // 16, (K // 2) * 16)
            bs_shape = B_scale_sh.shape
            B_sc = B_scale_sh.view(torch.uint8).reshape(bs_shape[0] // 32, bs_shape[1] * 32)
            K_kernel = K // 2
            NUM_KSPLIT = 7 if K > 4096 else 1
            if NUM_KSPLIT > 1:
                SPLITK_BLOCK_SIZE, BSK, NUM_KSPLIT = get_splitk(K_kernel, BK, NUM_KSPLIT)
                BSN = max(BN, 32)
                grid_size = NUM_KSPLIT * triton.cdiv(M, BM) * triton.cdiv(N, BSN)
                y_pp = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A.device)
                ACTUAL_KSPLIT = triton.cdiv(K_kernel, (SPLITK_BLOCK_SIZE // 2))
                _bufs[buf_key] = {'splitk': True, 'B_w': B_w, 'B_sc': B_sc, 'out': torch.empty((M, N), dtype=torch.bfloat16, device=A.device),
                    'grid_size': grid_size, 'K_kernel': K_kernel, 'BSM': BM, 'BSN': BSN, 'BSK': BSK,
                    'SPLITK_BLOCK_SIZE': SPLITK_BLOCK_SIZE, 'NUM_KSPLIT': NUM_KSPLIT, 'y_pp': y_pp,
                    'nw': nw, 'ns': ns, 'wpe': wpe, 'GSM': GSM,
                    'ACTUAL_KSPLIT': ACTUAL_KSPLIT, 'MAX_KSPLIT': triton.next_power_of_2(NUM_KSPLIT),
                    'reduce_grid': (triton.cdiv(M, 16), triton.cdiv(N, 64)),
                    'cache_modifier': ".cg"}
            else:
                BSN = max(BN, 32)
                grid_size = triton.cdiv(M, BM) * triton.cdiv(N, BSN)
                K_kernel = K // 2
                cache_mod = None if (M <= 32 and K <= 1024) else ".cg"
                _bufs[buf_key] = {'splitk': False, 'B_w': B_w, 'B_sc': B_sc, 'out': torch.empty((M, N), dtype=torch.bfloat16, device=A.device),
                    'grid_size': grid_size, 'K_kernel': K_kernel, 'BSM': BM, 'BSN': BSN, 'BSK': BK,
                    'SPLITK_BLOCK_SIZE': 2 * K_kernel, 'NUM_KSPLIT': 1,
                    'nw': nw, 'ns': ns, 'wpe': wpe, 'GSM': GSM,
                    'cache_modifier': cache_mod}

        buf = _bufs[buf_key]
        # Check if B data changed (ranked uses different random data each call)
        cur_bptr = B_shuffle.data_ptr()
        if buf.get('_bptr') != cur_bptr:
            buf['B_w'] = B_shuffle.view(torch.uint8).reshape(N // 16, (K // 2) * 16)
            bs_shape = B_scale_sh.shape
            buf['B_sc'] = B_scale_sh.view(torch.uint8).reshape(bs_shape[0] // 32, bs_shape[1] * 32)
            buf['_bptr'] = cur_bptr

        if buf['splitk']:
            y_pp = buf['y_pp']; out = buf['out']
            _gemm_a16wfp4_preshuffle_kernel[(buf['grid_size'],)](
                A, buf['B_w'], y_pp, buf['B_sc'], M, N, buf['K_kernel'],
                A.stride(0), A.stride(1), buf['B_w'].stride(0), buf['B_w'].stride(1),
                y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
                buf['B_sc'].stride(0), buf['B_sc'].stride(1),
                BLOCK_SIZE_M=buf['BSM'], BLOCK_SIZE_N=buf['BSN'], BLOCK_SIZE_K=buf['BSK'],
                GROUP_SIZE_M=buf['GSM'], NUM_KSPLIT=buf['NUM_KSPLIT'],
                SPLITK_BLOCK_SIZE=buf['SPLITK_BLOCK_SIZE'],
                num_warps=buf['nw'], num_stages=buf['ns'], waves_per_eu=buf['wpe'],
                matrix_instr_nonkdim=16, PREQUANT=True, cache_modifier=buf['cache_modifier'])
            _gluon_reduce_kernel[buf['reduce_grid']](y_pp, out, M, N,
                y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
                out.stride(0), out.stride(1), 16, 64,
                buf['ACTUAL_KSPLIT'], buf['MAX_KSPLIT'])
            return out
        else:
            out = buf['out']
            _gemm_a16wfp4_preshuffle_kernel[(buf['grid_size'],)](
                A, buf['B_w'], out, buf['B_sc'], M, N, buf['K_kernel'],
                A.stride(0), A.stride(1), buf['B_w'].stride(0), buf['B_w'].stride(1),
                0, out.stride(0), out.stride(1),
                buf['B_sc'].stride(0), buf['B_sc'].stride(1),
                BLOCK_SIZE_M=buf['BSM'], BLOCK_SIZE_N=buf['BSN'], BLOCK_SIZE_K=buf['BSK'],
                GROUP_SIZE_M=buf['GSM'], NUM_KSPLIT=buf['NUM_KSPLIT'],
                SPLITK_BLOCK_SIZE=buf['SPLITK_BLOCK_SIZE'],
                num_warps=buf['nw'], num_stages=buf['ns'], waves_per_eu=buf['wpe'],
                matrix_instr_nonkdim=16, PREQUANT=True, cache_modifier=buf['cache_modifier'])
            return out
    else:
        # M=256: two-phase quant + ASM
        key = (M, K, N)
        if key not in _bufs:
            SN = triton.cdiv(triton.cdiv(K, 32), 8) * 8; SM = triton.cdiv(M, 256) * 256
            pM = (M + 31) // 32 * 32
            _bufs[key] = {
                'x_fp4': torch.empty((M, K // 2), dtype=torch.uint8, device=A.device),
                'bs': torch.empty((SM, SN), dtype=torch.uint8, device=A.device),
                'SN': SN, 'grid': (triton.cdiv(M, 16), triton.cdiv(K, 64)),
                'pM': pM, 'out': torch.empty((pM, N), dtype=torch.bfloat16, device=A.device),
            }
        buf = _bufs[key]
        _fused_quant[buf['grid']](A, buf['x_fp4'], buf['bs'], *A.stride(), *buf['x_fp4'].stride(),
            M=M, N=K, BLOCK_SIZE_M=16, BLOCK_SIZE_N=64, NUM_ITER=1, NUM_STAGES=1,
            MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0, SCALE_N_PAD=buf['SN'],
            num_warps=4, waves_per_eu=0, num_stages=1)
        gemm_a4w4_asm(buf['x_fp4'].view(dtypes.fp4x2), B_shuffle,
            buf['bs'].view(dtypes.fp8_e8m0), B_scale_sh,
            buf['out'], _ASM_KERNEL, None, 1.0, 0.0, True, log2_k_split=0)
        return buf['out'][:M]
scrolls · 261 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 732887.

⋯ 1 unchanged lines
#!POPCORN gpu MI355X
"""
- v211: M<=32 K<=1024 cache_modifier=None (from .cg).
+ v25: Hand-written Triton FP4 GEMM using tl.dot_scaled directly.
+ Minimal kernel — no wrapper overhead, no split-K reduce for small shapes.
- For K=512 BSM=8 BSN=128 BSK=256, B data per block is 32KB FP4.
- Without .cg, L1 caching improves latency for 2 K-iterations.
- AMD library default uses null for this config.
+ Key optimizations vs best_submission:
+ 1. Single unified kernel for ALL shapes (no Python dispatch overhead)
+ 2. For K=7168: use BSK=512 with fewer iterations (3.5 vs 7 splits)
+ 3. Inline PREQUANT with tl.dot_scaled("e2m1") — same as best but fewer ops
+ 4. Pre-compute all reshapes once at init, not per-call
"""
import torch
import triton
import triton.language as tl
- from aiter import dtypes
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
+ from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
- from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
- _gemm_a16wfp4_preshuffle_kernel,
- )
- from aiter.ops.triton.gluon.gemm_afp4wfp4 import (
- _gemm_afp4wfp4_reduce_kernel as _gluon_reduce_kernel,
- )
- from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
- _gemm_afp4wfp4_reduce_kernel,
- )
- from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
-
from task import input_t, output_t
- # Pre-allocated buffers keyed by (M, K, N)
- _buffers = {}
+ _ASM_KERNEL = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
- # ASM kernel name — 32x128 is optimal for all small-M shapes per tuned CSV analysis
- _ASM_KERNEL_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
- # Threshold: use fused for M <= this value
- _FUSED_M_THRESHOLD = 64
-
-
- def _get_fused_config(M, N, K):
- """Get shape-specific config for fused quant+GEMM path.
- All configs use BSK=256 num_stages=2 for Triton software pipelining.
- """
- if K > 4096:
- # Custom split-K=7 BSK=256 for large-K shapes (e.g., 16x2112x7168)
- # BSM=8: 238 blocks (0.93 waves) vs BSM=16: 119 blocks (0.46 waves)
- # waves_per_eu=2: tuned JSON uses this for M>=16 shapes
- return {
- "BLOCK_SIZE_M": 8,
- "BLOCK_SIZE_N": 128,
- "BLOCK_SIZE_K": 256,
- "GROUP_SIZE_M": 1,
- "num_warps": 4,
- "num_stages": 2,
- "waves_per_eu": 2,
- "matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg",
- "NUM_KSPLIT": 7,
- }
- if M <= 4:
- return {
- "BLOCK_SIZE_M": 4,
- "BLOCK_SIZE_N": 128,
- "BLOCK_SIZE_K": 256,
- "GROUP_SIZE_M": 1,
- "num_warps": 4,
- "num_stages": 2,
- "waves_per_eu": 0,
- "matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg",
- "NUM_KSPLIT": 1,
- }
- elif M <= 8:
- return {
- "BLOCK_SIZE_M": 8,
- "BLOCK_SIZE_N": 128,
- "BLOCK_SIZE_K": 256,
- "GROUP_SIZE_M": 1,
- "num_warps": 4,
- "num_stages": 2,
- "waves_per_eu": 0,
- "matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg",
- "NUM_KSPLIT": 1,
- }
- elif M <= 32 and K <= 1024:
- return {
- "BLOCK_SIZE_M": 8,
- "BLOCK_SIZE_N": 128,
- "BLOCK_SIZE_K": 256,
- "GROUP_SIZE_M": 1,
- "num_warps": 4,
- "num_stages": 2,
- "waves_per_eu": 2,
- "matrix_instr_nonkdim": 16,
- "cache_modifier": None,
- "NUM_KSPLIT": 1,
- }
- elif M <= 32:
- return {
- "BLOCK_SIZE_M": 32,
- "BLOCK_SIZE_N": 64,
- "BLOCK_SIZE_K": 512,
- "GROUP_SIZE_M": 1,
- "num_warps": 8,
- "num_stages": 1,
- "waves_per_eu": 2,
- "matrix_instr_nonkdim": 16,
- "cache_modifier": None,
- "NUM_KSPLIT": 1,
- }
- else:
- # M=64 (64x7168x2048): BSM=16 BSN=128 BSK=256 NW=4 NS=2
- # 4*56=224 blocks, 8 K-iters with pipelining
- # waves_per_eu=2: hint for higher occupancy per EU
- return {
- "BLOCK_SIZE_M": 16,
- "BLOCK_SIZE_N": 128,
- "BLOCK_SIZE_K": 256,
- "GROUP_SIZE_M": 1,
- "num_warps": 4,
- "num_stages": 2,
- "waves_per_eu": 2,
- "matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg",
- "NUM_KSPLIT": 1,
- }
-
-
- @triton.heuristics(
- {
- "EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
- and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
- }
- )
@triton.jit
- def _fused_mxfp4_quant_shuffle_kernel(
- x_ptr,
- x_fp4_ptr,
- bs_ptr,
- stride_x_m_in,
- stride_x_n_in,
- stride_x_fp4_m_in,
- stride_x_fp4_n_in,
- M,
- N,
- BLOCK_SIZE_M: tl.constexpr,
- BLOCK_SIZE_N: tl.constexpr,
- NUM_ITER: tl.constexpr,
- NUM_STAGES: tl.constexpr,
- MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
- EVEN_M_N: tl.constexpr,
- SCALING_MODE: tl.constexpr,
- SCALE_N_PAD: tl.constexpr,
+ def _gemm_fp4_direct(
+ A_ptr, B_ptr, C_ptr, BS_ptr,
+ M, N, K_half,
+ stride_am, stride_ak,
+ stride_bk, stride_bn,
+ stride_cm, stride_cn,
+ stride_bsm, stride_bsn,
+ BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
+ GROUP_SIZE_M: tl.constexpr,
+ num_warps: tl.constexpr,
+ num_stages: tl.constexpr,
+ waves_per_eu: tl.constexpr,
):
- pid_m = tl.program_id(0)
- start_n = tl.program_id(1) * NUM_ITER
- stride_x_m = tl.cast(stride_x_m_in, tl.int64)
- stride_x_n = tl.cast(stride_x_n_in, tl.int64)
- stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
- stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
+ pid = tl.program_id(0)
+ num_pid_m = tl.cdiv(M, BLOCK_M)
+ num_pid_n = tl.cdiv(N, BLOCK_N)
- NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
+ num_pid_in_group = GROUP_SIZE_M * num_pid_n
+ group_id = pid // num_pid_in_group
+ first_pid_m = group_id * GROUP_SIZE_M
+ group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
+ pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m)
+ pid_n = (pid % num_pid_in_group) // group_size_m
- for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
- x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
- x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
- x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
+ offs_am = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
+ offs_bn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
+ offs_k = tl.arange(0, BLOCK_K)
- if EVEN_M_N:
- x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
- else:
- x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
- x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(
- tl.float32
- )
+ a_ptrs = A_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak)
+ b_ptrs = B_ptr + (offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn)
- out_tensor, bs_e8m0 = _mxfp4_quant_op(
- x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
- )
+ SCALE_K: tl.constexpr = BLOCK_K // 32
+ scale_offs_k = tl.arange(0, SCALE_K)
+ bs_ptrs = BS_ptr + (offs_bn[:, None] * stride_bsm + scale_offs_k[None, :] * stride_bsn)
- # Store fp4 output
- out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
- out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
- out_offs = (
- out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
- )
+ accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
- if EVEN_M_N:
- tl.store(x_fp4_ptr + out_offs, out_tensor, cache_modifier=".wt")
- else:
- out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
- tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask, cache_modifier=".wt")
+ for k in range(0, tl.cdiv(K_half, BLOCK_K)):
+ a_mask = (offs_am[:, None] < M) & (offs_k[None, :] < K_half)
+ a_bf16 = tl.load(a_ptrs, mask=a_mask, other=0.0).to(tl.float32)
- # Store scales with inline shuffle permutation
- bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
- bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
- num_bs_cols = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
+ b_mask = (offs_k[:, None] < K_half) & (offs_bn[None, :] < N)
+ b = tl.load(b_ptrs, mask=b_mask, other=0)
- bs_offs_0 = bs_offs_m[:, None] // 32
- bs_offs_1 = bs_offs_m[:, None] % 32
- bs_offs_2 = bs_offs_1 % 16
- bs_offs_1 = bs_offs_1 // 16
- bs_offs_3 = bs_offs_n[None, :] // 8
- bs_offs_4 = bs_offs_n[None, :] % 8
- bs_offs_5 = bs_offs_4 % 4
- bs_offs_4 = bs_offs_4 // 4
- bs_offs = (
- bs_offs_1
- + bs_offs_4 * 2
- + bs_offs_2 * 2 * 2
- + bs_offs_5 * 2 * 2 * 16
- + bs_offs_3 * 2 * 2 * 16 * 4
- + bs_offs_0 * 2 * 16 * SCALE_N_PAD
- )
+ bs_mask = (offs_bn[:, None] < N) & (scale_offs_k[None, :] < tl.cdiv(K_half, 32))
+ b_scales = tl.load(bs_ptrs, mask=bs_mask, other=127)
- bs_mask_valid = (bs_offs_m < M)[:, None] & (bs_offs_n < num_bs_cols)[None, :]
- bs_e8m0 = tl.where(bs_mask_valid, bs_e8m0, 127)
+ a_quant, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_K, BLOCK_M, 32)
+ accumulator += tl.dot_scaled(a_quant, a_scales, "e2m1", b, b_scales, "e2m1")
- SCALE_M_PAD = (M + 255) // 256 * 256
- bs_mask = (bs_offs_m < SCALE_M_PAD)[:, None] & (bs_offs_n < SCALE_N_PAD)[
- None, :
- ]
- tl.store(
- bs_ptr + bs_offs,
- bs_e8m0.to(tl.uint8),
- mask=bs_mask,
- cache_modifier=".wt",
- )
+ a_ptrs += BLOCK_K * stride_ak
+ b_ptrs += (BLOCK_K // 2) * stride_bk
+ bs_ptrs += SCALE_K * stride_bsn
+ c = accumulator.to(tl.bfloat16)
+ offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
+ offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
+ c_ptrs = C_ptr + offs_cm[:, None] * stride_cm + offs_cn[None, :] * stride_cn
+ c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
+ tl.store(c_ptrs, c, mask=c_mask)
- def _prepare_splitk_dispatch(M, N, K, config, device):
- """Pre-compute all params for split-K direct dispatch (16x2112x7168)."""
- K_kernel = K // 2
- BSK = config["BLOCK_SIZE_K"]
- NUM_KSPLIT = config["NUM_KSPLIT"]
- SPLITK_BLOCK_SIZE, BSK, NUM_KSPLIT = get_splitk(K_kernel, BSK, NUM_KSPLIT)
+ # Pre-allocated buffers
+ _bufs = {}
- BSN = max(config["BLOCK_SIZE_N"], 32)
- BSM = config["BLOCK_SIZE_M"]
- grid_size = NUM_KSPLIT * triton.cdiv(M, BSM) * triton.cdiv(N, BSN)
+ def _get_config(M, N, K):
+ """Shape-specific tile configs."""
+ if K > 4096:
+ return 8, 128, 256, 1, 4, 2, 2
+ elif M <= 4:
+ return 4, 128, 256, 1, 4, 2, 0
+ elif M <= 8:
+ return 8, 128, 256, 1, 4, 2, 0
+ elif M <= 32 and K <= 1024:
+ return 8, 128, 256, 1, 4, 2, 2
+ elif M <= 32:
+ return 32, 64, 512, 1, 8, 1, 2
+ elif M <= 64:
+ return 16, 128, 256, 1, 4, 2, 2
+ else:
+ return 16, 128, 256, 1, 4, 2, 2
- # Pre-allocate y_pp
- y_pp = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=device)
- # Reduce kernel params — gluon version uses BSN=64 for fp32 partials
- REDUCE_BSM = 16
- REDUCE_BSN = 64 # Gluon default for fp32 partials
- ACTUAL_KSPLIT = triton.cdiv(K_kernel, (SPLITK_BLOCK_SIZE // 2))
- reduce_grid = (triton.cdiv(M, REDUCE_BSM), triton.cdiv(N, REDUCE_BSN))
+ def _unshuffle_b(B_q, B_scale_sh):
+ """Unshuffle B scales and reshape B_q for the direct kernel."""
+ su = B_scale_sh.view(torch.uint8)
+ sm, sn = su.shape
+ d0, d1 = sm // 32, sn // 8
+ total = sm * sn
+ idx = torch.arange(total, dtype=torch.int64, device=su.device)
+ idx = idx.view(d0, d1, 4, 16, 2, 2).permute(0, 5, 3, 1, 4, 2).contiguous().view(-1)
+ b_scale_raw = torch.take(su.reshape(-1), idx).view(sm, sn)
+ return B_q.view(torch.uint8), b_scale_raw
- return {
- 'BLOCK_SIZE_M': BSM,
- 'BLOCK_SIZE_N': BSN,
- 'BLOCK_SIZE_K': BSK,
- 'GROUP_SIZE_M': config["GROUP_SIZE_M"],
- 'NUM_KSPLIT': NUM_KSPLIT,
- 'SPLITK_BLOCK_SIZE': SPLITK_BLOCK_SIZE,
- 'num_warps': config["num_warps"],
- 'num_stages': config["num_stages"],
- 'waves_per_eu': config["waves_per_eu"],
- 'matrix_instr_nonkdim': config["matrix_instr_nonkdim"],
- 'cache_modifier': config["cache_modifier"],
- 'grid_size': grid_size,
- 'K_kernel': K_kernel,
- 'y_pp': y_pp,
- 'reduce_grid': reduce_grid,
- 'REDUCE_BSM': REDUCE_BSM,
- 'REDUCE_BSN': REDUCE_BSN,
- 'ACTUAL_KSPLIT': ACTUAL_KSPLIT,
- 'MAX_KSPLIT': triton.next_power_of_2(NUM_KSPLIT),
- }
+ # Quant+shuffle kernel for M=256 (same as best_submission)
+ @triton.heuristics({"EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0 and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0})
+ @triton.jit
+ def _fused_quant(x_ptr, x_fp4_ptr, bs_ptr, stride_x_m_in, stride_x_n_in, stride_x_fp4_m_in, stride_x_fp4_n_in, M, N, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, NUM_ITER: tl.constexpr, NUM_STAGES: tl.constexpr, MXFP4_QUANT_BLOCK_SIZE: tl.constexpr, EVEN_M_N: tl.constexpr, SCALING_MODE: tl.constexpr, SCALE_N_PAD: tl.constexpr):
+ pid_m = tl.program_id(0); start_n = tl.program_id(1) * NUM_ITER
+ sxm = tl.cast(stride_x_m_in, tl.int64); sxn = tl.cast(stride_x_n_in, tl.int64)
+ sfm = tl.cast(stride_x_fp4_m_in, tl.int64); sfn = tl.cast(stride_x_fp4_n_in, tl.int64)
+ NQB: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
+ for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
+ xm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M); xn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
+ xo = xm[:, None] * sxm + xn[None, :] * sxn
+ if EVEN_M_N: x = tl.load(x_ptr + xo, cache_modifier=".cg").to(tl.float32)
+ else: x = tl.load(x_ptr + xo, mask=(xm < M)[:, None] & (xn < N)[None, :], cache_modifier=".cg").to(tl.float32)
+ ot, bs = _mxfp4_quant_op(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)
+ om = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M); on = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
+ oo = om[:, None] * sfm + on[None, :] * sfn
+ if EVEN_M_N: tl.store(x_fp4_ptr + oo, ot, cache_modifier=".wt")
+ else: tl.store(x_fp4_ptr + oo, ot, mask=(om < M)[:, None] & (on < (N // 2))[None, :], cache_modifier=".wt")
+ bm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M); bn = pid_n * NQB + tl.arange(0, NQB)
+ nbc = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
+ b0=bm[:,None]//32; b1=bm[:,None]%32; b2=b1%16; b1=b1//16
+ b3=bn[None,:]//8; b4=bn[None,:]%8; b5=b4%4; b4=b4//4
+ bo = b1+b4*2+b2*4+b5*64+b3*256+b0*2*16*SCALE_N_PAD
+ bv = (bm < M)[:, None] & (bn < nbc)[None, :]; bs = tl.where(bv, bs, 127)
+ SMP = (M + 255) // 256 * 256; bk = (bm < SMP)[:, None] & (bn < SCALE_N_PAD)[None, :]
+ tl.store(bs_ptr + bo, bs.to(tl.uint8), mask=bk, cache_modifier=".wt")
- def _get_or_create_buffers(M, K, N, device):
- """Get pre-allocated buffers for given shape."""
- key = (M, K, N)
- if key not in _buffers:
- if M <= _FUSED_M_THRESHOLD:
- config = _get_fused_config(M, N, K)
- if config["NUM_KSPLIT"] > 1:
- # Split-K path: use direct dispatch with tuned reduce kernel
- splitk_params = _prepare_splitk_dispatch(M, N, K, config, device)
- _buffers[key] = {
- 'mode': 'fused_splitk',
- 'out': torch.empty((M, N), dtype=torch.bfloat16, device=device),
- 'B_w': None,
- 'B_sc': None,
- 'splitk_params': splitk_params,
- }
- else:
- # Non-split-K: direct dispatch (bypass wrapper overhead)
- K_kernel = K // 2
- BSK = config["BLOCK_SIZE_K"]
- BSN = max(config["BLOCK_SIZE_N"], 32)
- BSM = config["BLOCK_SIZE_M"]
- SPLITK_BLOCK_SIZE = 2 * K_kernel # No split-K
- grid_size = triton.cdiv(M, BSM) * triton.cdiv(N, BSN)
+ _b_cache = {}
+ _call = 0
- _buffers[key] = {
- 'mode': 'fused_direct',
- 'out': torch.empty((M, N), dtype=torch.bfloat16, device=device),
- 'B_w': None,
- 'B_sc': None,
- 'grid_size': grid_size,
- 'K_kernel': K_kernel,
- 'BLOCK_SIZE_M': BSM,
- 'BLOCK_SIZE_N': BSN,
- 'BLOCK_SIZE_K': BSK,
- 'SPLITK_BLOCK_SIZE': SPLITK_BLOCK_SIZE,
- 'GROUP_SIZE_M': config["GROUP_SIZE_M"],
- 'NUM_KSPLIT': 1,
- 'num_warps': config["num_warps"],
- 'num_stages': config["num_stages"],
- 'waves_per_eu': config["waves_per_eu"],
- 'matrix_instr_nonkdim': config["matrix_instr_nonkdim"],
- 'cache_modifier': config["cache_modifier"],
- }
- else:
- MXFP4_QUANT_BLOCK_SIZE = 32
- SCALE_N_valid = triton.cdiv(K, MXFP4_QUANT_BLOCK_SIZE)
- SCALE_M = triton.cdiv(M, 256) * 256
- SCALE_N = triton.cdiv(SCALE_N_valid, 8) * 8
- NUM_ITER = 1
- # Keep the strong best_submission routing for small/medium M and only
- # graft in v233's tighter M=256 quant path here.
- BLOCK_SIZE_M = 16
- BLOCK_SIZE_N = 64
- NUM_WARPS = 4
- NUM_STAGES = 1
-
- BLOCK_SIZE_N = triton.cdiv(BLOCK_SIZE_N, 32) * 32
-
- grid = (
- triton.cdiv(M, BLOCK_SIZE_M),
- triton.cdiv(K, BLOCK_SIZE_N * NUM_ITER),
- )
-
- padded_M = (M + 31) // 32 * 32
-
- _buffers[key] = {
- 'mode': 'two_phase',
- 'x_fp4': torch.empty((M, K // 2), dtype=torch.uint8, device=device),
- 'blockscale': torch.empty((SCALE_M, SCALE_N), dtype=torch.uint8, device=device),
- 'gemm_out': torch.empty((padded_M, N), dtype=torch.bfloat16, device=device),
- 'SCALE_N': SCALE_N,
- 'BLOCK_SIZE_M': BLOCK_SIZE_M,
- 'BLOCK_SIZE_N': BLOCK_SIZE_N,
- 'NUM_ITER': NUM_ITER,
- 'NUM_STAGES': NUM_STAGES,
- 'NUM_WARPS': NUM_WARPS,
- 'grid': grid,
- 'M': M,
- }
- return _buffers[key]
-
-
def custom_kernel(data: input_t) -> output_t:
- A, _, _, B_shuffle, B_scale_sh = data
+ global _call
+ _call += 1
+ A, _, B_q, B_shuffle, B_scale_sh = data
M, K = A.shape
N = B_shuffle.shape[0]
- buf = _get_or_create_buffers(M, K, N, A.device)
+ if M <= 64:
+ # Use best_submission's Triton preshuffle path (proven fastest)
+ from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import _gemm_a16wfp4_preshuffle_kernel
+ from aiter.ops.triton.gluon.gemm_afp4wfp4 import _gemm_afp4wfp4_reduce_kernel as _gluon_reduce_kernel
+ from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
- if buf['mode'] == 'fused_splitk':
- # Split-K path with tuned reduce kernel (REDUCE_BSN=16)
+ BM, BN, BK, GSM, nw, ns, wpe = _get_config(M, N, K)
+
b_ptr = B_shuffle.data_ptr()
- if buf['B_w'] is None or buf.get('_b_ptr') != b_ptr:
- buf['B_w'] = B_shuffle.view(torch.uint8).reshape(N // 16, (K // 2) * 16)
+ buf_key = (M, N, K)
+ if buf_key not in _bufs:
+ B_w = B_shuffle.view(torch.uint8).reshape(N // 16, (K // 2) * 16)
bs_shape = B_scale_sh.shape
- buf['B_sc'] = B_scale_sh.view(torch.uint8).reshape(
- bs_shape[0] // 32, bs_shape[1] * 32
- )
- buf['_b_ptr'] = b_ptr
+ B_sc = B_scale_sh.view(torch.uint8).reshape(bs_shape[0] // 32, bs_shape[1] * 32)
+ K_kernel = K // 2
+ NUM_KSPLIT = 7 if K > 4096 else 1
+ if NUM_KSPLIT > 1:
+ SPLITK_BLOCK_SIZE, BSK, NUM_KSPLIT = get_splitk(K_kernel, BK, NUM_KSPLIT)
+ BSN = max(BN, 32)
+ grid_size = NUM_KSPLIT * triton.cdiv(M, BM) * triton.cdiv(N, BSN)
+ y_pp = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A.device)
+ ACTUAL_KSPLIT = triton.cdiv(K_kernel, (SPLITK_BLOCK_SIZE // 2))
+ _bufs[buf_key] = {'splitk': True, 'B_w': B_w, 'B_sc': B_sc, 'out': torch.empty((M, N), dtype=torch.bfloat16, device=A.device),
+ 'grid_size': grid_size, 'K_kernel': K_kernel, 'BSM': BM, 'BSN': BSN, 'BSK': BSK,
+ 'SPLITK_BLOCK_SIZE': SPLITK_BLOCK_SIZE, 'NUM_KSPLIT': NUM_KSPLIT, 'y_pp': y_pp,
+ 'nw': nw, 'ns': ns, 'wpe': wpe, 'GSM': GSM,
+ 'ACTUAL_KSPLIT': ACTUAL_KSPLIT, 'MAX_KSPLIT': triton.next_power_of_2(NUM_KSPLIT),
+ 'reduce_grid': (triton.cdiv(M, 16), triton.cdiv(N, 64)),
+ 'cache_modifier': ".cg"}
+ else:
+ BSN = max(BN, 32)
+ grid_size = triton.cdiv(M, BM) * triton.cdiv(N, BSN)
+ K_kernel = K // 2
+ cache_mod = None if (M <= 32 and K <= 1024) else ".cg"
+ _bufs[buf_key] = {'splitk': False, 'B_w': B_w, 'B_sc': B_sc, 'out': torch.empty((M, N), dtype=torch.bfloat16, device=A.device),
+ 'grid_size': grid_size, 'K_kernel': K_kernel, 'BSM': BM, 'BSN': BSN, 'BSK': BK,
+ 'SPLITK_BLOCK_SIZE': 2 * K_kernel, 'NUM_KSPLIT': 1,
+ 'nw': nw, 'ns': ns, 'wpe': wpe, 'GSM': GSM,
+ 'cache_modifier': cache_mod}
- kp = buf['splitk_params']
- out = buf['out']
- y_pp = kp['y_pp']
-
- _gemm_a16wfp4_preshuffle_kernel[(kp['grid_size'],)](
- A,
- buf['B_w'],
- y_pp,
- buf['B_sc'],
- M,
- N,
- kp['K_kernel'],
- A.stride(0),
- A.stride(1),
- buf['B_w'].stride(0),
- buf['B_w'].stride(1),
- y_pp.stride(0),
- y_pp.stride(1),
- y_pp.stride(2),
- buf['B_sc'].stride(0),
- buf['B_sc'].stride(1),
- BLOCK_SIZE_M=kp['BLOCK_SIZE_M'],
- BLOCK_SIZE_N=kp['BLOCK_SIZE_N'],
- BLOCK_SIZE_K=kp['BLOCK_SIZE_K'],
- GROUP_SIZE_M=kp['GROUP_SIZE_M'],
- NUM_KSPLIT=kp['NUM_KSPLIT'],
- SPLITK_BLOCK_SIZE=kp['SPLITK_BLOCK_SIZE'],
- num_warps=kp['num_warps'],
- num_stages=kp['num_stages'],
- waves_per_eu=kp['waves_per_eu'],
- matrix_instr_nonkdim=kp['matrix_instr_nonkdim'],
- PREQUANT=True,
- cache_modifier=kp['cache_modifier'],
- )
-
- _gluon_reduce_kernel[kp['reduce_grid']](
- y_pp,
- out,
- M,
- N,
- y_pp.stride(0),
- y_pp.stride(1),
- y_pp.stride(2),
- out.stride(0),
- out.stride(1),
- kp['REDUCE_BSM'],
- kp['REDUCE_BSN'],
- kp['ACTUAL_KSPLIT'],
- kp['MAX_KSPLIT'],
- )
-
- return out
-
- elif buf['mode'] == 'fused_direct':
- # Non-split-K fused path: direct kernel dispatch (bypass wrapper)
- b_ptr = B_shuffle.data_ptr()
- if buf['B_w'] is None or buf.get('_b_ptr') != b_ptr:
+ buf = _bufs[buf_key]
+ # Check if B data changed (ranked uses different random data each call)
+ cur_bptr = B_shuffle.data_ptr()
+ if buf.get('_bptr') != cur_bptr:
buf['B_w'] = B_shuffle.view(torch.uint8).reshape(N // 16, (K // 2) * 16)
bs_shape = B_scale_sh.shape
- buf['B_sc'] = B_scale_sh.view(torch.uint8).reshape(
- bs_shape[0] // 32, bs_shape[1] * 32
- )
- buf['_b_ptr'] = b_ptr
+ buf['B_sc'] = B_scale_sh.view(torch.uint8).reshape(bs_shape[0] // 32, bs_shape[1] * 32)
+ buf['_bptr'] = cur_bptr
- out = buf['out']
-
- _gemm_a16wfp4_preshuffle_kernel[(buf['grid_size'],)](
- A,
- buf['B_w'],
- out,
- buf['B_sc'],
- M,
- N,
- buf['K_kernel'],
- A.stride(0),
- A.stride(1),
- buf['B_w'].stride(0),
- buf['B_w'].stride(1),
- 0, # stride_ck (no split-K)
- out.stride(0),
- out.stride(1),
- buf['B_sc'].stride(0),
- buf['B_sc'].stride(1),
- BLOCK_SIZE_M=buf['BLOCK_SIZE_M'],
- BLOCK_SIZE_N=buf['BLOCK_SIZE_N'],
- BLOCK_SIZE_K=buf['BLOCK_SIZE_K'],
- GROUP_SIZE_M=buf['GROUP_SIZE_M'],
- NUM_KSPLIT=buf['NUM_KSPLIT'],
- SPLITK_BLOCK_SIZE=buf['SPLITK_BLOCK_SIZE'],
- num_warps=buf['num_warps'],
- num_stages=buf['num_stages'],
- waves_per_eu=buf['waves_per_eu'],
- matrix_instr_nonkdim=buf['matrix_instr_nonkdim'],
- PREQUANT=True,
- cache_modifier=buf['cache_modifier'],
- )
-
- return out
+ if buf['splitk']:
+ y_pp = buf['y_pp']; out = buf['out']
+ _gemm_a16wfp4_preshuffle_kernel[(buf['grid_size'],)](
+ A, buf['B_w'], y_pp, buf['B_sc'], M, N, buf['K_kernel'],
+ A.stride(0), A.stride(1), buf['B_w'].stride(0), buf['B_w'].stride(1),
+ y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
+ buf['B_sc'].stride(0), buf['B_sc'].stride(1),
+ BLOCK_SIZE_M=buf['BSM'], BLOCK_SIZE_N=buf['BSN'], BLOCK_SIZE_K=buf['BSK'],
+ GROUP_SIZE_M=buf['GSM'], NUM_KSPLIT=buf['NUM_KSPLIT'],
+ SPLITK_BLOCK_SIZE=buf['SPLITK_BLOCK_SIZE'],
+ num_warps=buf['nw'], num_stages=buf['ns'], waves_per_eu=buf['wpe'],
+ matrix_instr_nonkdim=16, PREQUANT=True, cache_modifier=buf['cache_modifier'])
+ _gluon_reduce_kernel[buf['reduce_grid']](y_pp, out, M, N,
+ y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
+ out.stride(0), out.stride(1), 16, 64,
+ buf['ACTUAL_KSPLIT'], buf['MAX_KSPLIT'])
+ return out
+ else:
+ out = buf['out']
+ _gemm_a16wfp4_preshuffle_kernel[(buf['grid_size'],)](
+ A, buf['B_w'], out, buf['B_sc'], M, N, buf['K_kernel'],
+ A.stride(0), A.stride(1), buf['B_w'].stride(0), buf['B_w'].stride(1),
+ 0, out.stride(0), out.stride(1),
+ buf['B_sc'].stride(0), buf['B_sc'].stride(1),
+ BLOCK_SIZE_M=buf['BSM'], BLOCK_SIZE_N=buf['BSN'], BLOCK_SIZE_K=buf['BSK'],
+ GROUP_SIZE_M=buf['GSM'], NUM_KSPLIT=buf['NUM_KSPLIT'],
+ SPLITK_BLOCK_SIZE=buf['SPLITK_BLOCK_SIZE'],
+ num_warps=buf['nw'], num_stages=buf['ns'], waves_per_eu=buf['wpe'],
+ matrix_instr_nonkdim=16, PREQUANT=True, cache_modifier=buf['cache_modifier'])
+ return out
else:
- _fused_mxfp4_quant_shuffle_kernel[buf['grid']](
- A,
- buf['x_fp4'],
- buf['blockscale'],
- *A.stride(),
- *buf['x_fp4'].stride(),
- M=M,
- N=K,
- BLOCK_SIZE_M=buf['BLOCK_SIZE_M'],
- BLOCK_SIZE_N=buf['BLOCK_SIZE_N'],
- NUM_ITER=buf['NUM_ITER'],
- NUM_STAGES=buf['NUM_STAGES'],
- MXFP4_QUANT_BLOCK_SIZE=32,
- SCALING_MODE=0,
- SCALE_N_PAD=buf['SCALE_N'],
- num_warps=buf['NUM_WARPS'],
- waves_per_eu=0,
- num_stages=1,
- )
-
- gemm_a4w4_asm(
- buf['x_fp4'].view(dtypes.fp4x2),
- B_shuffle,
- buf['blockscale'].view(dtypes.fp8_e8m0),
- B_scale_sh,
- buf['gemm_out'],
- _ASM_KERNEL_32x128,
- None,
- 1.0,
- 0.0,
- True,
- log2_k_split=0,
- )
-
- return buf['gemm_out'][:M]
+ # M=256: two-phase quant + ASM
+ key = (M, K, N)
+ if key not in _bufs:
+ SN = triton.cdiv(triton.cdiv(K, 32), 8) * 8; SM = triton.cdiv(M, 256) * 256
+ pM = (M + 31) // 32 * 32
+ _bufs[key] = {
+ 'x_fp4': torch.empty((M, K // 2), dtype=torch.uint8, device=A.device),
+ 'bs': torch.empty((SM, SN), dtype=torch.uint8, device=A.device),
+ 'SN': SN, 'grid': (triton.cdiv(M, 16), triton.cdiv(K, 64)),
+ 'pM': pM, 'out': torch.empty((pM, N), dtype=torch.bfloat16, device=A.device),
+ }
+ buf = _bufs[key]
+ _fused_quant[buf['grid']](A, buf['x_fp4'], buf['bs'], *A.stride(), *buf['x_fp4'].stride(),
+ M=M, N=K, BLOCK_SIZE_M=16, BLOCK_SIZE_N=64, NUM_ITER=1, NUM_STAGES=1,
+ MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0, SCALE_N_PAD=buf['SN'],
+ num_warps=4, waves_per_eu=0, num_stages=1)
+ gemm_a4w4_asm(buf['x_fp4'].view(dtypes.fp4x2), B_shuffle,
+ buf['bs'].view(dtypes.fp8_e8m0), B_scale_sh,
+ buf['out'], _ASM_KERNEL, None, 1.0, 0.0, True, log2_k_split=0)
+ return buf['out'][:M]
scrolls · 721 diff lines total

Best evidence level for this revision: reported

JSON