Skip to content
KernelIndex
Search⌘K

submission 732851

allen98. · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6b27ba8f981fdf998b3ef4949babc42900a77b15d5eae7912c8a7fd5efe1b855
license declaredunknown
license concludedunknown
authorsallen98.
imported2026-08-15

Techniques

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

fp4FP4 GEMM v31 — Hybrid: HIP fused for M≤32 (ANY K) + Triton for M=33-127 + ASM for M≥128.
num-warps = 4num_warps=4,num_stages=1,waves_per_eu=4,matrix_instr_nonkdim=16,NUM_KSPLIT=1),
shared-memory__shared__ float red[4][64*4+16];
split-kint64_t sm, const std::string& knl_name, const std::string& co_path, int64_t splitK);
stages = 1num_warps=4,num_stages=1,waves_per_eu=4,matrix_instr_nonkdim=16,NUM_KSPLIT=1),
tile-k = 512const int nk = (K + 511) / 512; // BK=512 outer loop
tile-m = 16(4, 2880, 512): dict(BLOCK_SIZE_M=16,BLOCK_SIZE_N=64,BLOCK_SIZE_K=512,GROUP_SIZE_M=1,
tile-n = 64(4, 2880, 512): dict(BLOCK_SIZE_M=16,BLOCK_SIZE_N=64,BLOCK_SIZE_K=512,GROUP_SIZE_M=1,
vector-width = uint4uint4 a0,a1,a2,a3;

Kernel source

kernel.py1394 lines
"""
FP4 GEMM v31 — Hybrid: HIP fused for M≤32 (ANY K) + Triton for M=33-127 + ASM for M≥128.
Key change from v30:
  - v31: Removed K threshold for HIP path. gemm_fused_16x16_kern now handles ALL K values.
    For shape 2 (M=16 K=7168): single HIP kernel replaces 2 Triton launches.
    Profiling showed 83% launch overhead for shape 2 — this eliminates it.
  - ASM path for M≥128: hand-tuned 186% MFMA util beats Triton for M=256
"""
import functools, hashlib, os, sys, subprocess, shutil, tempfile, csv, glob

os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")

_IS_PROFILE_SUB = os.environ.get("PROFILE_SUBPROCESS") == "1"
_DO_PROFILE = os.environ.get("DO_PROFILE", "0") == "1"  # Enable with DO_PROFILE=1 env var
_ORIG_DIR = os.environ.get("ORIG_DIR", os.path.dirname(os.path.abspath(__file__)))
if _ORIG_DIR not in sys.path:
    sys.path.insert(0, _ORIG_DIR)

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

SCALE_GROUP_SIZE = 32

# ═══════════════════════════════════════════════════════════════════
# HIP C++ kernels: 16x16 fused MFMA (M≤32 K≤1024) + ASM path (M≥128)
# ═══════════════════════════════════════════════════════════════════

CPP_SRC = r"""
#include <torch/extension.h>
void quantize_a_hip(torch::Tensor a, torch::Tensor a_q, torch::Tensor a_scale);
torch::Tensor gemm_fused_16x16_hip(torch::Tensor A, torch::Tensor Bq,
    torch::Tensor Bsc, int64_t N, torch::Tensor out);
void quant_a_shuffled(torch::Tensor A, torch::Tensor Aq, torch::Tensor Asc, int64_t sm);
torch::Tensor quant_and_asm(torch::Tensor A, torch::Tensor Bsh,
    torch::Tensor Bsc, int64_t N,
    torch::Tensor Aq, torch::Tensor Asc, torch::Tensor out,
    int64_t sm, const std::string& knl_name, const std::string& co_path, int64_t splitK);
void warmup_asm(const std::string& co_path, const std::string& knl_name);
void fast_hip16(int64_t A, int64_t Bq, int64_t Bsc, int64_t C,
    int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBq, int64_t BscSN);
void fast_ext_splitk(int64_t A, int64_t Bq, int64_t Bsc, int64_t Cfp, int64_t Cbf,
    int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBq, int64_t BscSN, int64_t ks);
void dispatch_hip(int64_t A, int64_t Bq, int64_t Bsc, int64_t C, int64_t Cfp,
    int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBq, int64_t BscSN, int64_t ks);
void fast_reduce_splitk(int64_t src, int64_t dst, int64_t M, int64_t N, int64_t ksplit);
void launch_fused_npar(int64_t A, int64_t Bsh, int64_t Bsc, int64_t C,
    int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBsh, int64_t BscSN);
void load_triton_hsaco(const std::string& path, const std::string& name);
void launch_triton_kernel(const std::string& key,
    int64_t p0, int64_t p1, int64_t p2, int64_t p3,
    int64_t a0, int64_t a1, int64_t a2, int64_t a3, int64_t a4,
    int64_t a5, int64_t a6, int64_t a7, int64_t a8, int64_t a9, int64_t a10, int64_t a11,
    int64_t gridX, int64_t gridY, int64_t blockX, int64_t shared_mem);
void launch_triton_fp4gemm(const std::string& key,
    int64_t Aq, int64_t Bt, int64_t C, int64_t Asc, int64_t Bsc,
    int64_t M, int64_t N, int64_t K, int64_t Ks,
    int64_t sam, int64_t sbn, int64_t sck, int64_t scm, int64_t sasm,
    int64_t gridX, int64_t shared_mem);
void launch_nofuse_triton(int64_t Aq, int64_t Bt, int64_t C,
    int64_t Asc, int64_t Bsc,
    int64_t M, int64_t N, int64_t K, int64_t Ks,
    int64_t sam, int64_t sbn, int64_t scm, int64_t sasm);
"""

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

namespace {

constexpr int QUANT_GROUP = 32;
constexpr int GROUPS_PER_BLOCK = 4;
constexpr int BLOCK_THREADS = QUANT_GROUP * GROUPS_PER_BLOCK;

__device__ inline uint32_t float_to_bits(float x) {
    union { float f; uint32_t u; } v; v.f = x; return v.u;
}
__device__ inline float bits_to_float(uint32_t x) {
    union { uint32_t u; float f; } v; v.u = x; return v.f;
}

__device__ inline uint8_t encode_mxfp4(float x) {
    constexpr float MAX_NORMAL = 6.0f, MIN_NORMAL = 1.0f;
    constexpr int32_t DENORM_EXP = (127-1)+(23-1)+1;
    constexpr int32_t DENORM_MASK_INT = DENORM_EXP << 23;
    constexpr uint32_t NORMAL_ROUND_BIAS_BITS = 0xC11FFFFFu;
    const uint32_t bits = float_to_bits(x);
    const uint32_t sign = bits & 0x80000000u;
    const float abs_x = bits_to_float(bits ^ sign);
    bool saturate = abs_x >= MAX_NORMAL;
    bool denormal = (!saturate) && (abs_x < MIN_NORMAL);
    uint8_t value = 0x7;
    if (denormal) {
        float df = abs_x + bits_to_float((uint32_t)DENORM_MASK_INT);
        value = (uint8_t)(float_to_bits(df) - (uint32_t)DENORM_MASK_INT);
    } else if (!saturate) {
        uint32_t ab = bits ^ sign;
        uint32_t mo = (ab >> 22) & 1u;
        int32_t r = (int32_t)ab + (int32_t)NORMAL_ROUND_BIAS_BITS + (int32_t)mo;
        value = (uint8_t)((uint32_t)r >> 22);
    }
    return value | (sign ? 0x8u : 0u);
}

__global__ __launch_bounds__(BLOCK_THREADS)
void quantize_a_kernel(const __hip_bfloat16* __restrict__ a, uint8_t* __restrict__ a_q,
    uint8_t* __restrict__ a_scale, int m, int k) {
    const int row = blockIdx.y;
    const int gib = threadIdx.x / QUANT_GROUP;
    const int lane = threadIdx.x % QUANT_GROUP;
    const int k_block = blockIdx.x * GROUPS_PER_BLOCK + gib;
    if (k_block >= k / QUANT_GROUP) return;
    const int col = k_block * QUANT_GROUP + lane;
    float value = (row < m && col < k) ? (float)a[row * k + col] : 0.0f;
    float abs_val = fabsf(value);
    for (int off = 16; off > 0; off >>= 1) {
        float o = __shfl_xor(abs_val, off); abs_val = fmaxf(abs_val, o);
    }
    uint32_t ab = float_to_bits(abs_val);
    ab = (ab + 0x00200000u) & 0xFF800000u;
    int ef = (ab >> 23u) & 0xFFu;
    int se = (ab == 0u) ? -127 : max(-127, min(127, ef - 127 - 2));
    float qs = (se <= 126) ? bits_to_float((uint32_t)(127 - se) << 23) : 0.0f;
    uint8_t sb = (uint8_t)(se + 127);
    if (lane == 0) a_scale[row * (k / QUANT_GROUP) + k_block] = sb;
    uint8_t nib = encode_mxfp4(value * qs);
    uint8_t pn = __shfl_xor(nib, 1);
    if ((lane & 1) == 0 && row < m)
        a_q[row * (k / 2) + k_block * 16 + (lane / 2)] = nib | (pn << 4);
}

typedef int __attribute__((ext_vector_type(8))) i32x8;
typedef float __attribute__((ext_vector_type(4))) f32x4;
typedef __bf16 bf16v2_t __attribute__((ext_vector_type(2)));

__device__ __forceinline__ void quant_32_hw(const uint16_t* ap, int out[4], int32_t& spk) {
    uint16_t mx = 0;
    #pragma unroll
    for (int j = 0; j < 32; j++) mx = max(mx, (uint16_t)(ap[j] & 0x7FFF));
    uint32_t au = (((uint32_t)mx << 16) + 0x200000u) & 0xFF800000u;
    int ef = (au >> 23u) & 0xFFu;
    int su = (au == 0u) ? -127 : max(-127, min(127, ef - 127 - 2));
    float hs = (su >= -126) ? __uint_as_float((uint32_t)(su + 127) << 23) : 0.0f;
    #pragma unroll
    for (int j = 0; j < 4; j++) {
        out[j] = 0;
        out[j] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j], *reinterpret_cast<const bf16v2_t*>(&ap[j*8+0]), hs, 0);
        out[j] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j], *reinterpret_cast<const bf16v2_t*>(&ap[j*8+2]), hs, 1);
        out[j] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j], *reinterpret_cast<const bf16v2_t*>(&ap[j*8+4]), hs, 2);
        out[j] = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j], *reinterpret_cast<const bf16v2_t*>(&ap[j*8+6]), hs, 3);
    }
    spk = (int32_t)(uint8_t)(su + 127);
}

