Skip to content
KernelIndex
Search⌘K

submission 730077

johnny.t.shi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:693da62df42d122e18b819d997a1ac4b9abc5038e53c71dcc9f72b37564a5491
license declaredunknown
license concludedunknown
authorsjohnny.t.shi
imported2026-08-15

Techniques

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

split-k"""v911: NK=14 for split-K shape — maximize CU utilization.

Kernel source

v911_nk14.py176 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""v911: NK=14 for split-K shape — maximize CU utilization.
NK=7 → 119 blocks (46% CU). NK=14 → 238 blocks (93% CU).
Each block does 1 K-iter instead of 2. More parallelism for memory-bound shape.
Trade: reduce sums 14 partials instead of 7."""
import os
os.environ['DISABLE_LLVM_OPT'] = 'disable-lsr,disable-machine-licm,disable-machine-sink'
os.environ['TRITON_HIP_USE_BLOCK_PINGPONG'] = '1'
import torch, triton, triton.language as tl
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid
from aiter.ops.triton.gluon.gemm_afp4wfp4 import _gemm_afp4wfp4_reduce_kernel as _gluon_reduce_kernel
import aiter.ops.triton.gemm_afp4wfp4 as _gm
_gs=_gm.get_splitk
from task import input_t, output_t

# ===== NX=2 XCD remap helper =====
@triton.jit
def _rx(p, n, NX: tl.constexpr):
    cs = tl.cdiv(n, NX)
    return (p % NX) * cs + p // NX

# ===== FORKED preshuffle kernel with NX=2 XCD remap =====
@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 _ps_nx2(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)

    if pid_unified < GRID_MN * NUM_KSPLIT:
        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 = _hw_cvt_quant(a_bf16.to(tl.float32), BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
                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)
            tl.store(c_ptrs, c, mask=c_mask)

