Skip to content
KernelIndex
Search⌘K

submission 735661

lmw-perfxlab · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:fcaec78fa61367b92420af5f832ffbc4918292216adaba468214f5f0a577a34c
license declaredunknown
license concludedunknown
authorslmw-perfxlab
imported2026-08-26

Techniques

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

fp4Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.
num-warps = 1num_warps=1
split-kand (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
stages = 2num_stages = 2
tile-k = 512BLOCK_SIZE_K = 512 #kwargs.get("BLOCK_SIZE_K", 256)
tile-m = 16BLOCK_SIZE_M = 16 #kwargs.get("BLOCK_SIZE_M", 32)
tile-n = 32BLOCK_SIZE_N = 32 #kwargs.get("BLOCK_SIZE_N", 32)

Kernel source

triton_a4w4_merge.py755 lines
# SPDX-License-Identifier: MIT
# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved.

import torch
import triton
import triton.language as tl
from task import input_t, output_t

import aiter
from aiter import QuantType, dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant 
from aiter.utility.fp4_utils import e8m0_shuffle
from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid, remap_xcd
from aiter.ops.triton.utils.gemm_config_utils import get_gemm_config

MXBLK = 32

GROUP_N = 64

@triton.jit
def mxfp4_quant_kernel_smallM(
    x_ptr, out_ptr, scale_ptr,
    stride_x_m, stride_x_n,
    stride_out_m, stride_out_n,
    stride_s_m, stride_s_n,
    M: tl.constexpr,
    N: tl.constexpr,
    GROUP_N_: tl.constexpr,
    MXBLK_: tl.constexpr,
):
    m_id = tl.program_id(0)
    g_id = tl.program_id(1)

    # ---------- 子块0 (0..31) ----------
    offs0 = g_id * GROUP_N_ + tl.arange(0, MXBLK_)
    mask0 = offs0 < N
    x0 = tl.load(
        x_ptr + m_id * stride_x_m + offs0 * stride_x_n,
        mask=mask0, other=0.0
    ).to(tl.float32)

    # 计算 scale0
    amax0 = tl.max(tl.abs(x0))
    amax0 = tl.where(amax0 == 0, 1e-8, amax0)
    amax_i32 = amax0.to(tl.int32, bitcast=True)
    amax_i32 = (amax_i32 + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    amax0 = amax_i32.to(tl.float32, bitcast=True)
    se0 = tl.log2(amax0).floor() - 2
    se0 = tl.clamp(se0, -127, 127)
    scale0 = tl.exp2(-se0)
    bs0 = (se0.to(tl.uint8) + 127)

    # 量化子块0
    q0 = x0 * scale0
    qbits0 = q0.to(tl.uint32, bitcast=True)
    s0 = qbits0 & 0x80000000
    qbits0 ^= s0
    qf0 = qbits0.to(tl.float32, bitcast=True)

    sat_mask0 = qf0 >= 6.0
    den_mask0 = (qf0 < 1.0) & (~sat_mask0)
    norm_mask0 = ~(sat_mask0 | den_mask0)

    denorm_exp = ((127 - 1) + (23 - 1) + 1) << 23
    denorm_float = tl.cast(denorm_exp, tl.float32, bitcast=True)
    dval0 = (qf0 + denorm_float).to(tl.uint32, bitcast=True) - denorm_exp
    dval0 = dval0.to(tl.uint8)

    nval0 = qbits0
    mant_odd0 = (nval0 >> (23 - 1)) & 1
    b1 = ((1 - 127) << 23) + (1 << 21) - 1
    nval0 = nval0 + b1 + mant_odd0
    nval0 = (nval0 >> (23 - 1)).to(tl.uint8)

    fp4_0 = tl.where(norm_mask0, nval0, tl.full([1], 0x7, tl.uint8))
    fp4_0 = tl.where(den_mask0, dval0, fp4_0)
    sign_lp0 = (s0 >> (23 + 8 - 1 - 2)).to(tl.uint8)
    fp4_0 |= sign_lp0          # shape: [32]

    # 打包子块0
    fp4_reshaped0 = tl.reshape(fp4_0, (MXBLK_ // 2, 2))   # [16, 2]
    ev0, od0 = tl.split(fp4_reshaped0)
    packed0 = (ev0 | (od0 << 4)).to(tl.uint8)              # [16]

    # 存储子块0的 packed 结果
    out_offs0 = g_id * (GROUP_N_ // 2) + tl.arange(0, MXBLK_ // 2)
    mask_out0 = out_offs0 < (N // 2)
    tl.store(
        out_ptr + m_id * stride_out_m + out_offs0 * stride_out_n,
        packed0,
        mask=mask_out0
    )

    # 存储 scale0
    s0_idx = g_id * 2
    if s0_idx < (N // MXBLK_):
        tl.store(scale_ptr + m_id * stride_s_m + s0_idx * stride_s_n, bs0)

    # ---------- 子块1 (32..63) ----------
    offs1 = g_id * GROUP_N_ + MXBLK_ + tl.arange(0, MXBLK_)
    mask1 = offs1 < N
    x1 = tl.load(
        x_ptr + m_id * stride_x_m + offs1 * stride_x_n,
        mask=mask1, other=0.0
    ).to(tl.float32)

    amax1 = tl.max(tl.abs(x1))
    amax1 = tl.where(amax1 == 0, 1e-8, amax1)
    amax_i32 = amax1.to(tl.int32, bitcast=True)
    amax_i32 = (amax_i32 + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    amax1 = amax_i32.to(tl.float32, bitcast=True)
    se1 = tl.log2(amax1).floor() - 2
    se1 = tl.clamp(se1, -127, 127)
    scale1 = tl.exp2(-se1)
    bs1 = (se1.to(tl.uint8) + 127)

    q1 = x1 * scale1
    qbits1 = q1.to(tl.uint32, bitcast=True)
    s1 = qbits1 & 0x80000000
    qbits1 ^= s1
    qf1 = qbits1.to(tl.float32, bitcast=True)

    sat_mask1 = qf1 >= 6.0
    den_mask1 = (qf1 < 1.0) & (~sat_mask1)
    norm_mask1 = ~(sat_mask1 | den_mask1)

    dval1 = (qf1 + denorm_float).to(tl.uint32, bitcast=True) - denorm_exp
    dval1 = dval1.to(tl.uint8)

    nval1 = qbits1
    mant_odd1 = (nval1 >> (23 - 1)) & 1
    nval1 = nval1 + b1 + mant_odd1
    nval1 = (nval1 >> (23 - 1)).to(tl.uint8)

    fp4_1 = tl.where(norm_mask1, nval1, tl.full([1], 0x7, tl.uint8))
    fp4_1 = tl.where(den_mask1, dval1, fp4_1)
    sign_lp1 = (s1 >> (23 + 8 - 1 - 2)).to(tl.uint8)
    fp4_1 |= sign_lp1

    fp4_reshaped1 = tl.reshape(fp4_1, (MXBLK_ // 2, 2))
    ev1, od1 = tl.split(fp4_reshaped1)
    packed1 = (ev1 | (od1 << 4)).to(tl.uint8)

    # 存储子块1的 packed 结果(紧接着子块0)
    out_offs1 = g_id * (GROUP_N_ // 2) + (MXBLK_ // 2) + tl.arange(0, MXBLK_ // 2)
    mask_out1 = out_offs1 < (N // 2)
    tl.store(
        out_ptr + m_id * stride_out_m + out_offs1 * stride_out_n,
        packed1,
        mask=mask_out1
    )

    # 存储 scale1
    s1_idx = g_id * 2 + 1
    if s1_idx < (N // MXBLK_):
        tl.store(scale_ptr + m_id * stride_s_m + s1_idx * stride_s_n, bs1)


def dynamic_mxfp4_quant_smallM(A: torch.Tensor):
    M, N = A.shape
    assert N % MXBLK == 0
    out = torch.empty((M, N // 2), dtype=torch.uint8, device=A.device)
    scale = torch.empty((M, N // MXBLK), dtype=torch.uint8, device=A.device)
    grid = (M, triton.cdiv(N, GROUP_N))
    mxfp4_quant_kernel_smallM[grid](
        A, out, scale,
        *A.stride(), *out.stride(), *scale.stride(),
        M=M, N=N, GROUP_N_=GROUP_N, MXBLK_=MXBLK,
        num_warps=1
    )
    return out, scale


_gemm_afp4wfp4_repr = make_kernel_repr(
    "_gemm_afp4wfp4_kernel",
    [
        "BLOCK_SIZE_M",
        "BLOCK_SIZE_N",
        "BLOCK_SIZE_K",
        "GROUP_SIZE_M",
        "num_warps",
        "num_stages",
        "waves_per_eu",
        "matrix_instr_nonkdim",
        "cache_modifier",
        "NUM_KSPLIT",
    ],
)

@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),
    }
)
@triton.jit(repr=_gemm_afp4wfp4_repr)
def _gemm_afp4wfp4_kernel(
    a_ptr, b_ptr, c_ptr, a_scales_ptr, b_scales_ptr,
    M, N, K,
    stride_am, stride_ak, stride_bk, stride_bn,
    stride_ck, stride_cm, stride_cn,
    stride_asm, stride_ask, 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, 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_asm > 0)
    tl.assume(stride_ask > 0)
    tl.assume(stride_bsk > 0)
    tl.assume(stride_bsn > 0)

    GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)

    pid_unified = tl.program_id(axis=0)
    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 = tl.arange(0, BLOCK_SIZE_K // 2)
        offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
        offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
        a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k_split[None, :] * stride_ak)
        b_ptrs = b_ptr + (offs_k_split[:, None] * stride_bk + offs_bn[None, :] * stride_bn)

        offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
        a_scale_ptrs = a_scales_ptr + offs_am[:, None] * stride_asm + offs_ks[None, :] * stride_ask
        b_scale_ptrs = b_scales_ptr + offs_bn[:, 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):
            a_scales = tl.load(a_scale_ptrs)
            b_scales = tl.load(b_scale_ptrs, cache_modifier=cache_modifier)

            if EVEN_K:
                a = tl.load(a_ptrs)
                b = tl.load(b_ptrs, cache_modifier=cache_modifier)
            else:
                a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * (BLOCK_SIZE_K // 2), other=0)
                b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k * (BLOCK_SIZE_K // 2), other=0, cache_modifier=cache_modifier)

            accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)

            a_ptrs += (BLOCK_SIZE_K // 2) * stride_ak
            b_ptrs += (BLOCK_SIZE_K // 2) * stride_bk
            a_scale_ptrs += (BLOCK_SIZE_K // SCALE_GROUP_SIZE) * stride_ask
            b_scale_ptrs += (BLOCK_SIZE_K // SCALE_GROUP_SIZE) * 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)

_gemm_afp4wfp4_preshuffle_scales_repr = make_kernel_repr(
    "_gemm_afp4wfp4_preshuffle_kernel",
    [
        "BLOCK_SIZE_M",
        "BLOCK_SIZE_N",
        "BLOCK_SIZE_K",
        "GROUP_SIZE_M",
        "num_warps",
        "num_stages",
        "waves_per_eu",
        "matrix_instr_nonkdim",
        "cache_modifier",
        "NUM_KSPLIT",
    ],
)

@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),
    }
)
@triton.jit(repr=_gemm_afp4wfp4_preshuffle_scales_repr)
def _gemm_afp4wfp4_kernel_preshuffle_scales(
    a_ptr, b_ptr, c_ptr, a_scales_ptr, b_scales_ptr,
    M, N, K,
    stride_am, stride_ak, stride_bk, stride_bn,
    stride_ck, stride_cm, stride_cn,
    stride_asm, stride_ask, 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, 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_asm > 0)
    tl.assume(stride_ask > 0)
    tl.assume(stride_bsk > 0)
    tl.assume(stride_bsn > 0)

    GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)

    pid_unified = tl.program_id(axis=0)
    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)
    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 = tl.arange(0, BLOCK_SIZE_K // 2)
        offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
        offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
        a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k_split[None, :] * stride_ak)
        b_ptrs = b_ptr + (offs_k_split[:, None] * stride_bk + offs_bn[None, :] * stride_bn)

        offs_asn = (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_asn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk

        if BLOCK_SIZE_M < 32:
            offs_ks_non_shufl = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
            a_scale_ptrs = a_scales_ptr + offs_am[:, None] * stride_asm + offs_ks_non_shufl[None, :] * stride_ask
        else:
            offs_asm = (pid_m * (BLOCK_SIZE_M // 32) + tl.arange(0, (BLOCK_SIZE_M // 32))) % M
            a_scale_ptrs = a_scales_ptr + offs_asm[:, None] * stride_asm + offs_ks[None, :] * stride_ask

        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):
            if BLOCK_SIZE_M < 32:
                a_scales = tl.load(a_scale_ptrs)
            else:
                a_scales = (
                    tl.load(a_scale_ptrs)
                    .reshape(BLOCK_SIZE_M // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)
                    .permute(0, 5, 3, 1, 4, 2, 6)
                    .reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
                )
            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 = tl.load(a_ptrs)
                b = tl.load(b_ptrs, cache_modifier=cache_modifier)
            else:
                a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * (BLOCK_SIZE_K // 2), other=0)
                b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k * (BLOCK_SIZE_K // 2), other=0)

            accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)

            a_ptrs += (BLOCK_SIZE_K // 2) * stride_ak
            b_ptrs += (BLOCK_SIZE_K // 2) * stride_bk
            if BLOCK_SIZE_M < 32:
                a_scale_ptrs += (BLOCK_SIZE_K // SCALE_GROUP_SIZE) * stride_ask
            else:
                a_scale_ptrs += BLOCK_SIZE_K * stride_ask
            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, cache_modifier=".wt")

_gemm_afp4wfp4_preshuffle_repr = make_kernel_repr(
    "_gemm_afp4wfp4_preshuffle_kernel",
    [
        "BLOCK_SIZE_M",
        "BLOCK_SIZE_N",
        "BLOCK_SIZE_K",
        "GROUP_SIZE_M",
        "num_warps",
        "num_stages",
        "waves_per_eu",
        "matrix_instr_nonkdim",
        "cache_modifier",
        "NUM_KSPLIT",
    ],
)

@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),
    }
)
@triton.jit(repr=_gemm_afp4wfp4_preshuffle_repr)
def _gemm_afp4wfp4_preshuffle_kernel(
    a_ptr, b_ptr, c_ptr, a_scales_ptr, b_scales_ptr,
    M, N, K,
    stride_am, stride_ak, stride_bn, stride_bk,
    stride_ck, stride_cm, stride_cn,
    stride_asm, stride_ask, 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, 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_asm > 0)
    tl.assume(stride_ask > 0)
    tl.assume(stride_bsk > 0)
    tl.assume(stride_bsn > 0)

    GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)

    pid_unified = tl.program_id(axis=0)
    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)
    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 = tl.arange(0, BLOCK_SIZE_K // 2)
        offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
        offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
        offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr

        offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
        a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k_split[None, :] * stride_ak)
        b_ptrs = b_ptr + (offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk)

        offs_asn = (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_asn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk

        if BLOCK_SIZE_M < 32:
            offs_ks_non_shufl = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
            a_scale_ptrs = a_scales_ptr + offs_am[:, None] * stride_asm + offs_ks_non_shufl[None, :] * stride_ask
        else:
            offs_asm = (pid_m * (BLOCK_SIZE_M // 32) + tl.arange(0, (BLOCK_SIZE_M // 32))) % M
            a_scale_ptrs = a_scales_ptr + offs_asm[:, None] * stride_asm + offs_ks[None, :] * stride_ask

        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):
            if BLOCK_SIZE_M < 32:
                a_scales = tl.load(a_scale_ptrs)
            else:
                a_scales = (
                    tl.load(a_scale_ptrs)
                    .reshape(BLOCK_SIZE_M // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)
                    .permute(0, 5, 3, 1, 4, 2, 6)
                    .reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
                )

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

            accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)

            a_ptrs += (BLOCK_SIZE_K // 2) * stride_ak
            b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
            if BLOCK_SIZE_M < 32:
                a_scale_ptrs += (BLOCK_SIZE_K // SCALE_GROUP_SIZE) * stride_ask
            else:
                a_scale_ptrs += BLOCK_SIZE_K * stride_ask
            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, cache_modifier=".wt")

_gemm_afp4wfp4_reduce_repr = make_kernel_repr(
    "_gemm_afp4wfp4_reduce_kernel",
    [
        "BLOCK_SIZE_M",
        "BLOCK_SIZE_N",
        "ACTUAL_KSPLIT",
        "MAX_KSPLIT",
    ],
)

@triton.heuristics({})
@triton.jit(repr=_gemm_afp4wfp4_reduce_repr)
def _gemm_afp4wfp4_reduce_kernel(
    c_in_ptr, c_out_ptr,
    M, N,
    stride_c_in_k, stride_c_in_m, stride_c_in_n,
    stride_c_out_m, stride_c_out_n,
    BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
    ACTUAL_KSPLIT: tl.constexpr, MAX_KSPLIT: tl.constexpr,
):
    pid_m = tl.program_id(axis=0)
    pid_n = tl.program_id(axis=1)

    offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
    offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
    offs_k = tl.arange(0, MAX_KSPLIT)
    c_in_ptrs = (
        c_in_ptr
        + (offs_k[:, None, None] * stride_c_in_k)
        + (offs_m[None, :, None] * stride_c_in_m)
        + (offs_n[None, None, :] * stride_c_in_n)
    )

    if ACTUAL_KSPLIT == MAX_KSPLIT:
        c = tl.load(c_in_ptrs)
    else:
        c = tl.load(c_in_ptrs, mask=offs_k[:, None, None] < ACTUAL_KSPLIT)
    c = tl.sum(c, axis=0)

    c = c.to(c_out_ptr.type.element_ty)

    c_out_ptrs = c_out_ptr + (offs_m[:, None] * stride_c_out_m) + (offs_n[None, :] * stride_c_out_n)
    tl.store(c_out_ptrs, c)


def _get_config(M: int, N: int, K: int, shuffle: bool = False):
    K = 2 * K
    if shuffle:
        return get_gemm_config(
            "GEMM-AFP4WFP4_PRESHUFFLED",
            M, N, K,
            bounds=(4, 8, 16, 31, 32, 64, 128, 256, 512, 1024, 2048, 4096, 8192),
        )
    else:
        return get_gemm_config("GEMM-AFP4WFP4", M, N, K)


def custom_kernel(data: input_t) -> output_t:
    """
    Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.
    Replaced Aiter generic gemm call with the direct Triton preshuffle kernel.
    """
    def _quant_mxfp4(x, shuffle=False):
        x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
        if shuffle:
            bs_e8m0 = e8m0_shuffle(bs_e8m0)
        return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)
    
    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    B = B.contiguous()
    # 保留张量的真实物理内存排布,规避前端类型检查
    M, K = A.shape
    N, _ = B.shape

    A_q, A_scale_sh = dynamic_mxfp4_quant_smallM(A)
    # A_q, A_scale_sh = _quant_mxfp4(A, False)
    
    # K dimension handling (fp4 packs 2 elements per uint8, logical physical length K_fp4 = K // 2)
    K_fp4 = K // 2

    # Fetching configuration and resolving parameters
    # config = _get_config(M, N, K_fp4, shuffle=True)
    # if hasattr(config, "kwargs"):
    #     kwargs = config.kwargs
    #     num_warps = getattr(config, "num_warps", 4)
    #     num_stages = getattr(config, "num_stages", 4)
    # elif isinstance(config, dict):
    #     kwargs = config
    #     num_warps = kwargs.get("num_warps", 4)
    #     num_stages = kwargs.get("num_stages", 4)
    # else:
    #     kwargs = {}
    num_warps = 2
    num_stages = 2
    # BLOCK_SIZE_M 8
    # BLOCK_SIZE_N 64
    # BLOCK_SIZE_K 512
    # GROUP_SIZE_M 4
    # num_warps 2
    # num_stages 2
    # waves_per_eu 4
    # matrix_instr_nonkdim 16 
    # cache_modifier CG 
    # NUM_KSPLIT_1
    BLOCK_SIZE_M = 16 #kwargs.get("BLOCK_SIZE_M", 32)
    BLOCK_SIZE_N = 32 #kwargs.get("BLOCK_SIZE_N", 32)
    BLOCK_SIZE_K = 512 #kwargs.get("BLOCK_SIZE_K", 256)
    # if BLOCK_SIZE_M >= 32 and BLOCK_SIZE_M % 32 != 0:
    #     BLOCK_SIZE_M = (BLOCK_SIZE_M // 32 + 1) * 32
    # if BLOCK_SIZE_N % 32 != 0:
    #     BLOCK_SIZE_N = max(32, (BLOCK_SIZE_N // 32 + 1) * 32)
    # if BLOCK_SIZE_K % 256 != 0:
    #     BLOCK_SIZE_K = max(256, (BLOCK_SIZE_K // 256 + 1) * 256)

    GROUP_SIZE_M = 4 #kwargs.get("GROUP_SIZE_M", 4)
    NUM_KSPLIT = 1 #kwargs.get("NUM_KSPLIT", 1)
    waves_per_eu = 4 #kwargs.get("waves_per_eu", 0)
    matrix_instr_nonkdim = 16 #kwargs.get("matrix_instr_nonkdim", 16)
    cache_modifier = "" #kwargs.get("cache_modifier", "")

    # shuffle_a = BLOCK_SIZE_M >= 32
    # A_q, A_scale_sh = _quant_mxfp4(A, False)
    # Preparing splits and outputs
    SPLITK_BLOCK_SIZE = max(1, K // NUM_KSPLIT)
    SPLITK_BLOCK_SIZE = triton.cdiv(SPLITK_BLOCK_SIZE, BLOCK_SIZE_K) * BLOCK_SIZE_K

    if NUM_KSPLIT == 1:
        C = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
        stride_ck = 0
        stride_cm, stride_cn = C.stride()
    else:
        C = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A.device)
        stride_ck, stride_cm, stride_cn = C.stride()

    # Extracting layout strides
    stride_am, stride_ak = A_q.stride()
    B_shuffle = B_shuffle.contiguous()
    stride_bn, stride_bk = B_shuffle.stride()
    
    # Forcing architectural scales strides
    A_scale_sh = A_scale_sh.contiguous()
    B_scale_sh = B_scale_sh.contiguous()
    stride_asm, stride_ask = A_scale_sh.stride()
    stride_bsn, stride_bsk = B_scale_sh.stride()
    # print(A_q, A_scale_sh)

    # 补偿宏块物理跨度,使前端步长与内核寻址游标对齐
    stride_bn *= 16
    stride_bsn *= 32
    # if shuffle_a:
    #     stride_asm *= 32

    grid = lambda META: (
        triton.cdiv(M, META["BLOCK_SIZE_M"]) * triton.cdiv(N, META["BLOCK_SIZE_N"]) * META["NUM_KSPLIT"],
    )

    A_q_uint8 = A_q.contiguous().view(torch.uint8)
    A_scale_sh_uint8 = A_scale_sh.contiguous().view(torch.uint8)
    B_shuffle_uint8 = B_shuffle.contiguous().view(torch.uint8)
    B_scale_sh_uint8 = B_scale_sh.contiguous().view(torch.uint8)
    # Executing Triton Kernel
    _gemm_afp4wfp4_preshuffle_kernel[grid](
        A_q_uint8,
        B_shuffle_uint8,
        C,
        A_scale_sh_uint8,
        B_scale_sh_uint8,
        M, N, K_fp4,
        stride_am, stride_ak,
        stride_bn, stride_bk,
        stride_ck, stride_cm, stride_cn,
        stride_asm, stride_ask,
        stride_bsn, stride_bsk,
        BLOCK_SIZE_M=BLOCK_SIZE_M,
        BLOCK_SIZE_N=BLOCK_SIZE_N,
        BLOCK_SIZE_K=BLOCK_SIZE_K,
        GROUP_SIZE_M=GROUP_SIZE_M,
        NUM_KSPLIT=NUM_KSPLIT,
        SPLITK_BLOCK_SIZE=SPLITK_BLOCK_SIZE,
        num_warps=num_warps,
        num_stages=num_stages,
        waves_per_eu=waves_per_eu,
        matrix_instr_nonkdim=matrix_instr_nonkdim,
        cache_modifier=cache_modifier,
    )

    # Optional reduction when KSPLIT > 1
    if NUM_KSPLIT > 1:
        C_out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
        grid_reduce = lambda META: (triton.cdiv(M, META["BLOCK_SIZE_M"]), triton.cdiv(N, META["BLOCK_SIZE_N"]))
        _gemm_afp4wfp4_reduce_kernel[grid_reduce](
            C, C_out,
            M, N,
            stride_ck, stride_cm, stride_cn,
            C_out.stride(0), C_out.stride(1),
            BLOCK_SIZE_M=BLOCK_SIZE_M,
            BLOCK_SIZE_N=BLOCK_SIZE_N,
            ACTUAL_KSPLIT=NUM_KSPLIT,
            MAX_KSPLIT=NUM_KSPLIT
        )
        return C_out
    else:
        return C
scrolls · 755 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON