Skip to content
KernelIndex
Search⌘K

submission 564882

trungnob · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v_23.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-564882?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
10.1µs
#245 of 1143
2026-03-16

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ec3af872bf871deb44a1e888b12ecfdc49a9adecd626f308e402e11e6726bd37
license declaredunknown
license concludedunknown
authorstrungnob
imported2026-08-15

Techniques

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

fp4MXFP4 GEMM hybrid:
shared-memory__shared__ float s_abs[WAVE_SIZE];
split-k"SPLITK_BLOCK_SIZE": 2 * 7168 // 4,

Kernel source

submission_v_23.py1082 lines
"""
MXFP4 GEMM hybrid:
- monkey-patched aiter.gemm_a16wfp4 for shapes where direct BF16-A/FP4-B wins
- HIP wavefront quant + tuned ASM GEMM where the direct path loses
"""
import os
from typing import Optional

os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"

import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
import aiter
import aiter.ops.triton.utils._triton.arch_info as arch_info
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm, get_padded_m
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
    _gemm_afp4wfp4_reduce_kernel,
)
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid
from aiter.ops.triton.utils.common_utils import deserialize_str
from aiter.ops.triton.utils.gemm_config_utils import get_gemm_config
from task import input_t, output_t


TUNED_CONFIGS = {
    (4, 2880, 512): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x256E", 3),
    (8, 2112, 7168): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 2),
    (16, 2112, 7168): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 1),
    (16, 3072, 1536): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0),
    (32, 4096, 512): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_96x384E", 1),
    (32, 2880, 512): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0),
    (64, 7168, 2048): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 1),
    (64, 3072, 1536): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0),
    (256, 2880, 512): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x384E", 2),
    (256, 3072, 1536): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 1),
}

DIRECT_SHAPES = {
    (4, 2880, 512),
    (32, 4096, 512),
    (32, 2880, 512),
    (64, 7168, 2048),
}

OUTLIER_RAW_SHAPE = (16, 2112, 7168)
OUTLIER_RAW_CONFIG = {
    "BLOCK_SIZE_M": 16,
    "BLOCK_SIZE_N": 128,
    # v7 candidate: keep the fused outlier route but halve the K-loop count.
    "BLOCK_SIZE_K": 256,
    "GROUP_SIZE_M": 1,
    "NUM_KSPLIT": 4,
    "SPLITK_BLOCK_SIZE": 2 * 7168 // 4,
    "num_warps": 8,
    "num_stages": 2,
    "waves_per_eu": 2,
    "matrix_instr_nonkdim": 16,
    "cache_modifier": ".cg",
}


HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <stdint.h>

#define SCALE_GROUP_SIZE 32
#define WAVE_SIZE 256
#define GROUPS_PER_BLOCK (WAVE_SIZE / 32)

