Skip to content
KernelIndex
Search⌘K

submission 653415

Nicky Pochinkov · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission-v1774680234.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-653415?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
11.0µs
#308 of 1143
2026-03-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:22e1bc72e4d5a23aead6eadd59f7337560654fb8c0c7a85b051c032374816b20
license declaredunknown
license concludedunknown
authorsNicky Pochinkov
imported2026-08-26

Techniques

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

shared-memoryextern __shared__ char lds[];
split-kint snb, int split_k_count)
stages = 1BLOCK_SIZE_M=q['BSM'], BLOCK_SIZE_N=q['BSN'], NUM_STAGES=1, num_warps=q['NW'], waves_per_eu=0, num_stages=1)
tile-n = 32M<=16: 128 threads, BN=32 (2x CTAs for better CU utilization)

Kernel source

submission-v1774680234.py728 lines
"""Attempt 613: Adaptive BN for fused_qgemm_k.

Based on 611 (XCD swizzle). Key change:
- fused_qgemm_k adapts to M dimension at launch time:
  M<=16: 128 threads, BN=32 (2x CTAs for better CU utilization)
  M>16:  256 threads, BN=64 (4 warps for better per-CU utilization)
  Uses blockDim.x to determine BN at runtime.
- Case 1 (m=4): 45->90 CTAs, ~0.3us improvement
"""
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"

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

_fp4x2 = dtypes.fp4x2
_e8m0 = dtypes.fp8_e8m0
_bf16 = torch.bfloat16

# ============================================================
# HIP MFMA kernels
# ============================================================

