Skip to content
KernelIndex
Search⌘K

submission 615003

trulyspinach · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-615003?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
12.7µs
#392 of 1143
2026-03-23

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2e3159a6fdcffc06e707786c1d52af48db99ec480c3ff86d65a432fbd31e68a6
license declaredunknown
license concludedunknown
authorstrulyspinach
imported2026-08-26

Techniques

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

shared-memory__shared__ float absmax[ROWS_PER_BLOCK * PAIRS_PER_ROW];
split-kdef _get_splitk(k_packed: int, block_size_k: int, num_ksplit: int):

Kernel source

submission.py905 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

import os

os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")

from task import input_t, output_t

import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline

from aiter import dtypes
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
from aiter.ops.triton.quant import dynamic_mxfp4_quant
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
from aiter.utility.fp4_utils import e8m0_shuffle

_M_BOUNDS = (4, 8, 16, 31, 32, 64, 128, 256, 512, 1024, 2048, 4096, 8192)
_REFERENCE_SHAPES: set[tuple[int, int, int]] = {(64, 7168, 2048), (256, 3072, 1536)}
_B_TRITON_CACHE: dict[tuple[int, int, int, tuple[int, ...], tuple[int, ...]], tuple[torch.Tensor, torch.Tensor]] = {}
_FUSED_CONFIGS = {
    (2880, 512): {
        "M_LEQ_8": {
            "BLOCK_SIZE_M": 8,
            "BLOCK_SIZE_N": 64,
            "BLOCK_SIZE_K": 256,
            "GROUP_SIZE_M": 4,
            "num_warps": 2,
            "num_stages": 2,
            "waves_per_eu": 1,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": None,
            "NUM_KSPLIT": 1,
        },
        "M_LEQ_31": {
            "BLOCK_SIZE_M": 16,
            "BLOCK_SIZE_N": 64,
            "BLOCK_SIZE_K": 256,
            "GROUP_SIZE_M": 4,
            "num_warps": 2,
            "num_stages": 2,
            "waves_per_eu": 1,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": None,
            "NUM_KSPLIT": 1,
        },
        "M_LEQ_32": {
            "BLOCK_SIZE_M": 32,
            "BLOCK_SIZE_N": 64,
            "BLOCK_SIZE_K": 256,
            "GROUP_SIZE_M": 4,
            "num_warps": 2,
            "num_stages": 2,
            "waves_per_eu": 1,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": None,
            "NUM_KSPLIT": 1,
        },
    },
    (4096, 512): {
        "M_LEQ_32": {
            "BLOCK_SIZE_M": 32,
            "BLOCK_SIZE_N": 64,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1,
            "num_warps": 4,
            "num_stages": 1,
            "waves_per_eu": 1,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": None,
            "NUM_KSPLIT": 1,
        },
    },
}
_OVERRIDE_CONFIGS = {
    (2112, 7168): {
        "M_LEQ_31": {
            "BLOCK_SIZE_M": 16,
            "BLOCK_SIZE_N": 32,
            "BLOCK_SIZE_K": 1024,
            "GROUP_SIZE_M": 1,
            "num_warps": 2,
            "num_stages": 2,
            "waves_per_eu": 2,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 1,
        },
    },
}

CPP_SRC = r"""
void quant_mxfp4(torch::Tensor input, torch::Tensor out_q, torch::Tensor out_scale);
void quant_mxfp4_shuffled(torch::Tensor input, torch::Tensor out_q, torch::Tensor out_scale);
void shuffle_e8m0_scales(torch::Tensor input, torch::Tensor output);
"""