__device__ __forceinline__ uint8_t fp32_to_e2m1_exact(float qv) {
    constexpr uint32_t EXP_BIAS_FP32 = 127;
    constexpr uint32_t EXP_BIAS_FP4 = 1;
    constexpr float max_normal = 6.0f;
    constexpr float min_normal = 1.0f;

    uint32_t bits = __float_as_uint(qv);
    uint32_t sign = bits & 0x80000000u;
    bits ^= sign;
    float absv = __uint_as_float(bits);

    uint8_t result;
    if (absv >= max_normal) {
        result = 0x7;
    } else if (absv < min_normal) {
        constexpr uint32_t denorm_exp =
            (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (23 - 1) + 1;
        constexpr uint32_t denorm_mask_int = denorm_exp << 23;
        float denorm_mask_float = __uint_as_float(denorm_mask_int);
        float denormal_x = absv + denorm_mask_float;
        uint32_t denormal_bits = __float_as_uint(denormal_x);
        denormal_bits -= denorm_mask_int;
        result = static_cast<uint8_t>(denormal_bits & 0xFF);
    } else {
        uint32_t normal_x = bits;
        uint32_t mant_odd = (normal_x >> (23 - 1)) & 1;
        int32_t val_to_add =
            ((static_cast<int32_t>(EXP_BIAS_FP4 - EXP_BIAS_FP32)) << 23) + (1 << 21) - 1;
        normal_x += static_cast<uint32_t>(val_to_add);
        normal_x += mant_odd;
        normal_x >>= (23 - 1);
        result = static_cast<uint8_t>(normal_x & 0xFF);
    }

    uint8_t sign_lp = static_cast<uint8_t>(
        (sign >> (23 + 8 - 1 - 2)) & 0x8
    );
    return result | sign_lp;
}

__device__ __forceinline__ uint8_t compute_e8m0_scale_exact(float amax) {
    if (amax == 0.0f) {
        return 127;
    }
    uint32_t bits = __float_as_uint(amax);
    bits = (bits + 0x200000u) & 0xFF800000u;
    int scale_unbiased = static_cast<int>((bits >> 23) & 0xFF) - 127 - 2;
    if (scale_unbiased < -127) scale_unbiased = -127;
    if (scale_unbiased > 127) scale_unbiased = 127;
    return static_cast<uint8_t>(scale_unbiased + 127);
}

__global__ __launch_bounds__(WAVE_SIZE)
void fused_quant_wave_kernel(
    const __hip_bfloat16* __restrict__ A,
    uint8_t* __restrict__ A_fp4,
    uint8_t* __restrict__ A_scale,
    int M,
    int K,
    int scaleN,
    int scaleN_pad
) {
    __shared__ float s_abs[WAVE_SIZE];
    __shared__ uint8_t s_nib[WAVE_SIZE];

    int row = blockIdx.x;
    int group_block = blockIdx.y;
    int lane = threadIdx.x;

    if (row >= M) {
        return;
    }

    int half = lane >> 5;
    int lane32 = lane & 31;
    int group = group_block * GROUPS_PER_BLOCK + half;
    bool valid_group = group < scaleN;

    float v = 0.0f;
    if (valid_group) {
        int k_idx = group * SCALE_GROUP_SIZE + lane32;
        v = __bfloat162float(A[row * K + k_idx]);
    }

    s_abs[lane] = fabsf(v);
    __syncthreads();

    int half_base = half * 32;
    for (int offset = 16; offset > 0; offset >>= 1) {
        if (lane32 < offset) {
            float other = s_abs[half_base + lane32 + offset];
            s_abs[half_base + lane32] = fmaxf(s_abs[half_base + lane32], other);
        }
        __syncthreads();
    }

    uint8_t scale_e8m0 = compute_e8m0_scale_exact(s_abs[half_base]);
    float scale_unbiased = static_cast<float>(static_cast<int>(scale_e8m0) - 127);
    float inv_scale = exp2f(-scale_unbiased);
    s_nib[lane] = fp32_to_e2m1_exact(v * inv_scale);
    __syncthreads();

    if (valid_group && lane32 < 16) {
        int nib_base = half_base + lane32 * 2;
        uint8_t lo = s_nib[nib_base];
        uint8_t hi = s_nib[nib_base + 1];
        A_fp4[row * (K / 2) + group * 16 + lane32] = lo | (hi << 4);
    }

    if (valid_group && lane32 == 0) {
        int d0 = row / 32;
        int rem_r = row % 32;
        int d1 = rem_r / 16;
        int d2 = rem_r % 16;
        int d3 = group / 8;
        int rem_c = group % 8;
        int d4 = rem_c / 4;
        int d5 = rem_c % 4;
        int sn8 = scaleN_pad / 8;
        int shuffle_idx = d0 * (sn8 * 256) + d3 * 256 + d5 * 64 + d2 * 4 + d4 * 2 + d1;
        A_scale[shuffle_idx] = scale_e8m0;
    }
}

void fused_mxfp4_quant_out(
    torch::Tensor A,
    torch::Tensor A_fp4,
    torch::Tensor A_scale,
    int scaleN,
    int scaleN_pad
) {
    int M = static_cast<int>(A.size(0));
    int K = static_cast<int>(A.size(1));
    dim3 blocks(M, (scaleN + GROUPS_PER_BLOCK - 1) / GROUPS_PER_BLOCK);
    dim3 threads(WAVE_SIZE);

    fused_quant_wave_kernel<<<blocks, threads>>>(
        reinterpret_cast<const __hip_bfloat16*>(A.data_ptr()),
        A_fp4.data_ptr<uint8_t>(),
        A_scale.data_ptr<uint8_t>(),
        M,
        K,
        scaleN,
        scaleN_pad
    );
}
"""


CPP_SRC = r"""
#include <torch/extension.h>

void fused_mxfp4_quant_out(
    torch::Tensor A,
    torch::Tensor A_fp4,
    torch::Tensor A_scale,
    int scaleN,
    int scaleN_pad
);
"""


_module = None
_hip_buf_cache = {}
_index_cache = {}
_raw_scale_buf_cache = {}
_view_scale_cache = {}
_direct_out_cache = {}
_fused_pp_cache = {}
_patched = False
_status_printed = set()


def _status(key, message: str) -> None:
    if key not in _status_printed:
        _status_printed.add(key)
        print(f"[CODEX-HYBRID] {message}", flush=True)


@triton.jit
def _unshuffle_e8m0_kernel(
    scale_sh_ptr,
    raw_ptr,
    sm: tl.constexpr,
    sn: tl.constexpr,
    rows: tl.constexpr,
    cols: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    pid = tl.program_id(0)
    offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    mask = offsets < (rows * cols)

    sn8: tl.constexpr = sn // 8
    r = offsets // cols
    c = offsets % cols

    d0 = r // 32
    rem_r = r % 32
    d1 = rem_r // 16
    d2 = rem_r % 16
    d3 = c // 8
    rem_c = c % 8
    d4 = rem_c // 4
    d5 = rem_c % 4

    shuffled_idx = d0 * (sn8 * 256) + d3 * 256 + d5 * 64 + d2 * 4 + d4 * 2 + d1
    val = tl.load(scale_sh_ptr + shuffled_idx, mask=mask)
    tl.store(raw_ptr + offsets, val, mask=mask)


def fast_unshuffle_e8m0(scale_sh: torch.Tensor, rows: int, cols: int) -> torch.Tensor:
    sm, sn = scale_sh.shape
    device_index = scale_sh.device.index if scale_sh.device.index is not None else 0
    key = (device_index, rows, cols)

    raw_flat = _raw_scale_buf_cache.get(key)
    if raw_flat is None:
        raw_flat = torch.empty(rows * cols, dtype=torch.uint8, device=scale_sh.device)
        _raw_scale_buf_cache[key] = raw_flat

    block_size = 1024
    grid = (triton.cdiv(rows * cols, block_size),)
    _unshuffle_e8m0_kernel[grid](
        scale_sh.view(torch.uint8),
        raw_flat,
        sm,
        sn,
        rows,
        cols,
        BLOCK_SIZE=block_size,
    )
    return raw_flat.view(rows, cols).view(scale_sh.dtype)


_gemm_a16wfp4_fused_repr = make_kernel_repr(
    "_gemm_a16wfp4_fused_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),
        "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
        * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
    }
)
@triton.jit(repr=_gemm_a16wfp4_fused_repr)
def _gemm_a16wfp4_fused_kernel(
    a_ptr,
    b_ptr,
    c_ptr,
    b_scales_ptr,
    M,
    N,
    K,
    stride_am,
    stride_ak,
    stride_bk,
    stride_bn,
    stride_ck,
    stride_cm,
    stride_cn,
    scale_sh_cols,
    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,
    ATOMIC_ADD: 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(scale_sh_cols > 0)

    pid_unified = tl.program_id(axis=0)
    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

    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 = tl.arange(0, BLOCK_SIZE_K // 2)
        offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
        offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
        b_ptrs = b_ptr + (offs_k_split[:, None] * stride_bk + offs_bn[None, :] * stride_bn)

        ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)) + tl.arange(
            0, BLOCK_SIZE_K // SCALE_GROUP_SIZE
        )
        offs_bn_2d = offs_bn[:, None]
        ks_2d = ks[None, :]
        sn8 = scale_sh_cols // 8

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

        for _ in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
            shuffle_idx = (
                (offs_bn_2d // 32) * (sn8 * 256)
                + (ks_2d // 8) * 256
                + ((ks_2d % 8) % 4) * 64
                + (offs_bn_2d % 16) * 4
                + ((ks_2d % 8) // 4) * 2
                + ((offs_bn_2d % 32) // 16)
            )
            b_scales = tl.load(b_scales_ptr + shuffle_idx)

            if EVEN_K:
                a_bf16 = tl.load(a_ptrs)
                b = tl.load(b_ptrs, cache_modifier=cache_modifier)
            else:
                a_bf16 = tl.load(
                    a_ptrs,
                    mask=offs_k_bf16[None, :] < 2 * K,
                    other=0,
                )
                b = tl.load(
                    b_ptrs,
                    mask=offs_k[:, None] < K,
                    other=0,
                    cache_modifier=cache_modifier,
                )

            a, a_scales = _mxfp4_quant_op(a_bf16, 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
            ks_2d += BLOCK_SIZE_K // SCALE_GROUP_SIZE

        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)
        if ATOMIC_ADD:
            tl.atomic_add(c_ptrs, c, mask=c_mask, sem="relaxed")
        else:
            tl.store(c_ptrs, c, mask=c_mask)


_gemm_a16wfp4_shuf_repr = make_kernel_repr(
    "_gemm_a16wfp4_shuffled_scales_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),
        "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
        * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
    }
)
@triton.jit(repr=_gemm_a16wfp4_shuf_repr)
def _gemm_a16wfp4_shuffled_scales_kernel(
    a_ptr,
    b_ptr,
    c_ptr,
    b_scales_ptr,
    M,
    N,
    K,
    stride_am,
    stride_ak,
    stride_bk,
    stride_bn,
    stride_ck,
    stride_cm,
    stride_cn,
    scale_sh_cols,
    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,
    ATOMIC_ADD: 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(scale_sh_cols > 0)

    pid_unified = tl.program_id(axis=0)
    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

    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 = tl.arange(0, BLOCK_SIZE_K // 2)
        offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
        offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
        b_ptrs = b_ptr + (
            offs_k_split[:, None] * stride_bk + offs_bn[None, :] * stride_bn
        )

        ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)) + tl.arange(
            0, BLOCK_SIZE_K // SCALE_GROUP_SIZE
        )
        offs_bn_2d = offs_bn[:, None]
        ks_2d = ks[None, :]
        sn8 = scale_sh_cols // 8

        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):
            shuffle_idx = (
                (offs_bn_2d // 32) * (sn8 * 256)
                + (ks_2d // 8) * 256
                + ((ks_2d % 8) % 4) * 64
                + (offs_bn_2d % 16) * 4
                + ((ks_2d % 8) // 4) * 2
                + ((offs_bn_2d % 32) // 16)
            )
            b_scales = tl.load(b_scales_ptr + shuffle_idx)

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

            a, a_scales = _mxfp4_quant_op(a_bf16, 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
            ks_2d += BLOCK_SIZE_K // SCALE_GROUP_SIZE

        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)
        if ATOMIC_ADD:
            tl.atomic_add(c_ptrs, c, mask=c_mask, sem="relaxed")
        else:
            tl.store(c_ptrs, c, mask=c_mask)


def _a16_shuf_get_config(M: int, N: int, K: int):
    return get_gemm_config("GEMM-A16WFP4", M, N, 2 * K)


def _a16_shuf_get_splitk(K: int, BLOCK_SIZE_K: int, NUM_KSPLIT: int):
    num_ksplit_step = 2
    block_size_k_step = 2
    splitk_block_size = (
        triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
    )
    while NUM_KSPLIT > 1 and BLOCK_SIZE_K > 16:
        if (
            K % (splitk_block_size // 2) == 0
            and splitk_block_size % BLOCK_SIZE_K == 0
            and K % (BLOCK_SIZE_K // 2) == 0
        ):
            break
        if K % (splitk_block_size // 2) != 0 and NUM_KSPLIT > 1:
            NUM_KSPLIT = NUM_KSPLIT // num_ksplit_step
        elif splitk_block_size % BLOCK_SIZE_K != 0:
            if NUM_KSPLIT > 1:
                NUM_KSPLIT = NUM_KSPLIT // num_ksplit_step
            elif BLOCK_SIZE_K > 16:
                BLOCK_SIZE_K = BLOCK_SIZE_K // block_size_k_step
        elif K % (BLOCK_SIZE_K // 2) != 0 and BLOCK_SIZE_K > 16:
            BLOCK_SIZE_K = BLOCK_SIZE_K // block_size_k_step
        else:
            break

        splitk_block_size = (
            triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
        )

    return splitk_block_size, BLOCK_SIZE_K, NUM_KSPLIT


def gemm_a16wfp4_shuffled_scales(
    x: torch.Tensor,
    w: torch.Tensor,
    w_scales_sh: torch.Tensor,
    atomic_add: Optional[bool] = False,
    dtype: Optional[torch.dtype] = torch.bfloat16,
    y: Optional[torch.Tensor] = None,
    config: Optional[str] = None,
) -> torch.Tensor:
    assert arch_info.is_fp4_avail(), "MXFP4 is not available on your device"

    M, K_bf16 = x.shape
    N, K_packed = w.shape
    assert K_packed * 2 == K_bf16, f"expected packed K/2 weights, got x={x.shape} w={w.shape}"

    w = w.view(torch.uint8).T
    w_scales_sh_u8 = w_scales_sh.view(torch.uint8).contiguous()

    if config is None:
        config, _ = _a16_shuf_get_config(M, N, K_packed)
    else:
        config = deserialize_str(config)

    if y is None:
        if atomic_add:
            y = torch.zeros((M, N), dtype=dtype, device=x.device)
        else:
            y = torch.empty((M, N), dtype=dtype, device=x.device)

    if config["NUM_KSPLIT"] > 1 and not atomic_add:
        splitk_block_size, block_size_k, num_ksplit = _a16_shuf_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

    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["NUM_KSPLIT"] = 1
    config["BLOCK_SIZE_K"] = max(config["BLOCK_SIZE_K"], 64)

    if config["NUM_KSPLIT"] > 1 and not atomic_add:
        y_pp = torch.empty(
            (config["NUM_KSPLIT"], M, N), dtype=torch.float32, device=y.device
        )
    else:
        config["SPLITK_BLOCK_SIZE"] = 2 * K_packed
        y_pp = None

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

    _gemm_a16wfp4_shuffled_scales_kernel[grid](
        x,
        w,
        y if y_pp is None else y_pp,
        w_scales_sh_u8.view(-1),
        M,
        N,
        K_packed,
        x.stride(0),
        x.stride(1),
        w.stride(0),
        w.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),
        w_scales_sh_u8.shape[1],
        ATOMIC_ADD=atomic_add,
        **config,
    )

    if config["NUM_KSPLIT"] > 1 and not atomic_add:
        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),
        )
        _gemm_afp4wfp4_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 _get_module():
    global _module
    from torch.utils.cpp_extension import load_inline
    if _module is None:
        _module = load_inline(
            name="fused_quant_wave_combo_hybrid_v2",
            cpp_sources=[CPP_SRC],
            cuda_sources=[HIP_SRC],
            functions=["fused_mxfp4_quant_out"],
            verbose=False,
            extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++17", "-O3"],
        )
    return _module


def _hip_quant_buffers(a: torch.Tensor):
    m, k = a.shape
    scale_n_valid = (k + 31) // 32
    scale_n_pad = ((scale_n_valid + 7) // 8) * 8
    scale_m_pad = ((m + 255) // 256) * 256
    key = (a.device.index if a.device.index is not None else 0, m, k, scale_m_pad, scale_n_pad)
    cached = _hip_buf_cache.get(key)
    if cached is None:
        a_fp4_u8 = torch.empty((m, k // 2), dtype=torch.uint8, device=a.device)
        a_scale_u8 = torch.full((scale_m_pad * scale_n_pad,), 127, dtype=torch.uint8, device=a.device)
        _hip_buf_cache[key] = (a_fp4_u8, a_scale_u8, scale_n_valid, scale_m_pad, scale_n_pad)
    return _hip_buf_cache[key]


def _ensure_patch() -> None:
    global _patched
    if _patched:
        return

    import triton._utils as tu

    tu.type_canonicalisation_dict["float4_e2m1fn_x2"] = "u8"
    tu.type_canonicalisation_dict["float4_e2m1fn"] = "u8"
    tu.type_canonicalisation_dict["float8_e8m0fnu"] = "u8"
    tu.type_canonicalisation_dict["float8_e8m0fnuz"] = "u8"
    tu.BITWIDTH_DICT["float4_e2m1fn_x2"] = 8
    tu.BITWIDTH_DICT["float4_e2m1fn"] = 8
    tu.BITWIDTH_DICT["float8_e8m0fnu"] = 8
    tu.BITWIDTH_DICT["float8_e8m0fnuz"] = 8
    _patched = True


def _unshuffle_e8m0(scale_sh: torch.Tensor, rows: int, cols: int) -> torch.Tensor:
    sm, sn = scale_sh.shape
    device_index = scale_sh.device.index if scale_sh.device.index is not None else 0
    key = (device_index, sm, sn, rows, cols)

    idx = _index_cache.get(key)
    raw_flat = _raw_scale_buf_cache.get(key)
    if idx is None or raw_flat is None:
        # Precompute the inverse-shuffle gather plan once per shape. The tensor
        # values change every ranked iteration, but the permutation does not.
        sn8 = sn // 8
        row_ids = torch.arange(rows, dtype=torch.int64).unsqueeze(1)
        col_ids = torch.arange(cols, dtype=torch.int64).unsqueeze(0)

        d0 = row_ids // 32
        rem_r = row_ids % 32
        d1 = rem_r // 16
        d2 = rem_r % 16
        d3 = col_ids // 8
        rem_c = col_ids % 8
        d4 = rem_c // 4
        d5 = rem_c % 4

        idx = (
            d0 * (sn8 * 256)
            + d3 * 256
            + d5 * 64
            + d2 * 4
            + d4 * 2
            + d1
        ).reshape(-1).to(scale_sh.device)
        raw_flat = torch.empty(rows * cols, dtype=torch.uint8, device=scale_sh.device)
        _index_cache[key] = idx
        _raw_scale_buf_cache[key] = raw_flat

    scale_u8_flat = scale_sh.view(torch.uint8).reshape(-1)
    torch.index_select(scale_u8_flat, 0, idx, out=raw_flat)
    return raw_flat.view(rows, cols).view(scale_sh.dtype)


def _unshuffle_e8m0_view(scale_sh: torch.Tensor, rows: int, cols: int) -> torch.Tensor:
    device_index = scale_sh.device.index if scale_sh.device.index is not None else 0
    key = (id(scale_sh), scale_sh.data_ptr(), rows, cols, device_index)
    cached = _view_scale_cache.get(key)
    if cached is not None:
        return cached[1]

    sm, sn = scale_sh.shape
    scale_u8 = scale_sh.view(torch.uint8).contiguous()
    raw_u8 = (
        scale_u8.view(sm // 32, sn // 8, 4, 16, 2, 2)
        .permute(0, 5, 3, 1, 4, 2)
        .contiguous()
        .view(sm, sn)
    )
    raw = raw_u8[:rows, :cols].contiguous().view(scale_sh.dtype)
    # Keep the source tensor alive so Python object ids cannot be recycled into
    # stale cache hits across later ranked inputs.
    _view_scale_cache[key] = (scale_sh, raw)
    return raw


def _direct_output_buffer(a: torch.Tensor, n: int) -> torch.Tensor:
    key = (a.device.index, a.shape[0], n)
    cached = _direct_out_cache.get(key)
    if cached is None:
        cached = torch.empty((a.shape[0], n), dtype=torch.bfloat16, device=a.device)
        _direct_out_cache[key] = cached
    return cached


def _fused_pp_buffer(device: torch.device, num_ksplit: int, m: int, n: int) -> torch.Tensor:
    device_index = device.index if device.index is not None else 0
    key = (device_index, num_ksplit, m, n)
    cached = _fused_pp_cache.get(key)
    if cached is None:
        cached = torch.empty((num_ksplit, m, n), dtype=torch.float32, device=device)
        _fused_pp_cache[key] = cached
    return cached


def gemm_a16wfp4_fused(
    A: torch.Tensor,
    B_q: torch.Tensor,
    B_scale_sh: torch.Tensor,
    config_override: dict | None = None,
    atomic_add: bool = False,
    y: torch.Tensor | None = None,
):
    M, K_bf16 = A.shape
    B_q_u8 = B_q.view(torch.uint8).contiguous()
    B_scale_sh_u8 = B_scale_sh.view(torch.uint8).contiguous()
    N, K_packed = B_q_u8.shape

    if config_override is None:
        config, _ = get_gemm_config("GEMM-A16WFP4", M, N, K_bf16)
    else:
        config = dict(config_override)

    if config["NUM_KSPLIT"] > 1 and not atomic_add:
        splitk_block_size, block_size_k, num_ksplit = _a16_shuf_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

    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["NUM_KSPLIT"] = 1
    config["BLOCK_SIZE_K"] = max(config["BLOCK_SIZE_K"], 64)

    if y is None:
        y = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
    if config["NUM_KSPLIT"] > 1 and not atomic_add:
        y_pp = _fused_pp_buffer(A.device, config["NUM_KSPLIT"], M, N)
    else:
        y_pp = None

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

    B = B_q_u8.t()
    _gemm_a16wfp4_fused_kernel[grid](
        A,
        B,
        y if y_pp is None else y_pp,
        B_scale_sh_u8.view(-1),
        M,
        N,
        K_packed,
        A.stride(0),
        A.stride(1),
        B.stride(0),
        B.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),
        B_scale_sh_u8.shape[1],
        ATOMIC_ADD=atomic_add,
        **config,
    )

    if config["NUM_KSPLIT"] > 1 and not atomic_add:
        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),
        )
        _gemm_afp4wfp4_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 _direct_path(A: torch.Tensor, B: torch.Tensor, B_q: torch.Tensor, B_scale_sh: torch.Tensor):
    n = B.shape[0]
    out = _direct_output_buffer(A, n)
    return gemm_a16wfp4_shuffled_scales(A, B_q, B_scale_sh, dtype=torch.bfloat16, y=out)


def _direct_raw_outlier_path(A: torch.Tensor, B: torch.Tensor, B_q: torch.Tensor, B_scale_sh: torch.Tensor):
    from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4

    n = B.shape[0]
    k = A.shape[1]
    _ensure_patch()
    out = _direct_output_buffer(A, n)
    B_scale_raw = fast_unshuffle_e8m0(B_scale_sh, n, k // 32)
    return gemm_a16wfp4(
        A,
        B_q.view(torch.uint8),
        B_scale_raw,
        dtype=torch.bfloat16,
        y=out,
        config=OUTLIER_RAW_CONFIG,
    )


def _direct_fused_outlier_path(A: torch.Tensor, B: torch.Tensor, B_q: torch.Tensor, B_scale_sh: torch.Tensor):
    n = B.shape[0]
    out = _direct_output_buffer(A, n)
    return gemm_a16wfp4_fused(
        A,
        B_q,
        B_scale_sh,
        config_override=OUTLIER_RAW_CONFIG,
        atomic_add=False,
        y=out,
    )


def _hip_path(a: torch.Tensor, b: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor):
    m, k = a.shape
    n = b.shape[0]

    try:
        mod = _get_module()
        a_fp4_u8, a_scale_u8, scale_n_valid, scale_m_pad, scale_n_pad = _hip_quant_buffers(a)
        mod.fused_mxfp4_quant_out(a, a_fp4_u8, a_scale_u8, scale_n_valid, scale_n_pad)
        a_q = a_fp4_u8.view(dtypes.fp4x2)
        a_scale_sh = a_scale_u8.view(scale_m_pad, scale_n_pad).view(dtypes.fp8_e8m0)
    except Exception as exc:
        _status((m, n, k, "hip-fallback"), f"HIP quant failed for {(m, n, k)}: {repr(exc)}")
        from aiter.ops.triton.quant import dynamic_mxfp4_quant
        from aiter.utility.fp4_utils import e8m0_shuffle

        x_fp4, bs_e8m0 = dynamic_mxfp4_quant(a)
        a_q = x_fp4.view(dtypes.fp4x2)
        a_scale_sh = e8m0_shuffle(bs_e8m0)

    config = TUNED_CONFIGS.get((m, n, k))
    if config is not None:
        knl_name, split_k = config
        padded_m = get_padded_m(m, n, k, 0)
        out = torch.empty(padded_m, n, dtype=torch.bfloat16, device=a.device)
        gemm_a4w4_asm(
            a_q, b_shuffle, a_scale_sh, b_scale_sh,
            out, knl_name, None, 1.0, 0.0, True, split_k,
        )
        return out[:m]

    return aiter.gemm_a4w4(
        a_q, b_shuffle, a_scale_sh, 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
    shape = (A.shape[0], B.shape[0], A.shape[1])
    _ensure_patch()

    if shape == OUTLIER_RAW_SHAPE:
        _status((shape, "route"), f"route fused-outlier for {shape}")
        try:
            return _direct_fused_outlier_path(A, B, B_q, B_scale_sh)
        except Exception as exc:
            _status((shape, "direct-fused-fail"), f"fused outlier failed for {shape}: {repr(exc)}")
        try:
            return _direct_raw_outlier_path(A, B, B_q, B_scale_sh)
        except Exception as exc:
            _status((shape, "direct-raw-fail"), f"direct raw custom failed for {shape}: {repr(exc)}")

    if shape in DIRECT_SHAPES:
        _status((shape, "route"), f"route direct for {shape}")
        try:
            return _direct_path(A, B, B_q, B_scale_sh)
        except Exception as exc:
            _status((shape, "direct-fail"), f"direct failed for {shape}: {repr(exc)}")

    _status((shape, "route"), f"route HIP for {shape}")
    return _hip_path(A, B, B_shuffle, B_scale_sh)
scrolls · 1082 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