_HIP_SRC = r"""
#include <hip/hip_runtime.h>

typedef int v8i __attribute__((ext_vector_type(8)));
typedef float v4f __attribute__((ext_vector_type(4)));
typedef unsigned int u4v __attribute__((ext_vector_type(4)));

__device__ __forceinline__ int eidx(int r, int c, int sn) {
    int A=r>>5, B=(r&31)>>4, C=r&15, D=c>>3, E=(c&7)>>2, F=c&3;
    return A*32*sn + D*256 + F*64 + C*4 + E*2 + B;
}

__device__ __forceinline__ unsigned char qfp4(float x) {
    unsigned int u;
    __builtin_memcpy(&u, &x, 4);
    unsigned char s = (unsigned char)((u >> 28) & 8u);
    u &= 0x7FFFFFFFu;
    float a;
    __builtin_memcpy(&a, &u, 4);
    unsigned char r;
    if (a >= 6.f) {
        r = 7u;
    } else if (a < 1.f) {
        const unsigned int dm = 149u << 23;
        float dmf;
        __builtin_memcpy(&dmf, &dm, 4);
        a += dmf;
        unsigned int t;
        __builtin_memcpy(&t, &a, 4);
        r = (unsigned char)((t - dm) & 7u);
    } else {
        unsigned int mo = (u >> 22) & 1u;
        u += 0xC11FFFFFu + mo;
        r = (unsigned char)((u >> 22) & 7u);
    }
    return r | s;
}

// ---- Fused quant+GEMM: adaptive threads/BN ----
// M<=16: 128 threads (2 warps), BN=32, more CTAs for CU utilization
// M>16:  256 threads (4 warps), BN=64, more warps for latency hiding
__global__ __launch_bounds__(256, 2)
void fused_qgemm_k(
    const unsigned short* __restrict__ A_bf16,
    const unsigned char* __restrict__ bq,
    const unsigned char* __restrict__ bsc,
    unsigned short* __restrict__ out,
    int M, int N, int K, int snb)
{
    const int tid = threadIdx.x;
    const int nwarps = (int)blockDim.x >> 6;
    const int warp_id = tid >> 6;
    const int tid_w = tid & 63;
    const int grp = tid_w >> 4;
    const int ln = tid_w & 15;
    const int rb = blockIdx.x * 16;
    const int bn = nwarps << 4;
    const int cb = blockIdx.y * bn;
    const int Kh = K >> 1;
    const int Ksc = K >> 5;
    const int Kh_pad = Kh + 4;

    extern __shared__ char lds[];
    unsigned char* aq_lds = (unsigned char*)lds;
    unsigned char* asc_lds = aq_lds + 16 * Kh_pad;

    // Phase 1: Cooperative quantization of A (vectorized loads)
    int bm_eff = M - rb;
    if (bm_eff > 16) bm_eff = 16;
    int total_groups = bm_eff > 0 ? bm_eff * Ksc : 0;
    for (int gi = tid; gi < total_groups; gi += (int)blockDim.x) {
        int row = gi / Ksc;
        int gk = gi - row * Ksc;
        int ar = rb + row;
        const u4v* vp = (const u4v*)(A_bf16 + (long long)ar * K + gk * 32);
        u4v t0 = vp[0], t1 = vp[1], t2 = vp[2], t3 = vp[3];
        unsigned int v[16];
        v[0]=t0[0]; v[1]=t0[1]; v[2]=t0[2]; v[3]=t0[3];
        v[4]=t1[0]; v[5]=t1[1]; v[6]=t1[2]; v[7]=t1[3];
        v[8]=t2[0]; v[9]=t2[1]; v[10]=t2[2]; v[11]=t2[3];
        v[12]=t3[0]; v[13]=t3[1]; v[14]=t3[2]; v[15]=t3[3];
        unsigned int mx = 0;
        #pragma unroll
        for (int i = 0; i < 16; i++) {
            unsigned int lo = (v[i] & 0x7FFF) << 16;
            unsigned int hi = v[i] & 0x7FFF0000;
            mx = mx > lo ? mx : lo;
            mx = mx > hi ? mx : hi;
        }
        float amax; __builtin_memcpy(&amax, &mx, 4);
        unsigned int ai; __builtin_memcpy(&ai, &amax, 4);
        ai = (ai + 0x200000u) & 0xFF800000u;
        int exp_raw = (int)((ai >> 23) & 0xFFu) - 127;
        int sub = exp_raw - 2;
        sub = sub < -127 ? -127 : (sub > 127 ? 127 : sub);
        asc_lds[row * Ksc + gk] = (unsigned char)(sub + 127);
        unsigned int qs_bits = (unsigned int)(127 - sub) << 23;
        float qs; __builtin_memcpy(&qs, &qs_bits, 4);
        unsigned char* dst = aq_lds + row * Kh_pad + gk * 16;
        #pragma unroll
        for (int i = 0; i < 16; i++) {
            unsigned int f0i = (v[i] & 0xFFFF) << 16;
            unsigned int f1i = v[i] & 0xFFFF0000;
            float f0, f1; __builtin_memcpy(&f0, &f0i, 4); __builtin_memcpy(&f1, &f1i, 4);
            dst[i] = qfp4(f0 * qs) | (qfp4(f1 * qs) << 4);
        }
    }
    __syncthreads();

    // Phase 2: Simple MFMA GEMM loop (no SW pipelining)
    const int mc = cb + warp_id * 16 + ln;
    const bool rv = (rb + ln) < M;
    const bool cv = mc < N;
    v4f acc = {0.f, 0.f, 0.f, 0.f};

    // Precompute eidx row base for B scales
    int eA = mc >> 5, eB_row = (mc & 31) >> 4, eC = mc & 15;
    int ebase = eA * 32 * snb + eC * 4 + eB_row;

    #pragma unroll
    for (int ks = 0; ks < Ksc; ks += 4) {
        int gk = ks + grp;
        v8i ad = {}, bd = {};
        int sa = 127, sb = 127;
        if (rv && gk < Ksc) {
            const int* p = (const int*)(aq_lds + ln * Kh_pad + (gk << 4));
            ad[0]=p[0]; ad[1]=p[1]; ad[2]=p[2]; ad[3]=p[3];
            sa = (int)asc_lds[ln * Ksc + gk];
        }
        if (cv && gk < Ksc) {
            const int* p = (const int*)(bq + (long long)mc * Kh + (gk << 4));
            bd[0]=p[0]; bd[1]=p[1]; bd[2]=p[2]; bd[3]=p[3];
            int eD = gk >> 3, eE = (gk & 7) >> 2, eF = gk & 3;
            sb = (int)bsc[ebase + eD * 256 + eF * 64 + eE * 2];
        }
        acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(ad, bd, acc, 4, 4, 0, sa, 0, sb);
    }
    for (int j = 0; j < 4; j++) {
        int or_ = rb + grp * 4 + j, oc = cb + warp_id * 16 + ln;
        if (or_ < M && oc < N) {
            float fv = acc[j]; unsigned int fb;
            __builtin_memcpy(&fb, &fv, 4);
            out[(long long)or_ * N + oc] = (unsigned short)((fb + 0x7FFFu + ((fb >> 16) & 1u)) >> 16);
        }
    }
}

// ---- Fused quant+wide4: BM=16, BN=128, 256 threads + split-K ----
// Lower occupancy (1 vs 2) gives compiler more VGPRs for SW pipeline.
__global__ __launch_bounds__(256, 1)
void fused_wide4_k(
    const unsigned short* __restrict__ A_bf16,
    const unsigned char* __restrict__ bq,
    const unsigned char* __restrict__ bsc,
    unsigned short* __restrict__ out_bf16,
    float* __restrict__ out_f32,
    int M, int N, int K,
    int snb, int split_k_count)
{
    const int tid = threadIdx.x;
    const int warp_id = tid >> 6;
    const int tid_w = tid & 63;
    const int grp = tid_w >> 4;
    const int ln = tid_w & 15;

    const int rb = blockIdx.x * 16;
    // XCD-aware: grid=(M_tiles, split_K, N_tiles)
    // blockIdx.y = split-K partition, blockIdx.z = N-tile
    // Groups all partitions of same N-tile on consecutive CUs (same XCD)
    const int cb = blockIdx.z * 128;
    const int Kh = K >> 1;
    const int Ksc = K >> 5;

    int tiles_total = Ksc >> 2;
    int tiles_per = tiles_total / split_k_count;
    int tile_s = blockIdx.y * tiles_per;
    int tile_e = tile_s + tiles_per;
    if ((int)blockIdx.y == split_k_count - 1) tile_e = tiles_total;
    int ks_s = tile_s << 2;
    int ks_e = tile_e << 2;
    int ksc_len = ks_e - ks_s;
    if (ksc_len <= 0) return;

    int kh_pad = (ksc_len << 4) + 4;

    extern __shared__ char lds[];
    unsigned char* aq_lds = (unsigned char*)lds;
    unsigned char* asc_lds = aq_lds + 16 * kh_pad;

    int bm_eff = M - rb;
    if (bm_eff > 16) bm_eff = 16;
    int total_groups = bm_eff > 0 ? bm_eff * ksc_len : 0;
    for (int gi = tid; gi < total_groups; gi += 256) {
        int row = gi / ksc_len;
        int gk = gi - row * ksc_len;
        int ar = rb + row;
        int abs_k = (ks_s + gk) * 32;
        const u4v* vp = (const u4v*)(A_bf16 + (long long)ar * K + abs_k);
        u4v t0 = vp[0], t1 = vp[1], t2 = vp[2], t3 = vp[3];
        unsigned int v[16];
        v[0]=t0[0]; v[1]=t0[1]; v[2]=t0[2]; v[3]=t0[3];
        v[4]=t1[0]; v[5]=t1[1]; v[6]=t1[2]; v[7]=t1[3];
        v[8]=t2[0]; v[9]=t2[1]; v[10]=t2[2]; v[11]=t2[3];
        v[12]=t3[0]; v[13]=t3[1]; v[14]=t3[2]; v[15]=t3[3];
        unsigned int mx = 0;
        #pragma unroll
        for (int i = 0; i < 16; i++) {
            unsigned int lo = (v[i] & 0x7FFF) << 16;
            unsigned int hi = v[i] & 0x7FFF0000;
            mx = mx > lo ? mx : lo;
            mx = mx > hi ? mx : hi;
        }
        float amax; __builtin_memcpy(&amax, &mx, 4);
        unsigned int ai; __builtin_memcpy(&ai, &amax, 4);
        ai = (ai + 0x200000u) & 0xFF800000u;
        int exp_raw = (int)((ai >> 23) & 0xFFu) - 127;
        int sub = exp_raw - 2;
        sub = sub < -127 ? -127 : (sub > 127 ? 127 : sub);
        asc_lds[row * ksc_len + gk] = (unsigned char)(sub + 127);
        unsigned int qs_bits = (unsigned int)(127 - sub) << 23;
        float qs; __builtin_memcpy(&qs, &qs_bits, 4);
        unsigned char* dst = aq_lds + row * kh_pad + gk * 16;
        #pragma unroll
        for (int i = 0; i < 16; i++) {
            unsigned int f0i = (v[i] & 0xFFFF) << 16;
            unsigned int f1i = v[i] & 0xFFFF0000;
            float f0, f1; __builtin_memcpy(&f0, &f0i, 4); __builtin_memcpy(&f1, &f1i, 4);
            dst[i] = qfp4(f0 * qs) | (qfp4(f1 * qs) << 4);
        }
    }
    __syncthreads();

    const int mr = rb + ln;
    const bool rv = mr < M;

    int ni0 = warp_id * 2;
    int ni1 = warp_id * 2 + 1;
    int mc0 = cb + ni0 * 16 + ln;
    int mc1 = cb + ni1 * 16 + ln;
    bool cv0 = mc0 < N;
    bool cv1 = mc1 < N;
    long long bro0 = cv0 ? (long long)mc0 * Kh : 0;
    long long bro1 = cv1 ? (long long)mc1 * Kh : 0;

    int eA_b0 = (mc0 >> 5), eB_b0 = ((mc0 & 31) >> 4), eC_b0 = (mc0 & 15);
    int ebase_b0 = eA_b0 * 32 * snb + eC_b0 * 4 + eB_b0;
    int eA_b1 = (mc1 >> 5), eB_b1 = ((mc1 & 31) >> 4), eC_b1 = (mc1 & 15);
    int ebase_b1 = eA_b1 * 32 * snb + eC_b1 * 4 + eB_b1;

    v4f acc0 = {0.f,0.f,0.f,0.f};
    v4f acc1 = {0.f,0.f,0.f,0.f};

    v8i ad_c = {}, bd0_c = {}, bd1_c = {};
    int sa_c = 127, sb0_c = 127, sb1_c = 127;

    {
        int gk_local = grp;
        int abs_gk = ks_s + gk_local;
        int kb_local = gk_local << 4;
        int kb_global = abs_gk << 4;
        int eD = abs_gk >> 3, eE = (abs_gk & 7) >> 2, eF = abs_gk & 3;
        int eoff = eD * 256 + eF * 64 + eE * 2;
        if (rv && gk_local < ksc_len) {
            const int* p = (const int*)(aq_lds + ln * kh_pad + kb_local);
            ad_c[0]=p[0]; ad_c[1]=p[1]; ad_c[2]=p[2]; ad_c[3]=p[3];
            sa_c = (int)asc_lds[ln * ksc_len + gk_local];
        }
        if (cv0 && abs_gk < Ksc) {
            const int* p = (const int*)(bq + bro0 + kb_global);
            bd0_c[0]=p[0]; bd0_c[1]=p[1]; bd0_c[2]=p[2]; bd0_c[3]=p[3];
            sb0_c = (int)bsc[ebase_b0 + eoff];
        }
        if (cv1 && abs_gk < Ksc) {
            const int* p = (const int*)(bq + bro1 + kb_global);
            bd1_c[0]=p[0]; bd1_c[1]=p[1]; bd1_c[2]=p[2]; bd1_c[3]=p[3];
            sb1_c = (int)bsc[ebase_b1 + eoff];
        }
    }

    for (int ks = 0; ks < ksc_len; ks += 4) {
        v8i ad_n = {}, bd0_n = {}, bd1_n = {};
        int sa_n = 127, sb0_n = 127, sb1_n = 127;
        int next_ks = ks + 4;
        if (next_ks < ksc_len) {
            int gk_local = next_ks + grp;
            int abs_gk = ks_s + gk_local;
            int kb_local = gk_local << 4;
            int kb_global = abs_gk << 4;
            int eD = abs_gk >> 3, eE = (abs_gk & 7) >> 2, eF = abs_gk & 3;
            int eoff = eD * 256 + eF * 64 + eE * 2;
            if (rv && gk_local < ksc_len) {
                const int* p = (const int*)(aq_lds + ln * kh_pad + kb_local);
                ad_n[0]=p[0]; ad_n[1]=p[1]; ad_n[2]=p[2]; ad_n[3]=p[3];
                sa_n = (int)asc_lds[ln * ksc_len + gk_local];
            }
            if (cv0 && abs_gk < Ksc) {
                const int* p = (const int*)(bq + bro0 + kb_global);
                bd0_n[0]=p[0]; bd0_n[1]=p[1]; bd0_n[2]=p[2]; bd0_n[3]=p[3];
                sb0_n = (int)bsc[ebase_b0 + eoff];
            }
            if (cv1 && abs_gk < Ksc) {
                const int* p = (const int*)(bq + bro1 + kb_global);
                bd1_n[0]=p[0]; bd1_n[1]=p[1]; bd1_n[2]=p[2]; bd1_n[3]=p[3];
                sb1_n = (int)bsc[ebase_b1 + eoff];
            }
        }

        acc0 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(ad_c, bd0_c, acc0, 4, 4, 0, sa_c, 0, sb0_c);
        acc1 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(ad_c, bd1_c, acc1, 4, 4, 0, sa_c, 0, sb1_c);

        ad_c = ad_n; bd0_c = bd0_n; bd1_c = bd1_n;
        sa_c = sa_n; sb0_c = sb0_n; sb1_c = sb1_n;
    }

    if (split_k_count == 1) {
        for (int j = 0; j < 4; j++) {
            int or_ = rb + grp*4+j;
            if (or_ < M) {
                int oc0 = cb + ni0*16+ln;
                if (oc0 < N) {
                    float fv = acc0[j]; unsigned int fb;
                    __builtin_memcpy(&fb, &fv, 4);
                    out_bf16[(long long)or_*N+oc0] = (unsigned short)((fb+0x7FFFu+((fb>>16)&1u))>>16);
                }
                int oc1 = cb + ni1*16+ln;
                if (oc1 < N) {
                    float fv = acc1[j]; unsigned int fb;
                    __builtin_memcpy(&fb, &fv, 4);
                    out_bf16[(long long)or_*N+oc1] = (unsigned short)((fb+0x7FFFu+((fb>>16)&1u))>>16);
                }
            }
        }
    } else {
        for (int j = 0; j < 4; j++) {
            int or_ = rb + grp*4+j;
            if (or_ < M) {
                int oc0 = cb + ni0*16+ln;
                if (oc0 < N) atomicAdd(&out_f32[(long long)or_*N+oc0], acc0[j]);
                int oc1 = cb + ni1*16+ln;
                if (oc1 < N) atomicAdd(&out_f32[(long long)or_*N+oc1], acc1[j]);
            }
        }
    }
}

// f32->bf16 convert + zero f32 buffer
__global__ void cvt_zero_k(float* __restrict__ buf,
                            unsigned short* __restrict__ out, int total) {
    int idx = blockIdx.x * 256 + threadIdx.x;
    if (idx < total) {
        float fv = buf[idx]; unsigned int fb;
        __builtin_memcpy(&fb, &fv, 4);
        out[idx] = (unsigned short)((fb + 0x7FFFu + ((fb>>16)&1u)) >> 16);
        buf[idx] = 0.f;
    }
}

extern "C" int run_fused(const void* a, const void* bq, const void* bsc,
                          void* out, int M, int N, int K, int snb, int shmem) {
    int block_sz = (M <= 16) ? 128 : 256;
    int bn = (M <= 16) ? 32 : 64;
    dim3 g(((unsigned)M+15)/16, ((unsigned)N+bn-1)/bn);
    hipLaunchKernelGGL(fused_qgemm_k, g, dim3(block_sz), shmem, 0,
        (const unsigned short*)a, (const unsigned char*)bq,
        (const unsigned char*)bsc, (unsigned short*)out,
        M, N, K, snb);
    return 0;
}

extern "C" int run_fused_wide4(
    const void* a_bf16, const void* bq, const void* bsc,
    void* out_bf16, void* out_f32,
    int M, int N, int K, int snb, int split_k, int shmem)
{
    dim3 g(((unsigned)M+15)/16, split_k, ((unsigned)N+127)/128);
    hipLaunchKernelGGL(fused_wide4_k, g, dim3(256), shmem, 0,
        (const unsigned short*)a_bf16, (const unsigned char*)bq,
        (const unsigned char*)bsc, (unsigned short*)out_bf16,
        (float*)out_f32, M, N, K, snb, split_k);
    return 0;
}

extern "C" int run_cvt_zero(void* buf, void* out, int total) {
    int blocks = (total + 255) / 256;
    hipLaunchKernelGGL(cvt_zero_k, dim3(blocks), dim3(256), 0, 0,
        (float*)buf, (unsigned short*)out, total);
    return 0;
}
"""