__device__ __forceinline__ int32_t load_b_scale_16(const uint8_t* Bsc,
    int ng, int bb, int SNG, int N) {
    if (ng >= N) return 127;
    return (int32_t)Bsc[(ng/32)*(SNG*256)+(bb/8)*256+(bb%4)*64+(ng%16)*4+((bb%8)/4)*2+(ng%32)/16];
}

__global__ void __launch_bounds__(256, 3)
gemm_fused_16x16_kern(const uint16_t* __restrict__ A, const uint8_t* __restrict__ Bq,
    const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
    int M, int N, int K, int strA, int strBq, int BscSN) {
    const int mt = blockIdx.x, nt = blockIdx.y;
    const int wid = threadIdx.x >> 6, lid = threadIdx.x & 63;
    const int t_row = lid & 15, kpart = lid >> 4;
    const int SNG = BscSN >> 3;
    const int m_row = mt * 16 + t_row, a_off = m_row * strA, b_ng = nt * 16 + t_row;
    int ts = K >> 7, spw = (ts + 3) >> 2;
    int ks = wid * spw * 128, ke = min(ks + spw * 128, K);
    f32x4 c = {0,0,0,0};
    for (int kb = ks; kb < ke; kb += 128) {
        int ko = kb + kpart * 32, bk = (kb >> 1) + kpart * 16;
        uint4 a0,a1,a2,a3;
        if (m_row<M&&ko+31<K){const uint4*s=reinterpret_cast<const uint4*>(&A[a_off+ko]);a0=s[0];a1=s[1];a2=s[2];a3=s[3];}
        else{a0={0,0,0,0};a1={0,0,0,0};a2={0,0,0,0};a3={0,0,0,0};}
        int bi[4];
        if(b_ng<N&&bk+15<(K>>1)){uint4 bd=*reinterpret_cast<const uint4*>(&Bq[b_ng*strBq+bk]);bi[0]=((int*)&bd)[0];bi[1]=((int*)&bd)[1];bi[2]=((int*)&bd)[2];bi[3]=((int*)&bd)[3];}
        else{bi[0]=0;bi[1]=0;bi[2]=0;bi[3]=0;}
        int32_t bs=load_b_scale_16(Bsc,b_ng,(kb>>5)+kpart,SNG,N);
        uint16_t al[32];uint4*d=reinterpret_cast<uint4*>(al);d[0]=a0;d[1]=a1;d[2]=a2;d[3]=a3;
        int ai[4];int32_t as2;quant_32_hw(al,ai,as2);
        i32x8 am={ai[0],ai[1],ai[2],ai[3],0,0,0,0};
        i32x8 bm={bi[0],bi[1],bi[2],bi[3],0,0,0,0};
        c=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(am,bm,c,4,4,0,as2,0,bs);
    }
    __shared__ float red[4][64*4+16];
    #pragma unroll
    for(int i=0;i<4;i++) red[wid][lid*4+i]=c[i];
    __syncthreads();
    if(wid==0){for(int j=0;j<4;j++){
        float s=red[0][lid*4+j]+red[1][lid*4+j]+red[2][lid*4+j]+red[3][lid*4+j];
        int mo=mt*16+kpart*4+j,no=nt*16+t_row;
        if(mo<M&&no<N){uint32_t fp=__float_as_uint(s);fp+=0x7FFFu+((fp>>16)&1u);C[mo*N+no]=(uint16_t)(fp>>16u);}
    }}
}

// v78: N-parallel kernel matching Triton's warp layout {1,4} from translation doc
// 32 M-rows × 64 N-cols per CTA. 4 warps × 16 N-cols each.
// 2 M-reps × 4 K-reps = 8 MFMAs per BK=512 step. B loaded once, reused across M-reps.
// B_shuffle layout: [N/16][K/32][16][16] — coalesced 256-byte tiles.
__global__ void __launch_bounds__(256, 2)
gemm_fused_npar_kern(const uint16_t* __restrict__ A, const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
    int M, int N, int K, int strA, int strBsh, int BscSN) {
    const int mt = blockIdx.x, nt = blockIdx.y;
    const int wid = threadIdx.x >> 6, lid = threadIdx.x & 63;
    const int t_row = lid & 15, kpart = lid >> 4;
    const int SNG = BscSN >> 3;
    const int b_tile_n = nt * 4 + wid;
    const int b_lane = t_row;
    const int b_ng = nt * 64 + wid * 16 + t_row;
    const int nkt = K / 32;
    // 2 M-rep base rows (32-row M-tile)
    const int m_row0 = mt * 32 + t_row;
    const int m_row1 = mt * 32 + 16 + t_row;
    const int a_off0 = m_row0 * strA;
    const int a_off1 = m_row1 * strA;
    f32x4 c0 = {0,0,0,0};
    f32x4 c1 = {0,0,0,0};
    const int nk = (K + 511) / 512;  // BK=512 outer loop
    for (int ki = 0; ki < nk; ki++) {
        #pragma unroll
        for (int krep = 0; krep < 4; krep++) {  // 4 K-reps per BK=512 step
            int kbase = ki * 512 + krep * 128;
            int ko = kbase + kpart * 32;
            // Load B (shared between M-reps)
            int btk = kbase / 32 + kpart;
            int boff = b_tile_n * nkt * 256 + btk * 256 + b_lane * 16;
            int bi[4];
            if(b_ng<N && btk<nkt){uint4 bd=*reinterpret_cast<const uint4*>(&Bsh[boff]);
                bi[0]=((int*)&bd)[0];bi[1]=((int*)&bd)[1];bi[2]=((int*)&bd)[2];bi[3]=((int*)&bd)[3];}
            else{bi[0]=0;bi[1]=0;bi[2]=0;bi[3]=0;}
            int32_t bs = load_b_scale_16(Bsc, b_ng, kbase/32+kpart, SNG, N);
            i32x8 bm = {bi[0],bi[1],bi[2],bi[3],0,0,0,0};
            // M-rep 0: load A, quant, MFMA
            {
                uint16_t al[32]; int ai[4]; int32_t as2;
                if(m_row0<M&&ko+31<K){const uint4*s=reinterpret_cast<const uint4*>(&A[a_off0+ko]);
                    uint4*d=reinterpret_cast<uint4*>(al);d[0]=s[0];d[1]=s[1];d[2]=s[2];d[3]=s[3];}
                else{for(int j=0;j<32;j++)al[j]=(m_row0<M&&ko+j<K)?A[a_off0+ko+j]:0;}
                quant_32_hw(al,ai,as2);
                i32x8 am={ai[0],ai[1],ai[2],ai[3],0,0,0,0};
                c0=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(am,bm,c0,4,4,0,as2,0,bs);
            }
            // M-rep 1: reuse B, load different A rows
            {
                uint16_t al[32]; int ai[4]; int32_t as2;
                if(m_row1<M&&ko+31<K){const uint4*s=reinterpret_cast<const uint4*>(&A[a_off1+ko]);
                    uint4*d=reinterpret_cast<uint4*>(al);d[0]=s[0];d[1]=s[1];d[2]=s[2];d[3]=s[3];}
                else{for(int j=0;j<32;j++)al[j]=(m_row1<M&&ko+j<K)?A[a_off1+ko+j]:0;}
                quant_32_hw(al,ai,as2);
                i32x8 am={ai[0],ai[1],ai[2],ai[3],0,0,0,0};
                c1=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(am,bm,c1,4,4,0,as2,0,bs);
            }
        }
    }
    // Output M-rep 0
    int no = nt*64 + wid*16 + t_row;
    int mo0 = mt*32 + kpart*4;
    #pragma unroll
    for(int j=0;j<4;j++){
        if(mo0+j<M&&no<N){uint32_t fp=__float_as_uint(c0[j]);fp+=0x7FFFu+((fp>>16)&1u);
            C[(mo0+j)*N+no]=(uint16_t)(fp>>16u);}
    }
    // Output M-rep 1
    int mo1 = mt*32 + 16 + kpart*4;
    #pragma unroll
    for(int j=0;j<4;j++){
        if(mo1+j<M&&no<N){uint32_t fp=__float_as_uint(c1[j]);fp+=0x7FFFu+((fp>>16)&1u);
            C[(mo1+j)*N+no]=(uint16_t)(fp>>16u);}
    }
}

