Skip to content
KernelIndex
Search⌘K

submission 693471

nikxkilla · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-693471?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
#444 of 1143
2026-04-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:01f8afa7c3322fa771342e686691a4512c5622d25a7da4976451f574fb54bcef
license declaredunknown
license concludedunknown
authorsnikxkilla
imported2026-08-26

Techniques

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

fp4m.def("quant_into_shuffled_hip", &quant_into_shuffled_hip, "HIP small-m native fp4 quant");
shared-memory__shared__ float s_vals[GROUPS_PER_BLOCK][32];
split-k_CUSTOM_CSV = """cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio
stages = 1num_stages=1,

Kernel source

submission_v3.py672 lines
from task import input_t, output_t

import functools
import os
import tempfile
import textwrap

import triton
import triton.language as tl

if "PYTORCH_ROCM_ARCH" not in os.environ:
    os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"

_WS = {}
_CSV_PATH = None

_CUSTOM_CSV = """cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio
256,4,2880,512,21,0,0.0,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E,0.0,0.0,0.0
256,16,2112,7168,21,0,0.0,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E,0.0,0.0,0.0
256,32,4096,512,21,0,0.0,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E,0.0,0.0,0.0
256,32,2880,512,21,0,0.0,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E,0.0,0.0,0.0
256,64,7168,2048,21,0,6.8112,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E,275.88,1221.97,0.0
256,256,3072,1536,21,0,6.1771,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E,391.11,668.4,0.0
"""

_HIP_SRC = textwrap.dedent(
    r"""
    #include <torch/extension.h>
    #include <pybind11/pybind11.h>
    #include <hip/hip_runtime.h>
    #include <hip/amd_detail/amd_hip_bf16.h>
    #include <cstdint>
    #include <stdexcept>

    #define CHECK_CUDA(x) TORCH_CHECK(x.is_cuda(), #x " must be on ROCm device")
    #define CHECK_CONTIGUOUS(x) TORCH_CHECK(x.is_contiguous(), #x " must be contiguous")
    #define CHECK_BF16(x) TORCH_CHECK(x.scalar_type() == torch::kBFloat16, #x " must be bf16")
    #define CHECK_U8(x) TORCH_CHECK(x.scalar_type() == torch::kUInt8, #x " must be uint8")

    static __device__ inline uint32_t f32_as_u32(float x) {
      return __builtin_bit_cast(uint32_t, x);
    }

    static __device__ inline int shuffled_scale_index(
        int row_idx, int scale_n_idx, int scale_n_pad) {
      int a = row_idx / 32;
      int m1 = row_idx % 32;
      int b = m1 / 16;
      int c = m1 % 16;

      int d = scale_n_idx / 8;
      int n1 = scale_n_idx % 8;
      int e = n1 / 4;
      int f = n1 % 4;

      int gn = scale_n_pad / 8;
      return (((((a * gn + d) * 4 + f) * 16 + c) * 2 + e) * 2 + b);
    }

    static __device__ inline uint8_t quant_one_fp4_e2m1(float x, uint8_t scale_e8m0) {
      constexpr int EXP_BIAS_FP32 = 127;
      constexpr int EXP_BIAS_FP4 = 1;
      constexpr int EBITS_F32 = 8;
      constexpr int EBITS_FP4 = 2;
      constexpr int MBITS_F32 = 23;
      constexpr int MBITS_FP4 = 1;
      constexpr float MAX_NORMAL = 6.0f;
      constexpr float MIN_NORMAL = 1.0f;

      int scale_unbiased = static_cast<int>(scale_e8m0) - 127;
      float qx = x * exp2f(-static_cast<float>(scale_unbiased));

      uint32_t qx_bits = f32_as_u32(qx);
      uint32_t sign = qx_bits & 0x80000000u;
      uint32_t abs_bits = qx_bits ^ sign;
      float abs_f = __builtin_bit_cast(float, abs_bits);

      bool saturate_mask = abs_f >= MAX_NORMAL;
      bool denormal_mask = (!saturate_mask) && (abs_f < MIN_NORMAL);

      uint8_t fp4_val = 0x7u;
      if (denormal_mask) {
        constexpr int denorm_exp =
            (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1;
        constexpr uint32_t denorm_mask_int =
            static_cast<uint32_t>(denorm_exp) << MBITS_F32;
        float denorm_mask_float = __builtin_bit_cast(float, denorm_mask_int);
        uint32_t denormal_x = f32_as_u32(abs_f + denorm_mask_float);
        denormal_x -= denorm_mask_int;
        fp4_val = static_cast<uint8_t>(denormal_x);
      } else if (!saturate_mask) {
        uint32_t mant_odd = (abs_bits >> (MBITS_F32 - MBITS_FP4)) & 1u;
        constexpr int32_t val_to_add =
            ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1;
        int32_t normal_x = static_cast<int32_t>(abs_bits);
        normal_x += val_to_add;
        normal_x += static_cast<int32_t>(mant_odd);
        normal_x >>= (MBITS_F32 - MBITS_FP4);
        fp4_val = static_cast<uint8_t>(normal_x);
      }

      uint8_t sign_lp = static_cast<uint8_t>(
          sign >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4));
      return static_cast<uint8_t>(fp4_val | sign_lp);
    }

    static __device__ inline uint8_t encode_scale_from_amax(float amax) {
      if (!(amax > 0.0f)) {
        return static_cast<uint8_t>(0);
      }
      uint32_t amax_bits = f32_as_u32(amax);
      uint32_t rounded = (amax_bits + 0x200000u) & 0xFF800000u;
      float rounded_amax = __builtin_bit_cast(float, rounded);
      int scale_unbiased = static_cast<int>(floorf(log2f(rounded_amax)) - 2.0f);
      if (scale_unbiased < -127) scale_unbiased = -127;
      if (scale_unbiased > 127) scale_unbiased = 127;
      return static_cast<uint8_t>(scale_unbiased + 127);
    }

    template <int GROUPS_PER_BLOCK>
    __global__ __launch_bounds__(GROUPS_PER_BLOCK * 32)
    void quant_smallm_native_fp4_kernel(
        const __hip_bfloat16* __restrict__ x_ptr,
        uint8_t* __restrict__ x_fp4_ptr,
        uint8_t* __restrict__ bs_shuffled_ptr,
        int64_t stride_x_m,
        int64_t stride_x_n,
        int64_t stride_x_fp4_m,
        int64_t stride_x_fp4_n,
        int M,
        int K,
        int SCALE_N_VALID,
        int SCALE_N_PAD) {
      __shared__ float s_vals[GROUPS_PER_BLOCK][32];
      __shared__ float s_abs[GROUPS_PER_BLOCK][32];
      __shared__ uint8_t s_codes[GROUPS_PER_BLOCK][32];
      __shared__ uint8_t s_scale[GROUPS_PER_BLOCK];

      const int tid = threadIdx.x;
      const int grp = tid >> 5;
      const int lane = tid & 31;
      const int row = blockIdx.x * GROUPS_PER_BLOCK + grp;
      const int scale_n = blockIdx.y;
      const int k_idx = scale_n * 32 + lane;

      float v = 0.0f;
      if (grp < GROUPS_PER_BLOCK && row < M && scale_n < SCALE_N_VALID && k_idx < K) {
        v = static_cast<float>(x_ptr[row * stride_x_m + k_idx * stride_x_n]);
      }
      s_vals[grp][lane] = v;
      s_abs[grp][lane] = fabsf(v);
      __syncthreads();

      for (int offset = 16; offset > 0; offset >>= 1) {
        if (lane < offset) {
          float other = s_abs[grp][lane + offset];
          if (other > s_abs[grp][lane]) {
            s_abs[grp][lane] = other;
          }
        }
        __syncthreads();
      }

      if (lane == 0) {
        uint8_t scale_e8m0 = encode_scale_from_amax(s_abs[grp][0]);
        s_scale[grp] = scale_e8m0;
        if (row < M && scale_n < SCALE_N_VALID) {
          int lin = shuffled_scale_index(row, scale_n, SCALE_N_PAD);
          bs_shuffled_ptr[lin] = scale_e8m0;
        }
      }
      __syncthreads();

      if (row < M && scale_n < SCALE_N_VALID) {
        s_codes[grp][lane] = quant_one_fp4_e2m1(s_vals[grp][lane], s_scale[grp]) & 0x0Fu;
      } else {
        s_codes[grp][lane] = 0;
      }
      __syncthreads();

      if ((lane & 1) == 0 && row < M && scale_n < SCALE_N_VALID) {
        uint8_t packed = static_cast<uint8_t>(
            s_codes[grp][lane] | (s_codes[grp][lane + 1] << 4));
        int out_n = scale_n * 16 + (lane >> 1);
        x_fp4_ptr[row * stride_x_fp4_m + out_n * stride_x_fp4_n] = packed;
      }
    }

    void quant_into_shuffled_hip(
        torch::Tensor A,
        torch::Tensor A_q_u8,
        torch::Tensor A_scale_sh_u8) {
      CHECK_CUDA(A);
      CHECK_CUDA(A_q_u8);
      CHECK_CUDA(A_scale_sh_u8);
      CHECK_CONTIGUOUS(A);
      CHECK_CONTIGUOUS(A_q_u8);
      CHECK_CONTIGUOUS(A_scale_sh_u8);
      CHECK_BF16(A);
      CHECK_U8(A_q_u8);
      CHECK_U8(A_scale_sh_u8);
      TORCH_CHECK(A.dim() == 2, "A must be 2D");
      TORCH_CHECK(A_q_u8.dim() == 2, "A_q_u8 must be 2D");
      TORCH_CHECK(A_scale_sh_u8.dim() == 2, "A_scale_sh_u8 must be 2D");

      const int M = static_cast<int>(A.size(0));
      const int K = static_cast<int>(A.size(1));
      const int SCALE_N_VALID = (K + 31) / 32;
      const int SCALE_N_PAD = static_cast<int>(A_scale_sh_u8.size(1));

      TORCH_CHECK(M <= 32, "HIP small-m quant path only supports m <= 32");
      TORCH_CHECK(A_q_u8.size(0) == M, "A_q_u8 row mismatch");
      TORCH_CHECK(A_q_u8.size(1) * 2 >= K, "A_q_u8 packed width mismatch");

      constexpr int GROUPS_PER_BLOCK = 4;
      dim3 block(GROUPS_PER_BLOCK * 32);
      dim3 grid((M + GROUPS_PER_BLOCK - 1) / GROUPS_PER_BLOCK, SCALE_N_VALID);

      hipLaunchKernelGGL(
          (quant_smallm_native_fp4_kernel<GROUPS_PER_BLOCK>),
          grid,
          block,
          0,
          0,
          reinterpret_cast<const __hip_bfloat16*>(A.data_ptr<at::BFloat16>()),
          reinterpret_cast<uint8_t*>(A_q_u8.data_ptr<uint8_t>()),
          reinterpret_cast<uint8_t*>(A_scale_sh_u8.data_ptr<uint8_t>()),
          static_cast<int64_t>(A.stride(0)),
          static_cast<int64_t>(A.stride(1)),
          static_cast<int64_t>(A_q_u8.stride(0)),
          static_cast<int64_t>(A_q_u8.stride(1)),
          M,
          K,
          SCALE_N_VALID,
          SCALE_N_PAD);

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

    PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
      m.def("quant_into_shuffled_hip", &quant_into_shuffled_hip, "HIP small-m native fp4 quant");
    }
    """
)


@functools.lru_cache(maxsize=1)
def _load_hip_smallm_mod():
    from torch.utils.cpp_extension import load_inline

    return load_inline(
        name="mxfp4_quant_smallm_nativefp4_gfx950",
        cpp_sources="",
        cuda_sources=_HIP_SRC,
        functions=None,
        with_cuda=True,
        verbose=False,
        extra_cuda_cflags=["-O3", "-std=c++20", "--offload-arch=gfx950"],
        no_implicit_headers=False,
    )


def _ensure_csv():
    global _CSV_PATH

    if _CSV_PATH is None:
        f = tempfile.NamedTemporaryFile(
            mode="w",
            suffix=".csv",
            delete=False,
            prefix="aiter_a4w4_fused_scale_codex_",
        )
        f.write(_CUSTOM_CSV)
        f.close()
        _CSV_PATH = f.name
        os.environ["AITER_CONFIG_GEMM_A4W4"] = _CSV_PATH


@triton.jit
def _mxfp4_quant_op_local(
    x,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_M: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
    EXP_BIAS_FP32: tl.constexpr = 127
    EXP_BIAS_FP4: tl.constexpr = 1
    EBITS_F32: tl.constexpr = 8
    EBITS_FP4: tl.constexpr = 2
    MBITS_F32: tl.constexpr = 23
    MBITS_FP4: tl.constexpr = 1

    max_normal: tl.constexpr = 6
    min_normal: tl.constexpr = 1

    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)

    amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
    amax = amax.to(tl.int32, bitcast=True)
    amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    amax = amax.to(tl.float32, bitcast=True)
    scale_e8m0_unbiased = tl.log2(amax).floor() - 2
    scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)

    bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
    quant_scale = tl.exp2(-scale_e8m0_unbiased)
    qx = x * quant_scale
    qx = qx.to(tl.uint32, bitcast=True)

    s = qx & 0x80000000
    qx = qx ^ s

    qx_fp32 = qx.to(tl.float32, bitcast=True)
    saturate_mask = qx_fp32 >= max_normal
    denormal_mask = (~saturate_mask) & (qx_fp32 < min_normal)
    normal_mask = ~(saturate_mask | denormal_mask)

    denorm_exp: tl.constexpr = (
        (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
    )
    denorm_mask_int: tl.constexpr = denorm_exp << MBITS_F32
    denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)

    denormal_x = qx_fp32 + denorm_mask_float
    denormal_x = denormal_x.to(tl.uint32, bitcast=True)
    denormal_x -= denorm_mask_int
    denormal_x = denormal_x.to(tl.uint8)

    normal_x = qx
    mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1
    val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1
    normal_x += val_to_add
    normal_x += mant_odd
    normal_x = normal_x >> (MBITS_F32 - MBITS_FP4)
    normal_x = normal_x.to(tl.uint8)

    e2m1_value = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
    e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
    e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)

    sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
    sign_lp = sign_lp.to(tl.uint8)
    e2m1_value = e2m1_value | sign_lp

    e2m1_value = tl.reshape(
        e2m1_value, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2]
    )
    evens, odds = tl.split(e2m1_value)
    x_fp4 = evens | (odds << 4)
    x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)

    return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)


@triton.heuristics(
    {
        "EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
        and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
    }
)
@triton.jit
def _dynamic_mxfp4_quant_kernel_shuffled(
    x_ptr,
    x_fp4_ptr,
    bs_shuffled_ptr,
    stride_x_m_in,
    stride_x_n_in,
    stride_x_fp4_m_in,
    stride_x_fp4_n_in,
    M,
    N,
    SCALE_N_VALID,
    SCALE_N_PAD,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    NUM_ITER: tl.constexpr,
    NUM_STAGES: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
    EVEN_M_N: tl.constexpr,
):
    pid_m = tl.program_id(0)
    start_n = tl.program_id(1) * NUM_ITER

    stride_x_m = tl.cast(stride_x_m_in, tl.int64)
    stride_x_n = tl.cast(stride_x_n_in, tl.int64)
    stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
    stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)

    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE

    for pid_n in tl.range(start_n, start_n + NUM_ITER, num_stages=NUM_STAGES):
        x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n

        if EVEN_M_N:
            x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
        else:
            x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
            x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(tl.float32)

        out_tensor, bs_e8m0 = _mxfp4_quant_op_local(
            x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
        )

        out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        out_offs_n = pid_n * (BLOCK_SIZE_N // 2) + tl.arange(0, BLOCK_SIZE_N // 2)
        out_offs = out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n

        if EVEN_M_N:
            tl.store(x_fp4_ptr + out_offs, out_tensor)
        else:
            out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
            tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)

        bs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        bs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)

        a = bs_m[:, None] // 32
        m1 = bs_m[:, None] % 32
        b = m1 // 16
        c = m1 % 16

        d = bs_n[None, :] // 8
        n1 = bs_n[None, :] % 8
        e = n1 // 4
        f = n1 % 4

        gn = SCALE_N_PAD // 8
        lin = (((((a * gn + d) * 4 + f) * 16 + c) * 2 + e) * 2 + b)

        bs_mask = (bs_m[:, None] < M) & (bs_n[None, :] < SCALE_N_VALID)
        tl.store(bs_shuffled_ptr + tl.cast(lin, tl.int64), bs_e8m0, mask=bs_mask)


def _get_ws(device, m, n, k):
    import torch
    from aiter import dtypes

    key = (device.index if device.index is not None else 0, m, n, k)
    ws = _WS.get(key)
    if ws is None:
        scale_n = (k + 31) // 32
        scale_n_pad = ((scale_n + 7) // 8) * 8
        scale_m_pad = ((m + 255) // 256) * 256

        ascale_cm_buf_u8 = torch.empty((scale_n, m), dtype=torch.uint8, device=device)
        ascale_sh_u8 = torch.empty(
            (scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device
        )
        ascale_sh_u8.fill_(127)

        ws = {
            "aq_u8": torch.empty((m, k // 2), dtype=torch.uint8, device=device),
            "ascale_cm_buf_u8": ascale_cm_buf_u8,
            "ascale_cm_u8": ascale_cm_buf_u8.T,
            "ascale_sh_u8": ascale_sh_u8,
            "out": torch.empty(
                (((m + 31) // 32) * 32, n), dtype=dtypes.bf16, device=device
            ),
        }
        _WS[key] = ws
    return ws


ACTIVE_SWEEP = "s7"

_QUANT_SWEEPS = {
    "s0": {
        "small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 256, "NUM_WARPS": 4, "NUM_STAGES": 1},
        "mid":   {"NUM_ITER": 4, "BLOCK_SIZE_M": 32,           "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
        "large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32,           "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
    },
    "s1": {
        "small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 256, "NUM_WARPS": 2, "NUM_STAGES": 1},
        "mid":   {"NUM_ITER": 4, "BLOCK_SIZE_M": 32,           "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
        "large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32,           "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
    },
    "s2": {
        "small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 1},
        "mid":   {"NUM_ITER": 4, "BLOCK_SIZE_M": 32,           "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
        "large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32,           "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
    },
    "s3": {
        "small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 256, "NUM_WARPS": 4, "NUM_STAGES": 1},
        "mid":   {"NUM_ITER": 2, "BLOCK_SIZE_M": 32,           "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
        "large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32,           "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
    },
    "s4": {
        "small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 256, "NUM_WARPS": 4, "NUM_STAGES": 1},
        "mid":   {"NUM_ITER": 4, "BLOCK_SIZE_M": 64,           "BLOCK_SIZE_N": 64,  "NUM_WARPS": 4, "NUM_STAGES": 2},
        "large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32,           "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
    },
    "s5": {
        "small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 256, "NUM_WARPS": 4, "NUM_STAGES": 1},
        "mid":   {"NUM_ITER": 4, "BLOCK_SIZE_M": 32,           "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
        "large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32,           "BLOCK_SIZE_N": 64,  "NUM_WARPS": 4, "NUM_STAGES": 2},
    },
    "s6": {
        "small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 256, "NUM_WARPS": 4, "NUM_STAGES": 1},
        "mid":   {"NUM_ITER": 4, "BLOCK_SIZE_M": 32,           "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
        "large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 16,           "BLOCK_SIZE_N": 64,  "NUM_WARPS": 4, "NUM_STAGES": 2},
    },
    "s7": {
        "small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 256, "NUM_WARPS": 4, "NUM_STAGES": 1},
        "mid":   {"NUM_ITER": 4, "BLOCK_SIZE_M": 32,           "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
        "large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32,           "BLOCK_SIZE_N": 128, "NUM_WARPS": 8, "NUM_STAGES": 2},
    },
}

ACTIVE_LARGE = "L1"

_LARGE_K_CFGS = {
    "L0": {
        "NUM_ITER": 4,
        "BLOCK_SIZE_M": 16,
        "BLOCK_SIZE_N": 64,
        "NUM_WARPS": 4,
        "NUM_STAGES": 2,
    },
    "L1": {
        "NUM_ITER": 2,
        "BLOCK_SIZE_M": 16,
        "BLOCK_SIZE_N": 128,
        "NUM_WARPS": 4,
        "NUM_STAGES": 2,
    },
    "L2": {
        "NUM_ITER": 4,
        "BLOCK_SIZE_M": 32,
        "BLOCK_SIZE_N": 64,
        "NUM_WARPS": 4,
        "NUM_STAGES": 2,
    },
}


def _resolve_block_m(m, block_m):
    import triton

    if block_m == "pow2":
        return triton.next_power_of_2(m)
    if block_m == "pow2_cap32":
        return min(32, triton.next_power_of_2(m))
    return block_m


def _pick_quant_cfg(M, K):
    if K <= 1024:
        src = _QUANT_SWEEPS["s2"]["small"]
    elif K <= 4096:
        src = _QUANT_SWEEPS["s3"]["mid"]
    else:
        src = _LARGE_K_CFGS[ACTIVE_LARGE]

    cfg = dict(src)
    cfg["BLOCK_SIZE_M"] = _resolve_block_m(M, cfg["BLOCK_SIZE_M"])
    return cfg


def _quant_into_raw(A, aq_u8, ascale_u8):
    from aiter.ops.triton._triton_kernels.quant.quant import _dynamic_mxfp4_quant_kernel

    M, K = A.shape
    cfg = _pick_quant_cfg(M, K)

    grid = (
        triton.cdiv(M, cfg["BLOCK_SIZE_M"]),
        triton.cdiv(K, cfg["BLOCK_SIZE_N"] * cfg["NUM_ITER"]),
    )

    _dynamic_mxfp4_quant_kernel[grid](
        A,
        aq_u8,
        ascale_u8,
        *A.stride(),
        *aq_u8.stride(),
        *ascale_u8.stride(),
        M=M,
        N=K,
        MXFP4_QUANT_BLOCK_SIZE=32,
        SCALING_MODE=0,
        NUM_ITER=cfg["NUM_ITER"],
        BLOCK_SIZE_M=cfg["BLOCK_SIZE_M"],
        BLOCK_SIZE_N=cfg["BLOCK_SIZE_N"],
        NUM_STAGES=cfg["NUM_STAGES"],
        num_warps=cfg["NUM_WARPS"],
        waves_per_eu=0,
        num_stages=1,
    )

    return aq_u8, ascale_u8


def _quant_into_shuffled_triton(A, aq_u8, ascale_sh_u8):
    M, K = A.shape
    cfg = _pick_quant_cfg(M, K)
    scale_n_valid = (K + 31) // 32
    scale_n_pad = ascale_sh_u8.shape[1]

    grid = (
        triton.cdiv(M, cfg["BLOCK_SIZE_M"]),
        triton.cdiv(K, cfg["BLOCK_SIZE_N"] * cfg["NUM_ITER"]),
    )

    _dynamic_mxfp4_quant_kernel_shuffled[grid](
        A,
        aq_u8,
        ascale_sh_u8,
        *A.stride(),
        *aq_u8.stride(),
        M=M,
        N=K,
        SCALE_N_VALID=scale_n_valid,
        SCALE_N_PAD=scale_n_pad,
        MXFP4_QUANT_BLOCK_SIZE=32,
        NUM_ITER=cfg["NUM_ITER"],
        BLOCK_SIZE_M=cfg["BLOCK_SIZE_M"],
        BLOCK_SIZE_N=cfg["BLOCK_SIZE_N"],
        NUM_STAGES=cfg["NUM_STAGES"],
        num_warps=cfg["NUM_WARPS"],
        waves_per_eu=0,
        num_stages=1,
    )

    return aq_u8, ascale_sh_u8


def _quant_into_shuffled_smallm_hip(A, aq_u8, ascale_sh_u8):
    mod = _load_hip_smallm_mod()
    mod.quant_into_shuffled_hip(A, aq_u8, ascale_sh_u8)
    return aq_u8, ascale_sh_u8


def _quant_into_shuffled(A, aq_u8, ascale_sh_u8):
    if A.shape[0] <= 32:
        return _quant_into_shuffled_smallm_hip(A, aq_u8, ascale_sh_u8)
    return _quant_into_shuffled_triton(A, aq_u8, ascale_sh_u8)


def custom_kernel(data: input_t) -> output_t:
    _ensure_csv()

    import aiter
    from aiter import dtypes

    A, _, _, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    B_shuffle = B_shuffle.contiguous()
    B_scale_sh = B_scale_sh.contiguous()
    m, k = A.shape
    n = B_shuffle.shape[0]

    ws = _get_ws(A.device, m, n, k)

    A_q_u8, A_scale_sh_u8 = _quant_into_shuffled(A, ws["aq_u8"], ws["ascale_sh_u8"])
    A_scale = A_scale_sh_u8.view(dtypes.fp8_e8m0)
    A_q = A_q_u8.view(dtypes.fp4x2)

    return aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
scrolls · 672 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