_CPP_SRC = """
#include <torch/extension.h>
extern "C" int run_fused(const void*, const void*, const void*, void*, int, int, int, int, int);
extern "C" int run_fused_wide4(const void*, const void*, const void*, void*, void*, int, int, int, int, int, int);
extern "C" int run_cvt_zero(void*, void*, int);

int64_t launch_fused(torch::Tensor a, torch::Tensor bq, torch::Tensor bsc,
                      torch::Tensor out,
                      int64_t M, int64_t N, int64_t K,
                      int64_t snb, int64_t shmem) {
    return (int64_t)run_fused(
        (const void*)a.data_ptr<at::BFloat16>(),
        (const void*)bq.data_ptr<uint8_t>(),
        (const void*)bsc.data_ptr<uint8_t>(),
        (void*)out.data_ptr<at::BFloat16>(),
        (int)M, (int)N, (int)K, (int)snb, (int)shmem);
}

int64_t launch_fused_wide4(torch::Tensor a, torch::Tensor bq, torch::Tensor bsc,
                            torch::Tensor out_bf16, torch::Tensor out_f32,
                            int64_t M, int64_t N, int64_t K,
                            int64_t snb, int64_t sk, int64_t shmem) {
    return (int64_t)run_fused_wide4(
        (const void*)a.data_ptr<at::BFloat16>(),
        (const void*)bq.data_ptr<uint8_t>(),
        (const void*)bsc.data_ptr<uint8_t>(),
        (void*)out_bf16.data_ptr<at::BFloat16>(),
        (void*)out_f32.data_ptr<float>(),
        (int)M, (int)N, (int)K, (int)snb, (int)sk, (int)shmem);
}

int64_t launch_cvt_zero(torch::Tensor buf, torch::Tensor out, int64_t total) {
    return (int64_t)run_cvt_zero(buf.data_ptr<float>(), out.data_ptr<at::BFloat16>(), (int)total);
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("launch_fused", &launch_fused);
    m.def("launch_fused_wide4", &launch_fused_wide4);
    m.def("launch_cvt_zero", &launch_cvt_zero);
}
"""