// v80: Non-fused FP4 GEMM reproducing Triton's exact ISA structure
// Key findings from ISA dump:
//   - isTransposed=true: MFMA src0=B_data, src1=A_data (SWAPPED)
//   - 5th operand = B_scale, 6th = A_scale
//   - opSel 0,1,2,3 via op_sel + op_sel_hi for 4 K-reps
//   - 4 K-reps per BK=512, 2 M-reps, 8 MFMAs per K-step
//   - Bt layout: [K/2, N] with stride(0)=1 (K contiguous), stride(1)=K/2
__global__ void __launch_bounds__(256, 1)
gemm_nofuse_triton(const uint8_t* __restrict__ Aq, const uint8_t* __restrict__ Bt,
    uint16_t* __restrict__ C,
    const uint8_t* __restrict__ Asc, const uint8_t* __restrict__ Bsc,
    int M, int N, int K, int Ks, int sam, int sbn, int scm, int sasm) {
    // K = Kp = original_K / 2 (packed FP4 byte count per row)
    const int tid = threadIdx.x;
    const int wid = tid >> 6;
    const int lid = tid & 63;
    const int t_row = lid & 15;
    const int kpart = lid >> 4;
    // Grid mapping (same as Triton)
    int npn = (N + 63) / 64;
    int pm = blockIdx.x / npn;
    int pn = blockIdx.x % npn;
    // M/N tile offsets with wrapping
    int m0 = (pm * 32 + t_row) % M;                  // M-rep 0 row
    int m1 = (pm * 32 + 16 + t_row) % M;             // M-rep 1 row
    int n_col = (pn * 64 + wid * 16 + t_row) % N;    // N column for this thread
    // Accumulators (matches Triton a[0:3], a[4:7])
    f32x4 acc0 = {0,0,0,0};
    f32x4 acc1 = {0,0,0,0};
    // BK=512 → 4 K-reps of 128 FP4 each. K is packed byte count.
    // nk = K / 256 (each BK step = 256 packed bytes = 512 FP4 elements)
    int nk = K / 256;  // K here = Kp = original_K / 2
    for (int ki = 0; ki < nk; ki++) {
        int k_byte_base = ki * 256;  // byte offset for this BK step
        // Pack 4 A scales for M-rep 0 (bytes for k-rep 0,1,2,3)
        int32_t packed_asc0 = 0, packed_asc1 = 0;
        // Pack 4 B scales
        int32_t packed_bsc = 0;
        {
            uint8_t a0s[4], a1s[4], bs[4];
            #pragma unroll
            for (int kr = 0; kr < 4; kr++) {
                int k_group = ki * 16 + kr * 4 + kpart; // scale group index
                a0s[kr] = Asc[m0 * sasm + k_group];
                a1s[kr] = Asc[m1 * sasm + k_group];
                // B scale uses shuffled layout: _bsc_off(n_col, k_group, Ks)
                int bsc_off = (n_col/32)*(Ks*32) + (k_group/8)*256 + (k_group%4)*64
                    + (n_col%16)*4 + ((k_group/4)%2)*2 + ((n_col/16)%2);
                bs[kr] = Bsc[bsc_off];
            }
            __builtin_memcpy(&packed_asc0, a0s, 4);
            __builtin_memcpy(&packed_asc1, a1s, 4);
            __builtin_memcpy(&packed_bsc, bs, 4);
        }
        // 4 K-reps manually unrolled (opSel must be compile-time constant)
        #define KREP(KR) { \
            int k_byte = k_byte_base + (KR) * 64 + kpart * 16; \
            int b_off = k_byte + n_col * sbn; \
            uint4 bd = *reinterpret_cast<const uint4*>(&Bt[b_off]); \
            i32x8 bm = {((int*)&bd)[0],((int*)&bd)[1],((int*)&bd)[2],((int*)&bd)[3],0,0,0,0}; \
            uint4 ad0 = *reinterpret_cast<const uint4*>(&Aq[m0 * sam + k_byte]); \
            i32x8 am0 = {((int*)&ad0)[0],((int*)&ad0)[1],((int*)&ad0)[2],((int*)&ad0)[3],0,0,0,0}; \
            uint4 ad1 = *reinterpret_cast<const uint4*>(&Aq[m1 * sam + k_byte]); \
            i32x8 am1 = {((int*)&ad1)[0],((int*)&ad1)[1],((int*)&ad1)[2],((int*)&ad1)[3],0,0,0,0}; \
            acc0 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4( \
                bm, am0, acc0, 4, 4, (KR), packed_bsc, (KR), packed_asc0); \
            acc1 = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4( \
                bm, am1, acc1, 4, 4, (KR), packed_bsc, (KR), packed_asc1); \
        }
        KREP(0) KREP(1) KREP(2) KREP(3)
        #undef KREP
    }
    // Output with isTransposed=true: a[j] = C[N=kpart*4+j, M=t_row]
    // M-row comes from t_row (src1=A dim), N-col from kpart*4+j (src0=B dim)
    int out_m0 = pm * 32 + t_row;              // M-rep 0
    int out_m1 = pm * 32 + 16 + t_row;         // M-rep 1
    int out_n_base = pn * 64 + wid * 16 + kpart * 4;
    if (out_m0 < M) {
        #pragma unroll
        for (int j = 0; j < 4; j++) {
            if (out_n_base + j < N) {
                uint32_t fp = __float_as_uint(acc0[j]); fp += 0x7FFFu + ((fp >> 16) & 1u);
                C[out_m0 * scm + out_n_base + j] = (uint16_t)(fp >> 16u);
            }
        }
    }
    if (out_m1 < M) {
        #pragma unroll
        for (int j = 0; j < 4; j++) {
            if (out_n_base + j < N) {
                uint32_t fp = __float_as_uint(acc1[j]); fp += 0x7FFFu + ((fp >> 16) & 1u);
                C[out_m1 * scm + out_n_base + j] = (uint16_t)(fp >> 16u);
            }
        }
    }
}

// ── ASM GEMM launcher ──
static std::unordered_map<std::string, std::pair<hipModule_t, hipFunction_t>> _asm_cache;

struct __attribute__((packed)) AsmArgs {
    void* D; char _0[8]; void* C; char _1[8];
    void* A; char _2[8]; void* B; char _3[8];
    float alpha; char _4[12]; float beta; char _5[12];
    unsigned sD0; char _6[12]; unsigned sD1; char _7[12];
    unsigned sC0; char _8[12]; unsigned sC1; char _9[12];
    unsigned sA0; char _10[12]; unsigned sA1; char _11[12];
    unsigned sB0; char _12[12]; unsigned sB1; char _13[12];
    unsigned M; char _14[12]; unsigned N; char _15[12];
    unsigned K; char _16[12];
    void* SA; char _17[8]; void* SB; char _18[8];
    unsigned sSA0; char _19[12]; unsigned sSA1; char _20[12];
    unsigned sSB0; char _21[12]; unsigned sSB1; char _22[12];
    int log2ks;
};

hipFunction_t get_asm_fn(const std::string& co_path, const std::string& knl_name) {
    auto key = co_path + "::" + knl_name;
    auto it = _asm_cache.find(key);
    if (it != _asm_cache.end()) return it->second.second;
    hipModule_t mod; hipFunction_t fn;
    hipModuleLoad(&mod, co_path.c_str());
    hipModuleGetFunction(&fn, mod, knl_name.c_str());
    _asm_cache[key] = {mod, fn};
    return fn;
}

__global__ void __launch_bounds__(64, 16)
quant_a_kern(const uint16_t* __restrict__ A, uint8_t* __restrict__ Aq,
    uint8_t* __restrict__ Asc, int M, int K, int strA, int sm) {
#if defined(__gfx950__)
    int idx = blockIdx.x * blockDim.x + threadIdx.x;
    int ns = K / 32, total = sm * ns;
    if (idx >= total) return;
    int row = idx / ns, blk = idx % ns, ko = blk * 32;
    uint16_t v[32];
    if (row < M) {
        const uint4* s = reinterpret_cast<const uint4*>(&A[row * strA + ko]);
        uint4* d = reinterpret_cast<uint4*>(v); d[0]=s[0];d[1]=s[1];d[2]=s[2];d[3]=s[3];
    } else { for(int j=0;j<32;j++) v[j]=0; }
    uint16_t mx=0;
    for(int j=0;j<32;j++) mx=max(mx,(uint16_t)(v[j]&0x7FFF));
    uint32_t au=(((uint32_t)mx<<16)+0x200000u)&0xFF800000u;
    int ef=(au>>23u)&0xFFu;
    int su=(au==0u)?-127:max(-127,min(127,ef-127-2));
    float hs=(su>=-126)?__uint_as_float((uint32_t)(su+127)<<23):0.0f;
    uint8_t sv=(uint8_t)(su+127);
    if(row<M){
        unsigned int p[4];
        for(int i=0;i<4;i++){p[i]=0;
            p[i]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p[i],*reinterpret_cast<bf16v2_t*>(&v[i*8+0]),hs,0);
            p[i]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p[i],*reinterpret_cast<bf16v2_t*>(&v[i*8+2]),hs,1);
            p[i]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p[i],*reinterpret_cast<bf16v2_t*>(&v[i*8+4]),hs,2);
            p[i]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p[i],*reinterpret_cast<bf16v2_t*>(&v[i*8+6]),hs,3);
        }
        *reinterpret_cast<uint4*>(&Aq[row*(K/2)+blk*16])=*reinterpret_cast<uint4*>(p);
    }
    int sn=K/32;
    int i0=row/32,i1=(row%32)/16,i2=row%16;
    int i3=blk/8,i4=(blk%8)/4,i5=blk%4;
    Asc[i0*(32*sn)+i3*256+i5*64+i2*4+i4*2+i1]=sv;
#endif
}

}  // namespace

void quantize_a_hip(torch::Tensor a, torch::Tensor a_q, torch::Tensor a_scale) {
    int m=(int)a.size(0), k=(int)a.size(1);
    int nkg=k/QUANT_GROUP;
    dim3 g((nkg+GROUPS_PER_BLOCK-1)/GROUPS_PER_BLOCK, m);
    quantize_a_kernel<<<g, BLOCK_THREADS>>>(
        reinterpret_cast<const __hip_bfloat16*>(a.data_ptr<at::BFloat16>()),
        a_q.data_ptr<uint8_t>(), a_scale.data_ptr<uint8_t>(), m, k);
}

torch::Tensor gemm_fused_16x16_hip(torch::Tensor A, torch::Tensor Bq,
    torch::Tensor Bsc, int64_t N, torch::Tensor out) {
    int M=(int)A.size(0),K=(int)A.size(1);
    gemm_fused_16x16_kern<<<dim3((M+15)/16,((int)N+15)/16),256>>>(
        reinterpret_cast<const uint16_t*>(A.data_ptr()),
        reinterpret_cast<const uint8_t*>(Bq.data_ptr()),
        reinterpret_cast<const uint8_t*>(Bsc.data_ptr()),
        reinterpret_cast<uint16_t*>(out.data_ptr()),
        M,(int)N,K,(int)A.stride(0),(int)(Bq.stride(0)*Bq.element_size()),(int)Bsc.size(1));
    return out;
}

void quant_a_shuffled(torch::Tensor A, torch::Tensor Aq, torch::Tensor Asc, int64_t sm) {
    int M=(int)A.size(0),K=(int)A.size(1);
    int ns=K/32, total=(int)sm*ns;
    quant_a_kern<<<(total+63)/64, 64>>>(
        reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
        Aq.data_ptr<uint8_t>(), Asc.data_ptr<uint8_t>(), M, K, (int)A.stride(0), (int)sm);
}