CUDA_SRC = r"""
#include <cmath>
#include <cstdint>
#include <stdexcept>

#include <hip/hip_runtime.h>
#include <torch/extension.h>

__device__ __forceinline__ float bf16_to_float(uint16_t x) {
    uint32_t bits = static_cast<uint32_t>(x) << 16;
    return __uint_as_float(bits);
}

__device__ __forceinline__ uint8_t float_to_mxfp4(float x) {
    uint32_t qx = __float_as_uint(x);
    uint32_t s = qx & 0x80000000u;
    qx ^= s;
    float qx_fp32 = __uint_as_float(qx);

    bool saturate_mask = qx_fp32 >= 6.0f;
    bool denormal_mask = (!saturate_mask) && (qx_fp32 < 1.0f);

    constexpr uint32_t MBITS_F32 = 23u;
    constexpr uint32_t MBITS_FP4 = 1u;
    constexpr uint32_t EXP_BIAS_FP32 = 127u;
    constexpr uint32_t EXP_BIAS_FP4 = 1u;
    constexpr uint32_t denorm_mask_int =
        ((EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1u) << MBITS_F32;
    float denorm_mask_float = __uint_as_float(denorm_mask_int);

    uint8_t denormal_x = 0;
    if (denormal_mask) {
        uint32_t denormal_bits = __float_as_uint(qx_fp32 + denorm_mask_float);
        denormal_x = static_cast<uint8_t>(denormal_bits - denorm_mask_int);
    }

    uint8_t normal_x = 0;
    if (!saturate_mask && !denormal_mask) {
        uint32_t mant_odd = (qx >> (MBITS_F32 - MBITS_FP4)) & 1u;
        uint32_t val_to_add =
            ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1u << 21) - 1u;
        uint32_t rounded = qx + val_to_add + mant_odd;
        normal_x = static_cast<uint8_t>(rounded >> (MBITS_F32 - MBITS_FP4));
    }

    uint8_t value = saturate_mask ? 0x7u : normal_x;
    value = denormal_mask ? denormal_x : value;
    value |= static_cast<uint8_t>(s >> 28);
    return value;
}

template <bool SHUFFLED_SCALE>
__global__ void quant_mxfp4_kernel(
    const uint16_t* input,
    uint8_t* out_q,
    uint8_t* out_scale,
    int M,
    int N,
    int stride_in_m,
    int stride_in_n,
    int stride_q_m,
    int stride_q_n,
    int stride_s_m,
    int stride_s_n) {
    constexpr int ROWS_PER_BLOCK = 16;
    constexpr int PAIRS_PER_ROW = 16;
    constexpr int GROUP_SIZE = 32;

    int tid = threadIdx.x;
    int row_local = tid / PAIRS_PER_ROW;
    int pair_idx = tid % PAIRS_PER_ROW;
    int row = blockIdx.x * ROWS_PER_BLOCK + row_local;
    int group = blockIdx.y;
    int groups = N / GROUP_SIZE;

    __shared__ float absmax[ROWS_PER_BLOCK * PAIRS_PER_ROW];
    __shared__ float quant_scale[ROWS_PER_BLOCK];

    bool active = row < M && group < groups;
    int col0 = group * GROUP_SIZE + pair_idx * 2;
    int col1 = col0 + 1;

    float x0 = 0.0f;
    float x1 = 0.0f;
    if (active) {
        x0 = bf16_to_float(input[row * stride_in_m + col0 * stride_in_n]);
        x1 = bf16_to_float(input[row * stride_in_m + col1 * stride_in_n]);
    }

    int shared_idx = row_local * PAIRS_PER_ROW + pair_idx;
    absmax[shared_idx] = fmaxf(fabsf(x0), fabsf(x1));
    __syncthreads();

    for (int offset = PAIRS_PER_ROW / 2; offset > 0; offset >>= 1) {
        if (pair_idx < offset) {
            absmax[shared_idx] = fmaxf(absmax[shared_idx], absmax[shared_idx + offset]);
        }
        __syncthreads();
    }

    if (pair_idx == 0) {
        float qscale = 0.0f;
        uint8_t scale_byte = 0;
        if (active) {
            float amax = absmax[row_local * PAIRS_PER_ROW];
            uint32_t rounded_bits = (__float_as_uint(amax) + 0x200000u) & 0xFF800000u;
            if (rounded_bits == 0u) {
                scale_byte = 0;
                qscale = __uint_as_float(254u << 23);
            } else {
                int exponent = static_cast<int>((rounded_bits >> 23) & 0xFFu) - 127;
                int scale_unbiased = exponent - 2;
                scale_unbiased = scale_unbiased < -127 ? -127 : scale_unbiased;
                scale_unbiased = scale_unbiased > 127 ? 127 : scale_unbiased;
                scale_byte = static_cast<uint8_t>(scale_unbiased + 127);
                qscale = __uint_as_float(static_cast<uint32_t>(127 - scale_unbiased) << 23);
            }
            if constexpr (SHUFFLED_SCALE) {
                int row_block = row / 32;
                int row_half = (row % 32) / 16;
                int row_inner = row % 16;
                int group_block = group / 8;
                int group_half = (group % 8) / 4;
                int group_inner = group % 4;
                int out_cols = stride_s_m;
                int rest =
                    group_block * 256 + group_inner * 64 + row_inner * 4 + group_half * 2 + row_half;
                int out_row = row_block * 32 + rest / out_cols;
                int out_col = rest % out_cols;
                out_scale[out_row * stride_s_m + out_col * stride_s_n] = scale_byte;
            } else {
                out_scale[row * stride_s_m + group * stride_s_n] = scale_byte;
            }
        }
        quant_scale[row_local] = qscale;
    }
    __syncthreads();

    if (!active) {
        return;
    }

    uint8_t q0 = float_to_mxfp4(x0 * quant_scale[row_local]);
    uint8_t q1 = float_to_mxfp4(x1 * quant_scale[row_local]);
    out_q[row * stride_q_m + (group * PAIRS_PER_ROW + pair_idx) * stride_q_n] =
        static_cast<uint8_t>(q0 | (q1 << 4));
}

void quant_mxfp4(torch::Tensor input, torch::Tensor out_q, torch::Tensor out_scale) {
    int M = input.size(0);
    int N = input.size(1);
    dim3 blocks((M + 15) / 16, N / 32);
    dim3 threads(256);
    quant_mxfp4_kernel<false><<<blocks, threads>>>(
        reinterpret_cast<const uint16_t*>(input.data_ptr()),
        out_q.data_ptr<uint8_t>(),
        out_scale.data_ptr<uint8_t>(),
        M,
        N,
        input.stride(0),
        input.stride(1),
        out_q.stride(0),
        out_q.stride(1),
        out_scale.stride(0),
        out_scale.stride(1)
    );

    hipError_t err = hipGetLastError();
    if (err != hipSuccess) {
        throw std::runtime_error(hipGetErrorString(err));
    }
}

void quant_mxfp4_shuffled(torch::Tensor input, torch::Tensor out_q, torch::Tensor out_scale) {
    int M = input.size(0);
    int N = input.size(1);
    dim3 blocks((M + 15) / 16, N / 32);
    dim3 threads(256);
    quant_mxfp4_kernel<true><<<blocks, threads>>>(
        reinterpret_cast<const uint16_t*>(input.data_ptr()),
        out_q.data_ptr<uint8_t>(),
        out_scale.data_ptr<uint8_t>(),
        M,
        N,
        input.stride(0),
        input.stride(1),
        out_q.stride(0),
        out_q.stride(1),
        out_scale.stride(0),
        out_scale.stride(1)
    );

    hipError_t err = hipGetLastError();
    if (err != hipSuccess) {
        throw std::runtime_error(hipGetErrorString(err));
    }
}

__global__ void shuffle_e8m0_scales_kernel(
    const uint8_t* input,
    uint8_t* output,
    int M,
    int G,
    int stride_in_m,
    int stride_in_g,
    int stride_out_m,
    int stride_out_g) {
    int row = blockIdx.x;
    int group = blockIdx.y * blockDim.x + threadIdx.x;
    if (row >= M || group >= G) {
        return;
    }

    int row_block = row / 32;
    int row_half = (row % 32) / 16;
    int row_inner = row % 16;
    int group_block = group / 8;
    int group_half = (group % 8) / 4;
    int group_inner = group % 4;
    int out_cols = stride_out_m;
    int rest = group_block * 256 + group_inner * 64 + row_inner * 4 + group_half * 2 + row_half;
    int out_row = row_block * 32 + rest / out_cols;
    int out_col = rest % out_cols;

    output[out_row * stride_out_m + out_col * stride_out_g] =
        input[row * stride_in_m + group * stride_in_g];
}

void shuffle_e8m0_scales(torch::Tensor input, torch::Tensor output) {
    int M = input.size(0);
    int G = input.size(1);
    dim3 blocks(M, (G + 127) / 128);
    dim3 threads(128);
    shuffle_e8m0_scales_kernel<<<blocks, threads>>>(
        input.data_ptr<uint8_t>(),
        output.data_ptr<uint8_t>(),
        M,
        G,
        input.stride(0),
        input.stride(1),
        output.stride(0),
        output.stride(1)
    );

    hipError_t err = hipGetLastError();
    if (err != hipSuccess) {
        throw std::runtime_error(hipGetErrorString(err));
    }
}
"""