_mfma_mod = None
_mfma_ok = False

try:
    from torch.utils.cpp_extension import load_inline as _li
    _mfma_mod = _li(
        name='mfma_v613',
        cpp_sources=[_CPP_SRC],
        cuda_sources=[_HIP_SRC],
        verbose=False,
        extra_cuda_cflags=['-Ofast', '-ffast-math', '-munsafe-fp-atomics', '-std=c++17'],
    )
    _mfma_ok = True
    print(f"[613] MFMA compiled OK", file=sys.stderr)
except Exception as e:
    print(f"[610] compile fail: {e}", file=sys.stderr)

# ============================================================
# Triton fused quant+shuffle kernel (for m>32 path)
# ============================================================

@triton.jit
def _mxfp4_quant_op(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; MBITS_F32: tl.constexpr = 23; MBITS_FP4: tl.constexpr = 1
    EBITS_F32: tl.constexpr = 8; EBITS_FP4: tl.constexpr = 2; 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 = (not saturate_mask) & (qx_fp32 < min_normal); normal_mask = not (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 >>= (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 _fused_quant_shuffle_kernel(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, sn, M, N, 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, SCALING_MODE: 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, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
        x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M); x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
        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(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_row = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M); bs_col = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
        A_idx = bs_row // 32; B_idx = (bs_row & 31) >> 4; C_idx = bs_row & 15; D_idx = bs_col // 8; E_idx = (bs_col & 7) >> 2; F_idx = bs_col & 3
        shuffled_flat = (A_idx * (32 * sn))[:, None] + (D_idx * 256 + F_idx * 64)[None, :] + (C_idx * 4)[:, None] + (E_idx * 2)[None, :] + B_idx[:, None]
        bs_valid = (bs_row < M)[:, None] & (bs_col < ((N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE))[None, :]
        tl.store(bs_shuffled_ptr + shuffled_flat, bs_e8m0, mask=bs_valid)


# ============================================================
# Buffer management
# ============================================================

_shapes = [(4,2880,512),(16,2112,7168),(32,4096,512),(32,2880,512),(64,7168,2048),(256,3072,1536)]
_fused_bufs = {}
_quant_bufs = {}
_wide_bufs = {}
_bsc_cache = {}

_KN = f"_ZN5aiter{len('f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128')}f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_asm = None
try: _asm = torch.ops.aiter.gemm_a4w4_asm
except: pass
_gemm_bufs = {}


def _setup_fused(m, n, k):
    key = (m, n, k)
    if key in _fused_bufs: return _fused_bufs[key]
    ob = torch.empty((m, n), dtype=_bf16, device='cuda')
    ksc = k // 32
    snb = ((ksc + 7) // 8) * 8
    kh_pad = k // 2 + 4
    shared = 16 * kh_pad + 16 * ksc
    entry = {'out': ob, 'snb': snb, 'shared': shared}
    _fused_bufs[key] = entry
    return entry


def _setup_quant(m, k):
    key = (m, k)
    if key in _quant_bufs: return _quant_bufs[key]
    x_fp4 = torch.empty((m, k // 2), dtype=torch.uint8, device='cuda')
    sc = (k + 31) // 32
    sm = ((m + 255) // 256) * 256
    sn = ((sc + 7) // 8) * 8
    bs = torch.empty((sm * sn,), dtype=torch.uint8, device='cuda')
    BSM = triton.next_power_of_2(m) if m <= 32 else 32
    NW = 1 if m <= 32 else 4
    if k <= 1024:
        BSN = min(256, max(32, triton.next_power_of_2(k)))
        BSM = min(8, triton.next_power_of_2(m)) if m <= 32 else BSM
        NW = 4
    else:
        BSN = 128
        if m <= 16:
            BSM = 4
            NW = 4
    grid = (triton.cdiv(m, BSM), triton.cdiv(k, BSN))
    entry = {
        'grid': grid, 'x_fp4': x_fp4, 'bs': bs,
        'sn': sn, 'BSM': BSM, 'BSN': BSN, 'NW': NW,
        'fp4_s0': k // 2, 'fp4_s1': 1,
        'aq': x_fp4.view(_fp4x2), 'ash': bs.view(sm, sn).view(_e8m0),
    }
    _quant_bufs[key] = entry
    return entry


def _compute_split_k(m, n, k):
    base_blocks = ((m + 15) // 16) * ((n + 127) // 128)
    if base_blocks >= 200: return 1
    k_tiles = k // 128
    if k_tiles < 2: return 1
    target = 200
    best_sk = 1
    best_diff = abs(base_blocks - target)
    for d in range(2, k_tiles + 1):
        if k_tiles % d != 0: continue
        tiles_per = k_tiles // d
        if tiles_per < 2: break
        total = base_blocks * d
        if total > target * 2: break
        diff = abs(total - target)
        if diff < best_diff:
            best_diff = diff
            best_sk = d
    return best_sk


def _setup_wide(m, n, k):
    key = (m, n, k)
    if key in _wide_bufs: return _wide_bufs[key]
    sk = _compute_split_k(m, n, k)
    out_bf16 = torch.empty((m, n), dtype=_bf16, device='cuda')
    out_f32 = torch.zeros((m, n), dtype=torch.float32, device='cuda') if sk > 1 else torch.empty(1, dtype=torch.float32, device='cuda')
    ksc = k // 32
    snb = ((ksc + 7) // 8) * 8
    ksc_per_part = ksc // sk if sk > 0 else ksc
    kh_per = ksc_per_part * 16
    kh_pad = kh_per + 4
    shmem_fw = 16 * kh_pad + 16 * ksc_per_part
    entry = {'out_bf16': out_bf16, 'out_f32': out_f32, 'snb': snb, 'split_k': sk, 'shmem_fw': shmem_fw}
    _wide_bufs[key] = entry
    return entry


def _setup_gemm(m, n, k):
    key = (m, n, k)
    if key in _gemm_bufs: return _gemm_bufs[key]
    if _asm is None: return None
    mp = ((m + 31) // 32) * 32
    ob = torch.empty((mp, n), dtype=_bf16, device='cuda')
    entry = (ob[:m], ob, _KN)
    _gemm_bufs[key] = entry
    return entry


def _get_bsc(bsh, n, k):
    ksc = k // 32
    sn = ((ksc + 7) // 8) * 8
    sm = ((n + 255) // 256) * 256
    needed = sm * sn
    u8 = bsh.view(torch.uint8).contiguous().view(-1)
    actual = u8.numel()
    if actual >= needed: return u8[:needed]
    key = (n, k)
    if key in _bsc_cache:
        buf = _bsc_cache[key]
        buf[:actual] = u8
        return buf
    buf = torch.full((needed,), 127, dtype=torch.uint8, device='cuda')
    buf[:actual] = u8
    _bsc_cache[key] = buf
    return buf


# Pre-allocate all paths
for _m, _n, _k in _shapes:
    _setup_quant(_m, _k)
    _setup_gemm(_m, _n, _k)
    if _mfma_ok:
        if _m <= 32 and _k <= 1024:
            _setup_fused(_m, _n, _k)
        elif _m <= 32:
            _setup_wide(_m, _n, _k)

# Warm up Triton quant
for _m, _n, _k in _shapes:
    q = _quant_bufs[(_m, _k)]
    _d = torch.randn(_m, _k, dtype=_bf16, device='cuda')
    _fused_quant_shuffle_kernel[q['grid']](_d, q['x_fp4'], q['bs'], *_d.stride(), q['fp4_s0'], q['fp4_s1'],
        q['sn'], M=_m, N=_k, MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0, NUM_ITER=1,
        BLOCK_SIZE_M=q['BSM'], BLOCK_SIZE_N=q['BSN'], NUM_STAGES=1, num_warps=q['NW'], waves_per_eu=0, num_stages=1)
torch.cuda.synchronize()


# ============================================================
# Dispatch logic
# ============================================================

def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n = B.shape[0]

    # Fused quant+GEMM (m<=32, k<=1024)
    if _mfma_ok and m <= 32 and k <= 1024:
        fb = _fused_bufs.get((m, n, k))
        if fb is None: fb = _setup_fused(m, n, k)
        bq_u8 = B_q.view(torch.uint8) if B_q.dtype != torch.uint8 else B_q
        bsc_u8 = _get_bsc(B_scale_sh, n, k)
        _mfma_mod.launch_fused(A, bq_u8, bsc_u8, fb['out'],
                                m, n, k, fb['snb'], fb['shared'])
        return fb['out']

    # Fused quant+wide4 MFMA (m<=32, k>1024)
    if _mfma_ok and m <= 32 and k > 1024:
        wb = _wide_bufs.get((m, n, k))
        if wb is None: wb = _setup_wide(m, n, k)
        bq_u8 = B_q.view(torch.uint8) if B_q.dtype != torch.uint8 else B_q
        bsc_u8 = _get_bsc(B_scale_sh, n, k)
        sk = wb['split_k']
        _mfma_mod.launch_fused_wide4(A, bq_u8, bsc_u8,
                                      wb['out_bf16'], wb['out_f32'],
                                      m, n, k, wb['snb'], sk, wb['shmem_fw'])
        if sk > 1: _mfma_mod.launch_cvt_zero(wb['out_f32'], wb['out_bf16'], m * n)
        return wb['out_bf16']

    # Quantize A (for aiter ASM path — m>32)
    q = _quant_bufs.get((m, k))
    if q is None: q = _setup_quant(m, k)
    _fused_quant_shuffle_kernel[q['grid']](A, q['x_fp4'], q['bs'], k, 1, q['fp4_s0'], q['fp4_s1'],
        q['sn'], M=m, N=k, MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0, NUM_ITER=1,
        BLOCK_SIZE_M=q['BSM'], BLOCK_SIZE_N=q['BSN'], NUM_STAGES=1, num_warps=q['NW'], waves_per_eu=0, num_stages=1)

    # aiter ASM (m>32)
    d = _gemm_bufs.get((m, n, k))
    if d is None: d = _setup_gemm(m, n, k)
    if d is not None:
        _asm(q['aq'], B_shuffle, q['ash'], B_scale_sh, d[1], d[2], None, 1.0, 0.0, True, None)
        return d[0]
    return aiter.gemm_a4w4(q['aq'], B_shuffle, q['ash'], B_scale_sh, dtype=_bf16, bpreshuffle=True)
scrolls · 728 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