Skip to content
KernelIndex
Search⌘K

submission 754911

steve · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v1385_backup.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754911?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
8.39µs
#62 of 1143
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:fbfa0f8dabf7dafcc3405d8ef4b85e3d4590535dea1057e4e622e7d83482747d
license declaredunknown
license concludedunknown
authorssteve
imported2026-08-15

Techniques

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

shared-memoryextern __shared__ uint8_t lds[];
split-kmxfp4_gemm_splitk_fused(
vector-width = uint4const uint4* src4 = reinterpret_cast<const uint4*>(src + offset);

Kernel source

submission_v1385_backup.py1617 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""v1385: -mcumode for CU-mode scheduling."""
from task import input_t, output_t
import torch
import ctypes
import os
from concurrent.futures import ThreadPoolExecutor
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"

HIP_SOURCE = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/hip_fp16.h>
#include <hip/hip_bfloat16.h>
#include <cstdint>

using mfma_acc_t = __attribute__((ext_vector_type(4))) float;
using mfma_input_t = __attribute__((ext_vector_type(8))) uint32_t;
using u32x4 = __attribute__((ext_vector_type(4))) uint32_t;

#define MXFP4_BLOCK_SIZE 32

__device__ __forceinline__ mfma_acc_t
mfma_fp4_16x16(mfma_input_t a, mfma_input_t b, mfma_acc_t c,
               uint32_t sa, uint32_t sb) {
#if defined(__gfx950__)
    return __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
        a, b, c, 4, 4, 0, sa, 0, sb);
#else
    return c;
#endif
}

#ifdef UNIT_HEX
// Inline ASM: 6 MFMAs for hex kernel in single block (no compiler waits)
__device__ __forceinline__ void
asm_mfma_hex(u32x4 af, uint32_t sa,
             u32x4 bf0, u32x4 bf1, u32x4 bf2, u32x4 bf3, u32x4 bf4, u32x4 bf5,
             uint32_t sb0, uint32_t sb1, uint32_t sb2, uint32_t sb3, uint32_t sb4, uint32_t sb5,
             mfma_acc_t& acc0, mfma_acc_t& acc1, mfma_acc_t& acc2,
             mfma_acc_t& acc3, mfma_acc_t& acc4, mfma_acc_t& acc5) {
#if defined(__gfx950__)
    asm volatile(
        "s_setprio 2\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a0], %[af], %[b0], %[a0], %[sa], %[s0] cbsz:4 blgp:4\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a1], %[af], %[b1], %[a1], %[sa], %[s1] cbsz:4 blgp:4\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a2], %[af], %[b2], %[a2], %[sa], %[s2] cbsz:4 blgp:4\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a3], %[af], %[b3], %[a3], %[sa], %[s3] cbsz:4 blgp:4\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a4], %[af], %[b4], %[a4], %[sa], %[s4] cbsz:4 blgp:4\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a5], %[af], %[b5], %[a5], %[sa], %[s5] cbsz:4 blgp:4\n\t"
        "s_setprio 0\n\t"
        : [a0] "+v"(acc0), [a1] "+v"(acc1), [a2] "+v"(acc2),
          [a3] "+v"(acc3), [a4] "+v"(acc4), [a5] "+v"(acc5)
        : [af] "v"(af), [sa] "v"(sa),
          [b0] "v"(bf0), [b1] "v"(bf1), [b2] "v"(bf2),
          [b3] "v"(bf3), [b4] "v"(bf4), [b5] "v"(bf5),
          [s0] "v"(sb0), [s1] "v"(sb1), [s2] "v"(sb2),
          [s3] "v"(sb3), [s4] "v"(sb4), [s5] "v"(sb5)
        : "memory"
    );
#endif
}

// Inline ASM: 6 MFMAs interleaved with 6 dwordx4 NT B loads for hex
// MFMA captures inputs at issue (4 cycles), so bf VGPRs safe to reuse as load dests
__device__ __forceinline__ void
asm_mfma_hex_loadb(
    u32x4 af, uint32_t sa,
    u32x4& bf0, u32x4& bf1, u32x4& bf2, u32x4& bf3, u32x4& bf4, u32x4& bf5,
    uint32_t sb0, uint32_t sb1, uint32_t sb2, uint32_t sb3, uint32_t sb4, uint32_t sb5,
    mfma_acc_t& acc0, mfma_acc_t& acc1, mfma_acc_t& acc2,
    mfma_acc_t& acc3, mfma_acc_t& acc4, mfma_acc_t& acc5,
    const void* p0, const void* p1, const void* p2,
    const void* p3, const void* p4, const void* p5) {
#if defined(__gfx950__)
    asm volatile(
        "s_setprio 2\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a0], %[af], %[b0], %[a0], %[sa], %[s0] cbsz:4 blgp:4\n\t"
        "global_load_dwordx4 %[b0], %[p0], off nt\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a1], %[af], %[b1], %[a1], %[sa], %[s1] cbsz:4 blgp:4\n\t"
        "global_load_dwordx4 %[b1], %[p1], off nt\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a2], %[af], %[b2], %[a2], %[sa], %[s2] cbsz:4 blgp:4\n\t"
        "global_load_dwordx4 %[b2], %[p2], off nt\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a3], %[af], %[b3], %[a3], %[sa], %[s3] cbsz:4 blgp:4\n\t"
        "global_load_dwordx4 %[b3], %[p3], off nt\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a4], %[af], %[b4], %[a4], %[sa], %[s4] cbsz:4 blgp:4\n\t"
        "global_load_dwordx4 %[b4], %[p4], off nt\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a5], %[af], %[b5], %[a5], %[sa], %[s5] cbsz:4 blgp:4\n\t"
        "global_load_dwordx4 %[b5], %[p5], off nt\n\t"
        "s_setprio 0\n\t"
        : [a0] "+v"(acc0), [a1] "+v"(acc1), [a2] "+v"(acc2),
          [a3] "+v"(acc3), [a4] "+v"(acc4), [a5] "+v"(acc5),
          [b0] "+v"(bf0), [b1] "+v"(bf1), [b2] "+v"(bf2),
          [b3] "+v"(bf3), [b4] "+v"(bf4), [b5] "+v"(bf5)
        : [af] "v"(af), [sa] "v"(sa),
          [s0] "v"(sb0), [s1] "v"(sb1), [s2] "v"(sb2),
          [s3] "v"(sb3), [s4] "v"(sb4), [s5] "v"(sb5),
          [p0] "v"(p0), [p1] "v"(p1), [p2] "v"(p2),
          [p3] "v"(p3), [p4] "v"(p4), [p5] "v"(p5)
        : "memory"
    );
#endif
}
#endif // UNIT_HEX

#if defined(UNIT_QUAD)
// Inline ASM: 4 MFMAs for quad kernel (no compiler waits)
__device__ __forceinline__ void
asm_mfma_quad(u32x4 af, uint32_t sa,
              u32x4 bf0, u32x4 bf1, u32x4 bf2, u32x4 bf3,
              uint32_t sb0, uint32_t sb1, uint32_t sb2, uint32_t sb3,
              mfma_acc_t& acc0, mfma_acc_t& acc1, mfma_acc_t& acc2, mfma_acc_t& acc3) {
#if defined(__gfx950__)
    asm volatile(
        "s_setprio 2\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a0], %[af], %[b0], %[a0], %[sa], %[s0] cbsz:4 blgp:4\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a1], %[af], %[b1], %[a1], %[sa], %[s1] cbsz:4 blgp:4\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a2], %[af], %[b2], %[a2], %[sa], %[s2] cbsz:4 blgp:4\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a3], %[af], %[b3], %[a3], %[sa], %[s3] cbsz:4 blgp:4\n\t"
        "s_setprio 0\n\t"
        : [a0] "+v"(acc0), [a1] "+v"(acc1), [a2] "+v"(acc2), [a3] "+v"(acc3)
        : [af] "v"(af), [sa] "v"(sa),
          [b0] "v"(bf0), [b1] "v"(bf1), [b2] "v"(bf2), [b3] "v"(bf3),
          [s0] "v"(sb0), [s1] "v"(sb1), [s2] "v"(sb2), [s3] "v"(sb3)
        : "memory"
    );
#endif
}

// Inline ASM: 4 MFMAs interleaved with 4 dwordx4 NT B loads
// MFMA captures inputs at issue (4 cycles), so bf VGPRs safe to reuse as load dests
__device__ __forceinline__ void
asm_mfma_quad_loadb(
    u32x4 af, uint32_t sa,
    u32x4& bf0, u32x4& bf1, u32x4& bf2, u32x4& bf3,
    uint32_t sb0, uint32_t sb1, uint32_t sb2, uint32_t sb3,
    mfma_acc_t& acc0, mfma_acc_t& acc1, mfma_acc_t& acc2, mfma_acc_t& acc3,
    const void* p0, const void* p1, const void* p2, const void* p3) {
#if defined(__gfx950__)
    asm volatile(
        "s_setprio 2\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a0], %[af], %[b0], %[a0], %[sa], %[s0] cbsz:4 blgp:4\n\t"
        "global_load_dwordx4 %[b0], %[p0], off nt\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a1], %[af], %[b1], %[a1], %[sa], %[s1] cbsz:4 blgp:4\n\t"
        "global_load_dwordx4 %[b1], %[p1], off nt\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a2], %[af], %[b2], %[a2], %[sa], %[s2] cbsz:4 blgp:4\n\t"
        "global_load_dwordx4 %[b2], %[p2], off nt\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a3], %[af], %[b3], %[a3], %[sa], %[s3] cbsz:4 blgp:4\n\t"
        "global_load_dwordx4 %[b3], %[p3], off nt\n\t"
        "s_setprio 0\n\t"
        : [a0] "+v"(acc0), [a1] "+v"(acc1), [a2] "+v"(acc2), [a3] "+v"(acc3),
          [b0] "+v"(bf0), [b1] "+v"(bf1), [b2] "+v"(bf2), [b3] "+v"(bf3)
        : [af] "v"(af), [sa] "v"(sa),
          [s0] "v"(sb0), [s1] "v"(sb1), [s2] "v"(sb2), [s3] "v"(sb3),
          [p0] "v"(p0), [p1] "v"(p1), [p2] "v"(p2), [p3] "v"(p3)
        : "memory"
    );
#endif
}

#endif // UNIT_QUAD (quad ASM functions)

// Non-temporal bf16 store via builtin
__device__ __forceinline__ void
store_nt_bf16(hip_bfloat16* __restrict__ ptr, hip_bfloat16 val) {
    __builtin_nontemporal_store(*reinterpret_cast<uint16_t*>(&val),
                                reinterpret_cast<uint16_t*>(ptr));
}

// Non-temporal float store via builtin
__device__ __forceinline__ void
store_nt_f32(float* __restrict__ ptr, float val) {
    __builtin_nontemporal_store(*reinterpret_cast<uint32_t*>(&val),
                                reinterpret_cast<uint32_t*>(ptr));
}

// Single dwordx4 NT load into u32x4 via ASM
__device__ __forceinline__ void
load_b_nt_u32x4(const uint8_t* __restrict__ B_shuffle, size_t base, int lid,
                u32x4& bf) {
#if defined(__gfx950__)
    const void* ptr = (const void*)(B_shuffle + base + (size_t)lid * 16);
    asm volatile("global_load_dwordx4 %0, %1, off nt"
                 : "=v"(bf) : "v"(ptr) : "memory");
#endif
}

__device__ __forceinline__ uint32_t
read_bscale_sh(const uint8_t* __restrict__ B_scale_sh,
               int nts, int pos, int grp, int sd4, int sp1d8) {
    int row = nts + pos, col = sd4 * 4 + grp;
    int r0 = row >> 5, r1 = (row >> 4) & 1, c0 = col >> 3, c1 = (col >> 2) & 1;
    return (uint32_t)__builtin_nontemporal_load(&B_scale_sh[r0 * sp1d8 * 256 + c0 * 256 + grp * 64 + pos * 4 + c1 * 2 + r1]);
}

__device__ __forceinline__ void
load_b_nt(const uint8_t* __restrict__ B_shuffle, size_t base, int lid,
          mfma_input_t& bf) {
    const uint32_t* bp = reinterpret_cast<const uint32_t*>(
        B_shuffle + base + (size_t)lid * 16);
    bf[0] = __builtin_nontemporal_load(bp + 0);
    bf[1] = __builtin_nontemporal_load(bp + 1);
    bf[2] = __builtin_nontemporal_load(bp + 2);
    bf[3] = __builtin_nontemporal_load(bp + 3);
}

// Packed single-pass quant: uint4 loads + v_pk_max_u16 absmax (from v104)
__device__ __forceinline__ void
quant_bf16x32_to_fp4(const uint16_t* __restrict__ src, int offset,
                     uint8_t* __restrict__ dst, uint32_t& scale_out) {
    const uint4* src4 = reinterpret_cast<const uint4*>(src + offset);

    uint4 d[4];
    #pragma unroll
    for (int j = 0; j < 4; j++) d[j] = src4[j];

#if defined(__gfx950__)
    uint32_t pk_max = 0;
    #pragma unroll
    for (int j = 0; j < 4; j++) {
        uint32_t a0 = d[j].x & 0x7FFF7FFFu;
        uint32_t a1 = d[j].y & 0x7FFF7FFFu;
        uint32_t a2 = d[j].z & 0x7FFF7FFFu;
        uint32_t a3 = d[j].w & 0x7FFF7FFFu;
        asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a0), "v"(pk_max));
        asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a1), "v"(pk_max));
        asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a2), "v"(pk_max));
        asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a3), "v"(pk_max));
    }
    uint32_t lo16 = pk_max & 0xFFFFu;
    uint32_t hi16 = pk_max >> 16;
    uint32_t max16;
    asm volatile("v_max_u32 %0, %1, %2" : "=v"(max16) : "v"(lo16), "v"(hi16));
    float absmax = __uint_as_float(max16 << 16);
#else
    float absmax = 0.0f;
    #pragma unroll
    for (int j = 0; j < 4; j++) {
        float v0, v1;
        v0 = __uint_as_float((d[j].x & 0xFFFFu) << 16);
        v1 = __uint_as_float(d[j].x & 0xFFFF0000u);
        absmax = fmaxf(absmax, fmaxf(fabsf(v0), fabsf(v1)));
        v0 = __uint_as_float((d[j].y & 0xFFFFu) << 16);
        v1 = __uint_as_float(d[j].y & 0xFFFF0000u);
        absmax = fmaxf(absmax, fmaxf(fabsf(v0), fabsf(v1)));
        v0 = __uint_as_float((d[j].z & 0xFFFFu) << 16);
        v1 = __uint_as_float(d[j].z & 0xFFFF0000u);
        absmax = fmaxf(absmax, fmaxf(fabsf(v0), fabsf(v1)));
        v0 = __uint_as_float((d[j].w & 0xFFFFu) << 16);
        v1 = __uint_as_float(d[j].w & 0xFFFF0000u);
        absmax = fmaxf(absmax, fmaxf(fabsf(v0), fabsf(v1)));
    }
#endif

    uint32_t bits = __float_as_uint(absmax);
    uint32_t rounded = (bits + 0x200000u) & 0xFF800000u;
    uint32_t E = (rounded >> 23) & 0xFFu;
    uint32_t e8m0 = (E > 2u) ? (E - 2u) : 0u;
    scale_out = e8m0;

#if defined(__gfx950__)
    float scale_f = (e8m0 > 0u) ? __uint_as_float(e8m0 << 23) : 1.0f;
    #pragma unroll
    for (int j = 0; j < 4; j++) {
        uint32_t pk0, pk1, pk2, pk3;
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
                     : "=v"(pk0) : "v"(d[j].x), "v"(scale_f));
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
                     : "=v"(pk1) : "v"(d[j].y), "v"(scale_f));
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
                     : "=v"(pk2) : "v"(d[j].z), "v"(scale_f));
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
                     : "=v"(pk3) : "v"(d[j].w), "v"(scale_f));
        uint32_t lo, hi, packed;
        asm volatile("v_perm_b32 %0, %1, %2, %3" : "=v"(lo) : "v"(pk1), "v"(pk0), "v"(0x0C0C0400u));
        asm volatile("v_perm_b32 %0, %1, %2, %3" : "=v"(hi) : "v"(pk3), "v"(pk2), "v"(0x04000C0Cu));
        asm volatile("v_or_b32 %0, %1, %2" : "=v"(packed) : "v"(lo), "v"(hi));
        reinterpret_cast<uint32_t*>(dst)[j] = packed;
    }
#else
    (void)dst;
#endif
}

// Compute-only quant: takes pre-loaded d[4], produces fp4 + scale (for deep pipeline)
__device__ __forceinline__ void
quant_compute_fp4(uint4 d[4], uint8_t* __restrict__ dst, uint32_t& scale_out) {
#if defined(__gfx950__)
    uint32_t pk_max = 0;
    #pragma unroll
    for (int j = 0; j < 4; j++) {
        uint32_t a0 = d[j].x & 0x7FFF7FFFu;
        uint32_t a1 = d[j].y & 0x7FFF7FFFu;
        uint32_t a2 = d[j].z & 0x7FFF7FFFu;
        uint32_t a3 = d[j].w & 0x7FFF7FFFu;
        asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a0), "v"(pk_max));
        asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a1), "v"(pk_max));
        asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a2), "v"(pk_max));
        asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a3), "v"(pk_max));
    }
    uint32_t lo16 = pk_max & 0xFFFFu;
    uint32_t hi16 = pk_max >> 16;
    uint32_t max16;
    asm volatile("v_max_u32 %0, %1, %2" : "=v"(max16) : "v"(lo16), "v"(hi16));
    float absmax = __uint_as_float(max16 << 16);

    uint32_t bits = __float_as_uint(absmax);
    uint32_t rounded = (bits + 0x200000u) & 0xFF800000u;
    uint32_t E = (rounded >> 23) & 0xFFu;
    uint32_t e8m0 = (E > 2u) ? (E - 2u) : 0u;
    scale_out = e8m0;

    float scale_f = (e8m0 > 0u) ? __uint_as_float(e8m0 << 23) : 1.0f;
    #pragma unroll
    for (int j = 0; j < 4; j++) {
        uint32_t pk0, pk1, pk2, pk3;
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
                     : "=v"(pk0) : "v"(d[j].x), "v"(scale_f));
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
                     : "=v"(pk1) : "v"(d[j].y), "v"(scale_f));
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
                     : "=v"(pk2) : "v"(d[j].z), "v"(scale_f));
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
                     : "=v"(pk3) : "v"(d[j].w), "v"(scale_f));
        uint32_t lo, hi, packed;
        asm volatile("v_perm_b32 %0, %1, %2, %3" : "=v"(lo) : "v"(pk1), "v"(pk0), "v"(0x0C0C0400u));
        asm volatile("v_perm_b32 %0, %1, %2, %3" : "=v"(hi) : "v"(pk3), "v"(pk2), "v"(0x04000C0Cu));
        asm volatile("v_or_b32 %0, %1, %2" : "=v"(packed) : "v"(lo), "v"(hi));
        reinterpret_cast<uint32_t*>(dst)[j] = packed;
    }
#else
    scale_out = 0x7F;
    (void)dst;
#endif
}

// Compute-only quant returning u32x4 directly (no stack intermediary)
__device__ __forceinline__ void
quant_compute_u32x4(uint4 d[4], u32x4& af_out, uint32_t& scale_out) {
#if defined(__gfx950__)
    uint32_t pk_max = 0;
    #pragma unroll
    for (int j = 0; j < 4; j++) {
        uint32_t a0 = d[j].x & 0x7FFF7FFFu;
        uint32_t a1 = d[j].y & 0x7FFF7FFFu;
        uint32_t a2 = d[j].z & 0x7FFF7FFFu;
        uint32_t a3 = d[j].w & 0x7FFF7FFFu;
        asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a0), "v"(pk_max));
        asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a1), "v"(pk_max));
        asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a2), "v"(pk_max));
        asm volatile("v_pk_max_u16 %0, %1, %2" : "=v"(pk_max) : "v"(a3), "v"(pk_max));
    }
    uint32_t lo16 = pk_max & 0xFFFFu;
    uint32_t hi16 = pk_max >> 16;
    uint32_t max16;
    asm volatile("v_max_u32 %0, %1, %2" : "=v"(max16) : "v"(lo16), "v"(hi16));
    float absmax = __uint_as_float(max16 << 16);

    uint32_t bits = __float_as_uint(absmax);
    uint32_t rounded = (bits + 0x200000u) & 0xFF800000u;
    uint32_t E = (rounded >> 23) & 0xFFu;
    uint32_t e8m0 = (E > 2u) ? (E - 2u) : 0u;
    scale_out = e8m0;

    float scale_f = (e8m0 > 0u) ? __uint_as_float(e8m0 << 23) : 1.0f;
    #pragma unroll
    for (int j = 0; j < 4; j++) {
        uint32_t pk0, pk1, pk2, pk3;
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
                     : "=v"(pk0) : "v"(d[j].x), "v"(scale_f));
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
                     : "=v"(pk1) : "v"(d[j].y), "v"(scale_f));
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
                     : "=v"(pk2) : "v"(d[j].z), "v"(scale_f));
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
                     : "=v"(pk3) : "v"(d[j].w), "v"(scale_f));
        uint32_t lo, hi, packed;
        asm volatile("v_perm_b32 %0, %1, %2, %3" : "=v"(lo) : "v"(pk1), "v"(pk0), "v"(0x0C0C0400u));
        asm volatile("v_perm_b32 %0, %1, %2, %3" : "=v"(hi) : "v"(pk3), "v"(pk2), "v"(0x04000C0Cu));
        asm volatile("v_or_b32 %0, %1, %2" : "=v"(packed) : "v"(lo), "v"(hi));
        af_out[j] = packed;
    }
#else
    af_out = {0,0,0,0};
    scale_out = 0x7F;
#endif
}

// Full quant from bf16 source, returns u32x4 directly
__device__ __forceinline__ void
quant_bf16x32_to_u32x4(const uint16_t* __restrict__ src, int offset,
                        u32x4& af_out, uint32_t& scale_out) {
    const uint4* src4 = reinterpret_cast<const uint4*>(src + offset);
    uint4 d[4];
    #pragma unroll
    for (int j = 0; j < 4; j++) d[j] = src4[j];
    quant_compute_u32x4(d, af_out, scale_out);
}

// Single MFMA with u32x4 inputs (for narrow kernel)
#if defined(UNIT_NARROW)
__device__ __forceinline__ void
asm_mfma_single(u32x4 af, uint32_t sa, u32x4 bf, uint32_t sb,
                mfma_acc_t& acc) {
#if defined(__gfx950__)
    asm volatile(
        "s_setprio 2\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a0], %[af], %[b0], %[a0], %[sa], %[s0] cbsz:4 blgp:4\n\t"
        "s_setprio 0\n\t"
        : [a0] "+v"(acc)
        : [af] "v"(af), [sa] "v"(sa),
          [b0] "v"(bf), [s0] "v"(sb)
        : "memory"
    );
#endif
}
#endif // UNIT_NARROW

// Inline ASM: 2 MFMAs (no compiler waits) — shared by wide + splitk
#if defined(UNIT_WIDE) || defined(UNIT_SPLITK)
__device__ __forceinline__ void
asm_mfma_pair(u32x4 af, uint32_t sa,
              u32x4 bf0, u32x4 bf1,
              uint32_t sb0, uint32_t sb1,
              mfma_acc_t& acc0, mfma_acc_t& acc1) {
#if defined(__gfx950__)
    asm volatile(
        "s_setprio 2\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a0], %[af], %[b0], %[a0], %[sa], %[s0] cbsz:4 blgp:4\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a1], %[af], %[b1], %[a1], %[sa], %[s1] cbsz:4 blgp:4\n\t"
        "s_setprio 0\n\t"
        : [a0] "+v"(acc0), [a1] "+v"(acc1)
        : [af] "v"(af), [sa] "v"(sa),
          [b0] "v"(bf0), [b1] "v"(bf1),
          [s0] "v"(sb0), [s1] "v"(sb1)
        : "memory"
    );
#endif
}
// Inline ASM: 2 MFMAs interleaved with 2 dwordx4 NT B loads
__device__ __forceinline__ void
asm_mfma_pair_loadb(
    u32x4 af, uint32_t sa,
    u32x4& bf0, u32x4& bf1,
    uint32_t sb0, uint32_t sb1,
    mfma_acc_t& acc0, mfma_acc_t& acc1,
    const void* p0, const void* p1) {
#if defined(__gfx950__)
    asm volatile(
        "s_setprio 2\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a0], %[af], %[b0], %[a0], %[sa], %[s0] cbsz:4 blgp:4\n\t"
        "global_load_dwordx4 %[b0], %[p0], off nt\n\t"
        "v_mfma_scale_f32_16x16x128_f8f6f4 %[a1], %[af], %[b1], %[a1], %[sa], %[s1] cbsz:4 blgp:4\n\t"
        "global_load_dwordx4 %[b1], %[p1], off nt\n\t"
        "s_setprio 0\n\t"
        : [a0] "+v"(acc0), [a1] "+v"(acc1),
          [b0] "+v"(bf0), [b1] "+v"(bf1)
        : [af] "v"(af), [sa] "v"(sa),
          [s0] "v"(sb0), [s1] "v"(sb1),
          [p0] "v"(p0), [p1] "v"(p1)
        : "memory"
    );
#endif
}
#endif // UNIT_WIDE || UNIT_SPLITK

#ifdef UNIT_NARROW
// ---- Shape 1: Narrow 16-wide fused (M=4) ----
__global__ void __launch_bounds__(256, 2)
mxfp4_fused_narrow(
    const uint16_t* __restrict__ A_bf16,
    const uint8_t*  __restrict__ B_shuffle,
    const uint8_t*  __restrict__ B_scale_sh,
    hip_bfloat16*   __restrict__ output,
    int M, int N, int K, int sp1d8
) {
    constexpr int NW = 4, T = 16, MB = 64;
    const int Kp = K / 2;
    const int KPW = Kp / NW, KI = KPW / MB, KB = Kp / MB;
    const int tid = threadIdx.x, wid = tid / 64, lid = tid % 64;
    const int ntx = blockIdx.x, nts = ntx * T, mts = blockIdx.y * T;
    if (nts >= N) return;
    const int grp = lid / 16, pos = lid % 16, mr = mts + pos;
    const bool valid = (mr < M);
    mfma_acc_t acc = {0,0,0,0};

    if (KI >= 2) {
        int ki = 0;
        int kb0 = wid*KI+ki, keo0 = wid*(K/NW)+ki*128;
        int sd40 = keo0/(MXFP4_BLOCK_SIZE*4);
        uint32_t sa0=0x7F;
        u32x4 af0={0,0,0,0};
        // Issue B load first, overlap with A quant
        u32x4 bf0={0,0,0,0};
        load_b_nt_u32x4(B_shuffle, (size_t)ntx*KB*1024+(size_t)kb0*1024, lid, bf0);
        if (valid) {
            quant_bf16x32_to_u32x4(A_bf16, mr*K+keo0+grp*32, af0, sa0);
        }
        asm volatile("s_waitcnt vmcnt(0)":::"memory");
        uint32_t sb0 = read_bscale_sh(B_scale_sh, nts, pos, grp, sd40, sp1d8);
        for (ki = 1; ki < KI; ki++) {
            int kb1=wid*KI+ki, keo1=wid*(K/NW)+ki*128;
            int sd41=keo1/(MXFP4_BLOCK_SIZE*4);
            asm_mfma_single(af0, sa0, bf0, sb0, acc);
            // Issue B load first, overlap with A quant
            load_b_nt_u32x4(B_shuffle, (size_t)ntx*KB*1024+(size_t)kb1*1024, lid, bf0);
            af0={0,0,0,0}; sa0=0x7F;
            if (valid) {
                quant_bf16x32_to_u32x4(A_bf16, mr*K+keo1+grp*32, af0, sa0);
            }
            asm volatile("s_waitcnt vmcnt(0)":::"memory");
            sb0 = read_bscale_sh(B_scale_sh, nts, pos, grp, sd41, sp1d8);
        }
        asm_mfma_single(af0, sa0, bf0, sb0, acc);
    } else if (KI == 1) {
        int kb=wid, keo=wid*(K/NW), sd4=keo/(MXFP4_BLOCK_SIZE*4);
        uint32_t sa=0x7F;
        u32x4 af={0,0,0,0};
        // Issue B load first, overlap with A quant
        u32x4 bf={0,0,0,0};
        load_b_nt_u32x4(B_shuffle, (size_t)ntx*KB*1024+(size_t)kb*1024, lid, bf);
        if (valid) {
            quant_bf16x32_to_u32x4(A_bf16, mr*K+keo+grp*32, af, sa);
        }
        asm volatile("s_waitcnt vmcnt(0)":::"memory");
        uint32_t sb = read_bscale_sh(B_scale_sh, nts, pos, grp, sd4, sp1d8);
        asm_mfma_single(af, sa, bf, sb, acc);
    }

    extern __shared__ uint8_t lds[];
    float* lr = reinterpret_cast<float*>(lds);
    #pragma unroll
    for (int i=0;i<4;i++) lr[wid*T*T+(grp*4+i)*T+pos]=acc[i];
    __syncthreads();
    {
        int ml=wid*4+grp, mg=mts+ml, ng=nts+pos;
        if (mg<M&&ng<N) {
            float s=0;
            #pragma unroll
            for (int w=0;w<NW;w++) s+=lr[w*T*T+ml*T+pos];
            output[mg*N+ng]=hip_bfloat16(s);
        }
    }
}
#endif // UNIT_NARROW

#ifdef UNIT_WIDE
// ---- Shapes 3,4: Wide 32-N fused (M=32) ----
__global__ void __launch_bounds__(256, 2)
mxfp4_fused_wide(
    const uint16_t* __restrict__ A_bf16,
    const uint8_t*  __restrict__ B_shuffle,
    const uint8_t*  __restrict__ B_scale_sh,
    hip_bfloat16*   __restrict__ output,
    int M, int N, int K, int sp1d8
) {
    constexpr int NW = 4, T = 16, MB = 64;
    const int Kp = K / 2;
    const int KPW = Kp / NW, KI = KPW / MB, KB = Kp / MB;
    const int tid = threadIdx.x, wid = tid / 64, lid = tid % 64;
    const int ntx0 = blockIdx.x * 2, ntx1 = ntx0 + 1;
    const int nts0 = ntx0 * T, nts1 = ntx1 * T;
    const int mts = blockIdx.y * T;
    if (nts0 >= N) return;
    const bool has_sub1 = (nts1 < N);
    const int grp = lid / 16, pos = lid % 16, mr = mts + pos;
    const bool valid = (mr < M);

    mfma_acc_t acc0 = {0,0,0,0}, acc1 = {0,0,0,0};

    if (KI >= 2) {
        // === DEEP PIPELINE: separate A load/quant, A double-buffer like quad/hex ===
        int ki = 0;
        int kb0 = wid*KI+ki, keo0 = wid*(K/NW)+ki*128;
        int sd40 = keo0/(MXFP4_BLOCK_SIZE*4);
        uint32_t sa0=0x7F;
        u32x4 af0={0,0,0,0};

        // Prologue: Load A[0] data into registers (oldest in vmcnt)
        uint4 ad[4] = {};
        if (valid) {
            const uint4* a_src = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo0+grp*32);
            #pragma unroll
            for (int j = 0; j < 4; j++) ad[j] = a_src[j];
        }
        // Issue B[0] loads (newer in vmcnt)
        u32x4 bfa={0,0,0,0}, bfb={0,0,0,0};
        load_b_nt_u32x4(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb0*1024, lid, bfa);
        if (has_sub1) load_b_nt_u32x4(B_shuffle, (size_t)ntx1*KB*1024+(size_t)kb0*1024, lid, bfb);
        // Wait for A only (oldest), B loads overlap with quant ALU
        if (has_sub1) {
            asm volatile("s_waitcnt vmcnt(2)":::"memory");
        } else {
            asm volatile("s_waitcnt vmcnt(1)":::"memory");
        }
        if (valid) quant_compute_u32x4(ad, af0, sa0);
        // Scale reads while B loads still in flight
        uint32_t sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd40, sp1d8);
        uint32_t sbb = 0x7F;
        if (has_sub1) sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd40, sp1d8);
        // Pre-load A[1] for double-buffer
        uint4 ad_next[4] = {};
        if (valid && KI > 2) {
            int keo_next = wid*(K/NW)+1*128;
            const uint4* a_src_next = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo_next+grp*32);
            #pragma unroll
            for (int j = 0; j < 4; j++) ad_next[j] = a_src_next[j];
        }
        // Wait for B[0] + scale loads; leave A pre-loads in flight
        if (KI > 2) {
            asm volatile("s_waitcnt vmcnt(4)":::"memory");
        } else {
            asm volatile("s_waitcnt vmcnt(0)":::"memory");
        }
        for (ki = 1; ki < KI; ki++) {
            int kb1=wid*KI+ki, keo1=wid*(K/NW)+ki*128;
            int sd41=keo1/(MXFP4_BLOCK_SIZE*4);
            // Interleaved MFMA(old data) + B[ki] loads
            {
                const void* pa = (const void*)(B_shuffle + (size_t)ntx0*KB*1024 + (size_t)kb1*1024 + (size_t)lid*16);
                const void* pb = (const void*)(B_shuffle + (size_t)ntx1*KB*1024 + (size_t)kb1*1024 + (size_t)lid*16);
                asm_mfma_pair_loadb(af0, sa0, bfa, bfb, sba, sbb, acc0, acc1, pa, pb);
            }
            // Wait for pre-loaded A[ki] (2 B loads in flight, A loads older)
            asm volatile("s_waitcnt vmcnt(2)":::"memory");
            ad[0] = ad_next[0]; ad[1] = ad_next[1]; ad[2] = ad_next[2]; ad[3] = ad_next[3];
            // Quant A[ki] (VALU only, overlaps with B loads in flight)
            af0 = {0,0,0,0}; sa0 = 0x7F;
            if (valid) quant_compute_u32x4(ad, af0, sa0);
            // Scale reads
            sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd41, sp1d8);
            if (has_sub1) sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd41, sp1d8);
            // Pre-load A[ki+1]
            ad_next[0] = {}; ad_next[1] = {}; ad_next[2] = {}; ad_next[3] = {};
            if (valid && ki + 1 < KI) {
                int keo_next = wid*(K/NW) + (ki+1)*128;
                const uint4* a_src_next = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo_next+grp*32);
                #pragma unroll
                for (int j = 0; j < 4; j++) ad_next[j] = a_src_next[j];
            }
            // Wait for B[ki] + scales, leave A pre-loads in flight
            if (ki + 1 < KI) {
                asm volatile("s_waitcnt vmcnt(4)":::"memory");
            } else {
                asm volatile("s_waitcnt vmcnt(0)":::"memory");
            }
        }
        // Epilogue: final MFMAs
        asm_mfma_pair(af0, sa0, bfa, bfb, sba, sbb, acc0, acc1);
    } else if (KI == 1) {
        int kb=wid, keo=wid*(K/NW), sd4=keo/(MXFP4_BLOCK_SIZE*4);
        uint32_t sa=0x7F;
        u32x4 af={0,0,0,0};
        // Load A data first (oldest in vmcnt FIFO)
        uint4 ad[4] = {};
        if (valid) {
            const uint4* a_src = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo+grp*32);
            #pragma unroll
            for (int j = 0; j < 4; j++) ad[j] = a_src[j];
        }
        // Issue B loads (newer in vmcnt FIFO)
        u32x4 bf0={0,0,0,0}, bf1={0,0,0,0};
        load_b_nt_u32x4(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb*1024, lid, bf0);
        if (has_sub1) load_b_nt_u32x4(B_shuffle, (size_t)ntx1*KB*1024+(size_t)kb*1024, lid, bf1);
        // Wait for A data only (oldest), B loads overlap with quant VALU
        if (has_sub1) {
            asm volatile("s_waitcnt vmcnt(2)":::"memory");
        } else {
            asm volatile("s_waitcnt vmcnt(1)":::"memory");
        }
        // Quant A (VALU-only, overlaps with B loads in flight)
        if (valid) quant_compute_u32x4(ad, af, sa);
        // Scale reads while B may still be in flight
        uint32_t sb0 = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd4, sp1d8);
        uint32_t sb1 = 0x7F;
        if (has_sub1) sb1 = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd4, sp1d8);
        // Wait for B loads
        asm volatile("s_waitcnt vmcnt(0)":::"memory");
        // ASM MFMA pair: both N-tiles in one block
        asm_mfma_pair(af, sa, bf0, bf1, sb0, sb1, acc0, acc1);
    }

    extern __shared__ uint8_t lds[];
    float* lr = reinterpret_cast<float*>(lds);
    #pragma unroll
    for (int i=0;i<4;i++) lr[wid*T*T+(grp*4+i)*T+pos]=acc0[i];
    #pragma unroll
    for (int i=0;i<4;i++) lr[NW*T*T+wid*T*T+(grp*4+i)*T+pos]=acc1[i];
    __syncthreads();
    {
        int ml=wid*4+grp, mg=mts+ml, ng=nts0+pos;
        if (mg<M&&ng<N) {
            float s=0;
            #pragma unroll
            for (int w=0;w<NW;w++) s+=lr[w*T*T+ml*T+pos];
            store_nt_bf16(&output[mg*N+ng], hip_bfloat16(s));
        }
        if (has_sub1) {
            ng=nts1+pos;
            if (mg<M&&ng<N) {
                float s=0;
                #pragma unroll
                for (int w=0;w<NW;w++) s+=lr[NW*T*T+w*T*T+ml*T+pos];
                store_nt_bf16(&output[mg*N+ng], hip_bfloat16(s));
            }
        }
    }
}
#endif // UNIT_WIDE

#ifdef UNIT_QUAD
// ---- Shape 5: Quad 64-N fused (M=64, 4 N-tiles per CTA) with deep pipeline ----
__global__ void __launch_bounds__(256, 2)
mxfp4_fused_quad(
    const uint16_t* __restrict__ A_bf16,
    const uint8_t*  __restrict__ B_shuffle,
    const uint8_t*  __restrict__ B_scale_sh,
    hip_bfloat16*   __restrict__ output,
    int M, int N, int K, int sp1d8
) {
    __builtin_amdgcn_s_setprio(3);
    constexpr int NW = 4, T = 16, MB = 64;
    const int Kp = K / 2;
    const int KPW = Kp / NW, KI = KPW / MB, KB = Kp / MB;
    const int tid = threadIdx.x, wid = tid / 64, lid = tid % 64;
    const int ntx0 = blockIdx.x * 4;
    const int nts0 = ntx0 * T, nts1 = (ntx0+1) * T, nts2 = (ntx0+2) * T, nts3 = (ntx0+3) * T;
    const int mts = blockIdx.y * T;
    if (nts0 >= N) return;
    const bool has1 = (nts1 < N), has2 = (nts2 < N), has3 = (nts3 < N);
    const int grp = lid / 16, pos = lid % 16, mr = mts + pos;
    const bool valid = (mr < M);
    mfma_acc_t acc0 = {0,0,0,0}, acc1 = {0,0,0,0}, acc2 = {0,0,0,0}, acc3 = {0,0,0,0};
    if (KI >= 2 && has3) {
        // === FAST PATH: interleaved MFMA+B loads (all 4 tiles present) ===
        int ki = 0;
        int kb0 = wid*KI+ki, keo0 = wid*(K/NW)+ki*128;
        int sd40 = keo0/(MXFP4_BLOCK_SIZE*4);
        uint32_t sa0=0x7F;
        u32x4 af0={0,0,0,0};

        // Prologue: Load A[0]
        uint4 ad[4] = {};
        if (valid) {
            const uint4* a_src = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo0+grp*32);
            #pragma unroll
            for (int j = 0; j < 4; j++) ad[j] = a_src[j];
        }
        // Load B[0] for all 4 tiles via ASM
        u32x4 bf0_u={0,0,0,0}, bf1_u={0,0,0,0}, bf2_u={0,0,0,0}, bf3_u={0,0,0,0};
        load_b_nt_u32x4(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb0*1024, lid, bf0_u);
        load_b_nt_u32x4(B_shuffle, (size_t)(ntx0+1)*KB*1024+(size_t)kb0*1024, lid, bf1_u);
        load_b_nt_u32x4(B_shuffle, (size_t)(ntx0+2)*KB*1024+(size_t)kb0*1024, lid, bf2_u);
        load_b_nt_u32x4(B_shuffle, (size_t)(ntx0+3)*KB*1024+(size_t)kb0*1024, lid, bf3_u);
        // Wait for A (oldest, vmcnt(4) = 4 B loads still in flight)
        asm volatile("s_waitcnt vmcnt(4)":::"memory");
        if (valid) {
            quant_compute_u32x4(ad, af0, sa0);
        }
        // Issue scale reads while B loads still in flight (overlap with B completion)
        uint32_t sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd40, sp1d8);
        uint32_t sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd40, sp1d8);
        uint32_t sbc = read_bscale_sh(B_scale_sh, nts2, pos, grp, sd40, sp1d8);
        uint32_t sbd = read_bscale_sh(B_scale_sh, nts3, pos, grp, sd40, sp1d8);
        // Pre-load A[1] for double-buffer
        uint4 ad_next[4] = {};
        if (valid && KI > 2) {
            int keo_next = wid*(K/NW)+1*128;
            const uint4* a_src_next = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo_next+grp*32);
            #pragma unroll
            for (int j = 0; j < 4; j++) ad_next[j] = a_src_next[j];
        }
        // Wait for B + scale loads; leave A pre-loads in flight
        if (KI > 2) {
            asm volatile("s_waitcnt vmcnt(4)":::"memory");
        } else {
            asm volatile("s_waitcnt vmcnt(0)":::"memory");
        }
        for (ki = 1; ki < KI; ki++) {
            int kb1=wid*KI+ki, keo1=wid*(K/NW)+ki*128;
            int sd41=keo1/(MXFP4_BLOCK_SIZE*4);
            // Compute B pointers for next iteration loads
            const void* p0 = (const void*)(B_shuffle + (size_t)ntx0*KB*1024 + (size_t)kb1*1024 + (size_t)lid*16);
            const void* p1 = (const void*)(B_shuffle + (size_t)(ntx0+1)*KB*1024 + (size_t)kb1*1024 + (size_t)lid*16);
            const void* p2 = (const void*)(B_shuffle + (size_t)(ntx0+2)*KB*1024 + (size_t)kb1*1024 + (size_t)lid*16);
            const void* p3 = (const void*)(B_shuffle + (size_t)(ntx0+3)*KB*1024 + (size_t)kb1*1024 + (size_t)lid*16);
            // Interleaved MFMA(old data) + B loads(new addresses)
            asm_mfma_quad_loadb(af0, sa0, bf0_u, bf1_u, bf2_u, bf3_u,
                                sba, sbb, sbc, sbd,
                                acc0, acc1, acc2, acc3,
                                p0, p1, p2, p3);
            // Wait for pre-loaded A[ki]
            asm volatile("s_waitcnt vmcnt(4)":::"memory");
            ad[0] = ad_next[0]; ad[1] = ad_next[1]; ad[2] = ad_next[2]; ad[3] = ad_next[3];
            // Quant A[ki] (VALU only, overlaps with B loads in flight)
            af0 = {0,0,0,0}; sa0 = 0x7F;
            if (valid) quant_compute_u32x4(ad, af0, sa0);
            // Issue Bscale reads early (overlap with remaining B loads)
            sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd41, sp1d8);
            sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd41, sp1d8);
            sbc = read_bscale_sh(B_scale_sh, nts2, pos, grp, sd41, sp1d8);
            sbd = read_bscale_sh(B_scale_sh, nts3, pos, grp, sd41, sp1d8);
            // Pre-load A[ki+1] (issued AFTER bscale so vmcnt(4) retains only A loads)
            ad_next[0] = {}; ad_next[1] = {}; ad_next[2] = {}; ad_next[3] = {};
            if (valid && ki + 1 < KI) {
                int keo_next = wid*(K/NW) + (ki+1)*128;
                const uint4* a_src_next = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo_next+grp*32);
                #pragma unroll
                for (int j = 0; j < 4; j++) ad_next[j] = a_src_next[j];
            }
            // Wait for B[ki] + Bscale (4 B + 4 bscale drained, 4 A newest remain)
            if (ki + 1 < KI) {
                asm volatile("s_waitcnt vmcnt(4)":::"memory");
            } else {
                asm volatile("s_waitcnt vmcnt(0)":::"memory");
            }
        }
        // Epilogue: final MFMAs (no B loads)
        asm_mfma_quad(af0, sa0, bf0_u, bf1_u, bf2_u, bf3_u,
                      sba, sbb, sbc, sbd, acc0, acc1, acc2, acc3);
    } else if (KI >= 2) {
        // Fallback for edge tiles (has3=false)
        int ki = 0;
        int kb0 = wid*KI+ki, keo0 = wid*(K/NW)+ki*128;
        int sd40 = keo0/(MXFP4_BLOCK_SIZE*4);
        uint8_t ap[16]; uint32_t sa0=0x7F;
        mfma_input_t af0={0,0,0,0,0,0,0,0};
        uint4 ad[4] = {};
        if (valid) {
            const uint4* a_src = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo0+grp*32);
            #pragma unroll
            for (int j = 0; j < 4; j++) ad[j] = a_src[j];
        }
        mfma_input_t bfa={0,0,0,0,0,0,0,0}, bfb={0,0,0,0,0,0,0,0};
        mfma_input_t bfc={0,0,0,0,0,0,0,0}, bfd={0,0,0,0,0,0,0,0};
        load_b_nt(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb0*1024, lid, bfa);
        if (has1) load_b_nt(B_shuffle, (size_t)(ntx0+1)*KB*1024+(size_t)kb0*1024, lid, bfb);
        if (has2) load_b_nt(B_shuffle, (size_t)(ntx0+2)*KB*1024+(size_t)kb0*1024, lid, bfc);
        asm volatile("s_waitcnt vmcnt(0)":::"memory");
        if (valid) {
            quant_compute_fp4(ad, ap, sa0);
            uint4 c=*reinterpret_cast<const uint4*>(ap);
            af0[0]=c.x;af0[1]=c.y;af0[2]=c.z;af0[3]=c.w;
        }
        uint32_t sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd40, sp1d8);
        uint32_t sbb = 0x7F, sbc = 0x7F;
        if (has1) sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd40, sp1d8);
        if (has2) sbc = read_bscale_sh(B_scale_sh, nts2, pos, grp, sd40, sp1d8);
        for (ki = 1; ki < KI; ki++) {
            int kb1=wid*KI+ki, keo1=wid*(K/NW)+ki*128;
            int sd41=keo1/(MXFP4_BLOCK_SIZE*4);
            ad[0] = {}; ad[1] = {}; ad[2] = {}; ad[3] = {};
            if (valid) {
                const uint4* a_src = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo1+grp*32);
                #pragma unroll
                for (int j = 0; j < 4; j++) ad[j] = a_src[j];
            }
            {
                u32x4 af_u4 = {af0[0], af0[1], af0[2], af0[3]};
                u32x4 bfu0 = {bfa[0], bfa[1], bfa[2], bfa[3]};
                u32x4 bfu1 = {bfb[0], bfb[1], bfb[2], bfb[3]};
                u32x4 bfu2 = {bfc[0], bfc[1], bfc[2], bfc[3]};
                u32x4 zero4 = {0,0,0,0};
                asm_mfma_quad(af_u4, sa0, bfu0, bfu1, bfu2, zero4,
                              sba, sbb, sbc, 0x7F, acc0, acc1, acc2, acc3);
            }
            load_b_nt(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb1*1024, lid, bfa);
            if (has1) load_b_nt(B_shuffle, (size_t)(ntx0+1)*KB*1024+(size_t)kb1*1024, lid, bfb);
            if (has2) load_b_nt(B_shuffle, (size_t)(ntx0+2)*KB*1024+(size_t)kb1*1024, lid, bfc);
            asm volatile("s_waitcnt vmcnt(0)":::"memory");
            af0={0,0,0,0,0,0,0,0}; sa0=0x7F;
            if (valid) {
                quant_compute_fp4(ad, ap, sa0);
                uint4 c=*reinterpret_cast<const uint4*>(ap);
                af0[0]=c.x;af0[1]=c.y;af0[2]=c.z;af0[3]=c.w;
            }
            sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd41, sp1d8);
            if (has1) sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd41, sp1d8);
            if (has2) sbc = read_bscale_sh(B_scale_sh, nts2, pos, grp, sd41, sp1d8);
        }
        {
            u32x4 af_u4 = {af0[0], af0[1], af0[2], af0[3]};
            u32x4 bfu0 = {bfa[0], bfa[1], bfa[2], bfa[3]};
            u32x4 bfu1 = {bfb[0], bfb[1], bfb[2], bfb[3]};
            u32x4 bfu2 = {bfc[0], bfc[1], bfc[2], bfc[3]};
            u32x4 zero4 = {0,0,0,0};
            asm_mfma_quad(af_u4, sa0, bfu0, bfu1, bfu2, zero4,
                          sba, sbb, sbc, 0x7F, acc0, acc1, acc2, acc3);
        }
    } else if (KI == 1) {
        int kb=wid, keo=wid*(K/NW), sd4=keo/(MXFP4_BLOCK_SIZE*4);
        uint8_t ap[16]; uint32_t sa=0x7F;
        mfma_input_t af={0,0,0,0,0,0,0,0};
        if (valid) {
            quant_bf16x32_to_fp4(A_bf16, mr*K+keo+grp*32, ap, sa);
            uint4 c=*reinterpret_cast<const uint4*>(ap);
            af[0]=c.x;af[1]=c.y;af[2]=c.z;af[3]=c.w;
        }
        mfma_input_t bf0={0,0,0,0,0,0,0,0};
        load_b_nt(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb*1024, lid, bf0);
        uint32_t sb0 = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd4, sp1d8);
        acc0 = mfma_fp4_16x16(af, bf0, acc0, sa, sb0);
        if (has1) {
            mfma_input_t bf1={0,0,0,0,0,0,0,0};
            load_b_nt(B_shuffle, (size_t)(ntx0+1)*KB*1024+(size_t)kb*1024, lid, bf1);
            uint32_t sb1 = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd4, sp1d8);
            acc1 = mfma_fp4_16x16(af, bf1, acc1, sa, sb1);
        }
        if (has2) {
            mfma_input_t bf2={0,0,0,0,0,0,0,0};
            load_b_nt(B_shuffle, (size_t)(ntx0+2)*KB*1024+(size_t)kb*1024, lid, bf2);
            uint32_t sb2 = read_bscale_sh(B_scale_sh, nts2, pos, grp, sd4, sp1d8);
            acc2 = mfma_fp4_16x16(af, bf2, acc2, sa, sb2);
        }
        if (has3) {
            mfma_input_t bf3={0,0,0,0,0,0,0,0};
            load_b_nt(B_shuffle, (size_t)(ntx0+3)*KB*1024+(size_t)kb*1024, lid, bf3);
            uint32_t sb3 = read_bscale_sh(B_scale_sh, nts3, pos, grp, sd4, sp1d8);
            acc3 = mfma_fp4_16x16(af, bf3, acc3, sa, sb3);
        }
    }
    // Warp reduction via LDS: 4 N-tiles x 4 warps
    extern __shared__ uint8_t lds_quad[];
    float* lr = reinterpret_cast<float*>(lds_quad);
    constexpr int NW2 = 4, T2 = 16;
    #pragma unroll
    for (int i=0;i<4;i++) lr[wid*T2*T2+(grp*4+i)*T2+pos]=acc0[i];
    #pragma unroll
    for (int i=0;i<4;i++) lr[NW2*T2*T2+wid*T2*T2+(grp*4+i)*T2+pos]=acc1[i];
    #pragma unroll
    for (int i=0;i<4;i++) lr[2*NW2*T2*T2+wid*T2*T2+(grp*4+i)*T2+pos]=acc2[i];
    #pragma unroll
    for (int i=0;i<4;i++) lr[3*NW2*T2*T2+wid*T2*T2+(grp*4+i)*T2+pos]=acc3[i];
    __syncthreads();
    // Parallel output: warp w handles tile w
    {
        const int nts_w[4] = {nts0, nts1, nts2, nts3};
        const bool has_w[4] = {true, has1, has2, has3};
        if (has_w[wid]) {
            #pragma unroll
            for (int i=0;i<4;i++) {
                int ml=grp*4+i,mg=mts+ml;
                if (mg<M) {
                    int ng=nts_w[wid]+pos;
                    if (ng<N) { float s=0; for (int w=0;w<NW2;w++) s+=lr[wid*NW2*T2*T2+w*T2*T2+ml*T2+pos]; store_nt_bf16(&output[mg*N+ng], hip_bfloat16(s)); }
                }
            }
        }
    }
}
#endif // UNIT_QUAD

#ifdef UNIT_HEX
// ---- Shape 6: Hex 96-N fused (M>=128, 6 N-tiles per CTA) with deep pipeline ----
__global__ void __launch_bounds__(256, 1)  // occ=1: 512 VGPRs, same CU util
mxfp4_fused_hex(
    const uint16_t* __restrict__ A_bf16,
    const uint8_t*  __restrict__ B_shuffle,
    const uint8_t*  __restrict__ B_scale_sh,
    hip_bfloat16*   __restrict__ output,
    int M, int N, int K, int sp1d8
) {
    __builtin_amdgcn_s_setprio(3);
    constexpr int NW = 4, T = 16, MB = 64;
    const int Kp = K / 2;
    const int KPW = Kp / NW, KI = KPW / MB, KB = Kp / MB;
    const int tid = threadIdx.x, wid = tid / 64, lid = tid % 64;
    const int ntx0 = blockIdx.x * 6;
    const int nts[6] = {ntx0*T, (ntx0+1)*T, (ntx0+2)*T, (ntx0+3)*T, (ntx0+4)*T, (ntx0+5)*T};
    const int mts = blockIdx.y * T;
    if (nts[0] >= N) return;
    const bool has[6] = {true, (nts[1]<N), (nts[2]<N), (nts[3]<N), (nts[4]<N), (nts[5]<N)};
    const int grp = lid / 16, pos = lid % 16, mr = mts + pos;
    const bool valid = (mr < M);
    mfma_acc_t acc[6] = {{0,0,0,0},{0,0,0,0},{0,0,0,0},{0,0,0,0},{0,0,0,0},{0,0,0,0}};
    if (KI >= 2) {
        int ki = 0;
        int kb0 = wid*KI+ki, keo0 = wid*(K/NW)+ki*128;
        int sd40 = keo0/(MXFP4_BLOCK_SIZE*4);
        uint32_t sa0=0x7F;
        u32x4 af0={0,0,0,0};
        // Prologue: Load A data into registers (oldest in vmcnt)
        uint4 ad[4] = {};
        if (valid) {
            const uint4* a_src = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo0+grp*32);
            #pragma unroll
            for (int j = 0; j < 4; j++) ad[j] = a_src[j];
        }
        // Issue B loads for all 6 tiles via u32x4 ASM (1 vmcnt each)
        u32x4 bf[6] = {{0,0,0,0},{0,0,0,0},{0,0,0,0},{0,0,0,0},{0,0,0,0},{0,0,0,0}};
        load_b_nt_u32x4(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb0*1024, lid, bf[0]);
        #pragma unroll
        for (int t = 1; t < 6; t++)
            if (has[t]) load_b_nt_u32x4(B_shuffle, (size_t)(ntx0+t)*KB*1024+(size_t)kb0*1024, lid, bf[t]);
        // Wait for A only (oldest), B loads overlap with quant ALU
        if (has[5]) {
            asm volatile("s_waitcnt vmcnt(6)":::"memory");
        } else {
            asm volatile("s_waitcnt vmcnt(0)":::"memory");
        }
        if (valid) quant_compute_u32x4(ad, af0, sa0);
        // Issue scale reads while B loads still in flight (overlap with B completion)
        uint32_t sb[6];
        sb[0] = read_bscale_sh(B_scale_sh, nts[0], pos, grp, sd40, sp1d8);
        #pragma unroll
        for (int t = 1; t < 6; t++) {
            sb[t] = 0x7F;
            if (has[t]) sb[t] = read_bscale_sh(B_scale_sh, nts[t], pos, grp, sd40, sp1d8);
        }
        // Pre-load A[1] for double-buffer
        uint4 ad_next[4] = {};
        if (valid && KI > 2) {
            int keo_next = wid*(K/NW)+1*128;
            const uint4* a_src_next = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo_next+grp*32);
            #pragma unroll
            for (int j = 0; j < 4; j++) ad_next[j] = a_src_next[j];
        }
        // Wait for B + scale loads to complete
        if (KI > 2) {
            asm volatile("s_waitcnt vmcnt(4)":::"memory");  // leave 4 A pre-loads in flight
        } else {
            asm volatile("s_waitcnt vmcnt(0)":::"memory");  // no A pre-load, wait for all
        }
        for (ki = 1; ki < KI; ki++) {
            int kb1=wid*KI+ki, keo1=wid*(K/NW)+ki*128;
            int sd41=keo1/(MXFP4_BLOCK_SIZE*4);
            // Compute B pointers for interleaved loads
            const void* p[6];
            p[0] = (const void*)(B_shuffle + (size_t)ntx0*KB*1024 + (size_t)kb1*1024 + (size_t)lid*16);
            #pragma unroll
            for (int t = 1; t < 6; t++)
                p[t] = (const void*)(B_shuffle + (size_t)(ntx0+t)*KB*1024 + (size_t)kb1*1024 + (size_t)lid*16);
            // Interleaved MFMA(old data) + B loads(new addresses)
            asm_mfma_hex_loadb(af0, sa0,
                               bf[0], bf[1], bf[2], bf[3], bf[4], bf[5],
                               sb[0], sb[1], sb[2], sb[3], sb[4], sb[5],
                               acc[0], acc[1], acc[2], acc[3], acc[4], acc[5],
                               p[0], p[1], p[2], p[3], p[4], p[5]);
            // Wait for pre-loaded A[ki] (4 oldest vmcnt entries behind 6 B loads)
            asm volatile("s_waitcnt vmcnt(6)":::"memory");
            ad[0] = ad_next[0]; ad[1] = ad_next[1]; ad[2] = ad_next[2]; ad[3] = ad_next[3];
            // Quant A[ki] — VALU only, overlaps with B loads in flight
            af0 = {0,0,0,0}; sa0 = 0x7F;
            if (valid) quant_compute_u32x4(ad, af0, sa0);
            // Issue Bscale reads early (overlap with B loads still in flight)
            // vmcnt FIFO: [6 B loads] → after scale reads: [6 B, 6 scale]
            sb[0] = read_bscale_sh(B_scale_sh, nts[0], pos, grp, sd41, sp1d8);
            #pragma unroll
            for (int t = 1; t < 6; t++)
                if (has[t]) sb[t] = read_bscale_sh(B_scale_sh, nts[t], pos, grp, sd41, sp1d8);
            // Pre-load A[ki+1] LAST (newest in vmcnt FIFO, stays in flight across vmcnt)
            // vmcnt FIFO: [6 B, 6 scale, 4 A_next]
            ad_next[0] = {}; ad_next[1] = {}; ad_next[2] = {}; ad_next[3] = {};
            if (valid && ki + 1 < KI) {
                int keo_next = wid*(K/NW) + (ki+1)*128;
                const uint4* a_src_next = reinterpret_cast<const uint4*>(A_bf16 + mr*K+keo_next+grp*32);
                #pragma unroll
                for (int j = 0; j < 4; j++) ad_next[j] = a_src_next[j];
            }
            // Wait for B + scales (oldest 12), leave A pre-loads (newest 4) in flight
            if (ki + 1 < KI) {
                asm volatile("s_waitcnt vmcnt(4)":::"memory");
            } else {
                asm volatile("s_waitcnt vmcnt(0)":::"memory");
            }
        }
        // Epilogue: final MFMAs (no B loads)
        asm_mfma_hex(af0, sa0,
                     bf[0], bf[1], bf[2], bf[3], bf[4], bf[5],
                     sb[0], sb[1], sb[2], sb[3], sb[4], sb[5],
                     acc[0], acc[1], acc[2], acc[3], acc[4], acc[5]);
    } else if (KI == 1) {
        int kb=wid, keo=wid*(K/NW), sd4=keo/(MXFP4_BLOCK_SIZE*4);
        uint8_t ap[16]; uint32_t sa=0x7F;
        mfma_input_t af={0,0,0,0,0,0,0,0};
        if (valid) {
            quant_bf16x32_to_fp4(A_bf16, mr*K+keo+grp*32, ap, sa);
            uint4 c=*reinterpret_cast<const uint4*>(ap);
            af[0]=c.x;af[1]=c.y;af[2]=c.z;af[3]=c.w;
        }
        #pragma unroll
        for (int t = 0; t < 6; t++) {
            if (has[t]) {
                mfma_input_t bft={0,0,0,0,0,0,0,0};
                load_b_nt(B_shuffle, (size_t)(ntx0+t)*KB*1024+(size_t)kb*1024, lid, bft);
                uint32_t sbt = read_bscale_sh(B_scale_sh, nts[t], pos, grp, sd4, sp1d8);
                acc[t] = mfma_fp4_16x16(af, bft, acc[t], sa, sbt);
            }
        }
    }
    // Warp reduction via LDS: 6 N-tiles x 4 warps
    extern __shared__ uint8_t lds_hex[];
    float* lr = reinterpret_cast<float*>(lds_hex);
    constexpr int NW2 = 4, T2 = 16;
    #pragma unroll
    for (int t = 0; t < 6; t++) {
        #pragma unroll
        for (int i=0;i<4;i++) lr[t*NW2*T2*T2+wid*T2*T2+(grp*4+i)*T2+pos]=acc[t][i];
    }
    __syncthreads();
    // Parallel reduction: warp 0 handles tiles 0,1; warp 1 handles 2,3; warp 2 handles 4,5
    if (wid <= 2) {
        int t0 = wid * 2, t1 = t0 + 1;
        #pragma unroll
        for (int i=0;i<4;i++) {
            int ml=grp*4+i,mg=mts+ml;
            if (mg<M) {
                int ng0=nts[t0]+pos;
                if (has[t0] && ng0<N) { float s=0; for (int w=0;w<NW2;w++) s+=lr[t0*NW2*T2*T2+w*T2*T2+ml*T2+pos]; store_nt_bf16(&output[mg*N+ng0], hip_bfloat16(s)); }
                int ng1=nts[t1]+pos;
                if (has[t1] && ng1<N) { float s=0; for (int w=0;w<NW2;w++) s+=lr[t1*NW2*T2*T2+w*T2*T2+ml*T2+pos]; store_nt_bf16(&output[mg*N+ng1], hip_bfloat16(s)); }
            }
        }
    }
}
#endif // UNIT_HEX

#ifdef UNIT_SPLITK
// ---- Shape 2: Fused quant + wide split-K GEMM with 2-stage deep pipeline (from v104) ----
__global__ void __launch_bounds__(256, 2)
mxfp4_gemm_splitk_fused(
    const uint16_t* __restrict__ A_bf16,
    const uint8_t*  __restrict__ B_shuffle,
    const uint8_t*  __restrict__ B_scale_sh,
    float*          __restrict__ partial_out,
    int M, int N, int K, int sp1d8, int ksplit
) {
    __builtin_amdgcn_s_setprio(3);
    constexpr int NW = 4, T = 16, MB = 64;
    const int kslice = blockIdx.z;
    const int Kp = K / 2;
    const int KB = Kp / MB;
    const int K_slice = K / ksplit;
    const int Kp_slice = Kp / ksplit;
    const int KPW = Kp_slice / NW, KI = KPW / MB;
    const int ke_base = kslice * K_slice;
    const int kb_base = kslice * (Kp_slice / MB);
    const int tid = threadIdx.x, wid = tid / 64, lid = tid % 64;
    const int ntx0 = blockIdx.x * 2, ntx1 = ntx0 + 1;
    const int nts0 = ntx0 * T, nts1 = ntx1 * T;
    const int mts = blockIdx.y * T;
    if (nts0 >= N) return;
    const bool has_sub1 = (nts1 < N);
    const int grp = lid / 16, pos = lid % 16, mr = mts + pos;
    const bool valid = (mr < M);

    mfma_acc_t acc0 = {0,0,0,0}, acc1 = {0,0,0,0};

    if (KI == 2) {
        // ============ 2-STAGE DEEP PIPELINE for KI==2 ============
        int kb0 = kb_base + wid*2;
        int kb1 = kb0 + 1;
        int keo0 = ke_base + wid*(K_slice/NW);
        int keo1 = keo0 + 128;
        int sd40 = keo0/(MXFP4_BLOCK_SIZE*4);
        int sd41 = keo1/(MXFP4_BLOCK_SIZE*4);

        // Load A data for BOTH iterations upfront (oldest in vmcnt)
        uint4 ad0[4] = {}, ad1[4] = {};
        if (valid) {
            const uint4* a_src0 = reinterpret_cast<const uint4*>(A_bf16 + mr*K + keo0 + grp*32);
            #pragma unroll
            for (int j = 0; j < 4; j++) ad0[j] = a_src0[j];
            const uint4* a_src1 = reinterpret_cast<const uint4*>(A_bf16 + mr*K + keo1 + grp*32);
            #pragma unroll
            for (int j = 0; j < 4; j++) ad1[j] = a_src1[j];
        }

        // Pre-read all Bscale values upfront (avoids compiler vmcnt interference later)
        uint32_t sba0 = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd40, sp1d8);
        uint32_t sbb0 = 0x7F;
        if (has_sub1) sbb0 = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd40, sp1d8);
        uint32_t sba1 = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd41, sp1d8);
        uint32_t sbb1 = 0x7F;
        if (has_sub1) sbb1 = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd41, sp1d8);

        // Issue B loads for BOTH iterations via u32x4 (1 vmcnt each, 4 total or 2 if no sub1)
        u32x4 bfa0={0,0,0,0}, bfb0={0,0,0,0};
        u32x4 bfa1={0,0,0,0}, bfb1={0,0,0,0};
        load_b_nt_u32x4(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb0*1024, lid, bfa0);
        if (has_sub1) load_b_nt_u32x4(B_shuffle, (size_t)ntx1*KB*1024+(size_t)kb0*1024, lid, bfb0);
        load_b_nt_u32x4(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb1*1024, lid, bfa1);
        if (has_sub1) load_b_nt_u32x4(B_shuffle, (size_t)ntx1*KB*1024+(size_t)kb1*1024, lid, bfb1);

        // Wait for A loads only (oldest) — B loads remain in flight
        if (has_sub1) {
            asm volatile("s_waitcnt vmcnt(4)":::"memory");
        } else {
            asm volatile("s_waitcnt vmcnt(2)":::"memory");
        }

        // Compute quant iter 0 directly to u32x4 (no stack intermediary)
        uint32_t sa0 = 0x7F;
        u32x4 af0 = {0,0,0,0};
        if (valid) {
            quant_compute_u32x4(ad0, af0, sa0);
        }

        // Wait for iter 0 B loads (2 if has_sub1, 1 if not)
        if (has_sub1) {
            asm volatile("s_waitcnt vmcnt(2)":::"memory");
        } else {
            asm volatile("s_waitcnt vmcnt(1)":::"memory");
        }

        // ASM MFMA pair iter 0 (Bscale already in registers from pre-read)
        asm_mfma_pair(af0, sa0, bfa0, bfb0, sba0, sbb0, acc0, acc1);

        // Wait for iter 1 B loads (should be ready by now)
        asm volatile("s_waitcnt vmcnt(0)":::"memory");

        // Compute quant iter 1 directly to u32x4 (no stack intermediary)
        uint32_t sa1 = 0x7F;
        u32x4 af1 = {0,0,0,0};
        if (valid) {
            quant_compute_u32x4(ad1, af1, sa1);
        }

        // ASM MFMA pair iter 1 (Bscale already in registers from pre-read)
        asm_mfma_pair(af1, sa1, bfa1, bfb1, sba1, sbb1, acc0, acc1);

    } else if (KI >= 3) {
        // General case
        int ki = 0;
        int kb0 = kb_base+wid*KI+ki;
        int keo0 = ke_base+wid*(K_slice/NW)+ki*128;
        int sd40 = keo0/(MXFP4_BLOCK_SIZE*4);
        uint8_t ap[16]; uint32_t sa0=0x7F;
        mfma_input_t af0={0,0,0,0,0,0,0,0};
        if (valid) {
            quant_bf16x32_to_fp4(A_bf16, mr*K+keo0+grp*32, ap, sa0);
            uint4 c=*reinterpret_cast<const uint4*>(ap);
            af0[0]=c.x;af0[1]=c.y;af0[2]=c.z;af0[3]=c.w;
        }
        mfma_input_t bfa={0,0,0,0,0,0,0,0}, bfb={0,0,0,0,0,0,0,0};
        load_b_nt(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb0*1024, lid, bfa);
        uint32_t sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd40, sp1d8);
        uint32_t sbb = 0x7F;
        if (has_sub1) {
            load_b_nt(B_shuffle, (size_t)ntx1*KB*1024+(size_t)kb0*1024, lid, bfb);
            sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd40, sp1d8);
        }
        for (ki = 1; ki < KI; ki++) {
            int kb1=kb_base+wid*KI+ki;
            int keo1=ke_base+wid*(K_slice/NW)+ki*128;
            int sd41=keo1/(MXFP4_BLOCK_SIZE*4);
#if defined(__gfx950__)
            asm volatile("s_setprio 2" ::: "memory");
#endif
            acc0 = mfma_fp4_16x16(af0, bfa, acc0, sa0, sba);
            if (has_sub1) acc1 = mfma_fp4_16x16(af0, bfb, acc1, sa0, sbb);
#if defined(__gfx950__)
            asm volatile("s_setprio 0" ::: "memory");
#endif
            af0={0,0,0,0,0,0,0,0}; sa0=0x7F;
            if (valid) {
                quant_bf16x32_to_fp4(A_bf16, mr*K+keo1+grp*32, ap, sa0);
                uint4 c=*reinterpret_cast<const uint4*>(ap);
                af0[0]=c.x;af0[1]=c.y;af0[2]=c.z;af0[3]=c.w;
            }
            load_b_nt(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb1*1024, lid, bfa);
            if (has_sub1) load_b_nt(B_shuffle, (size_t)ntx1*KB*1024+(size_t)kb1*1024, lid, bfb);
            asm volatile("s_waitcnt vmcnt(0)":::"memory");
            sba = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd41, sp1d8);
            if (has_sub1) sbb = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd41, sp1d8);
        }
#if defined(__gfx950__)
        asm volatile("s_setprio 2" ::: "memory");
#endif
        acc0 = mfma_fp4_16x16(af0, bfa, acc0, sa0, sba);
        if (has_sub1) acc1 = mfma_fp4_16x16(af0, bfb, acc1, sa0, sbb);
#if defined(__gfx950__)
        asm volatile("s_setprio 0" ::: "memory");
#endif
    } else if (KI == 1) {
        int kb=kb_base+wid;
        int keo=ke_base+wid*(K_slice/NW);
        int sd4=keo/(MXFP4_BLOCK_SIZE*4);
        uint8_t ap[16]; uint32_t sa=0x7F;
        mfma_input_t af={0,0,0,0,0,0,0,0};
        if (valid) {
            quant_bf16x32_to_fp4(A_bf16, mr*K+keo+grp*32, ap, sa);
            uint4 c=*reinterpret_cast<const uint4*>(ap);
            af[0]=c.x;af[1]=c.y;af[2]=c.z;af[3]=c.w;
        }
        mfma_input_t bf0={0,0,0,0,0,0,0,0};
        load_b_nt(B_shuffle, (size_t)ntx0*KB*1024+(size_t)kb*1024, lid, bf0);
        uint32_t sb0 = read_bscale_sh(B_scale_sh, nts0, pos, grp, sd4, sp1d8);
#if defined(__gfx950__)
        asm volatile("s_setprio 2" ::: "memory");
#endif
        acc0 = mfma_fp4_16x16(af, bf0, acc0, sa, sb0);
#if defined(__gfx950__)
        asm volatile("s_setprio 0" ::: "memory");
#endif
        if (has_sub1) {
            mfma_input_t bf1={0,0,0,0,0,0,0,0};
            load_b_nt(B_shuffle, (size_t)ntx1*KB*1024+(size_t)kb*1024, lid, bf1);
            uint32_t sb1 = read_bscale_sh(B_scale_sh, nts1, pos, grp, sd4, sp1d8);
#if defined(__gfx950__)
            asm volatile("s_setprio 2" ::: "memory");
#endif
            acc1 = mfma_fp4_16x16(af, bf1, acc1, sa, sb1);
#if defined(__gfx950__)
            asm volatile("s_setprio 0" ::: "memory");
#endif
        }
    }

    extern __shared__ uint8_t lds[];
    float* lr = reinterpret_cast<float*>(lds);
    #pragma unroll
    for (int i=0;i<4;i++) lr[wid*T*T+(grp*4+i)*T+pos]=acc0[i];
    #pragma unroll
    for (int i=0;i<4;i++) lr[NW*T*T+wid*T*T+(grp*4+i)*T+pos]=acc1[i];
    __syncthreads();
    {
        int ml=wid*4+grp, mg=mts+ml, ng=nts0+pos;
        if (mg<M&&ng<N) {
            float s=0;
            #pragma unroll
            for (int w=0;w<NW;w++) s+=lr[w*T*T+ml*T+pos];
            partial_out[kslice*M*N+mg*N+ng] = s;
        }
        if (has_sub1) {
            ng=nts1+pos;
            if (mg<M&&ng<N) {
                float s=0;
                #pragma unroll
                for (int w=0;w<NW;w++) s+=lr[NW*T*T+w*T*T+ml*T+pos];
                partial_out[kslice*M*N+mg*N+ng] = s;
            }
        }
    }
}

__global__ void __launch_bounds__(256)
reduce_splitk(
    const float* __restrict__ partial_in,
    hip_bfloat16* __restrict__ output,
    int M, int N, int ksplit
) {
    const int idx = blockIdx.x * 256 + threadIdx.x;
    const int total = M * N;
    if (idx >= total) return;
    float s = 0.0f;
    #pragma unroll 7
    for (int k = 0; k < ksplit; k++) {
        s += partial_in[k * total + idx];
    }
    store_nt_bf16(&output[idx], hip_bfloat16(s));
}
#endif // UNIT_SPLITK

#ifdef UNIT_NARROW
extern "C" void ct_fused_narrow(void* A, void* B, void* Bs, void* out,
                                 int M, int N, int K, int sp1d8) {
    int nt=(N+15)/16, mt=(M+15)/16;
    hipLaunchKernelGGL(mxfp4_fused_narrow, dim3(nt,mt), dim3(256), 4*16*16*4, 0,
        (const uint16_t*)A, (const uint8_t*)B, (const uint8_t*)Bs,
        (hip_bfloat16*)out, M, N, K, sp1d8);
}
#endif // UNIT_NARROW

#ifdef UNIT_WIDE
extern "C" void ct_fused_wide(void* A, void* B, void* Bs, void* out,
                               int M, int N, int K, int sp1d8) {
    int nt=(N+31)/32, mt=(M+15)/16;
    hipLaunchKernelGGL(mxfp4_fused_wide, dim3(nt,mt), dim3(256), 2*4*16*16*4, 0,
        (const uint16_t*)A, (const uint8_t*)B, (const uint8_t*)Bs,
        (hip_bfloat16*)out, M, N, K, sp1d8);
}
#endif // UNIT_WIDE

#ifdef UNIT_QUAD
extern "C" void ct_fused_quad(void* A, void* B, void* Bs, void* out,
                               int M, int N, int K, int sp1d8) {
    int nt=(N+63)/64, mt=(M+15)/16;
    hipLaunchKernelGGL(mxfp4_fused_quad, dim3(nt,mt), dim3(256), 4*4*16*16*4, 0,
        (const uint16_t*)A, (const uint8_t*)B, (const uint8_t*)Bs,
        (hip_bfloat16*)out, M, N, K, sp1d8);
}
#endif // UNIT_QUAD

#ifdef UNIT_SPLITK
extern "C" void ct_splitk_fused(void* A, void* B, void* Bs,
                                 void* partial, void* out,
                                 int M, int N, int K, int sp1d8, int ksplit) {
    int nt=(N+31)/32, mt=(M+15)/16;
    hipLaunchKernelGGL(mxfp4_gemm_splitk_fused, dim3(nt,mt,ksplit), dim3(256), 2*4*16*16*4, 0,
        (const uint16_t*)A, (const uint8_t*)B,
        (const uint8_t*)Bs, (float*)partial, M, N, K, sp1d8, ksplit);
    int total = M * N;
    int nblk = (total + 255) / 256;
    hipLaunchKernelGGL(reduce_splitk, dim3(nblk), dim3(256), 0, 0,
        (const float*)partial, (hip_bfloat16*)out, M, N, ksplit);
}
#endif // UNIT_SPLITK

#ifdef UNIT_HEX
extern "C" void ct_fused_hex(void* A, void* B, void* Bs, void* out,
                               int M, int N, int K, int sp1d8) {
    int nt=(N+95)/96, mt=(M+15)/16;
    hipLaunchKernelGGL(mxfp4_fused_hex, dim3(nt,mt), dim3(256), 6*4*16*16*4, 0,
        (const uint16_t*)A, (const uint8_t*)B, (const uint8_t*)Bs,
        (hip_bfloat16*)out, M, N, K, sp1d8);
}
#endif // UNIT_HEX

"""

CPP_NARROW = r"""
#include <torch/extension.h>
extern "C" void ct_fused_narrow(void*, void*, void*, void*, int, int, int, int);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {}
"""

CPP_NARROW_HEX = r"""
#include <torch/extension.h>
extern "C" void ct_fused_narrow(void*, void*, void*, void*, int, int, int, int);
extern "C" void ct_fused_hex(void*, void*, void*, void*, int, int, int, int);
using fn_t = void(*)(void*,void*,void*,void*,int,int,int,int);
struct Entry { fn_t fn; void* out; int M, N, K, sp; };
static Entry g_e[8];
void setup(int id, int64_t fn, int64_t out, int M, int N, int K, int sp) {
    g_e[id] = {(fn_t)fn, (void*)out, M, N, K, sp};
}
void dispatch(int id, const torch::Tensor& A, const torch::Tensor& B, const torch::Tensor& Bs) {
    auto& e = g_e[id];
    e.fn(A.data_ptr(), B.data_ptr(), Bs.data_ptr(), e.out, e.M, e.N, e.K, e.sp);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("setup", &setup);
    m.def("dispatch", &dispatch);
}
"""

CPP_WIDE = r"""
#include <torch/extension.h>
extern "C" void ct_fused_wide(void*, void*, void*, void*, int, int, int, int);
using fn_t = void(*)(void*,void*,void*,void*,int,int,int,int);
struct Entry { fn_t fn; void* out; int M, N, K, sp; };
static Entry g_e[8];
void setup(int id, int64_t fn, int64_t out, int M, int N, int K, int sp) {
    g_e[id] = {(fn_t)fn, (void*)out, M, N, K, sp};
}
void dispatch(int id, const torch::Tensor& A, const torch::Tensor& B, const torch::Tensor& Bs) {
    auto& e = g_e[id];
    e.fn(A.data_ptr(), B.data_ptr(), Bs.data_ptr(), e.out, e.M, e.N, e.K, e.sp);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("setup", &setup);
    m.def("dispatch", &dispatch);
}
"""

CPP_QUAD = r"""
#include <torch/extension.h>
extern "C" void ct_fused_quad(void*, void*, void*, void*, int, int, int, int);
using fn_t = void(*)(void*,void*,void*,void*,int,int,int,int);
struct Entry { fn_t fn; void* out; int M, N, K, sp; };
static Entry g_e[8];
void setup(int id, int64_t fn, int64_t out, int M, int N, int K, int sp) {
    g_e[id] = {(fn_t)fn, (void*)out, M, N, K, sp};
}
void dispatch(int id, const torch::Tensor& A, const torch::Tensor& B, const torch::Tensor& Bs) {
    auto& e = g_e[id];
    e.fn(A.data_ptr(), B.data_ptr(), Bs.data_ptr(), e.out, e.M, e.N, e.K, e.sp);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("setup", &setup);
    m.def("dispatch", &dispatch);
}
"""

CPP_SPLITK = r"""
#include <torch/extension.h>
extern "C" void ct_splitk_fused(void*, void*, void*, void*, void*, int, int, int, int, int);
using fn_t = void(*)(void*,void*,void*,void*,void*,int,int,int,int,int);
struct Entry { fn_t fn; void* part; void* out; int M, N, K, sp, ks; };
static Entry g_e[8];
void setup(int id, int64_t fn, int64_t part, int64_t out, int M, int N, int K, int sp, int ks) {
    g_e[id] = {(fn_t)fn, (void*)part, (void*)out, M, N, K, sp, ks};
}
void dispatch(int id, const torch::Tensor& A, const torch::Tensor& B, const torch::Tensor& Bs) {
    auto& e = g_e[id];
    e.fn(A.data_ptr(), B.data_ptr(), Bs.data_ptr(), e.part, e.out, e.M, e.N, e.K, e.sp, e.ks);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("setup", &setup);
    m.def("dispatch", &dispatch);
}
"""

CPP_HEX = r"""
#include <torch/extension.h>
extern "C" void ct_fused_hex(void*, void*, void*, void*, int, int, int, int);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {}
"""

_mod_narrow_hex = None  # pybind11 module (has .setup/.dispatch)
_mod_wide = None
_mod_quad = None
_mod_splitk = None
_output_cache = {}
_partial_cache = {}
_b_cache = {}

# dlsym for getting kernel function addresses
_libc = ctypes.CDLL(None)
_libc.dlsym.argtypes = [ctypes.c_void_p, ctypes.c_char_p]
_libc.dlsym.restype = ctypes.c_void_p

_BASE_FLAGS = [
    "-O3", "--offload-arch=gfx950",
    "-mcumode",
    "-mllvm", "-amdgpu-early-inline-all=true",
    "-mllvm", "--amdgpu-kernarg-preload-count=16",
    "-DHIP_HCC_COMPAT_MODE=1",
]


def _compile_unit(name, cpp_src, defines, extra_flags=None):
    from torch.utils.cpp_extension import load_inline
    if isinstance(defines, str):
        defines = [defines]
    flags = _BASE_FLAGS + [f"-D{d}" for d in defines]
    if extra_flags:
        flags += extra_flags
    mod = load_inline(
        name=name, cpp_sources=cpp_src, cuda_sources=HIP_SOURCE,
        extra_cuda_cflags=flags,
        verbose=False,
    )
    # Return (pybind11_module, ctypes_handle_for_dlsym)
    return mod, ctypes.CDLL(mod.__file__)


def _get_modules():
    global _mod_narrow_hex, _mod_wide, _mod_quad, _mod_splitk
    if _mod_narrow_hex is not None:
        return
    inline_flags = ["-mllvm", "-amdgpu-function-calls=false"]
    with ThreadPoolExecutor(max_workers=4) as pool:
        f_nh = pool.submit(_compile_unit, "mxfp4_nh_v1401", CPP_NARROW_HEX,
                           ["UNIT_NARROW", "UNIT_HEX"], inline_flags)
        f_w = pool.submit(_compile_unit, "mxfp4_w_v1401", CPP_WIDE,
                          "UNIT_WIDE", inline_flags)
        f_q = pool.submit(_compile_unit, "mxfp4_q_v1401", CPP_QUAD,
                          "UNIT_QUAD", inline_flags)
        f_s = pool.submit(_compile_unit, "mxfp4_s_v1401", CPP_SPLITK,
                          "UNIT_SPLITK", inline_flags)
        _mod_narrow_hex = f_nh.result()
        _mod_wide = f_w.result()
        _mod_quad = f_q.result()
        _mod_splitk = f_s.result()


def _fn_addr(mod_ct, name):
    """Get function pointer address via dlsym."""
    return _libc.dlsym(mod_ct._handle, name.encode())


_dispatch_cache = {}  # (M, N, K) -> (mod_dispatch, shape_id, out_tensor)
_next_shape_id = [0]  # mutable counter


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]

    key = (M, N, K)
    entry = _dispatch_cache.get(key)
    if entry is not None:
        mod_dispatch, sid, out = entry
        mod_dispatch(sid, A, B_shuffle, B_scale_sh)
        return out

    # Cold path: first call for this shape — set up cached dispatch
    _get_modules()
    sp1d8 = (K // 32 + 7) // 8
    out = torch.empty(M, N, dtype=torch.bfloat16, device=A.device)
    o_ptr = out.data_ptr()
    sid = _next_shape_id[0]
    _next_shape_id[0] += 1

    if M >= 128:
        mod_py, mod_ct = _mod_narrow_hex
        fn = _fn_addr(mod_ct, "ct_fused_hex")
        mod_py.setup(sid, fn, o_ptr, M, N, K, sp1d8)
    elif M >= 64:
        mod_py, mod_ct = _mod_quad
        fn = _fn_addr(mod_ct, "ct_fused_quad")
        mod_py.setup(sid, fn, o_ptr, M, N, K, sp1d8)
    elif K > 4096:
        ksplit = 7 if K == 7168 else 2
        partial = torch.empty(ksplit, M, N, dtype=torch.float32, device=A.device)
        _partial_cache[(ksplit, M, N)] = partial
        mod_py, mod_ct = _mod_splitk
        fn = _fn_addr(mod_ct, "ct_splitk_fused")
        mod_py.setup(sid, fn, partial.data_ptr(), o_ptr, M, N, K, sp1d8, ksplit)
    elif M <= 8:
        mod_py, mod_ct = _mod_narrow_hex
        fn = _fn_addr(mod_ct, "ct_fused_narrow")
        mod_py.setup(sid, fn, o_ptr, M, N, K, sp1d8)
    else:
        mod_py, mod_ct = _mod_wide
        fn = _fn_addr(mod_ct, "ct_fused_wide")
        mod_py.setup(sid, fn, o_ptr, M, N, K, sp1d8)

    _dispatch_cache[key] = (mod_py.dispatch, sid, out)
    mod_py.dispatch(sid, A, B_shuffle, B_scale_sh)
    return out
scrolls · 1617 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