_INLINE_QUANT = load_inline(
    name="submission_inline_quant_v1",
    cpp_sources=[CPP_SRC],
    cuda_sources=[CUDA_SRC],
    functions=["quant_mxfp4", "quant_mxfp4_shuffled", "shuffle_e8m0_scales"],
    verbose=False,
    extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20"],
)


def _tensor_token(x: torch.Tensor) -> tuple[int, int]:
    return (id(x), int(x._version))


def _tensor_cache_key(x: torch.Tensor) -> tuple[int, int, int]:
    token_id, token_version = _tensor_token(x)
    return (token_id, token_version, int(x.data_ptr()))


def _inline_dynamic_mxfp4_quant(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    m, k = int(a.shape[0]), int(a.shape[1])
    a_q = torch.empty((m, k // 2), dtype=torch.uint8, device=a.device)
    a_scale = torch.empty((m, k // 32), dtype=torch.uint8, device=a.device)
    _INLINE_QUANT.quant_mxfp4(a, a_q, a_scale)
    return a_q, a_scale.view(dtypes.fp8_e8m0)


def _inline_dynamic_mxfp4_quant_shuffled(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    m, k = int(a.shape[0]), int(a.shape[1])
    a_q = torch.empty((m, k // 2), dtype=torch.uint8, device=a.device)
    out_m = ((m + 255) // 256) * 256
    out_g = (((k // 32) + 7) // 8) * 8
    a_scale_sh = torch.zeros((out_m, out_g), dtype=torch.uint8, device=a.device)
    _INLINE_QUANT.quant_mxfp4_shuffled(a, a_q, a_scale_sh)
    return a_q, a_scale_sh.view(dtypes.fp8_e8m0)


def _inline_e8m0_shuffle(a_scale: torch.Tensor) -> torch.Tensor:
    m, g = int(a_scale.shape[0]), int(a_scale.shape[1])
    out_m = ((m + 255) // 256) * 256
    out_g = ((g + 7) // 8) * 8
    shuffled = torch.zeros((out_m, out_g), dtype=torch.uint8, device=a_scale.device)
    _INLINE_QUANT.shuffle_e8m0_scales(a_scale.view(torch.uint8), shuffled)
    return shuffled.view(dtypes.fp8_e8m0)


def _select_config(configs: dict[tuple[int, int], dict], m: int, n: int, k: int) -> dict | None:
    override = configs.get((n, k))
    if override is None:
        return None
    for bound in _M_BOUNDS:
        key = f"M_LEQ_{bound}"
        if m <= bound and key in override:
            return dict(override[key])
    if "any" in override:
        return dict(override["any"])
    return None


def _get_config(m: int, n: int, k_packed: int) -> dict:
    override = _select_config(_OVERRIDE_CONFIGS, m, n, 2 * k_packed)
    if override is not None:
        return override
    config, _ = get_gemm_config(
        "GEMM-AFP4WFP4_PRESHUFFLED",
        m,
        n,
        2 * k_packed,
        bounds=_M_BOUNDS,
    )
    return config


def _get_splitk(k_packed: int, block_size_k: int, num_ksplit: int):
    splitk_block_size = (
        triton.cdiv((2 * triton.cdiv(k_packed, num_ksplit)), block_size_k) * block_size_k
    )
    while num_ksplit > 1 and block_size_k > 16:
        if (
            k_packed % (splitk_block_size // 2) == 0
            and splitk_block_size % block_size_k == 0
            and k_packed % (block_size_k // 2) == 0
        ):
            break
        if k_packed % (splitk_block_size // 2) != 0 and num_ksplit > 1:
            num_ksplit //= 2
        elif splitk_block_size % block_size_k != 0:
            if num_ksplit > 1:
                num_ksplit //= 2
            elif block_size_k > 16:
                block_size_k //= 2
        elif k_packed % (block_size_k // 2) != 0 and block_size_k > 16:
            block_size_k //= 2
        else:
            break
        splitk_block_size = (
            triton.cdiv((2 * triton.cdiv(k_packed, num_ksplit)), block_size_k)
            * block_size_k
        )

    num_ksplit = triton.cdiv(k_packed, (splitk_block_size // 2))
    return splitk_block_size, block_size_k, num_ksplit


@triton.heuristics(
    {
        "EVEN_K": lambda args: (args["k_packed"] % (args["BLOCK_SIZE_K"] // 2) == 0)
        and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
        and (args["k_packed"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
    }
)
@triton.jit
def _gemm_afp4wfp4_preshuffle_kernel(
    a_ptr,
    b_ptr,
    c_ptr,
    a_scales_ptr,
    b_scales_ptr,
    m,
    n,
    k_packed,
    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_bn > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)
    tl.assume(stride_asm > 0)
    tl.assume(stride_ask > 0)
    tl.assume(stride_bsn > 0)
    tl.assume(stride_bsk > 0)

    scale_group_size: tl.constexpr = 32
    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

    if (pid_k * SPLITK_BLOCK_SIZE // 2) < k_packed:
        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_iter 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:
                k_remaining = k_packed - k_iter * (BLOCK_SIZE_K // 2)
                a = tl.load(a_ptrs, mask=offs_k[None, :] < k_remaining, other=0)
                b = tl.load(
                    b_ptrs,
                    mask=offs_k_shuffle_arr[None, :] < (k_remaining * 16),
                    other=0,
                    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")


@triton.jit
def _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 _run_custom_preshuffle(
    a_q_u8: torch.Tensor,
    b_shuffled_u8: torch.Tensor,
    a_scale_triton_u8: torch.Tensor,
    b_scale_triton_u8: torch.Tensor,
    dtype: torch.dtype = torch.bfloat16,
    config_override: dict | None = None,
) -> torch.Tensor:
    m, k_packed = a_q_u8.shape
    n_16, k_16 = b_shuffled_u8.shape
    n = n_16 * 16
    assert k_16 == k_packed * 16, (a_q_u8.shape, b_shuffled_u8.shape)

    config = dict(config_override) if config_override is not None else _get_config(m, n, k_packed)
    if config["NUM_KSPLIT"] > 1:
        splitk_block_size, block_size_k, num_ksplit = _get_splitk(
            k_packed,
            config["BLOCK_SIZE_K"],
            config["NUM_KSPLIT"],
        )
        config["SPLITK_BLOCK_SIZE"] = splitk_block_size
        config["BLOCK_SIZE_K"] = block_size_k
        config["NUM_KSPLIT"] = num_ksplit
    else:
        config["SPLITK_BLOCK_SIZE"] = 2 * k_packed

    if config["BLOCK_SIZE_K"] >= 2 * k_packed:
        config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * k_packed)
        config["SPLITK_BLOCK_SIZE"] = 2 * k_packed

    config["BLOCK_SIZE_N"] = max(config["BLOCK_SIZE_N"], 32)
    if m < 32:
        assert config["BLOCK_SIZE_M"] <= 16
    else:
        assert config["BLOCK_SIZE_M"] >= 32

    y = torch.empty((m, n), dtype=dtype, device=a_q_u8.device)
    y_pp = None
    if config["NUM_KSPLIT"] > 1:
        y_pp = torch.empty((config["NUM_KSPLIT"], m, n), dtype=torch.float32, device=a_q_u8.device)

    grid = lambda meta: (  # noqa: E731
        (
            meta["NUM_KSPLIT"]
            * triton.cdiv(m, meta["BLOCK_SIZE_M"])
            * triton.cdiv(n, meta["BLOCK_SIZE_N"])
        ),
    )

    _gemm_afp4wfp4_preshuffle_kernel[grid](
        a_q_u8,
        b_shuffled_u8,
        y if y_pp is None else y_pp,
        a_scale_triton_u8,
        b_scale_triton_u8,
        m,
        n,
        k_packed,
        a_q_u8.stride(0),
        a_q_u8.stride(1),
        b_shuffled_u8.stride(0),
        b_shuffled_u8.stride(1),
        0 if y_pp is None else y_pp.stride(0),
        y.stride(0) if y_pp is None else y_pp.stride(1),
        y.stride(1) if y_pp is None else y_pp.stride(2),
        a_scale_triton_u8.stride(0),
        a_scale_triton_u8.stride(1),
        b_scale_triton_u8.stride(0),
        b_scale_triton_u8.stride(1),
        **config,
    )

    if y_pp is None:
        return y

    reduce_block_size_m = 16
    reduce_block_size_n = 64
    actual_ksplit = triton.cdiv(k_packed, (config["SPLITK_BLOCK_SIZE"] // 2))
    grid_reduce = (
        triton.cdiv(m, reduce_block_size_m),
        triton.cdiv(n, reduce_block_size_n),
    )
    _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),
        reduce_block_size_m,
        reduce_block_size_n,
        actual_ksplit,
        triton.next_power_of_2(config["NUM_KSPLIT"]),
    )
    return y

def _prepare_a_scale(a_scale: torch.Tensor, m: int) -> torch.Tensor:
    if m < 32:
        return a_scale.view(torch.uint8).contiguous()
    return _inline_e8m0_shuffle(a_scale).view(torch.uint8).reshape(-1, a_scale.shape[1] * 32).contiguous()


def _prepare_b_for_triton(b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor):
    key = (
        int(b_shuffle.device.index or 0),
        *_tensor_cache_key(b_shuffle),
        *_tensor_cache_key(b_scale_sh),
        tuple(int(dim) for dim in b_shuffle.shape),
        tuple(int(dim) for dim in b_scale_sh.shape),
    )
    cached = _B_TRITON_CACHE.get(key)
    if cached is not None:
        return cached
    b_shuffle_u8 = b_shuffle.view(torch.uint8).reshape(b_shuffle.shape[0] // 16, b_shuffle.shape[1] * 16)
    b_scale_u8 = b_scale_sh.view(torch.uint8).reshape(b_scale_sh.shape[0] // 32, b_scale_sh.shape[1] * 32)
    cached = (b_shuffle_u8.contiguous(), b_scale_u8.contiguous())
    _B_TRITON_CACHE[key] = cached
    return cached


def _run_fused_preshuffle(
    a: torch.Tensor,
    b_shuffled_u8: torch.Tensor,
    b_scale_triton_u8: torch.Tensor,
    config: dict,
):
    return gemm_a16wfp4_preshuffle(
        a,
        b_shuffled_u8,
        b_scale_triton_u8,
        dtype=dtypes.bf16,
        config=config,
    )


def _compute_uncached(
    a: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
) -> torch.Tensor:
    shape = (int(a.shape[0]), int(b_shuffle.shape[0]), int(a.shape[1]))
    try:
        b_shuffle_triton, b_scale_triton = _prepare_b_for_triton(b_shuffle, b_scale_sh)
        fused_config = _select_config(_FUSED_CONFIGS, *shape)
        if fused_config is not None:
            return _run_fused_preshuffle(
                a,
                b_shuffle_triton,
                b_scale_triton,
                fused_config,
            )
        if shape in _REFERENCE_SHAPES:
            a_q_u8, a_scale_sh = _inline_dynamic_mxfp4_quant_shuffled(a)
            return _reference(a_q_u8.view(dtypes.fp4x2), b_shuffle, a_scale_sh, b_scale_sh)
        a_q_u8, a_scale = _inline_dynamic_mxfp4_quant(a)
        if shape == (16, 2112, 7168):
            return _run_custom_preshuffle(
                a_q_u8,
                b_shuffle_triton,
                _prepare_a_scale(a_scale, a.shape[0]),
                b_scale_triton,
                dtype=torch.bfloat16,
                config_override=_OVERRIDE_CONFIGS[(2112, 7168)]["M_LEQ_31"],
            )
        return _run_custom_preshuffle(
            a_q_u8,
            b_shuffle_triton,
            _prepare_a_scale(a_scale, a.shape[0]),
            b_scale_triton,
            dtype=torch.bfloat16,
        )
    except Exception:
        a_q, a_scale = dynamic_mxfp4_quant(a)
        a_scale_sh = e8m0_shuffle(a_scale)
        return _reference(a_q, b_shuffle, a_scale_sh, b_scale_sh)


def _reference(a_q: torch.Tensor, b_shuffle: torch.Tensor, a_scale_sh: torch.Tensor, b_scale_sh: torch.Tensor):
    import aiter

    return aiter.gemm_a4w4(
        a_q.view(dtypes.fp4x2),
        b_shuffle,
        a_scale_sh.view(dtypes.fp8_e8m0),
        b_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )


def custom_kernel(data: input_t) -> output_t:
    a, _b, _b_q, b_shuffle, b_scale_sh = data
    a = a.contiguous()
    b_shuffle = b_shuffle.contiguous()
    b_scale_sh = b_scale_sh.contiguous()
    return _compute_uncached(a, b_shuffle, b_scale_sh)
scrolls · 905 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