torch::Tensor quant_and_asm(torch::Tensor A, torch::Tensor Bsh,
    torch::Tensor Bsc, int64_t N,
    torch::Tensor Aq, torch::Tensor Asc, torch::Tensor out,
    int64_t sm, const std::string& knl_name, const std::string& co_path, int64_t splitK) {
    int M=(int)A.size(0),K=(int)A.size(1);
    int ns=K/32, total=(int)sm*ns;
    quant_a_kern<<<(total+63)/64, 64>>>(
        reinterpret_cast<const uint16_t*>(A.data_ptr<at::BFloat16>()),
        Aq.data_ptr<uint8_t>(), Asc.data_ptr<uint8_t>(), M, K, (int)A.stride(0), (int)sm);
    auto fn = get_asm_fn(co_path, knl_name);
    AsmArgs a={};
    a.D=out.data_ptr();a.C=out.data_ptr();
    a.A=Aq.data_ptr();a.B=Bsh.data_ptr();
    a.alpha=1.0f;a.beta=0.0f;
    a.sC0=(unsigned)out.stride(0);a.sC1=1;
    a.sA0=(unsigned)(Aq.stride(0)*2);a.sA1=1;
    a.sB0=(unsigned)(Bsh.stride(0)*2);a.sB1=1;
    a.M=(unsigned)M;a.N=(unsigned)N;a.K=(unsigned)K;
    a.SA=Asc.data_ptr();a.SB=Bsc.data_ptr();
    a.sSA0=(unsigned)Asc.stride(0);a.sSA1=1;
    a.sSB0=(unsigned)Bsc.stride(0);a.sSB1=1;
    a.log2ks=(int)splitK;
    size_t asz=sizeof(a);
    void*cfg[]={HIP_LAUNCH_PARAM_BUFFER_POINTER,&a,HIP_LAUNCH_PARAM_BUFFER_SIZE,&asz,HIP_LAUNCH_PARAM_END};
    // Parse tile_m x tile_n from kernel name (e.g. "..._64x128E" -> tm=64,tn=128)
    unsigned tm=32, tn=128;
    auto pos=knl_name.rfind('x');
    if(pos!=std::string::npos) {
        tn=std::stoi(knl_name.substr(pos+1));
        // Find tile_m: scan backwards from 'x' to find the start of the number
        auto pos2=pos-1;
        while(pos2>0 && knl_name[pos2-1]>='0' && knl_name[pos2-1]<='9') pos2--;
        tm=std::stoi(knl_name.substr(pos2, pos-pos2));
    }
    unsigned mp=((unsigned)M+tm-1)/tm*tm;
    hipModuleLaunchKernel(fn,(N+tn-1)/tn,mp/tm,1,256,1,1,0,0,nullptr,(void**)cfg);
    return out;
}

void warmup_asm(const std::string& co_path, const std::string& knl_name) {
    get_asm_fn(co_path, knl_name);
}

// Raw-pointer fast path: avoids pybind11 tensor marshalling
void fast_hip16(int64_t A, int64_t Bq, int64_t Bsc, int64_t C,
    int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBq, int64_t BscSN) {
    gemm_fused_16x16_kern<<<dim3(((int)M+15)/16,((int)N+15)/16),256>>>(
        (const uint16_t*)A, (const uint8_t*)Bq, (const uint8_t*)Bsc, (uint16_t*)C,
        (int)M, (int)N, (int)K, (int)strA, (int)strBq, (int)BscSN);
}

// ext_splitk: on-the-fly quant + atomicAdd split-K (for M<=16 K>1024)
__device__ __forceinline__ void quant_32_inline(const uint16_t* ap, int out[4], int32_t& spk) {
    uint16_t mx=0;
    #pragma unroll
    for(int j=0;j<32;j++) mx=max(mx,(uint16_t)(ap[j]&0x7FFF));
    uint32_t au=(((uint32_t)mx<<16)+0x200000u)&0xFF800000u;
    int ef=(au>>23u)&0xFFu;
    int su=(au==0u)?-127:max(-127,min(127,ef-127-2));
    float hs=(su>=-126)?__uint_as_float((uint32_t)(su+127)<<23):0.0f;
    #pragma unroll
    for(int j=0;j<4;j++){out[j]=0;
        out[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j],*reinterpret_cast<const bf16v2_t*>(&ap[j*8+0]),hs,0);
        out[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j],*reinterpret_cast<const bf16v2_t*>(&ap[j*8+2]),hs,1);
        out[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j],*reinterpret_cast<const bf16v2_t*>(&ap[j*8+4]),hs,2);
        out[j]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(out[j],*reinterpret_cast<const bf16v2_t*>(&ap[j*8+6]),hs,3);
    }
    spk=(int32_t)(uint8_t)(su+127);
}

__global__ void __launch_bounds__(64, 8)
ext_splitk_16x16_kern(const uint16_t* __restrict__ A, const uint8_t* __restrict__ Bq,
    const uint8_t* __restrict__ Bsc, float* __restrict__ Cfp,
    int M, int N, int K, int strA, int strBq, int BscSN, int kper) {
#if defined(__gfx950__)
    const int mt=blockIdx.x,nt=blockIdx.y,ks=blockIdx.z;
    const int lid=threadIdx.x,t_row=lid&15,kpart=lid>>4;
    const int SNG=BscSN>>3,m_row=mt*16+t_row,b_ng=nt*16+t_row;
    int kst=(ks*kper/128)*128,ken=min(((ks+1)*kper+127)/128*128,K);
    if(kst>=ken) return;
    f32x4 c_acc={0,0,0,0};
    for(int kb=kst;kb<ken;kb+=128){
        int k_off=kb+kpart*32;
        uint16_t al[32]; int a_i[4]; int32_t a_spk;
        if(m_row<M&&k_off+31<K){const uint4*s=reinterpret_cast<const uint4*>(&A[m_row*strA+k_off]);
            uint4*d=reinterpret_cast<uint4*>(al);d[0]=s[0];d[1]=s[1];d[2]=s[2];d[3]=s[3];}
        else{for(int j=0;j<32;j++)al[j]=(m_row<M&&k_off+j<K)?A[m_row*strA+k_off+j]:0;}
        quant_32_inline(al,a_i,a_spk);
        int bk=kb/2+kpart*16; int b_i[4]; int32_t b_spk;
        if(b_ng<N&&bk+15<K/2){uint4 bd=*reinterpret_cast<const uint4*>(&Bq[b_ng*strBq+bk]);
            b_i[0]=((int*)&bd)[0];b_i[1]=((int*)&bd)[1];b_i[2]=((int*)&bd)[2];b_i[3]=((int*)&bd)[3];}
        else{b_i[0]=0;b_i[1]=0;b_i[2]=0;b_i[3]=0;}
        b_spk=load_b_scale_16(Bsc,b_ng,kb/32+kpart,SNG,N);
        i32x8 av={a_i[0],a_i[1],a_i[2],a_i[3],0,0,0,0};
        i32x8 bv={b_i[0],b_i[1],b_i[2],b_i[3],0,0,0,0};
        c_acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(av,bv,c_acc,4,4,0,a_spk,0,b_spk);
    }
    for(int j=0;j<4;j++){int mo=mt*16+kpart*4+j,no=nt*16+t_row;
        if(mo<M&&no<N) atomicAdd(&Cfp[mo*N+no],c_acc[j]);}
#endif
}

__global__ void __launch_bounds__(1024)
fp32_to_bf16_zero(float*src,uint16_t*dst,int n){
    int i=blockIdx.x*1024+threadIdx.x;
    if(i<n){float v=src[i];src[i]=0.0f;
        uint32_t fp=__float_as_uint(v);fp+=0x7FFFu+((fp>>16)&1u);dst[i]=(uint16_t)(fp>>16u);}
}

// Raw-pointer ext_splitk launch (2 kernels, ~3µs total vs Triton's ~9.6µs)
void fast_ext_splitk(int64_t A, int64_t Bq, int64_t Bsc, int64_t Cfp, int64_t Cbf,
    int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBq, int64_t BscSN, int64_t ks) {
    int kper=((K/128+(int)ks-1)/(int)ks)*128;
    int aks=(K+kper-1)/kper;
    ext_splitk_16x16_kern<<<dim3(((int)M+15)/16,((int)N+15)/16,aks),64>>>(
        (const uint16_t*)A,(const uint8_t*)Bq,(const uint8_t*)Bsc,(float*)Cfp,
        (int)M,(int)N,(int)K,(int)strA,(int)strBq,(int)BscSN,kper);
    int t=(int)M*(int)N;
    fp32_to_bf16_zero<<<(t+1023)/1024,1024>>>((float*)Cfp,(uint16_t*)Cbf,t);
}

// C++ dispatch for M<=32 K<=1024 (single kernel, minimal overhead)
void dispatch_hip(int64_t A, int64_t Bq, int64_t Bsc, int64_t C, int64_t Cfp,
    int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBq, int64_t BscSN, int64_t ks) {
    gemm_fused_16x16_kern<<<dim3(((int)M+15)/16,((int)N+15)/16),256>>>(
        (const uint16_t*)A,(const uint8_t*)Bq,(const uint8_t*)Bsc,(uint16_t*)C,
        (int)M,(int)N,(int)K,(int)strA,(int)strBq,(int)BscSN);
}

// v78: N-parallel with 32-row M-tiles matching Triton's {1,4} warp layout
void launch_fused_npar(int64_t A, int64_t Bsh, int64_t Bsc, int64_t C,
    int64_t M, int64_t N, int64_t K, int64_t strA, int64_t strBsh, int64_t BscSN) {
    gemm_fused_npar_kern<<<dim3(((int)M+31)/32,((int)N+63)/64),256>>>(
        (const uint16_t*)A,(const uint8_t*)Bsh,(const uint8_t*)Bsc,(uint16_t*)C,
        (int)M,(int)N,(int)K,(int)strA,(int)strBsh,(int)BscSN);
}

// v80: Non-fused launcher reproducing Triton's ISA structure
void launch_nofuse_triton(int64_t Aq, int64_t Bt, int64_t C,
    int64_t Asc, int64_t Bsc,
    int64_t M, int64_t N, int64_t K, int64_t Ks,
    int64_t sam, int64_t sbn, int64_t scm, int64_t sasm) {
    int npn = ((int)N + 63) / 64;
    int npm = ((int)M + 31) / 32;
    int grid = npm * npn;
    gemm_nofuse_triton<<<grid, 256>>>(
        (const uint8_t*)Aq, (const uint8_t*)Bt, (uint16_t*)C,
        (const uint8_t*)Asc, (const uint8_t*)Bsc,
        (int)M, (int)N, (int)K, (int)Ks, (int)sam, (int)sbn, (int)scm, (int)sasm);
}

// ── HIP reduce kernel: replaces Triton _reduce_kernel (saves 11µs launch) ──
// For split-K=8: sum 8 fp32 partials → bf16 output
__global__ void __launch_bounds__(256)
hip_reduce_splitk(float* __restrict__ src, uint16_t* __restrict__ dst,
    int M, int N, int ksplit) {
    int idx = blockIdx.x * blockDim.x + threadIdx.x;
    int total = M * N;
    if (idx >= total) return;
    float sum = 0.0f;
    for (int k = 0; k < ksplit; k++)
        sum += src[k * total + idx];
    // Convert to bf16 and ZERO source for next atomic_add call
    src[idx] = 0.0f;
    uint32_t fp = __float_as_uint(sum);
    fp += 0x7FFFu + ((fp >> 16) & 1u);
    dst[idx] = (uint16_t)(fp >> 16u);
}

void fast_reduce_splitk(int64_t src, int64_t dst, int64_t M, int64_t N, int64_t ksplit) {
    int total = (int)M * (int)N;
    hip_reduce_splitk<<<(total+255)/256, 256>>>((float*)src, (uint16_t*)dst,
        (int)M, (int)N, (int)ksplit);
}