# ===== v710 kernel for M>32 =====
@triton.jit
def _hw_cvt_quant(x, BLOCK_K: tl.constexpr, BLOCK_M: tl.constexpr, QUANT_BLOCK: tl.constexpr):
    NUM_BLOCKS: tl.constexpr = BLOCK_K // QUANT_BLOCK
    HALF_QB: tl.constexpr = QUANT_BLOCK // 2
    x_3d = x.reshape(BLOCK_M, NUM_BLOCKS, QUANT_BLOCK)
    amax = tl.max(tl.abs(x_3d), axis=-1, keep_dims=True)
    amax_i32 = amax.to(tl.int32, bitcast=True)
    amax_rounded = ((amax_i32 + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000)
    amax_exp = (amax_rounded >> 23).to(tl.int32)
    su = amax_exp - 129;su = tl.where(su < -126, -126, su);su = tl.where(su > 127, 127, su)
    bs_e8m0 = (su + 127).to(tl.uint8)
    # Fused CVT scale: pass 2^(su) so hardware does srcExp -= su → FP4(value * 2^(-su))
    cvt_exp = su + 127
    cvt_scale = (cvt_exp << 23).to(tl.float32, bitcast=True)
    cvt_scale_bc = tl.broadcast_to(cvt_scale, (BLOCK_M, NUM_BLOCKS, HALF_QB))
    cvt_scale_flat = cvt_scale_bc.reshape(BLOCK_M, BLOCK_K // 2)
    # Raw x pairs — NO pre-multiply! Hardware CVT handles scaling.
    x_flat = x_3d.reshape(BLOCK_M, BLOCK_K)
    x_pairs = x_flat.reshape(BLOCK_M, BLOCK_K // 2, 2)
    evens, odds = tl.split(x_pairs)
    fp4_packed = tl.inline_asm_elementwise("v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3","=v,v,v,v",args=[evens, odds, cvt_scale_flat],dtype=tl.int32,is_pure=True,pack=1)
    x_fp4 = fp4_packed.to(tl.uint8)
    return x_fp4.reshape(BLOCK_M, BLOCK_K // 2), bs_e8m0.reshape(BLOCK_M, NUM_BLOCKS)
# ===== Configs =====
def _get_ps_config(M, N, K):
    if K > 4096:
        # v911: NK=14 → 238 blocks (93% CU) with 1 K-iter each
        return {"BLOCK_SIZE_M":16,"BLOCK_SIZE_N":128,"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":14}
    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 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}
    elif M <= 64:
        return {"BLOCK_SIZE_M":8,"BLOCK_SIZE_N":128,"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}
    else:
        # M>64: BSK=256 (v895 confirmed -0.4µs on M=256), BSM=16, nw=4/wpe=2
        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}

_cf={}
def custom_kernel(data:input_t)->output_t:
    A,B,Bq,Bs,Bss=data;M,K=A.shape;N=Bs.shape[0]
    key=(M,K,N);c=_cf.get(key)
    if c is None:
        config = _get_ps_config(M, N, K)
        K_kernel = K // 2; BSK = config["BLOCK_SIZE_K"]; BSN = max(config["BLOCK_SIZE_N"], 32); BSM = config["BLOCK_SIZE_M"]
        NK = config["NUM_KSPLIT"]
        if NK > 1: SBS, BSK, NK = _gs(K_kernel, BSK, NK)
        else: SBS = 2 * K_kernel
        grid_mn = triton.cdiv(M, BSM) * triton.cdiv(N, BSN)
        grid_size = grid_mn * NK
        out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
        y_pp = torch.empty((NK, M, N), dtype=torch.float32, device=A.device) if NK > 1 else None
        AK = triton.cdiv(K_kernel, (SBS // 2)) if NK > 1 else 1
        c = ('ps', out, y_pp, None, None, grid_size, K_kernel, BSM, BSN, BSK,
             config["GROUP_SIZE_M"], NK, SBS, config["num_warps"], config["num_stages"],
             config["waves_per_eu"], config["matrix_instr_nonkdim"], config["cache_modifier"],
             AK, triton.next_power_of_2(NK) if NK > 1 else 1,
             (triton.cdiv(M,16), triton.cdiv(N,64)) if NK > 1 else None)
        _cf[key] = c
    _, out, y_pp, Bw, Bsc, grid_size, K_kernel, BSM, BSN, BSK, GSM, NK, SBS, nw, ns, wpe, mind, cm, AK, MK, rgr = c
    b_ptr = Bs.data_ptr()
    if Bw is None or _cf.get(('_bp', key)) != b_ptr:
        Bw = Bs.view(torch.uint8).reshape(N // 16, (K // 2) * 16)
        bs_shape = Bss.shape; Bsc = Bss.view(torch.uint8).reshape(bs_shape[0] // 32, bs_shape[1] * 32)
        c_list = list(c); c_list[3] = Bw; c_list[4] = Bsc; c = tuple(c_list); _cf[key] = c
        _cf[('_bp', key)] = b_ptr
    target = y_pp if NK > 1 else out
    sk_off = y_pp.stride(0) if NK > 1 else 0
    cm_s = (y_pp.stride(1) if NK > 1 else out.stride(0))
    cn_s = (y_pp.stride(2) if NK > 1 else out.stride(1))
    _ps_nx2[(grid_size,)](A, Bw, target, Bsc, M, N, K_kernel,
        A.stride(0), A.stride(1), Bw.stride(0), Bw.stride(1),
        sk_off, cm_s, cn_s, Bsc.stride(0), Bsc.stride(1),
        BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN, BLOCK_SIZE_K=BSK,
        GROUP_SIZE_M=GSM, NUM_KSPLIT=NK, SPLITK_BLOCK_SIZE=SBS,
        num_warps=nw, num_stages=ns, waves_per_eu=wpe,
        matrix_instr_nonkdim=mind, PREQUANT=True, cache_modifier=cm)
    if NK > 1:
        _gluon_reduce_kernel[rgr](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, AK, MK)
    return out
scrolls · 176 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 553074.

#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
-
- """
- v150: v120 + lean quantization op.
- Replace _mxfp4_quant_op with a custom Triton JIT function that:
- 1. Uses branchless FP4 E2M1 quantization (no 3-way if/else)
- 2. Uses bitcast for exp2 (from v138)
- 3. Minimizes total VALU instructions
- 4. Uses integer-only rounding for max (avoids log2/floor/exp2 chain)
- """
+ """v911: NK=14 for split-K shape — maximize CU utilization.
+ NK=7 → 119 blocks (46% CU). NK=14 → 238 blocks (93% CU).
+ Each block does 1 K-iter instead of 2. More parallelism for memory-bound shape.
+ Trade: reduce sums 14 partials instead of 7."""
+ import os
+ os.environ['DISABLE_LLVM_OPT'] = 'disable-lsr,disable-machine-licm,disable-machine-sink'
+ os.environ['TRITON_HIP_USE_BLOCK_PINGPONG'] = '1'
+ import torch, triton, triton.language as tl
+ from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid
+ from aiter.ops.triton.gluon.gemm_afp4wfp4 import _gemm_afp4wfp4_reduce_kernel as _gluon_reduce_kernel
+ import aiter.ops.triton.gemm_afp4wfp4 as _gm
+ _gs=_gm.get_splitk
from task import input_t, output_t
- import torch
- import triton
- import triton.language as tl
- import aiter
- from aiter import dtypes
- from aiter.ops.gemm_op_a4w4 import get_GEMM_config
- from aiter.ops.gemm_op_common import get_padded_m
-
- import aiter.ops.triton.gemm_afp4wfp4 as _gemm_mod
- _reduce_kernel = _gemm_mod._gemm_afp4wfp4_reduce_kernel
- _get_splitk_fn = _gemm_mod.get_splitk
-
- _fp4x2 = dtypes.fp4x2
- _fp8_e8m0 = dtypes.fp8_e8m0
- _bf16 = dtypes.bf16
-
-
+ # ===== NX=2 XCD remap helper =====
@triton.jit
- def _lean_mxfp4_quant_op(
- x, # [BLOCK_M, BLOCK_K] float32
- BLOCK_K: tl.constexpr,
- BLOCK_M: tl.constexpr,
- QUANT_BLOCK: tl.constexpr,
- ):
- """Lean MXFP4 quantization — fewer VALU instructions, branchless E2M1.
- Produces bit-identical output to aiter's _mxfp4_quant_op.
- """
- NUM_BLOCKS: tl.constexpr = BLOCK_K // QUANT_BLOCK
+ def _rx(p, n, NX: tl.constexpr):
+ cs = tl.cdiv(n, NX)
+ return (p % NX) * cs + p // NX
- x_3d = x.reshape(BLOCK_M, NUM_BLOCKS, QUANT_BLOCK)
-
- # --- Scale computation (integer-only, no log2/floor) ---
- amax = tl.max(tl.abs(x_3d), axis=-1, keep_dims=True)
- # Round amax up to next power-of-2 (clear mantissa, round up exponent)
- amax_i32 = amax.to(tl.int32, bitcast=True)
- amax_rounded = ((amax_i32 + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000)
- amax_f32 = amax_rounded.to(tl.float32, bitcast=True)
-
- # Extract exponent directly: E8M0 scale = exponent_of(amax_rounded) - 2 - 127 + 127
- # = exponent_field - 2 = (amax_rounded >> 23) - 2
- # But we need to handle amax=0 => scale should be 0 (biased), i.e. -127 unbiased
- amax_exp = (amax_rounded >> 23).to(tl.int32)
- # scale_unbiased = amax_exp - 127 - 2 = amax_exp - 129
- scale_unbiased = amax_exp - 129
- # Clamp to [-127, 127] — tl.clamp only supports float, use tl.where
- scale_unbiased = tl.where(scale_unbiased < -127, -127, scale_unbiased)
- scale_unbiased = tl.where(scale_unbiased > 127, 127, scale_unbiased)
- # Handle amax=0: amax_exp=0, scale_unbiased=-129 clamped to -127. OK.
-
- # E8M0 biased scale (uint8)
- bs_e8m0 = (scale_unbiased + 127).to(tl.uint8)
-
- # --- Inverse scale via bitcast (fast exp2) ---
- # quant_scale = exp2(-scale_unbiased) = bitcast(((-scale_unbiased) + 127) << 23)
- inv_exp = (127 - scale_unbiased)
- quant_scale = (inv_exp << 23).to(tl.float32, bitcast=True)
-
- # --- Quantize: scale input ---
- qx = x_3d * quant_scale
-
- # --- Convert to FP4 E2M1 (branchless) ---
- # Extract sign and abs
- qx_i32 = qx.to(tl.int32, bitcast=True)
- sign_bit = ((qx_i32 >> 31) & 0x8).to(tl.uint8) # sign at bit 3 for FP4
- qx_abs = (qx_i32 & 0x7FFFFFFF).to(tl.float32, bitcast=True)
-
- # Saturate: clamp abs to [0, 6.0] — values >= 6.0 become 0x7 (max E2M1 = 1.5 * 2^2 = 6.0)
- # After clamping, all values are in representable range, no saturation branch needed
- qx_clamped = tl.minimum(qx_abs, 6.0)
- qx_clamped_i32 = qx_clamped.to(tl.int32, bitcast=True)
-
- # Denormal path: values < 1.0 need special handling
- # E2M1 denormals: 0.0 (0b000), 0.5 (0b001)
- # Normal E2M1: 1.0 (0b010), 1.5 (0b011), 2.0 (0b100), 3.0 (0b101), 4.0 (0b110), 6.0 (0b111)
- #
- # For denormals (< 1.0): add magic number to round, extract low bits
- denorm_exp: tl.constexpr = (127 - 1) + (23 - 1) + 1
- denorm_magic: tl.constexpr = denorm_exp << 23
- denorm_magic_f: tl.constexpr = tl.cast(denorm_magic, tl.float32, bitcast=True)
- denormal_result = (qx_clamped + denorm_magic_f).to(tl.int32, bitcast=True) - denorm_magic
- denormal_result = denormal_result.to(tl.uint8)
-
- # Normal path (>= 1.0): round to nearest E2M1
- # IEEE float32 mantissa has 23 bits, E2M1 mantissa has 1 bit
- # So we need to round at bit 22 (keep only 1 mantissa bit)
- # Bias adjustment: subtract (127-1) from exponent to get E2M1 exponent
- qx_clamped_abs_i32 = qx_clamped_i32
- mant_odd = (qx_clamped_abs_i32 >> 22) & 1
- val_to_add: tl.constexpr = ((1 - 127) << 23) + (1 << 21) - 1
- normal_result = (qx_clamped_abs_i32 + val_to_add + mant_odd) >> 22
- normal_result = normal_result.to(tl.uint8)
-
- # Select: denormal if < 1.0, normal otherwise
- is_normal = qx_abs >= 1.0
- e2m1 = tl.where(is_normal, normal_result, denormal_result)
-
- # Apply sign
- e2m1 = e2m1 | sign_bit
-
- # Pack 2 FP4 values per byte
- e2m1 = tl.reshape(e2m1, [BLOCK_M, NUM_BLOCKS, QUANT_BLOCK // 2, 2])
- evens, odds = tl.split(e2m1)
- x_fp4 = evens | (odds << 4)
- x_fp4 = x_fp4.reshape(BLOCK_M, BLOCK_K // 2)
-
- return x_fp4, bs_e8m0.reshape(BLOCK_M, NUM_BLOCKS)
-
-
+ # ===== FORKED preshuffle kernel with NX=2 XCD remap =====
+ @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 _remap_xcd(pid, num_pids, NUM_XCDS: tl.constexpr):
- chunk_size = tl.cdiv(num_pids, NUM_XCDS)
- xcd = pid % NUM_XCDS
- pid_in_xcd = pid // NUM_XCDS
- return xcd * chunk_size + pid_in_xcd
+ def _ps_nx2(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)
- @triton.jit
- def _pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M: tl.constexpr):
- 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
- return pid_m, pid_n
-
-
- @triton.jit
- def _fused_quant_gemm_kernel(
- a_ptr, b_ptr, c_ptr, b_scales_ptr,
- M, N, K_real,
- stride_am, stride_ak,
- stride_bk, stride_bn,
- 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,
- QUANT_BLOCK: tl.constexpr,
- ):
- SCALE_GROUP_SIZE: tl.constexpr = 32
- K_packed = K_real // 2
- GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
- total_pids = GRID_MN * NUM_KSPLIT
- total_pids_padded = ((total_pids + 7) // 8) * 8
-
pid_unified = tl.program_id(axis=0)
- pid_unified = _remap_xcd(pid_unified, total_pids_padded, NUM_XCDS=8)
- if pid_unified < total_pids:
+ if pid_unified < GRID_MN * NUM_KSPLIT:
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)
+ 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
+ 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)
- tl.assume(pid_m >= 0)
- tl.assume(pid_n >= 0)
-
- if (pid_k * SPLITK_BLOCK_SIZE) < K_real:
- num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE, BLOCK_SIZE_K)
-
+ 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
- offs_k = pid_k * SPLITK_BLOCK_SIZE + tl.arange(0, BLOCK_SIZE_K)
- a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak
-
- offs_k_packed = pid_k * (SPLITK_BLOCK_SIZE // 2) + tl.arange(0, BLOCK_SIZE_K // 2)
- offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
- b_ptrs = b_ptr + offs_k_packed[:, None] * stride_bk + offs_bn[None, :] * stride_bn
-
- offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)) % N
- offs_ks_scale = (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_scale[None, :] * stride_bsk
-
+ 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 tl.range(0, num_k_iter):
- a_bf16 = tl.load(a_ptrs).to(tl.float32)
- a_fp4, a_scales = _lean_mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, QUANT_BLOCK)
-
- b_fp4 = tl.load(b_ptrs)
-
- b_scales = (
- tl.load(b_scale_ptrs)
+ 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)
- )
-
- accumulator = tl.dot_scaled(a_fp4, a_scales, "e2m1", b_fp4, b_scales, "e2m1", accumulator)
-
+ .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 = _hw_cvt_quant(a_bf16.to(tl.float32), BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
+ 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) * stride_bk
+ 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)
⋯ 1 unchanged lines
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
-
+ # ===== v710 kernel for M>32 =====
@triton.jit
- def _fused_quant_shuffle_kernel(
- x_ptr, x_fp4_ptr, bs_ptr,
- stride_x_m, stride_x_n,
- stride_x_fp4_m, stride_x_fp4_n,
- M, N, scale_n_valid,
- SCALE_N: tl.constexpr,
- BLOCK_SIZE_M: tl.constexpr,
- BLOCK_SIZE_N: tl.constexpr,
- NUM_ITER: tl.constexpr,
- NUM_STAGES: tl.constexpr,
- MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
- ):
- pid_m = tl.program_id(0)
- start_n = tl.program_id(1) * NUM_ITER
- NUM_QUANT_BLOCKS: 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):
- 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
- x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
- x = tl.load(x_ptr + x_offs, mask=x_mask, other=0.0).to(tl.float32)
-
- out_tensor, bs_e8m0 = _lean_mxfp4_quant_op(
- x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
- )
-
- 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
- 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)
-
- 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)
- m_idx = bs_offs_m[:, None]
- n_idx = bs_offs_n[None, :]
- i0 = m_idx // 32
- i1 = (m_idx // 16) % 2
- i2 = m_idx % 16
- i3 = n_idx // 8
- i4 = (n_idx // 4) % 2
- i5 = n_idx % 4
- shuffled_offset = (i0 * (SCALE_N * 32) + i3 * 256 + i5 * 64 + i2 * 4 + i4 * 2 + i1)
- bs_valid = (bs_offs_m < M)[:, None] & (bs_offs_n < scale_n_valid)[None, :]
- bs_e8m0 = tl.where(bs_valid, bs_e8m0, 127)
- bs_store_mask = (m_idx < (M + 255) // 256 * 256) & (n_idx < SCALE_N)
- tl.store(bs_ptr + shuffled_offset, bs_e8m0, mask=bs_store_mask)
-
-
- _cache_asm = {}
- _cache_fused = {}
- _gemm_asm = None
- _warmup_done = False
-
-
- def custom_kernel(data: input_t) -> output_t:
- global _gemm_asm, _warmup_done
-
- A, B, B_q, B_shuffle, B_scale_sh = data
- M, K = A.shape
- N = B_shuffle.shape[0]
-
- use_fused = (M <= 64)
-
- # Warmup: use ASM path to init aiter module
- if not _warmup_done:
- scale_n_valid = (K + 31) // 32
- SCALE_M = ((M + 255) // 256) * 256
- SCALE_N = ((scale_n_valid + 7) // 8) * 8
- BSM = triton.next_power_of_2(M) if M <= 32 else 16
- grid = (triton.cdiv(M, BSM), triton.cdiv(K, 32))
-
- x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)
- bs_sh = torch.full((SCALE_M, SCALE_N), 127, dtype=torch.uint8, device=A.device)
-
- _fused_quant_shuffle_kernel[grid](
- A, x_fp4, bs_sh,
- A.stride(0), A.stride(1),
- x_fp4.stride(0), x_fp4.stride(1),
- M, K, scale_n_valid,
- SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=32,
- NUM_ITER=1, NUM_STAGES=1, MXFP4_QUANT_BLOCK_SIZE=32,
- num_warps=1, waves_per_eu=0, num_stages=1,
- )
-
- result = aiter.gemm_a4w4(
- x_fp4.view(_fp4x2), B_shuffle,
- bs_sh.view(_fp8_e8m0), B_scale_sh,
- dtype=_bf16, bpreshuffle=True,
- )
- _warmup_done = True
- try:
- _gemm_asm = torch.ops.aiter.gemm_a4w4_asm
- except Exception:
- try:
- import aiter.jit.core as _jc
- _gemm_asm = getattr(_jc, 'gemm_a4w4_asm', None)
- except Exception:
- pass
- return result
-
- if use_fused:
- # --- Fused quant+GEMM: single kernel launch ---
- key = (M, K, N)
- c = _cache_fused.get(key)
- if c is None:
- K_packed = K // 2
- scale_n = (K + 31) // 32
- SCALE_N_B = ((scale_n + 7) // 8) * 8
-
- BLOCK_SIZE_M = 16
- BLOCK_SIZE_N = 64 if M <= 16 else 128
- BLOCK_SIZE_K = 512
-
- base_blocks = triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)
- target_ksplit = max(1, 256 // max(1, base_blocks))
-
- if target_ksplit > 1:
- SPLITK_BLOCK_SIZE, BLOCK_SIZE_K_adj, NUM_KSPLIT = _get_splitk_fn(
- K_packed, BLOCK_SIZE_K, target_ksplit
- )
- if BLOCK_SIZE_K_adj < 512:
- BLOCK_SIZE_K_adj = 512
- SPLITK_BLOCK_SIZE = 2 * K_packed
- NUM_KSPLIT = 1
- else:
- BLOCK_SIZE_K = BLOCK_SIZE_K_adj
- else:
- NUM_KSPLIT = 1
- SPLITK_BLOCK_SIZE = 2 * K_packed
-
- if NUM_KSPLIT > 1:
- y_pp = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A.device)
- else:
- y_pp = None
- SPLITK_BLOCK_SIZE = 2 * K_packed
-
- y = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
-
- total_blocks_raw = NUM_KSPLIT * triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)
- total_blocks = ((total_blocks_raw + 7) // 8) * 8
-
- bs_stride_n = 32 * SCALE_N_B
- bs_stride_k = 1
-
- c = (K_packed, SCALE_N_B, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K,
- NUM_KSPLIT, SPLITK_BLOCK_SIZE,
- y, y_pp, total_blocks, bs_stride_n, bs_stride_k)
- _cache_fused[key] = c
-
- (K_packed, SCALE_N_B, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K,
- NUM_KSPLIT, SPLITK_BLOCK_SIZE,
- y, y_pp, total_blocks, bs_stride_n, bs_stride_k) = c
-
- B_q_u8 = B_q.view(torch.uint8) if B_q.dtype != torch.uint8 else B_q
- B_q_T = B_q_u8.T
- B_scale_u8 = B_scale_sh.view(torch.uint8)
-
- out_tensor = y if NUM_KSPLIT == 1 else y_pp
-
- _fused_quant_gemm_kernel[(total_blocks,)](
- A, B_q_T, out_tensor, B_scale_u8,
- M, N, K,
- A.stride(0), A.stride(1),
- B_q_T.stride(0), B_q_T.stride(1),
- 0 if NUM_KSPLIT == 1 else y_pp.stride(0),
- y.stride(0) if NUM_KSPLIT == 1 else y_pp.stride(1),
- y.stride(1) if NUM_KSPLIT == 1 else y_pp.stride(2),
- bs_stride_n, bs_stride_k,
- BLOCK_SIZE_M=BLOCK_SIZE_M,
- BLOCK_SIZE_N=BLOCK_SIZE_N,
- BLOCK_SIZE_K=BLOCK_SIZE_K,
- GROUP_SIZE_M=8,
- NUM_KSPLIT=NUM_KSPLIT,
- SPLITK_BLOCK_SIZE=SPLITK_BLOCK_SIZE,
- QUANT_BLOCK=32,
- num_warps=8,
- num_stages=2,
- waves_per_eu=0,
- )
-
- if NUM_KSPLIT > 1:
- ACTUAL_KSPLIT = triton.cdiv(K_packed, (SPLITK_BLOCK_SIZE // 2))
- grid_reduce = (triton.cdiv(M, 16), triton.cdiv(N, 64))
- _reduce_kernel[grid_reduce](
- y_pp, y, M, N,
- y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
- y.stride(0), y.stride(1),
- 16, 64, ACTUAL_KSPLIT,
- triton.next_power_of_2(NUM_KSPLIT),
- )
-
- return y
-
+ def _hw_cvt_quant(x, BLOCK_K: tl.constexpr, BLOCK_M: tl.constexpr, QUANT_BLOCK: tl.constexpr):
+ NUM_BLOCKS: tl.constexpr = BLOCK_K // QUANT_BLOCK
+ HALF_QB: tl.constexpr = QUANT_BLOCK // 2
+ x_3d = x.reshape(BLOCK_M, NUM_BLOCKS, QUANT_BLOCK)
+ amax = tl.max(tl.abs(x_3d), axis=-1, keep_dims=True)
+ amax_i32 = amax.to(tl.int32, bitcast=True)
+ amax_rounded = ((amax_i32 + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000)
+ amax_exp = (amax_rounded >> 23).to(tl.int32)
+ su = amax_exp - 129;su = tl.where(su < -126, -126, su);su = tl.where(su > 127, 127, su)
+ bs_e8m0 = (su + 127).to(tl.uint8)
+ # Fused CVT scale: pass 2^(su) so hardware does srcExp -= su → FP4(value * 2^(-su))
+ cvt_exp = su + 127
+ cvt_scale = (cvt_exp << 23).to(tl.float32, bitcast=True)
+ cvt_scale_bc = tl.broadcast_to(cvt_scale, (BLOCK_M, NUM_BLOCKS, HALF_QB))
+ cvt_scale_flat = cvt_scale_bc.reshape(BLOCK_M, BLOCK_K // 2)
+ # Raw x pairs — NO pre-multiply! Hardware CVT handles scaling.
+ x_flat = x_3d.reshape(BLOCK_M, BLOCK_K)
+ x_pairs = x_flat.reshape(BLOCK_M, BLOCK_K // 2, 2)
+ evens, odds = tl.split(x_pairs)
+ fp4_packed = tl.inline_asm_elementwise("v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3","=v,v,v,v",args=[evens, odds, cvt_scale_flat],dtype=tl.int32,is_pure=True,pack=1)
+ x_fp4 = fp4_packed.to(tl.uint8)
+ return x_fp4.reshape(BLOCK_M, BLOCK_K // 2), bs_e8m0.reshape(BLOCK_M, NUM_BLOCKS)
+ # ===== Configs =====
+ def _get_ps_config(M, N, K):
+ if K > 4096:
+ # v911: NK=14 → 238 blocks (93% CU) with 1 K-iter each
+ return {"BLOCK_SIZE_M":16,"BLOCK_SIZE_N":128,"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":14}
+ 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 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}
+ elif M <= 64:
+ return {"BLOCK_SIZE_M":8,"BLOCK_SIZE_N":128,"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}
else:
- # --- ASM GEMM path ---
- key = (M, K, N)
- c = _cache_asm.get(key)
- if c is None:
- scale_n_valid = (K + 31) // 32
- SCALE_M = ((M + 255) // 256) * 256
- SCALE_N = ((scale_n_valid + 7) // 8) * 8
- padded_m = get_padded_m(M, N, K, 0)
+ # M>64: BSK=256 (v895 confirmed -0.4µs on M=256), BSM=16, nw=4/wpe=2
+ 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}
- BSM = triton.next_power_of_2(M) if M <= 32 else 16
- NW = 1
- BSN = 32
- NUM_ITER_Q = 2
- grid = (triton.cdiv(M, BSM), triton.cdiv(K, BSN * NUM_ITER_Q))
-
- ck_config = get_GEMM_config(M, N, K)
- kernel_name = ""
- split_k = 0
- if ck_config is not None:
- split_k = ck_config.get("splitK", 0) or 0
- kernel_name = ck_config["kernelName"]
-
- x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)
- bs_sh = torch.full((SCALE_M, SCALE_N), 127, dtype=torch.uint8, device=A.device)
- out = torch.empty((padded_m, N), dtype=torch.bfloat16, device=A.device)
-
- x_fp4_view = x_fp4.view(_fp4x2)
- bs_sh_view = bs_sh.view(_fp8_e8m0)
- out_view = out[:M] if M < padded_m else out
-
- c = (scale_n_valid, SCALE_N, BSM, BSN, NUM_ITER_Q, grid,
- x_fp4, bs_sh, out, x_fp4_view, bs_sh_view, out_view,
- kernel_name, split_k,
- A.stride(0), A.stride(1), x_fp4.stride(0), x_fp4.stride(1))
- _cache_asm[key] = c
-
- (scale_n_valid, SCALE_N, BSM, BSN, NUM_ITER_Q, grid,
- x_fp4, bs_sh, out, x_fp4_view, bs_sh_view, out_view,
- kernel_name, split_k,
- stride_a0, stride_a1, stride_fp4_0, stride_fp4_1) = c
-
- _fused_quant_shuffle_kernel[grid](
- A, x_fp4, bs_sh,
- stride_a0, stride_a1,
- stride_fp4_0, stride_fp4_1,
- M, K, scale_n_valid,
- SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN,
- NUM_ITER=NUM_ITER_Q, NUM_STAGES=NUM_ITER_Q, MXFP4_QUANT_BLOCK_SIZE=32,
- num_warps=1, waves_per_eu=0, num_stages=NUM_ITER_Q,
- )
-
- if _gemm_asm is not None:
- _gemm_asm(x_fp4_view, B_shuffle, bs_sh_view, B_scale_sh,
- out, kernel_name, None, 1.0, 0.0, True, split_k)
- return out_view
-
- return aiter.gemm_a4w4(
- x_fp4_view, B_shuffle, bs_sh_view, B_scale_sh,
- dtype=_bf16, bpreshuffle=True,
- )
+ _cf={}
+ def custom_kernel(data:input_t)->output_t:
+ A,B,Bq,Bs,Bss=data;M,K=A.shape;N=Bs.shape[0]
+ key=(M,K,N);c=_cf.get(key)
+ if c is None:
+ config = _get_ps_config(M, N, K)
+ K_kernel = K // 2; BSK = config["BLOCK_SIZE_K"]; BSN = max(config["BLOCK_SIZE_N"], 32); BSM = config["BLOCK_SIZE_M"]
+ NK = config["NUM_KSPLIT"]
+ if NK > 1: SBS, BSK, NK = _gs(K_kernel, BSK, NK)
+ else: SBS = 2 * K_kernel
+ grid_mn = triton.cdiv(M, BSM) * triton.cdiv(N, BSN)
+ grid_size = grid_mn * NK
+ out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
+ y_pp = torch.empty((NK, M, N), dtype=torch.float32, device=A.device) if NK > 1 else None
+ AK = triton.cdiv(K_kernel, (SBS // 2)) if NK > 1 else 1
+ c = ('ps', out, y_pp, None, None, grid_size, K_kernel, BSM, BSN, BSK,
+ config["GROUP_SIZE_M"], NK, SBS, config["num_warps"], config["num_stages"],
+ config["waves_per_eu"], config["matrix_instr_nonkdim"], config["cache_modifier"],
+ AK, triton.next_power_of_2(NK) if NK > 1 else 1,
+ (triton.cdiv(M,16), triton.cdiv(N,64)) if NK > 1 else None)
+ _cf[key] = c
+ _, out, y_pp, Bw, Bsc, grid_size, K_kernel, BSM, BSN, BSK, GSM, NK, SBS, nw, ns, wpe, mind, cm, AK, MK, rgr = c
+ b_ptr = Bs.data_ptr()
+ if Bw is None or _cf.get(('_bp', key)) != b_ptr:
+ Bw = Bs.view(torch.uint8).reshape(N // 16, (K // 2) * 16)
+ bs_shape = Bss.shape; Bsc = Bss.view(torch.uint8).reshape(bs_shape[0] // 32, bs_shape[1] * 32)
+ c_list = list(c); c_list[3] = Bw; c_list[4] = Bsc; c = tuple(c_list); _cf[key] = c
+ _cf[('_bp', key)] = b_ptr
+ target = y_pp if NK > 1 else out
+ sk_off = y_pp.stride(0) if NK > 1 else 0
+ cm_s = (y_pp.stride(1) if NK > 1 else out.stride(0))
+ cn_s = (y_pp.stride(2) if NK > 1 else out.stride(1))
+ _ps_nx2[(grid_size,)](A, Bw, target, Bsc, M, N, K_kernel,
+ A.stride(0), A.stride(1), Bw.stride(0), Bw.stride(1),
+ sk_off, cm_s, cn_s, Bsc.stride(0), Bsc.stride(1),
+ BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN, BLOCK_SIZE_K=BSK,
+ GROUP_SIZE_M=GSM, NUM_KSPLIT=NK, SPLITK_BLOCK_SIZE=SBS,
+ num_warps=nw, num_stages=ns, waves_per_eu=wpe,
+ matrix_instr_nonkdim=mind, PREQUANT=True, cache_modifier=cm)
+ if NK > 1:
+ _gluon_reduce_kernel[rgr](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, AK, MK)
+ return out
scrolls · 625 diff lines total

Best evidence level for this revision: reported

JSON