// ── Triton HSACO loader: bypass 11µs Python dispatch → 2.5µs HIP launch ──
static std::unordered_map<std::string, std::pair<hipModule_t, hipFunction_t>> _triton_cache;

void load_triton_hsaco(const std::string& path, const std::string& name) {
    auto key = path + "::" + name;
    if (_triton_cache.count(key)) return;
    hipModule_t mod; hipFunction_t fn;
    if (hipModuleLoad(&mod, path.c_str()) != hipSuccess) return;
    if (hipModuleGetFunction(&fn, mod, name.c_str()) != hipSuccess) return;
    _triton_cache[key] = {mod, fn};
}

// Generic Triton kernel launcher: 4 ptrs + up to 12 int args
void launch_triton_kernel(const std::string& key,
    int64_t p0, int64_t p1, int64_t p2, int64_t p3,
    int64_t a0, int64_t a1, int64_t a2, int64_t a3, int64_t a4,
    int64_t a5, int64_t a6, int64_t a7, int64_t a8, int64_t a9, int64_t a10, int64_t a11,
    int64_t gridX, int64_t gridY, int64_t blockX, int64_t shared_mem) {
    auto it = _triton_cache.find(key);
    if (it == _triton_cache.end()) return;
    void* vp0=(void*)p0; void* vp1=(void*)p1; void* vp2=(void*)p2; void* vp3=(void*)p3;
    int32_t i0=(int32_t)a0,i1=(int32_t)a1,i2=(int32_t)a2,i3=(int32_t)a3,i4=(int32_t)a4;
    int64_t i5=a5,i6=a6,i7=a7,i8=a8,i9=a9,i10=a10,i11=a11;
    void* args[] = {&vp0,&vp1,&vp2,&vp3,&i0,&i1,&i2,&i3,&i4,&i5,&i6,&i7,&i8,&i9,&i10,&i11};
    hipModuleLaunchKernel(it->second.second,
        (unsigned)gridX,(unsigned)gridY,1,(unsigned)blockX,1,1,
        (unsigned)shared_mem,0,args,nullptr);
}

// Launch Triton _gemm_fp4_kernel via cached HSACO with exact 96-byte kernarg layout
// Derived from ISA metadata dump: kernarg_segment_size=96, 5 ptrs + 9 scalars + padding + 2 unused ptrs
void launch_triton_fp4gemm(const std::string& key,
    int64_t Aq, int64_t Bt, int64_t C, int64_t Asc, int64_t Bsc,
    int64_t M, int64_t N, int64_t K, int64_t Ks,
    int64_t sam, int64_t sbn, int64_t sck, int64_t scm, int64_t sasm,
    int64_t gridX, int64_t shared_mem) {
    auto it = _triton_cache.find(key);
    if (it == _triton_cache.end()) return;
    // Pack 96-byte kernarg buffer matching Triton's compiled layout
    struct __attribute__((packed)) {
        void* a_ptr;      // 0
        void* b_ptr;      // 8
        void* c_ptr;      // 16
        void* asc_ptr;    // 24
        void* bsc_ptr;    // 32
        int32_t M;        // 40
        int32_t N;        // 44
        int32_t K;        // 48
        int32_t Ks;       // 52
        int32_t sam;      // 56
        int32_t sbn;      // 60
        int32_t sck;      // 64
        int32_t scm;      // 68
        int32_t sasm;     // 72
        int32_t _pad;     // 76
        void* _ptr0;      // 80  (unused by kernel, but kernarg_size=96)
        void* _ptr1;      // 88
    } args;
    args.a_ptr = (void*)Aq;
    args.b_ptr = (void*)Bt;
    args.c_ptr = (void*)C;
    args.asc_ptr = (void*)Asc;
    args.bsc_ptr = (void*)Bsc;
    args.M = (int32_t)M;
    args.N = (int32_t)N;
    args.K = (int32_t)K;
    args.Ks = (int32_t)Ks;
    args.sam = (int32_t)sam;
    args.sbn = (int32_t)sbn;
    args.sck = (int32_t)sck;
    args.scm = (int32_t)scm;
    args.sasm = (int32_t)sasm;
    args._pad = 0;
    args._ptr0 = nullptr;
    args._ptr1 = nullptr;
    size_t sz = sizeof(args);
    void* cfg[] = {HIP_LAUNCH_PARAM_BUFFER_POINTER, &args,
                   HIP_LAUNCH_PARAM_BUFFER_SIZE, &sz,
                   HIP_LAUNCH_PARAM_END};
    hipModuleLaunchKernel(it->second.second,
        (unsigned)gridX, 1, 1, 256, 1, 1,
        (unsigned)shared_mem, 0, nullptr, (void**)cfg);
}
"""

@functools.lru_cache(maxsize=1)
def _hip_module():
    d = hashlib.sha1(HIP_SRC.encode()).hexdigest()[:16]
    return load_inline(name=f"hip30_{d}", cpp_sources=[CPP_SRC], cuda_sources=[HIP_SRC],
        functions=["quantize_a_hip","gemm_fused_16x16_hip","fast_hip16","fast_ext_splitk",
                   "dispatch_hip","fast_reduce_splitk","load_triton_hsaco","launch_triton_kernel",
                   "launch_triton_fp4gemm","launch_nofuse_triton",
                   "quant_a_shuffled","quant_and_asm","warmup_asm","launch_fused_npar"],
        extra_cflags=["-O3"],
        extra_cuda_cflags=["-O3","--offload-arch=gfx950","-std=c++17"],
        with_cuda=True, verbose=False)


# ═══════════════════════════════════════════════════════════════════
# Triton kernels: fused quant+GEMM (M≤64) + standard GEMM (fallback)
# ═══════════════════════════════════════════════════════════════════

@triton.jit
def _mxfp4_quant_inline(x, BLOCK_SIZE_K: tl.constexpr, BLOCK_SIZE_M: tl.constexpr,
                          MXFP4_QUANT_BLOCK_SIZE: tl.constexpr):
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_K // 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)
    se = tl.log2(amax).floor() - 2
    se = tl.clamp(se, min=-127, max=127)
    bs = se.to(tl.uint8) + 127
    qs = tl.exp2(-se)
    qx = x * qs
    qx = qx.to(tl.uint32, bitcast=True)
    s = qx & 0x80000000
    qx = qx ^ s
    qf = qx.to(tl.float32, bitcast=True)
    sat = qf >= 6.0
    den = (not sat) & (qf < 1.0)
    nor = not (sat | den)
    dexp: tl.constexpr = (127-1)+(23-1)+1
    dmi: tl.constexpr = dexp << 23
    dmf: tl.constexpr = tl.cast(dmi, tl.float32, bitcast=True)
    dx = qf + dmf
    dx = dx.to(tl.uint32, bitcast=True)
    dx -= dmi
    dx = dx.to(tl.uint8)
    nx = qx
    mo = (nx >> 22) & 1
    nx += 0xC11FFFFF  # ((1-127)<<23) + (1<<21) - 1 as uint32
    nx += mo
    nx = nx >> 22
    nx = nx.to(tl.uint8)
    v = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
    v = tl.where(nor, nx, v)
    v = tl.where(den, dx, v)
    sl = s >> (23+8-1-2)
    sl = sl.to(tl.uint8)
    v = v | sl
    v = tl.reshape(v, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2])
    evens, odds = tl.split(v)
    xfp4 = evens | (odds << 4)
    return xfp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // 2), bs.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)

@triton.jit
def remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
    ppx = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
    tall = GRID_MN % NUM_XCDS
    tall = NUM_XCDS if tall == 0 else tall
    xcd = pid % NUM_XCDS
    lp = pid // NUM_XCDS
    if xcd < tall: pid = xcd * ppx + lp
    else: pid = tall * ppx + (xcd - tall) * (ppx - 1) + lp
    return pid

@triton.jit
def pid_grid(pid, npm, npn, GROUP_SIZE_M: tl.constexpr = 1):
    if GROUP_SIZE_M == 1:
        return pid // npn, pid % npn
    else:
        nig = GROUP_SIZE_M * npn
        gid = pid // nig
        fm = gid * GROUP_SIZE_M
        gsm = min(npm - fm, GROUP_SIZE_M)
        tl.assume(gsm >= 0)
        return fm + (pid % gsm), (pid % nig) // gsm

@triton.jit
def _bsc_off(row, col, Ks):
    return ((row//32)*(Ks*32) + (col//8)*256 + (col%4)*64
            + (row%16)*4 + ((col//4)%2)*2 + ((row//16)%2))

@triton.heuristics({"EVEN_K": lambda a: (a["K"]%(a["BLOCK_SIZE_K"]//2)==0)
    and (a["SPLITK_BLOCK_SIZE"]%a["BLOCK_SIZE_K"]==0)
    and (a["K"]%(a["SPLITK_BLOCK_SIZE"]//2)==0)})
@triton.jit
def _fused_gemm_fp4_kernel(
    a_ptr, b_ptr, c_ptr, bsc_ptr,
    M, N, K, actual_K, Ks,
    sam, sak, sbk, sbn, sck, scm, scn,
    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):
    tl.assume(sam>0);tl.assume(sak>0);tl.assume(sbk>0);tl.assume(sbn>0);tl.assume(scm>0);tl.assume(scn>0)
    GMN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
    pu = tl.program_id(0)
    pu = remap_xcd(pu, GMN * NUM_KSPLIT, NUM_XCDS=8)
    pk = pu % NUM_KSPLIT
    pid = pu // NUM_KSPLIT
    npm = tl.cdiv(M, BLOCK_SIZE_M); npn = tl.cdiv(N, BLOCK_SIZE_N)
    if NUM_KSPLIT == 1: pm, pn = pid_grid(pid, npm, npn, GROUP_SIZE_M=GROUP_SIZE_M)
    else: pm = pid // npn; pn = pid % npn
    tl.assume(pm>=0);tl.assume(pn>=0)
    SPK: tl.constexpr = BLOCK_SIZE_K // 32
    if (pk * SPLITK_BLOCK_SIZE // 2) < K:
        nki = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
        okp = tl.arange(0, BLOCK_SIZE_K // 2)
        okps = pk * (SPLITK_BLOCK_SIZE // 2) + okp
        oka = tl.arange(0, BLOCK_SIZE_K)
        okas = pk * SPLITK_BLOCK_SIZE + oka
        oam = (pm * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        obn = (pn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
        ap = a_ptr + (oam[:, None] * sam + okas[None, :] * sak)
        bp = b_ptr + (okps[:, None] * sbk + obn[None, :] * sbn)
        ksb = pk * (SPLITK_BLOCK_SIZE // 32)
        oksl = tl.arange(0, SPK)
        acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
        for ki in range(pk * nki, (pk + 1) * nki):
            if EVEN_K: abf = tl.load(ap)
            else: abf = tl.load(ap, mask=oka[None, :] < actual_K - ki * BLOCK_SIZE_K, other=0.0)
            af32 = abf.to(tl.float32)
            afp4, asc = _mxfp4_quant_inline(af32, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
            coks = ksb + oksl
            bso = _bsc_off(obn[:, None], coks[None, :], Ks)
            bscv = tl.load(bsc_ptr + bso, cache_modifier=".cg")
            if EVEN_K: b = tl.load(bp, cache_modifier=".cg")
            else: b = tl.load(bp, mask=okp[:, None] < K - ki * (BLOCK_SIZE_K // 2), other=0, cache_modifier=".cg")
            acc = tl.dot_scaled(afp4, asc, "e2m1", b, bscv, "e2m1", acc)
            ap += BLOCK_SIZE_K * sak
            bp += (BLOCK_SIZE_K // 2) * sbk
            ksb += SPK
        ocm = pm * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
        ocn = pn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
        c_mask = (ocm[:, None] < M) & (ocn[None, :] < N)
        c = acc.to(c_ptr.type.element_ty)
        cp = c_ptr + scm * ocm[:, None] + scn * ocn[None, :] + pk * sck
        tl.store(cp, c, mask=c_mask)

@triton.heuristics({"EVEN_K": lambda a: (a["K"]%(a["BLOCK_SIZE_K"]//2)==0)
    and (a["SPLITK_BLOCK_SIZE"]%a["BLOCK_SIZE_K"]==0)
    and (a["K"]%(a["SPLITK_BLOCK_SIZE"]//2)==0)})
@triton.jit
def _gemm_fp4_kernel(
    a_ptr, b_ptr, c_ptr, asc_ptr, bsc_ptr,
    M, N, K, Ks,
    sam, sak, sbk, sbn, sck, scm, scn, sasm, sask,
    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):
    tl.assume(sam>0);tl.assume(sak>0);tl.assume(sbk>0);tl.assume(sbn>0);tl.assume(scm>0);tl.assume(scn>0)
    tl.assume(sasm>0);tl.assume(sask>0)
    GMN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
    pu = tl.program_id(0)
    pu = remap_xcd(pu, GMN * NUM_KSPLIT, NUM_XCDS=8)
    pk = pu % NUM_KSPLIT
    pid = pu // NUM_KSPLIT
    npm = tl.cdiv(M, BLOCK_SIZE_M); npn = tl.cdiv(N, BLOCK_SIZE_N)
    if NUM_KSPLIT == 1: pm, pn = pid_grid(pid, npm, npn, GROUP_SIZE_M=GROUP_SIZE_M)
    else: pm = pid // npn; pn = pid % npn
    tl.assume(pm>=0);tl.assume(pn>=0)
    SPK: tl.constexpr = BLOCK_SIZE_K // 32
    if (pk * SPLITK_BLOCK_SIZE // 2) < K:
        nki = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
        ok = tl.arange(0, BLOCK_SIZE_K // 2)
        oks = pk * (SPLITK_BLOCK_SIZE // 2) + ok
        oam = (pm * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        obn = (pn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
        ap = a_ptr + (oam[:, None] * sam + oks[None, :] * sak)
        bp = b_ptr + (oks[:, None] * sbk + obn[None, :] * sbn)
        oaks = (pk * (SPLITK_BLOCK_SIZE // 32)) + tl.arange(0, SPK)
        ascp = asc_ptr + oam[:, None] * sasm + oaks[None, :] * sask
        ksb = pk * (SPLITK_BLOCK_SIZE // 32)
        oksl = tl.arange(0, SPK)
        acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
        for k in range(pk * nki, (pk + 1) * nki):
            ascv = tl.load(ascp)
            coks = ksb + oksl
            bso = _bsc_off(obn[:, None], coks[None, :], Ks)
            bscv = tl.load(bsc_ptr + bso, cache_modifier=".cg")
            if EVEN_K: a = tl.load(ap); b = tl.load(bp, cache_modifier=".cg")
            else:
                a = tl.load(ap, mask=ok[None, :] < K - k * (BLOCK_SIZE_K // 2), other=0)
                b = tl.load(bp, mask=ok[:, None] < K - k * (BLOCK_SIZE_K // 2), other=0, cache_modifier=".cg")
            acc = tl.dot_scaled(a, ascv, "e2m1", b, bscv, "e2m1", acc)
            ap += (BLOCK_SIZE_K // 2) * sak
            bp += (BLOCK_SIZE_K // 2) * sbk
            ascp += SPK * sask
            ksb += SPK
        c = acc.to(c_ptr.type.element_ty)
        ocm = pm * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
        ocn = pn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
        cp = c_ptr + scm * ocm[:, None] + scn * ocn[None, :] + pk * sck
        tl.store(cp, c, mask=(ocm[:, None] < M) & (ocn[None, :] < N))

@triton.jit
def _reduce_kernel(ci, co, M, N, sick, sicm, sicn, socm, socn,
    BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
    ACTUAL_KSPLIT: tl.constexpr, MAX_KSPLIT: tl.constexpr):
    pm = tl.program_id(0); pn = tl.program_id(1)
    om = (pm*BLOCK_SIZE_M+tl.arange(0,BLOCK_SIZE_M))%M
    on = (pn*BLOCK_SIZE_N+tl.arange(0,BLOCK_SIZE_N))%N
    ok = tl.arange(0,MAX_KSPLIT)
    ip = ci + ok[:,None,None]*sick + om[None,:,None]*sicm + on[None,None,:]*sicn
    if ACTUAL_KSPLIT==MAX_KSPLIT: c = tl.load(ip)
    else: c = tl.load(ip, mask=ok[:,None,None]<ACTUAL_KSPLIT)
    c = tl.sum(c, axis=0).to(co.type.element_ty)
    op = co + om[:,None]*socm + on[None,:]*socn
    tl.store(op, c)


# ═══════════════════════════════════════════════════════════════════
# Per-shape configs
# ═══════════════════════════════════════════════════════════════════

SHAPE_CONFIGS = {
    (4, 2880, 512): dict(BLOCK_SIZE_M=16,BLOCK_SIZE_N=64,BLOCK_SIZE_K=512,GROUP_SIZE_M=1,
                         num_warps=4,num_stages=1,waves_per_eu=4,matrix_instr_nonkdim=16,NUM_KSPLIT=1),
    # v125: Shape 2 num_warps=8 — matches CDNA4 blog 8-wave scheduling
    (16, 2112, 7168): dict(BLOCK_SIZE_M=16,BLOCK_SIZE_N=128,BLOCK_SIZE_K=256,GROUP_SIZE_M=1,
                           num_warps=8,num_stages=2,waves_per_eu=1,matrix_instr_nonkdim=16,NUM_KSPLIT=16),
    (32, 4096, 512): dict(BLOCK_SIZE_M=32,BLOCK_SIZE_N=64,BLOCK_SIZE_K=512,GROUP_SIZE_M=1,
                          num_warps=4,num_stages=1,waves_per_eu=4,matrix_instr_nonkdim=16,NUM_KSPLIT=1),
    (32, 2880, 512): dict(BLOCK_SIZE_M=32,BLOCK_SIZE_N=64,BLOCK_SIZE_K=512,GROUP_SIZE_M=1,
                          num_warps=4,num_stages=1,waves_per_eu=4,matrix_instr_nonkdim=16,NUM_KSPLIT=1),
    # Revert shape 5 to baseline (8 warps failed correctness)
    (64, 7168, 2048): dict(BLOCK_SIZE_M=32,BLOCK_SIZE_N=64,BLOCK_SIZE_K=512,GROUP_SIZE_M=4,
                           num_warps=4,num_stages=5,waves_per_eu=1,matrix_instr_nonkdim=16,NUM_KSPLIT=1),
}

CONFIGS = {
    4: dict(BLOCK_SIZE_M=16,BLOCK_SIZE_N=128,BLOCK_SIZE_K=256,GROUP_SIZE_M=1,
            num_warps=4,num_stages=2,waves_per_eu=3,matrix_instr_nonkdim=16,NUM_KSPLIT=16),
    16: dict(BLOCK_SIZE_M=16,BLOCK_SIZE_N=128,BLOCK_SIZE_K=256,GROUP_SIZE_M=1,
             num_warps=4,num_stages=2,waves_per_eu=3,matrix_instr_nonkdim=16,NUM_KSPLIT=16),
    32: dict(BLOCK_SIZE_M=32,BLOCK_SIZE_N=128,BLOCK_SIZE_K=256,GROUP_SIZE_M=1,
             num_warps=4,num_stages=2,waves_per_eu=3,matrix_instr_nonkdim=16,NUM_KSPLIT=1),
    64: dict(BLOCK_SIZE_M=64,BLOCK_SIZE_N=256,BLOCK_SIZE_K=256,GROUP_SIZE_M=1,
             num_warps=4,num_stages=3,waves_per_eu=2,matrix_instr_nonkdim=32,NUM_KSPLIT=1),
    256: dict(BLOCK_SIZE_M=128,BLOCK_SIZE_N=256,BLOCK_SIZE_K=256,GROUP_SIZE_M=2,
              num_warps=4,num_stages=3,waves_per_eu=2,matrix_instr_nonkdim=32,NUM_KSPLIT=1),
}

def get_config(M, N=None, K=None):
    if N is not None and K is not None:
        key = (M, N, K)
        if key in SHAPE_CONFIGS: return SHAPE_CONFIGS[key].copy()
    for t in sorted(CONFIGS.keys()):
        if M <= t: return CONFIGS[t].copy()
    return CONFIGS[256].copy()

def get_splitk(K, BSK, NSK):
    SPBS = triton.cdiv(2 * triton.cdiv(K, NSK), BSK) * BSK
    while NSK > 1 and BSK > 16:
        if K % (SPBS // 2) == 0 and SPBS % BSK == 0 and K % (BSK // 2) == 0: break
        elif K % (SPBS // 2) != 0 and NSK > 1: NSK //= 2
        elif SPBS % BSK != 0:
            if NSK > 1: NSK //= 2
            elif BSK > 16: BSK //= 2
        elif K % (BSK // 2) != 0 and BSK > 16: BSK //= 2
        else: break
        SPBS = triton.cdiv(2 * triton.cdiv(K, NSK), BSK) * BSK
    NSK = triton.cdiv(K, SPBS // 2)
    return SPBS, BSK, NSK


# ═══════════════════════════════════════════════════════════════════
# Hybrid dispatch
# ═══════════════════════════════════════════════════════════════════

ASM_BASE = "/home/runner/aiter/hsa/gfx950/f4gemm/"
ASM_CONFIGS_SHAPE = {
    # Per-shape best ASM tile (from discovery: 26 tiles available)
    # 64x128 for M=64: perfect M-tile fit, no padding. 128x128 for M=256 was worse (15.1 vs 12.5)
    (64, 7168, 2048): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",
                        "f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128.co", 0),
}
ASM_DEFAULT = ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
               "f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co", 0)
FUSED_M_THRESHOLD = 32  # Fused quant M≤32 ONLY. M=64 fused causes 20µs regression even with BLOCK_M=16.
# v31: Removed K threshold. HIP gemm_fused_16x16_kern handles ANY K (single kernel, LDS reduction).
# Key insight: For M=16 K=7168, Triton uses 2 launches (8µs overhead) vs HIP 1 launch (2µs overhead).
# GPU self-time is only 2.2µs — overhead dominates. Single HIP kernel: ~6µs vs Triton's ~12µs.

_OUT_CACHE = {}
_AQ_CACHE = {}
_HIP16_OUT = {}
_ASM_WS = {}
_warmed = False
_hip = None
_triton_hsaco_keys = {}  # {shape: (key, grid, block, shared_mem, num_args)}

def _warm_triton_hsaco(hip_mod):
    """Trigger Triton JIT for shape 5, find cached .hsaco, load via hipModule."""
    import glob as _glob
    # Shape 5: M=64, N=7168, K=2048 — trigger compilation with dummy data
    dev = torch.device('cuda')
    m, n, k = 64, 7168, 2048
    Kp = k >> 1
    Aq = torch.zeros(m, k//2, dtype=torch.uint8, device=dev)
    Asc = torch.zeros(m, k//32, dtype=torch.uint8, device=dev)
    Bq8 = torch.zeros(n, k//2, dtype=torch.uint8, device=dev)
    Bt = Bq8.T
    Bsc8 = torch.zeros(n, k//32, dtype=torch.uint8, device=dev)  # approximate
    out = torch.zeros(m, n, dtype=torch.bfloat16, device=dev)
    config = SHAPE_CONFIGS.get((64, 7168, 2048), get_config(64, 7168, 2048))
    config["SPLITK_BLOCK_SIZE"] = 2 * Kp
    grid_size = ((m + config["BLOCK_SIZE_M"] - 1) // config["BLOCK_SIZE_M"]) * \
                ((n + config["BLOCK_SIZE_N"] - 1) // config["BLOCK_SIZE_N"])
    try:
        _gemm_fp4_kernel[(grid_size,)](
            Aq, Bt, out, Asc, Bsc8,
            m, n, Kp, Bsc8.shape[1],
            Aq.stride(0), Aq.stride(1), Bt.stride(0), Bt.stride(1),
            0, out.stride(0), out.stride(1), Asc.stride(0), Asc.stride(1),
            SPLITK_BLOCK_SIZE=config["SPLITK_BLOCK_SIZE"],
            **{k2: config[k2] for k2 in ["BLOCK_SIZE_M","BLOCK_SIZE_N","BLOCK_SIZE_K",
               "GROUP_SIZE_M","NUM_KSPLIT","num_warps","num_stages","waves_per_eu","matrix_instr_nonkdim"]})
        torch.cuda.synchronize()
    except Exception as e:
        print(f"[HSACO] Triton warmup failed: {e}")
        return
    # Scan Triton cache for the compiled .hsaco
    cache_base = os.path.expanduser("~/.triton/cache")
    hsaco_path = None
    for root, _, files in os.walk(cache_base):
        for f in files:
            if f == '_gemm_fp4_kernel.hsaco':
                hsaco_path = os.path.join(root, f)
                break
        if hsaco_path:
            break
    if hsaco_path:
        hsaco_key = hsaco_path + "::_gemm_fp4_kernel"
        hip_mod.load_triton_hsaco(hsaco_path, "_gemm_fp4_kernel")
        custom_kernel._hsaco_key = hsaco_key
        print(f"[HSACO] Loaded: {hsaco_path}")
    else:
        print("[HSACO] .hsaco not found in cache")
    # Also compile .amdgcn → .co (ISA reproduction)
    amdgcn_path = None
    for root2, _, files2 in os.walk(cache_base):
        for f2 in files2:
            if f2 == '_gemm_fp4_kernel.amdgcn':
                amdgcn_path = os.path.join(root2, f2)
                break
        if amdgcn_path: break
    if amdgcn_path:
        import subprocess as _sp
        co_path = "/tmp/_gemm_fp4_isa.co"
        obj_path = "/tmp/_gemm_fp4_isa.o"
        # Find llvm-mc (try common ROCm paths)
        llvm_mc = None
        for p in ["/opt/rocm/llvm/bin/llvm-mc", "/usr/bin/llvm-mc",
                  "/opt/rocm/lib/llvm/bin/llvm-mc"]:
            if os.path.exists(p):
                llvm_mc = p; break
        lld = None
        for p in ["/opt/rocm/llvm/bin/ld.lld", "/usr/bin/ld.lld",
                  "/opt/rocm/lib/llvm/bin/ld.lld"]:
            if os.path.exists(p):
                lld = p; break
        if llvm_mc and lld:
            r1 = _sp.run([llvm_mc, "-triple=amdgcn-amd-amdhsa", "-mcpu=gfx950",
                          "-filetype=obj", amdgcn_path, "-o", obj_path],
                         capture_output=True, text=True)
            if r1.returncode == 0:
                r2 = _sp.run([lld, "-shared", obj_path, "-o", co_path],
                             capture_output=True, text=True)
                if r2.returncode == 0:
                    co_key = co_path + "::_gemm_fp4_kernel"
                    hip_mod.load_triton_hsaco(co_path, "_gemm_fp4_kernel")
                    custom_kernel._isa_co_key = co_key
                    print(f"[ISA-CO] Compiled & loaded: {co_path}")
                else:
                    print(f"[ISA-CO] ld.lld failed: {r2.stderr[:200]}")
            else:
                print(f"[ISA-CO] llvm-mc failed: {r1.stderr[:200]}")
        else:
            print(f"[ISA-CO] llvm-mc={llvm_mc} lld={lld} — tools not found")
    # Free dummy tensors
    del Aq, Asc, Bq8, Bt, Bsc8, out

def _get_aq(m, k, dev):
    key = (m, k, dev)
    if key not in _AQ_CACHE:
        _AQ_CACHE[key] = (torch.empty((m, k//2), dtype=torch.uint8, device=dev),
                           torch.empty((m, k//SCALE_GROUP_SIZE), dtype=torch.uint8, device=dev))
    return _AQ_CACHE[key]

def custom_kernel(data: input_t) -> output_t:
    global _warmed, _hip
    if not _warmed:
        _hip = _hip_module()
        knl, co, _ = ASM_DEFAULT
        _hip.warmup_asm(ASM_BASE + co, knl)
        for _, v in ASM_CONFIGS_SHAPE.items():
            _hip.warmup_asm(ASM_BASE + v[1], v[0])
        _warm_triton_hsaco(_hip)
        _warmed = True
    A = data[0]
    m, k = A.shape
    n = data[2].shape[0]  # B_q.shape[0]
    Kp = k >> 1
    hip = _hip
    B_q = data[2]
    B_shuffle = data[3]
    B_scale_sh = data[4]

    # ── Path 1: HIP fused 16x16 for M≤32 K≤1024 (single kernel, LDS K-reduction) ──
    # v31 REVERTED: HIP fused for large K regressed 12→20.7µs (uncoalesced B access).
    # Triton tl.dot_scaled uses coalesced tiled loads — must keep Triton for K>1024.
    if m <= 32 and k <= 1024:
        key_hip = (m, n)
        if key_hip not in _HIP16_OUT:
            _HIP16_OUT[key_hip] = (
                torch.empty(m, n, dtype=torch.bfloat16, device=A.device),
                int(B_q.stride(0) * B_q.element_size()),  # cache strBq
                int(B_scale_sh.size(1)),  # cache BscSN
            )
        out, strBq, BscSN = _HIP16_OUT[key_hip]
        hip.dispatch_hip(A.data_ptr(), B_q.data_ptr(), B_scale_sh.data_ptr(),
                         out.data_ptr(), 0, m, n, k, A.stride(0), strBq, BscSN, 0)
        return out

    # ── Path 1b: v78 N-parallel DISABLED — inline quant too slow for K>1024 (32.3µs vs 12.6µs Triton)
    # Need non-fused version (pre-quantized A) to eliminate quant overhead. See v78 results.
    if False:  # 32 < m <= 127 and n % 64 == 0:
        key_np = (m, n)
        if key_np not in _HIP16_OUT:
            out = torch.empty(m, n, dtype=torch.bfloat16, device=A.device)
            _HIP16_OUT[key_np] = (out, int(B_shuffle.stride(0) * B_shuffle.element_size()), int(B_scale_sh.size(1)))
        out, strBsh, BscSN = _HIP16_OUT[key_np]
        # Pass B_shuffle (coalesced [N/16][K/32][16][16]) instead of B_q (strided [N, K/2])
        hip.launch_fused_npar(A.data_ptr(), B_shuffle.data_ptr(), B_scale_sh.data_ptr(),
                              out.data_ptr(), m, n, k, A.stride(0), strBsh, BscSN)
        return out

    # ── Path 2: HIP ASM for M≥128 (hand-tuned 186% MFMA util, beats Triton) ──
    if m >= 128:
        key_asm = (m, n, k)
        if key_asm not in _ASM_WS:
            # Parse tile_m from kernel name for correct padding
            knl_tmp, _, _ = ASM_CONFIGS_SHAPE.get((m, n, k), ASM_DEFAULT)
            tile_m = 32  # default
            xpos = knl_tmp.rfind('x')
            if xpos > 0:
                i = xpos - 1
                while i > 0 and knl_tmp[i-1].isdigit(): i -= 1
                tile_m = int(knl_tmp[i:xpos])
            mp = ((m + tile_m - 1) // tile_m) * tile_m
            sm = mp
            _ASM_WS[key_asm] = {
                'aq': torch.empty(m, k // 2, dtype=torch.uint8, device=A.device),
                'asc': torch.empty(sm, k // 32, dtype=torch.uint8, device=A.device),
                'out': torch.empty(mp, n, dtype=torch.bfloat16, device=A.device),
                'sm': sm, 'mp': mp,
            }
        w = _ASM_WS[key_asm]
        knl, co, sk = ASM_CONFIGS_SHAPE.get((m, n, k), ASM_DEFAULT)
        hip.quant_and_asm(A, B_shuffle, B_scale_sh, n,
                          w['aq'], w['asc'], w['out'], w['sm'],
                          knl, ASM_BASE + co, sk)
        return w['out'][:m] if w['mp'] > m else w['out']

    # ── Path 2b: Launch Triton's ISA compiled to .co (shape 5 ONLY) ──
    if m == 64 and n == 7168 and k == 2048 and hasattr(custom_kernel, '_isa_co_key'):
        Aq, Asc = _get_aq(m, k, A.device)
        hip.quantize_a_hip(A, Aq, Asc)
        Bq8 = B_q.view(torch.uint8) if B_q.dtype != torch.uint8 else B_q
        Bt = Bq8.T
        Bsc8 = B_scale_sh.view(torch.uint8) if B_scale_sh.dtype != torch.uint8 else B_scale_sh
        key_out = ('isa_co', m, n)
        if key_out not in _OUT_CACHE:
            _OUT_CACHE[key_out] = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        out = _OUT_CACHE[key_out]
        Kp = k >> 1
        grid_size = ((m + 31) // 32) * ((n + 63) // 64)
        hip.launch_triton_fp4gemm(
            custom_kernel._isa_co_key,
            Aq.data_ptr(), Bt.data_ptr(), out.data_ptr(), Asc.data_ptr(), Bsc8.data_ptr(),
            m, n, Kp, Bsc8.shape[1],
            Aq.stride(0), Bt.stride(1), 0, out.stride(0), Asc.stride(0),
            grid_size, 104448)
        return out

    # ── Path 3: Triton GEMM (M=16-64 with K>1024, or M=33-127) ──
    Bq8 = B_q.view(torch.uint8) if B_q.dtype != torch.uint8 else B_q
    Bt = Bq8.T
    Bsc8 = B_scale_sh.view(torch.uint8) if B_scale_sh.dtype != torch.uint8 else B_scale_sh
    Ks_stride = Bsc8.shape[1]

    use_fused = (m <= FUSED_M_THRESHOLD)

    if not use_fused:
        Aq, Asc = _get_aq(m, k, A.device)
        hip.quantize_a_hip(A, Aq, Asc)

    config = get_config(m, n, k)

    if config["NUM_KSPLIT"] > 1:
        SPBS, BSK, NSK = get_splitk(Kp, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"])
        config["SPLITK_BLOCK_SIZE"] = SPBS
        config["BLOCK_SIZE_K"] = BSK
        config["NUM_KSPLIT"] = NSK
    else:
        config["SPLITK_BLOCK_SIZE"] = 2 * Kp

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

    config["BLOCK_SIZE_K"] = max(config["BLOCK_SIZE_K"], 128)
    NSK = config["NUM_KSPLIT"]

    out_key = (m, n, NSK, A.device)
    cached = _OUT_CACHE.get(out_key)
    if cached is not None:
        if NSK > 1: y_pp, y = cached
        else: y = cached; y_pp = None
    elif NSK > 1:
        y_pp = torch.empty((NSK, m, n), dtype=torch.float32, device=A.device)
        y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        _OUT_CACHE[out_key] = (y_pp, y)
    else:
        y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        _OUT_CACHE[out_key] = y
        y_pp = None

    out_t = y if NSK == 1 else y_pp
    grid = lambda META: (META["NUM_KSPLIT"] * triton.cdiv(m, META["BLOCK_SIZE_M"]) * triton.cdiv(n, META["BLOCK_SIZE_N"]),)

    if use_fused:
        _fused_gemm_fp4_kernel[grid](
            A, Bt, out_t, Bsc8,
            m, n, Kp, k, Ks_stride,
            A.stride(0), A.stride(1), Bt.stride(0), Bt.stride(1),
            0 if NSK == 1 else y_pp.stride(0), out_t.stride(-2), out_t.stride(-1),
            SPLITK_BLOCK_SIZE=config["SPLITK_BLOCK_SIZE"],
            **{k2: config[k2] for k2 in ["BLOCK_SIZE_M","BLOCK_SIZE_N","BLOCK_SIZE_K",
               "GROUP_SIZE_M","NUM_KSPLIT","num_warps","num_stages","waves_per_eu","matrix_instr_nonkdim"]})
    else:
        _gemm_fp4_kernel[grid](
            Aq, Bt, out_t, Asc, Bsc8,
            m, n, Kp, Ks_stride,
            Aq.stride(0), Aq.stride(1), Bt.stride(0), Bt.stride(1),
            0 if NSK == 1 else y_pp.stride(0), out_t.stride(-2), out_t.stride(-1),
            Asc.stride(0), Asc.stride(1),
            SPLITK_BLOCK_SIZE=config["SPLITK_BLOCK_SIZE"],
            **{k2: config[k2] for k2 in ["BLOCK_SIZE_M","BLOCK_SIZE_N","BLOCK_SIZE_K",
               "GROUP_SIZE_M","NUM_KSPLIT","num_warps","num_stages","waves_per_eu","matrix_instr_nonkdim"]})

    if NSK > 1:
        RBLM, RBLN = 16, 16
        ANSK = triton.cdiv(Kp, config["SPLITK_BLOCK_SIZE"] // 2)
        _reduce_kernel[(triton.cdiv(m, RBLM), triton.cdiv(n, RBLN))](
            y_pp, y, m, n,
            y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
            y.stride(0), y.stride(1),
            BLOCK_SIZE_M=RBLM, BLOCK_SIZE_N=RBLN,
            ACTUAL_KSPLIT=ANSK, MAX_KSPLIT=triton.next_power_of_2(NSK))

    return y


# ═══════════════════════════════════════════════════════════════════
# Profiling (DO_PROFILE=0 by default, set to 1 to enable)
# ═══════════════════════════════════════════════════════════════════

_COUNTER_GROUPS = [
    "TCC_HIT_sum TCC_MISS_sum TCC_EA0_RDREQ_sum TCC_EA0_WRREQ_sum",
    "SQ_LDS_BANK_CONFLICT SQ_INSTS_MFMA SQ_INSTS_VALU SQ_BUSY_CYCLES",
    "SQ_VALU_MFMA_BUSY_CYCLES SQ_INSTS_VMEM_RD SQ_WAIT_ANY SQ_ACTIVE_INST_VMEM",
]
_SHAPES = [(4,2880,512),(16,2112,7168),(32,4096,512),(32,2880,512),(64,7168,2048),(256,3072,1536)]

def _make_profile_inputs():
    try:
        from reference import generate_input
        return [generate_input(m=M, n=N, k=K, seed=42) for M, N, K in _SHAPES]
    except Exception as e:
        print(f"[PROFILE] generate_input failed: {e}"); return []

def _run_profiling_kernels():
    inputs = _make_profile_inputs()
    if not inputs: return
    for d in inputs: custom_kernel(d)
    torch.cuda.synchronize()
    for i, d in enumerate(inputs):
        print(f"[PROFILE] Shape M={_SHAPES[i][0]} N={_SHAPES[i][1]} K={_SHAPES[i][2]}")
        custom_kernel(d); torch.cuda.synchronize()

def _parse_results(outdir):
    dispatches = {}
    _OUR = ('splitk','gemm_fused_16x16','fp32_to_bf16','quant_a_kern','f4gemm',
            '_fused_gemm_fp4','_gemm_fp4','_reduce_kernel','quantize_a_kernel')
    for cf in sorted(glob.glob(os.path.join(outdir, "**/*counter_collection.csv"), recursive=True)):
        with open(cf) as f:
            for row in csv.DictReader(f):
                did, kn, cn, cv = row.get("Dispatch_Id",""), row.get("Kernel_Name",""), row.get("Counter_Name",""), row.get("Counter_Value","0")
                if not kn or not any(k in kn for k in _OUR): continue
                if did not in dispatches:
                    dispatches[did] = {'_k':kn,'_g':row.get("Grid_Size",""),'_b':row.get("Workgroup_Size",""),
                        '_v':row.get("VGPR_Count",""),'_a':row.get("Accum_VGPR_Count",""),'_s':row.get("Scratch_Size",""),'_l':row.get("LDS_Block_Size","")}
                dispatches[did][cn] = cv
    print("\n" + "="*80); print("ROCPROFV3 HARDWARE COUNTER RESULTS"); print("="*80)
    for did in sorted(dispatches.keys(), key=lambda x: int(x)):
        d = dispatches[did]
        short = d['_k'].split('::')[-1][:50] if '::' in d['_k'] else d['_k'][:50]
        print(f"\n--- Dispatch {did}: {short}")
        print(f"    Grid={d['_g']} Block={d['_b']} VGPR={d['_v']} AGPR={d['_a']} Scratch={d['_s']} LDS={d['_l']}")
        for k2 in sorted(d.keys()):
            if not k2.startswith('_'): print(f"    {k2:40s} = {d[k2]}")
    sys.stdout.flush()

def _run_profiling():
    print("[PROFILE] Starting rocprofv3...")
    tmpdir = tempfile.mkdtemp(prefix="rp_")
    cf = os.path.join(tmpdir, "counters.txt")
    with open(cf, "w") as f:
        for g in _COUNTER_GROUPS: f.write(f"pmc: {g}\n")
    runner = os.path.join(tmpdir, "runner.py")
    shutil.copy(os.path.abspath(__file__), runner)
    outdir = os.path.join(tmpdir, "out"); os.makedirs(outdir)
    env = os.environ.copy()
    env.update(PROFILE_SUBPROCESS="1", DO_PROFILE="0", ORIG_DIR=_ORIG_DIR)
    cmd = ["rocprofv3","--input",cf,"--output-directory",outdir,"--output-format","csv","--",sys.executable,runner]
    try:
        r = subprocess.run(cmd, env=env, capture_output=True, text=True, timeout=280)
        if r.stdout:
            for l in r.stdout.splitlines()[-30:]: print(f"  sub: {l}")
        if r.stderr:
            for l in r.stderr.splitlines()[-10:]: print(f"  err: {l}")
        _parse_results(outdir)
    except Exception as e:
        print(f"[PROFILE] Error: {e}")

if _IS_PROFILE_SUB:
    _run_profiling_kernels(); sys.exit(0)
elif _DO_PROFILE:
    _run_profiling()

# torch.compile failed: can't trace pybind11 C++ extensions.
# Would need torch.library custom op registration to work.
scrolls · 1394 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