Skip to content
KernelIndex
Search⌘K

submission 729970

aswinkumar · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:cfc8ea814dcc3a8216a43bcd0c6a0d0c8b6b32d0750d529c859694580861459b
license declaredunknown
license concludedunknown
authorsaswinkumar
imported2026-08-15

Techniques

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

fp4__device__ __forceinline__ void quant_a(uint32_t ap[16], v4i32& fp4, uint32_t& e8m0) {
shared-memory__shared__ float lds[4 * 256];
split-k__global__ void kern_splitk(
vector-width = uint4__device__ __forceinline__ uint4 lds_load_u128(lds_u32_ptr ptr) {

Kernel source

submission_v90.py1862 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
v99_afirst: v98 + A-first load reorder. Issue A loads → GLL+Bsc → quant A.
Pre-barrier: pairs of (buffer_load B[s], global_load Bsc[s]) for predictable vmcnt drain.
Landmarks at key vmcnt points for later binary patching.
Pure C++ via <<<>>>.
"""

import os
os.environ.setdefault('PYTORCH_ROCM_ARCH', 'gfx950')

import torch
from torch.utils.cpp_extension import load_inline

input_t = tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]
output_t = torch.Tensor

HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <torch/extension.h>
namespace py = pybind11;

typedef int   v4i32  __attribute__((__vector_size__(16)));
typedef float v4f32  __attribute__((__vector_size__(16)));
typedef const __attribute__((address_space(3))) uint32_t* lds_u32_ptr;

__device__ __forceinline__ float bf16_to_f32(uint16_t x) {
    return __uint_as_float((uint32_t)x << 16);
}
__device__ __forceinline__ uint16_t f32_to_bf16(float x) {
    uint32_t u = __float_as_uint(x);
    u += 0x7FFFu + ((u >> 16) & 1u);
    return (uint16_t)(u >> 16);
}
__device__ __forceinline__ uint4 lds_load_u128(lds_u32_ptr ptr) {
    uint4 value;
    asm volatile("ds_read_b128 %0, %1"
                 : "=v"(value)
                 : "v"(ptr)
                 : "memory");
    return value;
}
__device__ __forceinline__ int bsc_idx(int bn, int bkg, int pk) {
    return (bn>>5)*(pk*32)+(bkg>>3)*256+(bkg&3)*64+(bn&15)*4+((bkg&7)>>2)*2+((bn&31)>>4);
}
__device__ __forceinline__ void store_tile(
    uint16_t* C, v4f32& a, int mt, int nt, int gid, int lid, int M, int N
) {
    int on=nt*16+lid, om=mt*16+gid*4;
    if(on<N){ for(int r=0;r<4;r++) if(om+r<M) C[(om+r)*N+on]=f32_to_bf16(a[r]); }
}
__device__ __forceinline__ void store_tile_full(
    uint16_t* C, v4f32& a, int mt, int nt, int gid, int lid, int N
) {
    int on=nt*16+lid, om=mt*16+gid*4;
    C[(om+0)*N+on]=f32_to_bf16(a[0]);
    C[(om+1)*N+on]=f32_to_bf16(a[1]);
    C[(om+2)*N+on]=f32_to_bf16(a[2]);
    C[(om+3)*N+on]=f32_to_bf16(a[3]);
}

// --- Original quant (still used by kern_fused for non-LDS path) ---
__device__ __forceinline__ void quant_a(uint32_t ap[16], v4i32& fp4, uint32_t& e8m0) {
    uint32_t mx[16];
    for(int i=0;i<16;i++) mx[i]=ap[i]&0x7FFF7FFFu;
    for(int i=0;i<8;i++) asm volatile("v_pk_max_u16 %0,%1,%2":"=v"(mx[i]):"v"(mx[i]),"v"(mx[i+8]));
    for(int i=0;i<4;i++) asm volatile("v_pk_max_u16 %0,%1,%2":"=v"(mx[i]):"v"(mx[i]),"v"(mx[i+4]));
    for(int i=0;i<2;i++) asm volatile("v_pk_max_u16 %0,%1,%2":"=v"(mx[i]):"v"(mx[i]),"v"(mx[i+2]));
    asm volatile("v_pk_max_u16 %0,%1,%2":"=v"(mx[0]):"v"(mx[0]),"v"(mx[1]));
    uint32_t h=mx[0]>>16,l=mx[0]&0xFFFFu; uint16_t am=(uint16_t)(h>l?h:l);
    float sf;
    if(am==0u){e8m0=0u;sf=0.0f;}
    else{unsigned ai=__float_as_uint(bf16_to_f32(am));ai=(ai+0x200000u)&0xFF800000u;
        int su=(int)((ai>>23)&0xFFu)-129;su=max(-127,min(127,su));
        e8m0=(uint32_t)(su+127);sf=__uint_as_float(e8m0<<23);}
    for(int v=0;v<4;v++){uint32_t r=0;
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2":"+v"(r):"v"(ap[v*4+0]),"v"(sf));
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,0]":"+v"(r):"v"(ap[v*4+1]),"v"(sf));
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,0,1]":"+v"(r):"v"(ap[v*4+2]),"v"(sf));
        asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,1]":"+v"(r):"v"(ap[v*4+3]),"v"(sf));
        fp4[v]=(int)r;}
}

// --- Optimized quant: takes 4 uint4 directly, no intermediate array ---
__device__ __forceinline__ void quant_a_direct(
    uint4 c0, uint4 c1, uint4 c2, uint4 c3,
    v4i32& fp4, uint32_t& e8m0
) {
    // Absolute value (clear sign bits of packed bf16)
    const uint32_t ABSMASK = 0x7FFF7FFFu;
    uint32_t m0=c0.x&ABSMASK, m1=c0.y&ABSMASK, m2=c0.z&ABSMASK, m3=c0.w&ABSMASK;
    uint32_t m4=c1.x&ABSMASK, m5=c1.y&ABSMASK, m6=c1.z&ABSMASK, m7=c1.w&ABSMASK;
    uint32_t m8=c2.x&ABSMASK, m9=c2.y&ABSMASK, mA=c2.z&ABSMASK, mB=c2.w&ABSMASK;
    uint32_t mC=c3.x&ABSMASK, mD=c3.y&ABSMASK, mE=c3.z&ABSMASK, mF=c3.w&ABSMASK;

    // Tree reduction: 16→8→4→2→1 using pk_max_u16
    // NOTE: "+v" (read-write) because %0 is both source and destination
    asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m0):"v"(m8));
    asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m1):"v"(m9));
    asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m2):"v"(mA));
    asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m3):"v"(mB));
    asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m4):"v"(mC));
    asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m5):"v"(mD));
    asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m6):"v"(mE));
    asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m7):"v"(mF));
    asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m0):"v"(m4));
    asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m1):"v"(m5));
    asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m2):"v"(m6));
    asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m3):"v"(m7));
    asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m0):"v"(m2));
    asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m1):"v"(m3));
    asm volatile("v_pk_max_u16 %0,%0,%1":"+v"(m0):"v"(m1));

    // Extract scalar max from packed result
    uint32_t h=m0>>16, l=m0&0xFFFFu;
    uint16_t am=(uint16_t)(h>l?h:l);

    // Compute E8M0 scale
    float sf;
    if(am==0u){e8m0=0u;sf=0.0f;}
    else{unsigned ai=__float_as_uint(bf16_to_f32(am));ai=(ai+0x200000u)&0xFF800000u;
        int su=(int)((ai>>23)&0xFFu)-129;su=max(-127,min(127,su));
        e8m0=(uint32_t)(su+127);sf=__uint_as_float(e8m0<<23);}

    // Convert bf16 pairs to FP4 — operate directly on uint4 fields
    uint32_t r0=0, r1=0, r2=0, r3=0;
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2":"+v"(r0):"v"(c0.x),"v"(sf));
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,0]":"+v"(r0):"v"(c0.y),"v"(sf));
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,0,1]":"+v"(r0):"v"(c0.z),"v"(sf));
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,1]":"+v"(r0):"v"(c0.w),"v"(sf));
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2":"+v"(r1):"v"(c1.x),"v"(sf));
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,0]":"+v"(r1):"v"(c1.y),"v"(sf));
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,0,1]":"+v"(r1):"v"(c1.z),"v"(sf));
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,1]":"+v"(r1):"v"(c1.w),"v"(sf));
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2":"+v"(r2):"v"(c2.x),"v"(sf));
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,0]":"+v"(r2):"v"(c2.y),"v"(sf));
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,0,1]":"+v"(r2):"v"(c2.z),"v"(sf));
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,1]":"+v"(r2):"v"(c2.w),"v"(sf));
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2":"+v"(r3):"v"(c3.x),"v"(sf));
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,0]":"+v"(r3):"v"(c3.y),"v"(sf));
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,0,1]":"+v"(r3):"v"(c3.z),"v"(sf));
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2 op_sel:[0,0,1,1]":"+v"(r3):"v"(c3.w),"v"(sf));
    fp4[0]=(int)r0; fp4[1]=(int)r1; fp4[2]=(int)r2; fp4[3]=(int)r3;
}

// ==========================================================================
// KERNEL 1: Fused quant+GEMM (single wave, pipelined, for K<=1024)
// ==========================================================================
__launch_bounds__(64)
__global__ void kern_fused(
    const uint32_t* __restrict__ A, const uint8_t* __restrict__ Bs,
    const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
    const int M, const int N, const int K, const int kt, const int pk
) {
    const int lane=threadIdx.x, mt=blockIdx.x, nt=blockIdx.y;
    const int lid=lane&15, gid=lane>>4, mg=mt*16+lid, Kh=K>>1;
    const int bn=nt*16+lid, bb=nt*(kt*512);
    v4f32 acc={0,0,0,0};
    uint32_t ap[16];
    if(mg<M){const uint4* av=reinterpret_cast<const uint4*>(A+mg*Kh+(gid*32>>1));
        uint4 c0=av[0],c1=av[1],c2=av[2],c3=av[3];
        ap[0]=c0.x;ap[1]=c0.y;ap[2]=c0.z;ap[3]=c0.w;ap[4]=c1.x;ap[5]=c1.y;ap[6]=c1.z;ap[7]=c1.w;
        ap[8]=c2.x;ap[9]=c2.y;ap[10]=c2.z;ap[11]=c2.w;ap[12]=c3.x;ap[13]=c3.y;ap[14]=c3.z;ap[15]=c3.w;
    } else { for(int i=0;i<16;i++) ap[i]=0u; }
    v4i32 af; uint32_t ae; quant_a(ap,af,ae);
    uint4 bc=*reinterpret_cast<const uint4*>(Bs+bb+(gid/2)*512+(gid&1)*256+lid*16);
    v4i32 bf; bf[0]=(int)bc.x;bf[1]=(int)bc.y;bf[2]=(int)bc.z;bf[3]=(int)bc.w;
    uint32_t be=(uint32_t)Bsc[bsc_idx(bn,gid,pk)];
    for(int ki=1,Ki=K>>7;ki<Ki;ki++){int kb=ki<<7;
        uint32_t apn[16];
        if(mg<M){const uint4* avn=reinterpret_cast<const uint4*>(A+mg*Kh+((kb+gid*32)>>1));
            uint4 d0=avn[0],d1=avn[1],d2=avn[2],d3=avn[3];
            apn[0]=d0.x;apn[1]=d0.y;apn[2]=d0.z;apn[3]=d0.w;apn[4]=d1.x;apn[5]=d1.y;apn[6]=d1.z;apn[7]=d1.w;
            apn[8]=d2.x;apn[9]=d2.y;apn[10]=d2.z;apn[11]=d2.w;apn[12]=d3.x;apn[13]=d3.y;apn[14]=d3.z;apn[15]=d3.w;
        } else { for(int i=0;i<16;i++) apn[i]=0u; }
        uint4 bcn=*reinterpret_cast<const uint4*>(Bs+bb+(kb/64+gid/2)*512+(gid&1)*256+lid*16);
        v4i32 bfn; bfn[0]=(int)bcn.x;bfn[1]=(int)bcn.y;bfn[2]=(int)bcn.z;bfn[3]=(int)bcn.w;
        uint32_t ben=(uint32_t)Bsc[bsc_idx(bn,kb/32+gid,pk)];
        asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
            :"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));
        quant_a(apn,af,ae); bf=bfn; be=ben;
    }
    asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
        :"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));
    store_tile(C,acc,mt,nt,gid,lid,M,N);
}

// ==========================================================================
// KERNEL 1b: Template-specialized fused quant+GEMM
// Template params: cM, cN, cK are compile-time shape constants.
// Compiler can constant-fold Ki, Kh, kt, pk, address math, eliminate branches.
// For K<=1024: 1 K-tile per wave. For K>1024: loop over Ki/4 tiles per wave.
// ==========================================================================
template<int cM, int cN, int cK>
__launch_bounds__(256)
__global__ void kern_fused_spec(
    const uint32_t* __restrict__ A, const uint8_t* __restrict__ Bs,
    const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C
) {
    // All shape-dependent values are compile-time constants
    constexpr int cKi = cK >> 7;        // K-tiles
    constexpr int cKh = cK >> 1;        // K/2
    constexpr int ckt = cK / 64;
    constexpr int cpk = ((cK/32 + 7) / 8) * 8;

    const int tid = threadIdx.x;
    const int wave_id = tid >> 6;
    const int lane = tid & 63;
    const int mt = blockIdx.x, nt = blockIdx.y;
    const int lid = lane & 15, gid = lane >> 4;
    const int mg = mt*16 + lid;

    // Compile-time K-tile assignment per wave
    constexpr int tiles_per_wave = (cKi + 3) / 4;  // ceil(Ki/4)

    // Compile-time: which K-tiles this wave handles
    const int ki_start = wave_id * tiles_per_wave;
    const int ki_end_raw = ki_start + tiles_per_wave;
    const int ki_end = (ki_end_raw < cKi) ? ki_end_raw : cKi;

    const int bb = nt * (ckt * 512);
    const int bn = nt*16 + lid;

    v4f32 acc = {0,0,0,0};

    if (ki_start < cKi) {  // compile-time: waves 0-3 all active when Ki>=4
        // Prologue
        int kb = ki_start << 7;
        uint32_t ap[16];
        if (mg < cM) {  // compile-time M enables branch elimination
            const uint4* av = reinterpret_cast<const uint4*>(A + mg*cKh + ((kb+(gid<<5))>>1));
            uint4 c0=av[0],c1=av[1],c2=av[2],c3=av[3];
            ap[0]=c0.x;ap[1]=c0.y;ap[2]=c0.z;ap[3]=c0.w;
            ap[4]=c1.x;ap[5]=c1.y;ap[6]=c1.z;ap[7]=c1.w;
            ap[8]=c2.x;ap[9]=c2.y;ap[10]=c2.z;ap[11]=c2.w;
            ap[12]=c3.x;ap[13]=c3.y;ap[14]=c3.z;ap[15]=c3.w;
        } else { for(int i=0;i<16;i++) ap[i]=0u; }
        v4i32 af; uint32_t ae;
        quant_a(ap, af, ae);

        int b_off = bb + ((kb>>6)+(gid>>1))*512 + (gid&1)*256 + (lid<<4);
        uint4 bc = *reinterpret_cast<const uint4*>(Bs + b_off);
        v4i32 bf; bf[0]=(int)bc.x;bf[1]=(int)bc.y;bf[2]=(int)bc.z;bf[3]=(int)bc.w;
        uint32_t be = (uint32_t)Bsc[bsc_idx(bn, (kb>>5)+gid, cpk)];

        // Pipelined K-loop (compiler unrolls when tiles_per_wave is small constant)
        for (int ki = ki_start+1; ki < ki_end; ki++) {
            int kbn = ki << 7;
            uint32_t apn[16];
            if (mg < cM) {
                const uint4* avn = reinterpret_cast<const uint4*>(A + mg*cKh + ((kbn+(gid<<5))>>1));
                uint4 d0=avn[0],d1=avn[1],d2=avn[2],d3=avn[3];
                apn[0]=d0.x;apn[1]=d0.y;apn[2]=d0.z;apn[3]=d0.w;
                apn[4]=d1.x;apn[5]=d1.y;apn[6]=d1.z;apn[7]=d1.w;
                apn[8]=d2.x;apn[9]=d2.y;apn[10]=d2.z;apn[11]=d2.w;
                apn[12]=d3.x;apn[13]=d3.y;apn[14]=d3.z;apn[15]=d3.w;
            } else { for(int i=0;i<16;i++) apn[i]=0u; }
            int b_offn = bb + ((kbn>>6)+(gid>>1))*512 + (gid&1)*256 + (lid<<4);
            uint4 bcn = *reinterpret_cast<const uint4*>(Bs + b_offn);
            v4i32 bfn; bfn[0]=(int)bcn.x;bfn[1]=(int)bcn.y;bfn[2]=(int)bcn.z;bfn[3]=(int)bcn.w;
            uint32_t ben = (uint32_t)Bsc[bsc_idx(bn, (kbn>>5)+gid, cpk)];

            asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
                :"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));
            quant_a(apn, af, ae); bf=bfn; be=ben;
        }
        asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
            :"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));
    }

    // LDS reduction — transposed layout
    __shared__ float lds[4 * 256];
    int lds_w = wave_id * 256 + lane;
    lds[lds_w]=acc[0]; lds[lds_w+64]=acc[1]; lds[lds_w+128]=acc[2]; lds[lds_w+192]=acc[3];
    __syncthreads();

    if (wave_id == 0) {
        v4f32 sum;
        sum[0]=lds[lane]+lds[lane+256]+lds[lane+512]+lds[lane+768];
        sum[1]=lds[lane+64]+lds[lane+320]+lds[lane+576]+lds[lane+832];
        sum[2]=lds[lane+128]+lds[lane+384]+lds[lane+640]+lds[lane+896];
        sum[3]=lds[lane+192]+lds[lane+448]+lds[lane+704]+lds[lane+960];
        store_tile(C, sum, mt, nt, gid, lid, cM, cN);
    }
}

// ==========================================================================
// KERNEL 2: Warp-parallel A quantization
// ==========================================================================
#define QPG 16
__global__ void kern_quant(
    const uint32_t* __restrict__ A, unsigned char* __restrict__ fp4,
    unsigned char* __restrict__ sc,
    const int M, const int K, const int Ksc, const int Mpad, const int Kscpad
) {
    const int tid=blockIdx.x*blockDim.x+threadIdx.x;
    const int gid=tid/QPG, lig=tid%QPG;
    if(gid>=M*Ksc)return;
    const int row=gid/Ksc, grp=gid%Ksc;
    uint32_t p=*(A+row*(K/2)+grp*16+lig);
    uint32_t mx=p&0x7FFF7FFFu;
    for(int o=8;o>=1;o>>=1){uint32_t ot=__shfl_xor(mx,o);
        asm volatile("v_pk_max_u16 %0,%1,%2":"=v"(mx):"v"(mx),"v"(ot));}
    uint32_t h=mx>>16,l=mx&0xFFFFu;uint16_t am=(uint16_t)(h>l?h:l);
    float sf;unsigned char e8m0;
    if(am==0u){e8m0=0;sf=0.0f;}
    else{unsigned ai=__float_as_uint(bf16_to_f32(am));ai=(ai+0x200000u)&0xFF800000u;
        int su=(int)((ai>>23)&0xFFu)-129;su=max(-127,min(127,su));
        e8m0=(unsigned char)(su+127);sf=__uint_as_float((uint32_t)e8m0<<23);}
    uint32_t fp4b; asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0,%1,%2":"=v"(fp4b):"v"(p),"v"(sf));
    fp4[row*(K/2)+grp*16+lig]=(unsigned char)(fp4b&0xFFu);
    if(lig==0){int sn=Kscpad;
        int si=(row/32)*(32*sn)+(grp/8)*256+(grp%4)*64+(row%16)*4+((grp%8)/4)*2+((row%32)/16);
        if(si<Mpad*Kscpad)sc[si]=e8m0;}
}

// ==========================================================================
// KERNEL 3: Lean GEMM (pre-quantized A)
// ==========================================================================
__launch_bounds__(64)
__global__ void kern_gemm(
    const uint8_t* __restrict__ Afp4, const uint8_t* __restrict__ Asc,
    const uint8_t* __restrict__ Bs, const uint8_t* __restrict__ Bsc,
    uint16_t* __restrict__ C,
    const int M, const int N, const int K, const int kt, const int pkA, const int pkB
) {
    const int lane=threadIdx.x,mt=blockIdx.x,nt=blockIdx.y;
    const int lid=lane&15,gid=lane>>4,mg=mt*16+lid,Kh=K>>1;
    const int bn=nt*16+lid,bb=nt*(kt*512);
    v4f32 acc={0,0,0,0};
    v4i32 af,bf;uint32_t ae,be;
    if(mg<M){uint4 ac=*reinterpret_cast<const uint4*>(Afp4+mg*Kh+(gid*32)/2);
        af[0]=(int)ac.x;af[1]=(int)ac.y;af[2]=(int)ac.z;af[3]=(int)ac.w;}
    else{af[0]=af[1]=af[2]=af[3]=0;}
    ae=mg<M?(uint32_t)Asc[bsc_idx(mg,gid,pkA)]:0u;
    {uint4 bc=*reinterpret_cast<const uint4*>(Bs+bb+(gid/2)*512+(gid&1)*256+lid*16);
     bf[0]=(int)bc.x;bf[1]=(int)bc.y;bf[2]=(int)bc.z;bf[3]=(int)bc.w;
     be=(uint32_t)Bsc[bsc_idx(bn,gid,pkB)];}
    for(int ki=1,Ki=K>>7;ki<Ki;ki++){int kb=ki<<7;
        v4i32 afn,bfn;uint32_t aen,ben;
        if(mg<M){uint4 acn=*reinterpret_cast<const uint4*>(Afp4+mg*Kh+(kb+gid*32)/2);
            afn[0]=(int)acn.x;afn[1]=(int)acn.y;afn[2]=(int)acn.z;afn[3]=(int)acn.w;}
        else{afn[0]=afn[1]=afn[2]=afn[3]=0;}
        aen=mg<M?(uint32_t)Asc[bsc_idx(mg,kb/32+gid,pkA)]:0u;
        {uint4 bcn=*reinterpret_cast<const uint4*>(Bs+bb+(kb/64+gid/2)*512+(gid&1)*256+lid*16);
         bfn[0]=(int)bcn.x;bfn[1]=(int)bcn.y;bfn[2]=(int)bcn.z;bfn[3]=(int)bcn.w;
         ben=(uint32_t)Bsc[bsc_idx(bn,kb/32+gid,pkB)];}
        asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
            :"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));
        af=afn;ae=aen;bf=bfn;be=ben;}
    asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
        :"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));
    store_tile(C,acc,mt,nt,gid,lid,M,N);
}
// ==========================================================================
// KERNEL 3b: Lean GEMM with 32x32x64 MFMA (pre-quantized A)
// B-sharing across 32 M-rows happens WITHIN the wave — no LDS, no barriers!
// For M>=32 with K>1024. 4x fewer tiles, 2x less B bandwidth per output element.
// ==========================================================================
typedef float v16f32 __attribute__((__vector_size__(64)));
// ==========================================================================
// KERNEL 4: split-K GEMM with atomicAdd
// ==========================================================================
__launch_bounds__(64)
__global__ void kern_splitk(
    const uint8_t* __restrict__ Afp4, const uint8_t* __restrict__ Asc,
    const uint8_t* __restrict__ Bs, const uint8_t* __restrict__ Bsc,
    float* __restrict__ temp,
    const int M, const int N, const int K,
    const int kt, const int pkA, const int pkB,
    const int n_tiles, const int k_split
) {
    int bid=blockIdx.x,tid_t=bid/k_split,sid=bid%k_split;
    int mt=tid_t/n_tiles,nt=tid_t%n_tiles;
    const int lane=threadIdx.x,lid=lane&15,gid=lane>>4,mg=mt*16+lid,Kh=K>>1;
    const int bn=nt*16+lid,bb=nt*(kt*512);
    int Ki=K>>7,kip=Ki/k_split,kis=sid*kip,kie=(sid==k_split-1)?Ki:kis+kip;
    v4f32 acc={0,0,0,0};
    if(kis<kie){int kb0=kis<<7;
        v4i32 af,bf;uint32_t ae,be;
        if(mg<M){uint4 ac=*reinterpret_cast<const uint4*>(Afp4+mg*Kh+(kb0+gid*32)/2);
            af[0]=(int)ac.x;af[1]=(int)ac.y;af[2]=(int)ac.z;af[3]=(int)ac.w;}
        else{af[0]=af[1]=af[2]=af[3]=0;}
        ae=mg<M?(uint32_t)Asc[bsc_idx(mg,kb0/32+gid,pkA)]:0u;
        {uint4 bc=*reinterpret_cast<const uint4*>(Bs+bb+(kb0/64+gid/2)*512+(gid&1)*256+lid*16);
         bf[0]=(int)bc.x;bf[1]=(int)bc.y;bf[2]=(int)bc.z;bf[3]=(int)bc.w;
         be=(uint32_t)Bsc[bsc_idx(bn,kb0/32+gid,pkB)];}
        for(int ki=kis+1;ki<kie;ki++){int kb=ki<<7;
            v4i32 afn,bfn;uint32_t aen,ben;
            if(mg<M){uint4 acn=*reinterpret_cast<const uint4*>(Afp4+mg*Kh+(kb+gid*32)/2);
                afn[0]=(int)acn.x;afn[1]=(int)acn.y;afn[2]=(int)acn.z;afn[3]=(int)acn.w;}
            else{afn[0]=afn[1]=afn[2]=afn[3]=0;}
            aen=mg<M?(uint32_t)Asc[bsc_idx(mg,kb/32+gid,pkA)]:0u;
            {uint4 bcn=*reinterpret_cast<const uint4*>(Bs+bb+(kb/64+gid/2)*512+(gid&1)*256+lid*16);
             bfn[0]=(int)bcn.x;bfn[1]=(int)bcn.y;bfn[2]=(int)bcn.z;bfn[3]=(int)bcn.w;
             ben=(uint32_t)Bsc[bsc_idx(bn,kb/32+gid,pkB)];}
            asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
                :"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));
            af=afn;ae=aen;bf=bfn;be=ben;}
        asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
            :"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be));}
    int on=nt*16+lid,om=mt*16+gid*4;
    if(on<N){for(int r=0;r<4;r++)if(om+r<M) atomicAdd(&temp[(om+r)*N+on],acc[r]);}
}

__global__ void kern_convert(const float* src, uint16_t* dst, int n) {
    int i=blockIdx.x*blockDim.x+threadIdx.x; if(i>=n)return;
    uint32_t u=__float_as_uint(src[i]); u+=0x7FFFu+((u>>16)&1u); dst[i]=(uint16_t)(u>>16);
}
// ==========================================================================
// KERNEL 6b: Hybrid inline-A-quant + GLL-B for Shape 1.
// Each wave owns one K chunk, preloads its B chunk into private LDS, then
// quantizes A inline while feeding the MFMA loop. The only barrier is the
// final LDS reduction across split waves.
// ==========================================================================
template<int BATCH, int NUM_SPLITS>
__launch_bounds__(512, 2)
__global__ void kern_hybrid_s1(
    const uint32_t* __restrict__ A, const uint8_t* __restrict__ Bs,
    const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
    const int M, const int N, const int K, const int kt, const int pk
) {
    static_assert(BATCH <= 14, "supports up to 14 staged B tiles");
    constexpr int N_WAVES = NUM_SPLITS;

    const int tid = threadIdx.x;
    const int wave_id = tid >> 6;
    const int lane = tid & 63;
    const int lid = lane & 15;
    const int gid = lane >> 4;

    // 2D grid: blockIdx.x = mt (M-tile), blockIdx.y = nt (N-tile)
    // For Shape 1 (M=16): grid(1, 132) → mt always 0
    // For Shape 4 (M=64): grid(4, 448) → mt = 0..3
    const int mt = blockIdx.x;
    const int nt = blockIdx.y;
    const int Kh = K >> 1;
    const int Ki = K >> 7;
    const int Ki_chunk = Ki / NUM_SPLITS;
    const int k_start = wave_id * Ki_chunk * 128;

    extern __shared__ uint32_t lds_hybrid[];

    const int MY_B_DW = wave_id * BATCH * 256;
    const int my_b_byte = MY_B_DW * 4;
    const int B_LDS_DW_TOTAL = N_WAVES * BATCH * 256;
    float* red_buf = reinterpret_cast<float*>(&lds_hybrid[B_LDS_DW_TOTAL]);

    const int bn = nt * 16 + lid;
    const int bb = nt * (kt * 512);
    const int nt_total = (N + 15) / 16;
    const bool valid_wave = (nt < nt_total);

    uint32_t be0=0,be1=0,be2=0,be3=0,be4=0,be5=0,be6=0,be7=0;
    uint32_t be8=0,be9=0,be10=0,be11=0,be12=0,be13=0;
    if (valid_wave) {
        for (int s = 0; s < BATCH && s < Ki_chunk; s++) {
            int kb = k_start + (s << 7);
            const uint8_t* baddr = Bs + bb + (kb/64+gid/2)*512 + (gid&1)*256 + lid*16;
            uint32_t m0v = (uint32_t)(my_b_byte + s * 1024);
            asm volatile("s_mov_b32 m0, %0\n\t"
                         "s_nop 4\n\t"
                         "global_load_lds_dwordx4 %1, off nt"
                         :: "s"(__builtin_amdgcn_readfirstlane(m0v)), "v"(baddr) : "memory");
        }
        #define LOAD_BE_H(I) do { int _kb=k_start+((I)<<7); be##I=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; } while(0)
        LOAD_BE_H(0); if(Ki_chunk>1) LOAD_BE_H(1); if(Ki_chunk>2) LOAD_BE_H(2); if(Ki_chunk>3) LOAD_BE_H(3);
        if(Ki_chunk>4) LOAD_BE_H(4); if(Ki_chunk>5) LOAD_BE_H(5); if(Ki_chunk>6) LOAD_BE_H(6); if(Ki_chunk>7) LOAD_BE_H(7);
        if(Ki_chunk>8) LOAD_BE_H(8); if(Ki_chunk>9) LOAD_BE_H(9); if(Ki_chunk>10) LOAD_BE_H(10); if(Ki_chunk>11) LOAD_BE_H(11);
        if(Ki_chunk>12) LOAD_BE_H(12); if(Ki_chunk>13) LOAD_BE_H(13);
        #undef LOAD_BE_H
    }

    if (!valid_wave) return;

    asm volatile("s_waitcnt vmcnt(%c0)" :: "n"(BATCH) : "memory");

    const int mg = mt * 16 + lid;
    auto load_a_tile = [&](int ki_local, uint4& c0, uint4& c1, uint4& c2, uint4& c3) {
        int kb = k_start + (ki_local << 7);
        if (mg < M) {
            const uint4* av = reinterpret_cast<const uint4*>(A + mg*Kh + ((kb + (gid<<5)) >> 1));
            c0 = av[0]; c1 = av[1]; c2 = av[2]; c3 = av[3];
        } else { c0=c1=c2=c3=make_uint4(0,0,0,0); }
    };
    auto load_b_tile = [&](int slot, v4i32& bf_out) {
        int boff = MY_B_DW + slot * 256 + lane * 4;
        uint4 bc = lds_load_u128((lds_u32_ptr)(&lds_hybrid[boff]));
        bf_out[0]=(int)bc.x; bf_out[1]=(int)bc.y; bf_out[2]=(int)bc.z; bf_out[3]=(int)bc.w;
    };
    auto get_be = [&](int idx) -> uint32_t {
        switch (idx) {
            case 0: return be0; case 1: return be1; case 2: return be2; case 3: return be3;
            case 4: return be4; case 5: return be5; case 6: return be6; case 7: return be7;
            case 8: return be8; case 9: return be9; case 10: return be10; case 11: return be11;
            case 12: return be12; case 13: return be13;
            default: return 0u;
        }
    };

    v4f32 acc = {0,0,0,0};

    uint4 a0c0, a0c1, a0c2, a0c3;
    load_a_tile(0, a0c0, a0c1, a0c2, a0c3);
    v4i32 af; uint32_t ae;
    quant_a_direct(a0c0, a0c1, a0c2, a0c3, af, ae);

    v4i32 bf;
    load_b_tile(0, bf);

    for (int ki_local = 1; ki_local < Ki_chunk; ++ki_local) {
        int far_ki = ki_local + BATCH - 1;
        if (far_ki < Ki_chunk) {
            int kb = k_start + (far_ki << 7);
            const uint8_t* baddr = Bs + bb + (kb/64+gid/2)*512 + (gid&1)*256 + lid*16;
            uint32_t m0v = (uint32_t)(my_b_byte + ((ki_local - 1) % BATCH) * 1024);
            asm volatile("s_mov_b32 m0, %0\n\t"
                         "s_nop 4\n\t"
                         "global_load_lds_dwordx4 %1, off nt"
                         :: "s"(__builtin_amdgcn_readfirstlane(m0v)), "v"(baddr) : "memory");
        }

        uint4 next_c0, next_c1, next_c2, next_c3;
        load_a_tile(ki_local, next_c0, next_c1, next_c2, next_c3);

        v4i32 bfn;
        load_b_tile(ki_local % BATCH, bfn);

        uint32_t be_cur = get_be(ki_local - 1);
        asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
            :"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(be_cur));
        asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
        quant_a_direct(next_c0, next_c1, next_c2, next_c3, af, ae);
        bf = bfn;
    }

    asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
        :"+v"(acc):"v"(af),"v"(bf),"v"(ae),"v"(get_be(Ki_chunk - 1)));

    int red_off = wave_id * 256 + lane * 4;
    red_buf[red_off+0] = acc[0]; red_buf[red_off+1] = acc[1];
    red_buf[red_off+2] = acc[2]; red_buf[red_off+3] = acc[3];

    __syncthreads();

    if (wave_id == 0) {
        float s0=0, s1=0, s2=0, s3=0;
        for (int w = 0; w < NUM_SPLITS; w++) {
            int off = w * 256 + lane * 4;
            s0 += red_buf[off+0]; s1 += red_buf[off+1];
            s2 += red_buf[off+2]; s3 += red_buf[off+3];
        }
        v4f32 result; result[0]=s0; result[1]=s1; result[2]=s2; result[3]=s3;
        store_tile(C, result, mt, nt, gid, lid, M, N);
    }
}
// ==========================================================================
// Dedicated kernel: M=64 N=7168 K=2048 Ki=16 N_N=8 BATCH=10 SPLIT=10
__launch_bounds__(512)
__global__ void kern_s4(
    const uint32_t* __restrict__ A, const uint8_t* __restrict__ Bs,
    const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
    const int M, const int N, const int K, const int kt, const int pk
) {
    constexpr int N_M = 1, N_N = 8, BATCH = 10;
    constexpr int N_WAVES = 8;
    constexpr int A_TILE = 384;
    constexpr int A_FP4_STRIDE = 5;
    constexpr int Ki = 16;

    const int tid = threadIdx.x;
    const int wave_id = tid >> 6;
    const int lane = tid & 63;
    const int lid = lane & 15;
    const int gid = lane >> 4;
    const int mt = blockIdx.x;
    const int nt = blockIdx.y * N_N + wave_id;
    const int Kh = K >> 1;
    extern __shared__ uint32_t lds_gll[];
    const int B_LDS_DW = Ki * A_TILE;
    const int MY_B_DW = B_LDS_DW + wave_id * BATCH * 256;
    const int my_b_byte = MY_B_DW * 4;
    const int bn = nt * 16 + lid;
    const int bb = nt * (kt * 512);
    const int nt_total = (N + 15) / 16;
    const bool valid_wave = (nt < nt_total && mt < (M + 15) / 16);

    // === Quant A → LDS (2 tile(s) per wave) ===
    {
        const int mg_q = mt * 16 + lid;
        // Quant tile 0
        {
            const int ki_t = wave_id * 2 + 0;
            int kb = ki_t << 7;
            uint4 c0,c1,c2,c3;
            if (mg_q < M) {
                const uint4* av = reinterpret_cast<const uint4*>(A + mg_q*Kh + ((kb+(gid<<5))>>1));
                c0=av[0];c1=av[1];c2=av[2];c3=av[3];
            } else { c0=c1=c2=c3=make_uint4(0,0,0,0); }
            v4i32 fp4; uint32_t e8m0;
            quant_a_direct(c0, c1, c2, c3, fp4, e8m0);
            int base = ki_t * A_TILE;
            lds_gll[base+lane*A_FP4_STRIDE+0]=(uint32_t)fp4[0];
            lds_gll[base+lane*A_FP4_STRIDE+1]=(uint32_t)fp4[1];
            lds_gll[base+lane*A_FP4_STRIDE+2]=(uint32_t)fp4[2];
            lds_gll[base+lane*A_FP4_STRIDE+3]=(uint32_t)fp4[3];
            lds_gll[base+320+lane]=e8m0;
        }
        // Quant tile 1
        {
            const int ki_t = wave_id * 2 + 1;
            int kb = ki_t << 7;
            uint4 c0,c1,c2,c3;
            if (mg_q < M) {
                const uint4* av = reinterpret_cast<const uint4*>(A + mg_q*Kh + ((kb+(gid<<5))>>1));
                c0=av[0];c1=av[1];c2=av[2];c3=av[3];
            } else { c0=c1=c2=c3=make_uint4(0,0,0,0); }
            v4i32 fp4; uint32_t e8m0;
            quant_a_direct(c0, c1, c2, c3, fp4, e8m0);
            int base = ki_t * A_TILE;
            lds_gll[base+lane*A_FP4_STRIDE+0]=(uint32_t)fp4[0];
            lds_gll[base+lane*A_FP4_STRIDE+1]=(uint32_t)fp4[1];
            lds_gll[base+lane*A_FP4_STRIDE+2]=(uint32_t)fp4[2];
            lds_gll[base+lane*A_FP4_STRIDE+3]=(uint32_t)fp4[3];
            lds_gll[base+320+lane]=e8m0;
        }
    }
    // === Load B+Bsc: SPLIT=10 pre-barrier, 0 post-barrier ===
    uint32_t be0,be1,be2,be3,be4,be5,be6,be7,be8,be9,be10,be11,be12,be13,be14,be15;
    if (valid_wave) {
        const uint8_t* b_base = Bs + bb;
        uint64_t b_addr = (uint64_t)b_base;
        asm volatile(
            "s_mov_b32 s[16], %0\n\t"
            "s_mov_b32 s[17], %1\n\t"
            "s_mov_b32 s[18], 0xFFFFFFFF\n\t"
            "s_mov_b32 s[19], 0x20000"
            :: "s"(__builtin_amdgcn_readfirstlane((uint32_t)b_addr)),
               "s"(__builtin_amdgcn_readfirstlane((uint32_t)(b_addr >> 32)))
            : "s16", "s17", "s18", "s19"
        );
        // Pre-barrier: 10 interleaved B+Bsc pairs
        { uint32_t _m0v = (uint32_t)(my_b_byte + 0 * 1024);
          uint32_t _buf_off = (0*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 0 << 7; be0 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 1 * 1024);
          uint32_t _buf_off = (1*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 1 << 7; be1 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 2 * 1024);
          uint32_t _buf_off = (2*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 2 << 7; be2 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 3 * 1024);
          uint32_t _buf_off = (3*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 3 << 7; be3 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 4 * 1024);
          uint32_t _buf_off = (4*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 4 << 7; be4 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 5 * 1024);
          uint32_t _buf_off = (5*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 5 << 7; be5 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 6 * 1024);
          uint32_t _buf_off = (6*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 6 << 7; be6 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 7 * 1024);
          uint32_t _buf_off = (7*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 7 << 7; be7 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 8 * 1024);
          uint32_t _buf_off = (8*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 8 << 7; be8 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 9 * 1024);
          uint32_t _buf_off = (9*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 9 << 7; be9 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { int _kb=10<<7; be10=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
        { int _kb=11<<7; be11=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
        { int _kb=12<<7; be12=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
        { int _kb=13<<7; be13=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
        { int _kb=14<<7; be14=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
        { int _kb=15<<7; be15=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
    } else {
        be0=0; be1=0; be2=0; be3=0; be4=0; be5=0; be6=0; be7=0; be8=0; be9=0; be10=0; be11=0; be12=0; be13=0; be14=0; be15=0;
    }

    __syncthreads();
    if (!valid_wave) return;
    const int a_offset = 0;
    v4f32 acc = {0,0,0,0};
    asm volatile("s_waitcnt vmcnt(0)" ::: "memory");

    #define DB_READ_B_K(...) DB_READ_B_IMPL(__VA_ARGS__)
    #define DB_READ_A01_K(...) DB_READ_A01_IMPL(__VA_ARGS__)
    #define DB_READ_A23_K(...) DB_READ_A23_IMPL(__VA_ARGS__)
    #define DB_READ_AE_K(...) DB_READ_AE_IMPL(__VA_ARGS__)
    #define DB_ISSUE_K(...) DB_ISSUE_IMPL(__VA_ARGS__)
    #define DB_PACK_K(...) DB_PACK_IMPL(__VA_ARGS__)
    #define DB_MFMA_K(...) DB_MFMA_IMPL(__VA_ARGS__)
    #define DB_TILE_A_K(...) DB_TILE_A_IMPL(__VA_ARGS__)
    #define DB_TILE_B_K(...) DB_TILE_B_IMPL(__VA_ARGS__)
    #define DO_MFMA_RT_K(...) DO_MFMA_RT_IMPL(__VA_ARGS__)
    #define ISSUE_GLL_K(...) ISSUE_GLL_IMPL(__VA_ARGS__)

    #define DB_READ_B_IMPL(dst, addr) \
        asm volatile("ds_read_b128 %0, %1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
    #define DB_READ_A01_IMPL(dst, addr) \
        asm volatile("ds_read2_b32 %0, %1 offset1:1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
    #define DB_READ_A23_IMPL(dst, addr) \
        asm volatile("ds_read2_b32 %0, %1 offset0:2 offset1:3" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
    #define DB_READ_AE_IMPL(dst, addr) \
        asm volatile("ds_read_b32 %0, %1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
    #define DB_ISSUE_IMPL(KI, SLOT, bb, aa01, aa23, aae) do { \
        int _boff = MY_B_DW + (SLOT) * 256 + lane * 4; \
        DB_READ_B_IMPL(bb, _boff); \
        int _abase = a_offset + (KI) * A_TILE + lane * A_FP4_STRIDE; \
        DB_READ_A01_IMPL(aa01, _abase); \
        DB_READ_A23_IMPL(aa23, _abase); \
        DB_READ_AE_IMPL(aae, a_offset + (KI) * A_TILE + 320 + lane); \
    } while(0)
    #define DB_PACK_IMPL(bb, aa01, aa23, aae, bf_v, af_v, ae_v) do { \
        bf_v[0]=(int)bb.x; bf_v[1]=(int)bb.y; bf_v[2]=(int)bb.z; bf_v[3]=(int)bb.w; \
        af_v[0]=(int)(uint32_t)aa01; af_v[1]=(int)(uint32_t)(aa01>>32); \
        af_v[2]=(int)(uint32_t)aa23; af_v[3]=(int)(uint32_t)(aa23>>32); \
        ae_v = (uint32_t)aae; \
    } while(0)
    #define DB_MFMA_IMPL(af_v, bf_v, ae_v, be_v) \
        asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4" \
            :"+v"(acc):"v"(af_v),"v"(bf_v),"v"(ae_v),"v"(be_v))
    #define DB_TILE_A_IMPL(KI, SLOT, BE_VAR, NEXT_KI, NEXT_SLOT) do { \
        DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae); \
        DB_MFMA_IMPL(maf, mbf, mae, BE_VAR); \
        DB_ISSUE_IMPL(NEXT_KI, NEXT_SLOT, bb_a, a01_a, a23_a, ae_a); \
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory"); \
    } while(0)
    #define DB_TILE_B_IMPL(KI, SLOT, BE_VAR, NEXT_KI, NEXT_SLOT) do { \
        DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae); \
        DB_MFMA_IMPL(maf, mbf, mae, BE_VAR); \
        DB_ISSUE_IMPL(NEXT_KI, NEXT_SLOT, bb_b, a01_b, a23_b, ae_b); \
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory"); \
    } while(0)
    #define DO_MFMA_RT_IMPL(KI_RT, SLOT, BE_VAR) do { \
        v4i32 _bf; \
        { int _boff = MY_B_DW + (SLOT) * 256 + lane * 4; \
          uint4 _bc = lds_load_u128((lds_u32_ptr)(&lds_gll[_boff])); \
          _bf[0]=(int)_bc.x; _bf[1]=(int)_bc.y; _bf[2]=(int)_bc.z; _bf[3]=(int)_bc.w; } \
        v4i32 _af; uint32_t _ae; \
        { int _base = a_offset + (KI_RT) * A_TILE; \
          _af[0]=(int)lds_gll[_base+lane*A_FP4_STRIDE+0]; _af[1]=(int)lds_gll[_base+lane*A_FP4_STRIDE+1]; \
          _af[2]=(int)lds_gll[_base+lane*A_FP4_STRIDE+2]; _af[3]=(int)lds_gll[_base+lane*A_FP4_STRIDE+3]; \
          _ae=lds_gll[_base+320+lane]; } \
        asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4" \
            :"+v"(acc):"v"(_af),"v"(_bf),"v"(_ae),"v"(BE_VAR)); \
    } while(0)
    #define ISSUE_GLL_IMPL(KI_RT, SLOT) do { \
        uint32_t _buf_off = ((KI_RT)*2+gid/2)*512 + (gid&1)*256 + lid*16; \
        uint32_t _m0 = (uint32_t)(my_b_byte + (SLOT) * 1024); \
        const uint8_t* _bbase = Bs + bb; \
        uint64_t _baddr = (uint64_t)_bbase; \
        asm volatile( \
            "s_mov_b32 s[16], %2\n\t" \
            "s_mov_b32 s[17], %3\n\t" \
            "s_mov_b32 s[18], 0xFFFFFFFF\n\t" \
            "s_mov_b32 s[19], 0x20000\n\t" \
            "s_mov_b32 m0, %0\n\t" \
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds" \
            :: "s"(__builtin_amdgcn_readfirstlane(_m0)), "v"(_buf_off), \
               "s"(__builtin_amdgcn_readfirstlane((uint32_t)_baddr)), \
               "s"(__builtin_amdgcn_readfirstlane((uint32_t)(_baddr >> 32))) \
            : "memory", "s16", "s17", "s18", "s19" \
        ); \
    } while(0)

    {
        uint4 bb_a, bb_b; uint64_t a01_a, a23_a, a01_b, a23_b; uint32_t ae_a, ae_b;
        v4i32 maf, mbf; uint32_t mae;

        DB_ISSUE_IMPL(0, 0, bb_a, a01_a, a23_a, ae_a);
        DB_ISSUE_IMPL(1, 1, bb_b, a01_b, a23_b, ae_b);
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");

        // Tile 0: MFMA + GLL(10)→slot0 + prefetch 2
        DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
        DB_MFMA_IMPL(maf, mbf, mae, be0);
        ISSUE_GLL_IMPL(10, 0);
        DB_ISSUE_IMPL(2, 2, bb_a, a01_a, a23_a, ae_a);
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");

        // Tile 1: MFMA + GLL(11)→slot1 + prefetch 3
        DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
        DB_MFMA_IMPL(maf, mbf, mae, be1);
        ISSUE_GLL_IMPL(11, 1);
        DB_ISSUE_IMPL(3, 3, bb_b, a01_b, a23_b, ae_b);
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");

        // Tile 2: MFMA + GLL(12)→slot2 + prefetch 4
        DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
        DB_MFMA_IMPL(maf, mbf, mae, be2);
        ISSUE_GLL_IMPL(12, 2);
        DB_ISSUE_IMPL(4, 4, bb_a, a01_a, a23_a, ae_a);
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");

        // Tile 3: MFMA + GLL(13)→slot3 + prefetch 5
        DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
        DB_MFMA_IMPL(maf, mbf, mae, be3);
        ISSUE_GLL_IMPL(13, 3);
        DB_ISSUE_IMPL(5, 5, bb_b, a01_b, a23_b, ae_b);
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");

        // Tile 4: MFMA + GLL(14)→slot4 + prefetch 6
        DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
        DB_MFMA_IMPL(maf, mbf, mae, be4);
        ISSUE_GLL_IMPL(14, 4);
        DB_ISSUE_IMPL(6, 6, bb_a, a01_a, a23_a, ae_a);
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");

        // Tile 5: MFMA + GLL(15)→slot5 + prefetch 7
        DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
        DB_MFMA_IMPL(maf, mbf, mae, be5);
        ISSUE_GLL_IMPL(15, 5);
        DB_ISSUE_IMPL(7, 7, bb_b, a01_b, a23_b, ae_b);
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");

        DB_TILE_A_IMPL(6, 6, be6, 8, 8);

        DB_TILE_B_IMPL(7, 7, be7, 9, 9);

        // Tile 8: no more pre-loaded to prefetch
        DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
        DB_MFMA_IMPL(maf, mbf, mae, be8);
        asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");

        // Tile 9: last pre-loaded
        DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
        DB_MFMA_IMPL(maf, mbf, mae, be9);

        // Tiles 10..15: GLL loaded
        asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
        DO_MFMA_RT_IMPL(10, 0, be10);
        DO_MFMA_RT_IMPL(11, 1, be11);
        DO_MFMA_RT_IMPL(12, 2, be12);
        DO_MFMA_RT_IMPL(13, 3, be13);
        DO_MFMA_RT_IMPL(14, 4, be14);
        DO_MFMA_RT_IMPL(15, 5, be15);
    }
    #undef DB_READ_B_IMPL
    #undef DB_READ_B_K
    #undef DB_READ_A01_IMPL
    #undef DB_READ_A01_K
    #undef DB_READ_A23_IMPL
    #undef DB_READ_A23_K
    #undef DB_READ_AE_IMPL
    #undef DB_READ_AE_K
    #undef DB_ISSUE_IMPL
    #undef DB_ISSUE_K
    #undef DB_PACK_IMPL
    #undef DB_PACK_K
    #undef DB_MFMA_IMPL
    #undef DB_MFMA_K
    #undef DB_TILE_A_IMPL
    #undef DB_TILE_A_K
    #undef DB_TILE_B_IMPL
    #undef DB_TILE_B_K
    #undef DO_MFMA_RT_IMPL
    #undef DO_MFMA_RT_K
    #undef ISSUE_GLL_IMPL
    #undef ISSUE_GLL_K

    store_tile(C, acc, mt, nt, gid, lid, M, N);
}
// Dedicated kernel: M=256 N=3072 K=1536 Ki=12 N_N=12 BATCH=8 SPLIT=8
__launch_bounds__(768)
__global__ void kern_s5(
    const uint32_t* __restrict__ A, const uint8_t* __restrict__ Bs,
    const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
    const int M, const int N, const int K, const int kt, const int pk
) {
    constexpr int N_M = 1, N_N = 12, BATCH = 8;
    constexpr int N_WAVES = 12;
    constexpr int A_TILE = 384;
    constexpr int A_FP4_STRIDE = 5;
    constexpr int Ki = 12;

    const int tid = threadIdx.x;
    const int wave_id = tid >> 6;
    const int lane = tid & 63;
    const int lid = lane & 15;
    const int gid = lane >> 4;
    const int mt = blockIdx.x;
    const int nt = blockIdx.y * N_N + wave_id;
    const int Kh = K >> 1;
    extern __shared__ uint32_t lds_gll[];
    const int B_LDS_DW = Ki * A_TILE;
    const int MY_B_DW = B_LDS_DW + wave_id * BATCH * 256;
    const int my_b_byte = MY_B_DW * 4;
    const int bn = nt * 16 + lid;
    const int bb = nt * (kt * 512);
    const int nt_total = (N + 15) / 16;
    const bool valid_wave = (nt < nt_total && mt < (M + 15) / 16);

    // === Quant A → LDS (1 tile(s) per wave) ===
    {
        const int mg_q = mt * 16 + lid;
        // Quant tile 0
        {
            const int ki_t = wave_id * 1 + 0;
            int kb = ki_t << 7;
            uint4 c0,c1,c2,c3;
            if (mg_q < M) {
                const uint4* av = reinterpret_cast<const uint4*>(A + mg_q*Kh + ((kb+(gid<<5))>>1));
                c0=av[0];c1=av[1];c2=av[2];c3=av[3];
            } else { c0=c1=c2=c3=make_uint4(0,0,0,0); }
            v4i32 fp4; uint32_t e8m0;
            quant_a_direct(c0, c1, c2, c3, fp4, e8m0);
            int base = ki_t * A_TILE;
            lds_gll[base+lane*A_FP4_STRIDE+0]=(uint32_t)fp4[0];
            lds_gll[base+lane*A_FP4_STRIDE+1]=(uint32_t)fp4[1];
            lds_gll[base+lane*A_FP4_STRIDE+2]=(uint32_t)fp4[2];
            lds_gll[base+lane*A_FP4_STRIDE+3]=(uint32_t)fp4[3];
            lds_gll[base+320+lane]=e8m0;
        }
    }
    // === Load B+Bsc: SPLIT=8 pre-barrier, 0 post-barrier ===
    uint32_t be0,be1,be2,be3,be4,be5,be6,be7,be8,be9,be10,be11;
    if (valid_wave) {
        const uint8_t* b_base = Bs + bb;
        uint64_t b_addr = (uint64_t)b_base;
        asm volatile(
            "s_mov_b32 s[16], %0\n\t"
            "s_mov_b32 s[17], %1\n\t"
            "s_mov_b32 s[18], 0xFFFFFFFF\n\t"
            "s_mov_b32 s[19], 0x20000"
            :: "s"(__builtin_amdgcn_readfirstlane((uint32_t)b_addr)),
               "s"(__builtin_amdgcn_readfirstlane((uint32_t)(b_addr >> 32)))
            : "s16", "s17", "s18", "s19"
        );
        // Pre-barrier: 8 interleaved B+Bsc pairs
        { uint32_t _m0v = (uint32_t)(my_b_byte + 0 * 1024);
          uint32_t _buf_off = (0*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 0 << 7; be0 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 1 * 1024);
          uint32_t _buf_off = (1*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 1 << 7; be1 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 2 * 1024);
          uint32_t _buf_off = (2*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 2 << 7; be2 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 3 * 1024);
          uint32_t _buf_off = (3*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 3 << 7; be3 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 4 * 1024);
          uint32_t _buf_off = (4*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 4 << 7; be4 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 5 * 1024);
          uint32_t _buf_off = (5*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 5 << 7; be5 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 6 * 1024);
          uint32_t _buf_off = (6*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 6 << 7; be6 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { uint32_t _m0v = (uint32_t)(my_b_byte + 7 * 1024);
          uint32_t _buf_off = (7*2+gid/2)*512 + (gid&1)*256 + lid*16;
          asm volatile("s_mov_b32 m0, %0\n\t"
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds"
            :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off)
            : "memory", "s16", "s17", "s18", "s19");
          { int _kb = 7 << 7; be7 = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; } }
        { int _kb=8<<7; be8=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
        { int _kb=9<<7; be9=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
        { int _kb=10<<7; be10=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
        { int _kb=11<<7; be11=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; }
    } else {
        be0=0; be1=0; be2=0; be3=0; be4=0; be5=0; be6=0; be7=0; be8=0; be9=0; be10=0; be11=0;
    }

    __syncthreads();
    if (!valid_wave) return;
    const int a_offset = 0;
    v4f32 acc = {0,0,0,0};
    asm volatile("s_waitcnt vmcnt(0)" ::: "memory");

    #define DB_READ_B_K(...) DB_READ_B_IMPL(__VA_ARGS__)
    #define DB_READ_A01_K(...) DB_READ_A01_IMPL(__VA_ARGS__)
    #define DB_READ_A23_K(...) DB_READ_A23_IMPL(__VA_ARGS__)
    #define DB_READ_AE_K(...) DB_READ_AE_IMPL(__VA_ARGS__)
    #define DB_ISSUE_K(...) DB_ISSUE_IMPL(__VA_ARGS__)
    #define DB_PACK_K(...) DB_PACK_IMPL(__VA_ARGS__)
    #define DB_MFMA_K(...) DB_MFMA_IMPL(__VA_ARGS__)
    #define DB_TILE_A_K(...) DB_TILE_A_IMPL(__VA_ARGS__)
    #define DB_TILE_B_K(...) DB_TILE_B_IMPL(__VA_ARGS__)
    #define DO_MFMA_RT_K(...) DO_MFMA_RT_IMPL(__VA_ARGS__)
    #define ISSUE_GLL_K(...) ISSUE_GLL_IMPL(__VA_ARGS__)

    #define DB_READ_B_IMPL(dst, addr) \
        asm volatile("ds_read_b128 %0, %1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
    #define DB_READ_A01_IMPL(dst, addr) \
        asm volatile("ds_read2_b32 %0, %1 offset1:1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
    #define DB_READ_A23_IMPL(dst, addr) \
        asm volatile("ds_read2_b32 %0, %1 offset0:2 offset1:3" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
    #define DB_READ_AE_IMPL(dst, addr) \
        asm volatile("ds_read_b32 %0, %1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
    #define DB_ISSUE_IMPL(KI, SLOT, bb, aa01, aa23, aae) do { \
        int _boff = MY_B_DW + (SLOT) * 256 + lane * 4; \
        DB_READ_B_IMPL(bb, _boff); \
        int _abase = a_offset + (KI) * A_TILE + lane * A_FP4_STRIDE; \
        DB_READ_A01_IMPL(aa01, _abase); \
        DB_READ_A23_IMPL(aa23, _abase); \
        DB_READ_AE_IMPL(aae, a_offset + (KI) * A_TILE + 320 + lane); \
    } while(0)
    #define DB_PACK_IMPL(bb, aa01, aa23, aae, bf_v, af_v, ae_v) do { \
        bf_v[0]=(int)bb.x; bf_v[1]=(int)bb.y; bf_v[2]=(int)bb.z; bf_v[3]=(int)bb.w; \
        af_v[0]=(int)(uint32_t)aa01; af_v[1]=(int)(uint32_t)(aa01>>32); \
        af_v[2]=(int)(uint32_t)aa23; af_v[3]=(int)(uint32_t)(aa23>>32); \
        ae_v = (uint32_t)aae; \
    } while(0)
    #define DB_MFMA_IMPL(af_v, bf_v, ae_v, be_v) \
        asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4" \
            :"+v"(acc):"v"(af_v),"v"(bf_v),"v"(ae_v),"v"(be_v))
    #define DB_TILE_A_IMPL(KI, SLOT, BE_VAR, NEXT_KI, NEXT_SLOT) do { \
        DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae); \
        DB_MFMA_IMPL(maf, mbf, mae, BE_VAR); \
        DB_ISSUE_IMPL(NEXT_KI, NEXT_SLOT, bb_a, a01_a, a23_a, ae_a); \
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory"); \
    } while(0)
    #define DB_TILE_B_IMPL(KI, SLOT, BE_VAR, NEXT_KI, NEXT_SLOT) do { \
        DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae); \
        DB_MFMA_IMPL(maf, mbf, mae, BE_VAR); \
        DB_ISSUE_IMPL(NEXT_KI, NEXT_SLOT, bb_b, a01_b, a23_b, ae_b); \
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory"); \
    } while(0)
    #define DO_MFMA_RT_IMPL(KI_RT, SLOT, BE_VAR) do { \
        v4i32 _bf; \
        { int _boff = MY_B_DW + (SLOT) * 256 + lane * 4; \
          uint4 _bc = lds_load_u128((lds_u32_ptr)(&lds_gll[_boff])); \
          _bf[0]=(int)_bc.x; _bf[1]=(int)_bc.y; _bf[2]=(int)_bc.z; _bf[3]=(int)_bc.w; } \
        v4i32 _af; uint32_t _ae; \
        { int _base = a_offset + (KI_RT) * A_TILE; \
          _af[0]=(int)lds_gll[_base+lane*A_FP4_STRIDE+0]; _af[1]=(int)lds_gll[_base+lane*A_FP4_STRIDE+1]; \
          _af[2]=(int)lds_gll[_base+lane*A_FP4_STRIDE+2]; _af[3]=(int)lds_gll[_base+lane*A_FP4_STRIDE+3]; \
          _ae=lds_gll[_base+320+lane]; } \
        asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4" \
            :"+v"(acc):"v"(_af),"v"(_bf),"v"(_ae),"v"(BE_VAR)); \
    } while(0)
    #define ISSUE_GLL_IMPL(KI_RT, SLOT) do { \
        uint32_t _buf_off = ((KI_RT)*2+gid/2)*512 + (gid&1)*256 + lid*16; \
        uint32_t _m0 = (uint32_t)(my_b_byte + (SLOT) * 1024); \
        const uint8_t* _bbase = Bs + bb; \
        uint64_t _baddr = (uint64_t)_bbase; \
        asm volatile( \
            "s_mov_b32 s[16], %2\n\t" \
            "s_mov_b32 s[17], %3\n\t" \
            "s_mov_b32 s[18], 0xFFFFFFFF\n\t" \
            "s_mov_b32 s[19], 0x20000\n\t" \
            "s_mov_b32 m0, %0\n\t" \
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds" \
            :: "s"(__builtin_amdgcn_readfirstlane(_m0)), "v"(_buf_off), \
               "s"(__builtin_amdgcn_readfirstlane((uint32_t)_baddr)), \
               "s"(__builtin_amdgcn_readfirstlane((uint32_t)(_baddr >> 32))) \
            : "memory", "s16", "s17", "s18", "s19" \
        ); \
    } while(0)

    {
        uint4 bb_a, bb_b; uint64_t a01_a, a23_a, a01_b, a23_b; uint32_t ae_a, ae_b;
        v4i32 maf, mbf; uint32_t mae;

        DB_ISSUE_IMPL(0, 0, bb_a, a01_a, a23_a, ae_a);
        DB_ISSUE_IMPL(1, 1, bb_b, a01_b, a23_b, ae_b);
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");

        // Tile 0: MFMA + GLL(8)→slot0 + prefetch 2
        DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
        DB_MFMA_IMPL(maf, mbf, mae, be0);
        ISSUE_GLL_IMPL(8, 0);
        DB_ISSUE_IMPL(2, 2, bb_a, a01_a, a23_a, ae_a);
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");

        // Tile 1: MFMA + GLL(9)→slot1 + prefetch 3
        DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
        DB_MFMA_IMPL(maf, mbf, mae, be1);
        ISSUE_GLL_IMPL(9, 1);
        DB_ISSUE_IMPL(3, 3, bb_b, a01_b, a23_b, ae_b);
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");

        // Tile 2: MFMA + GLL(10)→slot2 + prefetch 4
        DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
        DB_MFMA_IMPL(maf, mbf, mae, be2);
        ISSUE_GLL_IMPL(10, 2);
        DB_ISSUE_IMPL(4, 4, bb_a, a01_a, a23_a, ae_a);
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");

        // Tile 3: MFMA + GLL(11)→slot3 + prefetch 5
        DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
        DB_MFMA_IMPL(maf, mbf, mae, be3);
        ISSUE_GLL_IMPL(11, 3);
        DB_ISSUE_IMPL(5, 5, bb_b, a01_b, a23_b, ae_b);
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");

        DB_TILE_A_IMPL(4, 4, be4, 6, 6);

        DB_TILE_B_IMPL(5, 5, be5, 7, 7);

        // Tile 6: no more pre-loaded to prefetch
        DB_PACK_IMPL(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
        DB_MFMA_IMPL(maf, mbf, mae, be6);
        asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");

        // Tile 7: last pre-loaded
        DB_PACK_IMPL(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
        DB_MFMA_IMPL(maf, mbf, mae, be7);

        // Tiles 8..11: GLL loaded
        asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
        DO_MFMA_RT_IMPL(8, 0, be8);
        DO_MFMA_RT_IMPL(9, 1, be9);
        DO_MFMA_RT_IMPL(10, 2, be10);
        DO_MFMA_RT_IMPL(11, 3, be11);
    }
    #undef DB_READ_B_IMPL
    #undef DB_READ_B_K
    #undef DB_READ_A01_IMPL
    #undef DB_READ_A01_K
    #undef DB_READ_A23_IMPL
    #undef DB_READ_A23_K
    #undef DB_READ_AE_IMPL
    #undef DB_READ_AE_K
    #undef DB_ISSUE_IMPL
    #undef DB_ISSUE_K
    #undef DB_PACK_IMPL
    #undef DB_PACK_K
    #undef DB_MFMA_IMPL
    #undef DB_MFMA_K
    #undef DB_TILE_A_IMPL
    #undef DB_TILE_A_K
    #undef DB_TILE_B_IMPL
    #undef DB_TILE_B_K
    #undef DO_MFMA_RT_IMPL
    #undef DO_MFMA_RT_K
    #undef ISSUE_GLL_IMPL
    #undef ISSUE_GLL_K

    store_tile(C, acc, mt, nt, gid, lid, M, N);
}
// KERNEL 5b: GLL batch — B via GLOBAL_LOAD_LDS, batch processing
// Issue D GLL loads to LDS, process D MFMAs while next batch's GLL flies.
// No B VGPRs in pipeline — only bf (4 VGPRs) read from LDS per MFMA.
// ==========================================================================
template<int N_M, int N_N, int BATCH>
__launch_bounds__(N_M * N_N * 64)
__global__ void kern_fused_gll(
    const uint32_t* __restrict__ A, const uint8_t* __restrict__ Bs,
    const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
    const int M, const int N, const int K, const int kt, const int pk
) {
    constexpr int N_WAVES = N_M * N_N;
    constexpr int A_TILE = 384;
    constexpr int A_FP4_STRIDE = 5;

    const int tid = threadIdx.x;
    const int wave_id = tid >> 6;
    const int lane = tid & 63;
    const int lid = lane & 15;
    const int gid = lane >> 4;

    const int mt_local = wave_id / N_N;
    const int nt_local = wave_id % N_N;
    const int mt = blockIdx.x * N_M + mt_local;
    const int nt = blockIdx.y * N_N + nt_local;
    const int mg = mt * 16 + lid;
    const int Kh = K >> 1;
    const int Ki = K >> 7;

    extern __shared__ uint32_t lds_gll[];

    // B LDS layout: after A data. Each wave gets BATCH slots of 256 dwords.
    const int B_LDS_DW = N_M * Ki * A_TILE;
    const int MY_B_DW = B_LDS_DW + wave_id * BATCH * 256;
    const int my_b_byte = MY_B_DW * 4;

    const int bn = nt * 16 + lid;
    const int bb = nt * (kt * 512);
    const int nt_total = (N + 15) / 16;
    const bool valid_wave = (nt < nt_total && mt < (M + 15) / 16);

    // REORDERED Phase 1+2: load A regs → issue GLL+Bsc → quant A → store LDS
    // A loads go to vmcnt first. Then GLL+Bsc go to vmcnt.
    // Compiler's vmcnt(0) before quant drains A (oldest). GLL+Bsc keep flying.
    // After landmark patching: vmcnt(0) → vmcnt(N) to drain only A loads.

    int total_work = N_M * Ki;
    int work_per_wave = (total_work + N_WAVES - 1) / N_WAVES;
    int w_start = wave_id * work_per_wave;
    int w_end = w_start + work_per_wave;
    if (w_end > total_work) w_end = total_work;
    int num_work = w_end - w_start;

    // Step 1: Issue A loads into registers (compiler tracks these as vmcnt)
    uint4 aq_c0_0={},aq_c1_0={},aq_c2_0={},aq_c3_0={};
    uint4 aq_c0_1={},aq_c1_1={},aq_c2_1={},aq_c3_1={};
    int aq_m_idx_0=0, aq_ki_0=0, aq_m_idx_1=0, aq_ki_1=0;

    if (num_work > 0) {
        int w = w_start;
        aq_m_idx_0 = w / Ki; aq_ki_0 = w % Ki;
        int mt_q = blockIdx.x * N_M + aq_m_idx_0;
        int mg_q = mt_q * 16 + lid;
        int kb = aq_ki_0 << 7;
        if (mg_q < M) {
            const uint4* av = reinterpret_cast<const uint4*>(A + mg_q*Kh + ((kb+(gid<<5))>>1));
            aq_c0_0=av[0]; aq_c1_0=av[1]; aq_c2_0=av[2]; aq_c3_0=av[3];
        }
    }
    if (num_work > 1) {
        int w = w_start + 1;
        aq_m_idx_1 = w / Ki; aq_ki_1 = w % Ki;
        int mt_q = blockIdx.x * N_M + aq_m_idx_1;
        int mg_q = mt_q * 16 + lid;
        int kb = aq_ki_1 << 7;
        if (mg_q < M) {
            const uint4* av = reinterpret_cast<const uint4*>(A + mg_q*Kh + ((kb+(gid<<5))>>1));
            aq_c0_1=av[0]; aq_c1_1=av[1]; aq_c2_1=av[2]; aq_c3_1=av[3];
        }
    }
    // vmcnt now has 4*num_work A loads (up to 8 for 2 tiles)

    // Step 2: Issue GLL+Bsc (interleaved pairs) — BEFORE quant!
    // These go to vmcnt AFTER A loads. GLL+Bsc fly during quant ALU.
    // Pre-barrier: INTERLEAVED (buffer_load B[s], global_load Bsc[s]) pairs
    // vmcnt queue: [B0,Bsc0, B1,Bsc1, ..., B11,Bsc11] = 24 entries for BATCH=12
    // Each pair = 2 vmcnt entries. vmcnt(N) drains oldest, keeping N newest.
    uint32_t be0=0,be1=0,be2=0,be3=0,be4=0,be5=0,be6=0,be7=0,be8=0,be9=0,be10=0,be11=0,be12=0,be13=0,be14=0,be15=0;
    if (valid_wave) {
        // Setup buffer descriptor for B in s[16:19]
        const uint8_t* b_base = Bs + bb;
        uint64_t b_addr = (uint64_t)b_base;
        asm volatile(
            "s_mov_b32 s[16], %0\n\t"
            "s_mov_b32 s[17], %1\n\t"
            "s_mov_b32 s[18], 0xFFFFFFFF\n\t"
            "s_mov_b32 s[19], 0x20000"
            :: "s"(__builtin_amdgcn_readfirstlane((uint32_t)b_addr)),
               "s"(__builtin_amdgcn_readfirstlane((uint32_t)(b_addr >> 32)))
            : "s16", "s17", "s18", "s19"
        );

        // Interleaved: issue B[s] then Bsc[s] as a pair, repeat for each slot
        #define ISSUE_PAIR(S) do { \
            int _kb = (S) << 7; \
            uint32_t _m0v = (uint32_t)(my_b_byte + (S) * 1024); \
            uint32_t _buf_off = (_kb/64+gid/2)*512 + (gid&1)*256 + lid*16; \
            asm volatile( \
                "s_mov_b32 m0, %0\n\t" \
                "buffer_load_dwordx4 %1, s[16:19], 0 offen lds" \
                :: "s"(__builtin_amdgcn_readfirstlane(_m0v)), "v"(_buf_off) \
                : "memory", "s16", "s17", "s18", "s19" \
            ); \
            be##S = (uint32_t)Bsc[bsc_idx(bn, (_kb>>5)+gid, pk)]; \
        } while(0)

        // Issue interleaved pairs for pre-loaded BATCH slots only
        if constexpr (BATCH > 0)              ISSUE_PAIR(0);
        if constexpr (BATCH > 1)  { if (Ki>1)  ISSUE_PAIR(1); }
        if constexpr (BATCH > 2)  { if (Ki>2)  ISSUE_PAIR(2); }
        if constexpr (BATCH > 3)  { if (Ki>3)  ISSUE_PAIR(3); }
        if constexpr (BATCH > 4)  { if (Ki>4)  ISSUE_PAIR(4); }
        if constexpr (BATCH > 5)  { if (Ki>5)  ISSUE_PAIR(5); }
        if constexpr (BATCH > 6)  { if (Ki>6)  ISSUE_PAIR(6); }
        if constexpr (BATCH > 7)  { if (Ki>7)  ISSUE_PAIR(7); }
        if constexpr (BATCH > 8)  { if (Ki>8)  ISSUE_PAIR(8); }
        if constexpr (BATCH > 9)  { if (Ki>9)  ISSUE_PAIR(9); }
        if constexpr (BATCH > 10) { if (Ki>10) ISSUE_PAIR(10); }
        if constexpr (BATCH > 11) { if (Ki>11) ISSUE_PAIR(11); }

        // Bsc-only for tiles beyond BATCH (no GLL — loaded mid-loop or not at all)
        #define LOAD_BE_ONLY(I) do { int _kb=(I)<<7; be##I=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; } while(0)
        if constexpr (BATCH <= 6)  { if (Ki>6)  LOAD_BE_ONLY(6);  }
        if constexpr (BATCH <= 7)  { if (Ki>7)  LOAD_BE_ONLY(7);  }
        if constexpr (BATCH <= 8)  { if (Ki>8)  LOAD_BE_ONLY(8);  }
        if constexpr (BATCH <= 9)  { if (Ki>9)  LOAD_BE_ONLY(9);  }
        if constexpr (BATCH <= 10) { if (Ki>10) LOAD_BE_ONLY(10); }
        if constexpr (BATCH <= 11) { if (Ki>11) LOAD_BE_ONLY(11); }
        if (Ki>12) LOAD_BE_ONLY(12);
        if (Ki>13) LOAD_BE_ONLY(13);
        if (Ki>14) LOAD_BE_ONLY(14);
        if (Ki>15) LOAD_BE_ONLY(15);
        #undef LOAD_BE_ONLY

        #undef ISSUE_PAIR
    }
    // vmcnt queue: [A_0(×4), A_1(×4), GLL_0,Bsc_0, ..., GLL_11,Bsc_11, Bsc_12..15]
    // For S4: 8 (A) + 24 (12 GLL+Bsc pairs) + 4 (Bsc only) = 36

    // Step 3: Quant A from preloaded registers
    // Compiler will insert vmcnt(0) before using aq_c0_0 etc.
    // LANDMARK 3: marks the vmcnt before quant tile 0 for patching
    // After patch: vmcnt(0) → vmcnt(28) = drain 8 A loads, keep 28 GLL+Bsc
    asm volatile("s_nop 3" ::: "memory"); // LANDMARK 3: before quant tile 0
    if (num_work > 0) {
        v4i32 fp4; uint32_t e8m0;
        quant_a_direct(aq_c0_0, aq_c1_0, aq_c2_0, aq_c3_0, fp4, e8m0);
        int base = aq_m_idx_0 * Ki * A_TILE + aq_ki_0 * A_TILE;
        lds_gll[base+lane*A_FP4_STRIDE+0]=(uint32_t)fp4[0];
        lds_gll[base+lane*A_FP4_STRIDE+1]=(uint32_t)fp4[1];
        lds_gll[base+lane*A_FP4_STRIDE+2]=(uint32_t)fp4[2];
        lds_gll[base+lane*A_FP4_STRIDE+3]=(uint32_t)fp4[3];
        lds_gll[base+320+lane]=e8m0;
    }
    asm volatile("s_nop 4" ::: "memory"); // LANDMARK 4: before quant tile 1
    if (num_work > 1) {
        v4i32 fp4; uint32_t e8m0;
        quant_a_direct(aq_c0_1, aq_c1_1, aq_c2_1, aq_c3_1, fp4, e8m0);
        int base = aq_m_idx_1 * Ki * A_TILE + aq_ki_1 * A_TILE;
        lds_gll[base+lane*A_FP4_STRIDE+0]=(uint32_t)fp4[0];
        lds_gll[base+lane*A_FP4_STRIDE+1]=(uint32_t)fp4[1];
        lds_gll[base+lane*A_FP4_STRIDE+2]=(uint32_t)fp4[2];
        lds_gll[base+lane*A_FP4_STRIDE+3]=(uint32_t)fp4[3];
        lds_gll[base+320+lane]=e8m0;
    }

    __syncthreads();

    if (!valid_wave) return;

    const int a_offset = mt_local * Ki * A_TILE;
    v4f32 acc = {0,0,0,0};

    // Runtime MFMA iteration: ki and be are runtime values
    #define DO_MFMA_RT(KI_RT, SLOT, BE_VAR) do { \
        v4i32 _bf; \
        { int _boff = MY_B_DW + (SLOT) * 256 + lane * 4; \
          uint4 _bc = lds_load_u128((lds_u32_ptr)(&lds_gll[_boff])); \
          _bf[0]=(int)_bc.x; _bf[1]=(int)_bc.y; _bf[2]=(int)_bc.z; _bf[3]=(int)_bc.w; } \
        v4i32 _af; uint32_t _ae; \
        { int _base = a_offset + (KI_RT) * A_TILE; \
          _af[0]=(int)lds_gll[_base+lane*A_FP4_STRIDE+0]; _af[1]=(int)lds_gll[_base+lane*A_FP4_STRIDE+1]; \
          _af[2]=(int)lds_gll[_base+lane*A_FP4_STRIDE+2]; _af[3]=(int)lds_gll[_base+lane*A_FP4_STRIDE+3]; \
          _ae=lds_gll[_base+320+lane]; } \
        asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4" \
            :"+v"(acc):"v"(_af),"v"(_bf),"v"(_ae),"v"(BE_VAR)); \
    } while(0)

    #define ISSUE_GLL_RT(KI_RT, SLOT) do { \
        int _kb = (KI_RT) << 7; \
        uint32_t _buf_off = (_kb/64+gid/2)*512 + (gid&1)*256 + lid*16; \
        uint32_t _m0 = (uint32_t)(my_b_byte + (SLOT) * 1024); \
        const uint8_t* _bbase = Bs + bb; \
        uint64_t _baddr = (uint64_t)_bbase; \
        asm volatile( \
            "s_mov_b32 s[16], %2\n\t" \
            "s_mov_b32 s[17], %3\n\t" \
            "s_mov_b32 s[18], 0xFFFFFFFF\n\t" \
            "s_mov_b32 s[19], 0x20000\n\t" \
            "s_mov_b32 m0, %0\n\t" \
            "buffer_load_dwordx4 %1, s[16:19], 0 offen lds" \
            :: "s"(__builtin_amdgcn_readfirstlane(_m0)), "v"(_buf_off), \
               "s"(__builtin_amdgcn_readfirstlane((uint32_t)_baddr)), \
               "s"(__builtin_amdgcn_readfirstlane((uint32_t)(_baddr >> 32))) \
            : "memory", "s16", "s17", "s18", "s19" \
        ); \
    } while(0)

    #define LOAD_BSC_RT(VAR, KI_RT) do { int _kb=(KI_RT)<<7; VAR=(uint32_t)Bsc[bsc_idx(bn,(_kb>>5)+gid,pk)]; } while(0)

    if (Ki > 16) {
        // === BATCHED LOOP for large Ki (e.g., Ki=56) ===
        // Only for Ki > 16 where we can't use the unrolled paths.
        // Uses BATCH GLL slots + be0..be7 (8 Bsc), reloaded per batch.
        // First batch already pre-loaded (BATCH GLL + 8 Bsc).
        // vmcnt(8): drain BATCH GLL, keep 8 Bsc.
        // Hardware: BATCH + 8. Drain to 8 → removes BATCH GLL entries.
        asm volatile("s_waitcnt vmcnt(8)" ::: "memory");

        for (int batch = 0; batch < Ki; batch += BATCH) {
            int remaining = Ki - batch;
            int count = remaining < BATCH ? remaining : BATCH;

            // 8 unrolled MFMA iterations using be0..be7, slots 0..7
            if (count > 0) DO_MFMA_RT(batch+0, 0, be0);
            if (count > 1) DO_MFMA_RT(batch+1, 1, be1);
            if (count > 2) DO_MFMA_RT(batch+2, 2, be2);
            if (count > 3) DO_MFMA_RT(batch+3, 3, be3);
            if (count > 4) DO_MFMA_RT(batch+4, 4, be4);
            if (count > 5) DO_MFMA_RT(batch+5, 5, be5);
            if (count > 6) DO_MFMA_RT(batch+6, 6, be6);
            if (count > 7) DO_MFMA_RT(batch+7, 7, be7);

            int next_batch = batch + BATCH;
            if (next_batch < Ki) {
                int next_count = Ki - next_batch;
                if (next_count > BATCH) next_count = BATCH;

                // Issue next batch's GLL → slots 0..7 (already consumed)
                for (int s = 0; s < next_count; s++) {
                    ISSUE_GLL_RT(next_batch + s, s);
                }

                // Reload Bsc for next batch
                if (next_count > 0) LOAD_BSC_RT(be0, next_batch+0);
                if (next_count > 1) LOAD_BSC_RT(be1, next_batch+1);
                if (next_count > 2) LOAD_BSC_RT(be2, next_batch+2);
                if (next_count > 3) LOAD_BSC_RT(be3, next_batch+3);
                if (next_count > 4) LOAD_BSC_RT(be4, next_batch+4);
                if (next_count > 5) LOAD_BSC_RT(be5, next_batch+5);
                if (next_count > 6) LOAD_BSC_RT(be6, next_batch+6);
                if (next_count > 7) LOAD_BSC_RT(be7, next_batch+7);

                // Drain GLL, keep 8 Bsc. Hardware: next_count GLL + <=8 Bsc.
                // vmcnt(8): drain GLL entries, keep Bsc.
                asm volatile("s_waitcnt vmcnt(8)" ::: "memory");
            }
        }
    } else if constexpr (BATCH >= 16) {
        // vmcnt: drain BATCH GLL, keep Ki Bsc (Ki <= BATCH)
        if (Ki == 12) { asm volatile("s_waitcnt vmcnt(12)" ::: "memory"); }
        else if (Ki == 16) { asm volatile("s_waitcnt vmcnt(16)" ::: "memory"); }
        else { asm volatile("s_waitcnt vmcnt(0)" ::: "memory"); }
        // ALL Ki tiles pre-loaded — straight-through
        if (Ki > 0)  DO_MFMA_RT(0, 0, be0);   if (Ki > 1)  DO_MFMA_RT(1, 1, be1);
        if (Ki > 2)  DO_MFMA_RT(2, 2, be2);   if (Ki > 3)  DO_MFMA_RT(3, 3, be3);
        if (Ki > 4)  DO_MFMA_RT(4, 4, be4);   if (Ki > 5)  DO_MFMA_RT(5, 5, be5);
        if (Ki > 6)  DO_MFMA_RT(6, 6, be6);   if (Ki > 7)  DO_MFMA_RT(7, 7, be7);
        if (Ki > 8)  DO_MFMA_RT(8, 8, be8);   if (Ki > 9)  DO_MFMA_RT(9, 9, be9);
        if (Ki > 10) DO_MFMA_RT(10, 10, be10); if (Ki > 11) DO_MFMA_RT(11, 11, be11);
        if (Ki > 12) DO_MFMA_RT(12, 12, be12); if (Ki > 13) DO_MFMA_RT(13, 13, be13);
        if (Ki > 14) DO_MFMA_RT(14, 14, be14); if (Ki > 15) DO_MFMA_RT(15, 15, be15);
    } else if constexpr (BATCH == 12) {
        // BATCH=12: DOUBLE-BUFFERED lgkmcnt(4) ds_reads
        // Two register sets. All ds_reads via inline asm for lgkmcnt control.
        if (Ki == 12) { asm volatile("s_waitcnt vmcnt(12)" ::: "memory"); }
        else { asm volatile("s_waitcnt vmcnt(0)" ::: "memory"); }

        #define DB_READ_B(dst, addr) \
            asm volatile("ds_read_b128 %0, %1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
        #define DB_READ_A01(dst, addr) \
            asm volatile("ds_read2_b32 %0, %1 offset1:1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
        #define DB_READ_A23(dst, addr) \
            asm volatile("ds_read2_b32 %0, %1 offset0:2 offset1:3" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")
        #define DB_READ_AE(dst, addr) \
            asm volatile("ds_read_b32 %0, %1" : "=v"(dst) : "v"((uint32_t)((addr)*4)) : "memory")

        #define DB_ISSUE(KI, SLOT, bb, aa01, aa23, aae) do { \
            int _boff = MY_B_DW + (SLOT) * 256 + lane * 4; \
            DB_READ_B(bb, _boff); \
            int _abase = a_offset + (KI) * A_TILE + lane * A_FP4_STRIDE; \
            DB_READ_A01(aa01, _abase); \
            DB_READ_A23(aa23, _abase); \
            DB_READ_AE(aae, a_offset + (KI) * A_TILE + 320 + lane); \
        } while(0)

        #define DB_PACK(bb, aa01, aa23, aae, bf_v, af_v, ae_v) do { \
            bf_v[0]=(int)bb.x; bf_v[1]=(int)bb.y; bf_v[2]=(int)bb.z; bf_v[3]=(int)bb.w; \
            af_v[0]=(int)(uint32_t)aa01; af_v[1]=(int)(uint32_t)(aa01>>32); \
            af_v[2]=(int)(uint32_t)aa23; af_v[3]=(int)(uint32_t)(aa23>>32); \
            ae_v = (uint32_t)aae; \
        } while(0)

        #define DB_MFMA(af_v, bf_v, ae_v, be_v) \
            asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4" \
                :"+v"(acc):"v"(af_v),"v"(bf_v),"v"(ae_v),"v"(be_v))

        #define DB_TILE_A(KI, SLOT, BE_VAR, NEXT_KI, NEXT_SLOT) do { \
            DB_PACK(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae); \
            DB_MFMA(maf, mbf, mae, BE_VAR); \
            DB_ISSUE(NEXT_KI, NEXT_SLOT, bb_a, a01_a, a23_a, ae_a); \
            asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory"); \
        } while(0)

        #define DB_TILE_B(KI, SLOT, BE_VAR, NEXT_KI, NEXT_SLOT) do { \
            DB_PACK(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae); \
            DB_MFMA(maf, mbf, mae, BE_VAR); \
            DB_ISSUE(NEXT_KI, NEXT_SLOT, bb_b, a01_b, a23_b, ae_b); \
            asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory"); \
        } while(0)

        uint4 bb_a, bb_b; uint64_t a01_a, a23_a, a01_b, a23_b; uint32_t ae_a, ae_b;
        v4i32 maf, mbf; uint32_t mae;

        asm volatile("s_nop 1" ::: "memory"); // LANDMARK 1: after vmcnt drain, before MFMA

        // Prologue: prefetch tiles 0 and 1
        DB_ISSUE(0, 0, bb_a, a01_a, a23_a, ae_a);
        DB_ISSUE(1, 1, bb_b, a01_b, a23_b, ae_b);
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");

        DB_TILE_A(0, 0, be0, 2, 2);
        DB_TILE_B(1, 1, be1, 3, 3);

        // Tile 2: MFMA + issue GLL[12..15] → slots 0..3 (consumed)
        DB_PACK(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
        DB_MFMA(maf, mbf, mae, be2);
        if (Ki > 12) ISSUE_GLL_RT(12, 0);
        if (Ki > 13) ISSUE_GLL_RT(13, 1);
        if (Ki > 14) ISSUE_GLL_RT(14, 2);
        if (Ki > 15) ISSUE_GLL_RT(15, 3);
        DB_ISSUE(4, 4, bb_a, a01_a, a23_a, ae_a);
        asm volatile("s_waitcnt lgkmcnt(4)" ::: "memory");

        DB_TILE_B(3, 3, be3, 5, 5);
        DB_TILE_A(4, 4, be4, 6, 6);
        DB_TILE_B(5, 5, be5, 7, 7);
        DB_TILE_A(6, 6, be6, 8, 8);
        DB_TILE_B(7, 7, be7, 9, 9);
        DB_TILE_A(8, 8, be8, 10, 10);
        DB_TILE_B(9, 9, be9, 11, 11);

        // Tile 10: no more pre-loaded tiles
        DB_PACK(bb_a, a01_a, a23_a, ae_a, mbf, maf, mae);
        DB_MFMA(maf, mbf, mae, be10);
        asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");

        // Tile 11
        DB_PACK(bb_b, a01_b, a23_b, ae_b, mbf, maf, mae);
        DB_MFMA(maf, mbf, mae, be11);

        // Tiles 12-15: GLL loaded, use DO_MFMA_RT
        asm volatile("s_nop 2" ::: "memory"); // LANDMARK 2: before mid-loop vmcnt drain
        if (Ki > 12) {
            asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
            DO_MFMA_RT(12, 0, be12);
        }
        if (Ki > 13) DO_MFMA_RT(13, 1, be13);
        if (Ki > 14) DO_MFMA_RT(14, 2, be14);
        if (Ki > 15) DO_MFMA_RT(15, 3, be15);

        #undef DB_READ_B
        #undef DB_READ_A01
        #undef DB_READ_A23
        #undef DB_READ_AE
        #undef DB_ISSUE
        #undef DB_PACK
        #undef DB_MFMA
        #undef DB_TILE_A
        #undef DB_TILE_B
    } else if constexpr (BATCH == 6) {
        // BATCH=6: 3+3 double-buffer for Ki<=12
        // Pre-loaded: 6 GLL (slots 0..5) + Ki Bsc
        // vmcnt: drain 6 GLL, keep Ki Bsc
        if (Ki == 12) { asm volatile("s_waitcnt vmcnt(12)" ::: "memory"); }
        else { asm volatile("s_waitcnt vmcnt(0)" ::: "memory"); }
        // First 3 MFMAs: ki=0..2 (slots 0..2)
        if (Ki > 0) DO_MFMA_RT(0, 0, be0);
        if (Ki > 1) DO_MFMA_RT(1, 1, be1);
        if (Ki > 2) DO_MFMA_RT(2, 2, be2);
        // Burst GLL[6..8] → slots 0..2 (consumed)
        if (Ki > 6) ISSUE_GLL_RT(6, 0);
        if (Ki > 7) ISSUE_GLL_RT(7, 1);
        if (Ki > 8) ISSUE_GLL_RT(8, 2);
        // Next 3 MFMAs: ki=3..5 (slots 3..5)
        if (Ki > 3) DO_MFMA_RT(3, 3, be3);
        if (Ki > 4) DO_MFMA_RT(4, 4, be4);
        if (Ki > 5) DO_MFMA_RT(5, 5, be5);
        // Burst GLL[9..11] → slots 3..5 (consumed)
        if (Ki > 9)  ISSUE_GLL_RT(9, 3);
        if (Ki > 10) ISSUE_GLL_RT(10, 4);
        if (Ki > 11) ISSUE_GLL_RT(11, 5);
        // Drain all GLL before consuming batch 1
        if (Ki > 6) {
            asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
            DO_MFMA_RT(6, 0, be6);
        }
        if (Ki > 7)  DO_MFMA_RT(7, 1, be7);
        if (Ki > 8)  DO_MFMA_RT(8, 2, be8);
        if (Ki > 9)  DO_MFMA_RT(9, 3, be9);
        if (Ki > 10) DO_MFMA_RT(10, 4, be10);
        if (Ki > 11) DO_MFMA_RT(11, 5, be11);
    } else {
        // BATCH=8: batch processing for Ki=8..16
        // vmcnt: drain BATCH GLL, keep Ki Bsc
        if (Ki == 12) { asm volatile("s_waitcnt vmcnt(12)" ::: "memory"); }
        else if (Ki == 16) { asm volatile("s_waitcnt vmcnt(16)" ::: "memory"); }
        else { asm volatile("s_waitcnt vmcnt(8)" ::: "memory"); }
        // First half: ki=0..3 (slots 0..3)
        if (Ki > 0) DO_MFMA_RT(0, 0, be0);   if (Ki > 1) DO_MFMA_RT(1, 1, be1);
        if (Ki > 2) DO_MFMA_RT(2, 2, be2);   if (Ki > 3) DO_MFMA_RT(3, 3, be3);
        // Burst GLL[8..11] → slots 0..3
        if (Ki > 8)  ISSUE_GLL_RT(8, 0);  if (Ki > 9)  ISSUE_GLL_RT(9, 1);
        if (Ki > 10) ISSUE_GLL_RT(10, 2); if (Ki > 11) ISSUE_GLL_RT(11, 3);
        // Second half: ki=4..7 (slots 4..7)
        if (Ki > 4) DO_MFMA_RT(4, 4, be4);   if (Ki > 5) DO_MFMA_RT(5, 5, be5);
        if (Ki > 6) DO_MFMA_RT(6, 6, be6);   if (Ki > 7) DO_MFMA_RT(7, 7, be7);
        // Burst GLL[12..15] → slots 4..7
        if (Ki > 12) ISSUE_GLL_RT(12, 4); if (Ki > 13) ISSUE_GLL_RT(13, 5);
        if (Ki > 14) ISSUE_GLL_RT(14, 6); if (Ki > 15) ISSUE_GLL_RT(15, 7);
        // Drain GLL, consume batch 1
        if (Ki > 8) {
            asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
            DO_MFMA_RT(8, 0, be8);
        }
        if (Ki > 9)  DO_MFMA_RT(9, 1, be9);
        if (Ki > 10) DO_MFMA_RT(10, 2, be10); if (Ki > 11) DO_MFMA_RT(11, 3, be11);
        if (Ki > 12) DO_MFMA_RT(12, 4, be12); if (Ki > 13) DO_MFMA_RT(13, 5, be13);
        if (Ki > 14) DO_MFMA_RT(14, 6, be14); if (Ki > 15) DO_MFMA_RT(15, 7, be15);
    }

    #undef LOAD_BSC_RT
    #undef ISSUE_GLL_RT
    #undef DO_MFMA_RT

    store_tile(C, acc, mt, nt, gid, lid, M, N);
}
// ==========================================================================
// MASTER C++ DISPATCH: Raw pointers, zero ATen overhead in hot path.
// All buffers pre-allocated once in init(), reused via raw pointers.
// ==========================================================================
static uint8_t*  s_fp4_ptr  = nullptr;
static uint8_t*  s_sc_ptr   = nullptr;
static float*    s_temp_ptr = nullptr;
static uint16_t* s_C_ptr    = nullptr;
static int s_fp4_cap=0, s_sc_cap=0, s_temp_cap=0, s_C_cap=0;

// One-time allocation of max-size buffers for all benchmark shapes
void preallocate(int64_t max_M, int64_t max_N, int64_t max_K) {
    int max_MN = (int)(max_M * max_N);
    int max_Kh = (int)(max_K / 2);
    int max_Ksc = (int)(max_K / 32);
    int max_Mpad = ((int)max_M + 255) / 256 * 256;
    int max_Kscpad = (max_Ksc + 7) / 8 * 8;

    int fp4_need = (int)(max_M * max_Kh);
    int sc_need  = max_Mpad * max_Kscpad;
    int temp_need = max_MN;
    int c_need   = max_MN;

    if (fp4_need > s_fp4_cap) {
        if (s_fp4_ptr) hipFree(s_fp4_ptr);
        hipMalloc(&s_fp4_ptr, fp4_need); s_fp4_cap = fp4_need;
    }
    if (sc_need > s_sc_cap) {
        if (s_sc_ptr) hipFree(s_sc_ptr);
        hipMalloc(&s_sc_ptr, sc_need); s_sc_cap = sc_need;
    }
    if (temp_need * 4 > s_temp_cap) {
        if (s_temp_ptr) hipFree(s_temp_ptr);
        hipMalloc(&s_temp_ptr, temp_need * sizeof(float)); s_temp_cap = temp_need * 4;
    }
    if (c_need * 2 > s_C_cap) {
        if (s_C_ptr) hipFree(s_C_ptr);
        hipMalloc(&s_C_ptr, c_need * sizeof(uint16_t)); s_C_cap = c_need * 2;
    }
}

torch::Tensor dispatch(
    torch::Tensor A, torch::Tensor Bs, torch::Tensor Bsc,
    int64_t M, int64_t N, int64_t K
) {
    const int NUM_CUS = 304;
    int mt = (int)((M+15)/16), nt = (int)((N+15)/16);
    int tiles = mt * nt;
    int Ki = (int)(K >> 7);
    int kt = (int)(K / 64);
    int pk = (((int)(K/32)+7)/8)*8;

    // Output: cached tensor, only recreated if shape changes
    static torch::Tensor s_C_tensor;
    static int64_t s_C_M = 0, s_C_N = 0;
    if (M != s_C_M || N != s_C_N) {
        s_C_tensor = torch::from_blob(s_C_ptr, {M, N},
            torch::TensorOptions().dtype(torch::kBFloat16).device(A.device()));
        s_C_M = M; s_C_N = N;
    }
    uint16_t* C_ptr = (uint16_t*)s_C_tensor.data_ptr();

    // Helper macro: launch specialized fused kernel
    #define LAUNCH_FUSED(cM, cN, cK) do { \
        dim3 grid((unsigned)((cM+15)/16), (unsigned)((cN+15)/16)); \
        kern_fused_spec<cM, cN, cK><<<grid, 256>>>( \
            (const uint32_t*)A.data_ptr(), (const uint8_t*)Bs.data_ptr(), \
            (const uint8_t*)Bsc.data_ptr(), C_ptr); \
    } while(0)
    // Dispatch to specialized kernel if shape matches benchmark shapes
    // Specialized fused: K<=1024 shapes + M=16 K=7168 (underutilized split-k)
    if (M==4 && N==2880 && K==512)        { LAUNCH_FUSED(4, 2880, 512); }
    else if (M==16 && N==2112 && K==7168) {
        // Hybrid inline-A + GLL-B + nt — sweep winner: sp8_b2
        constexpr int BATCH_S1 = 2;
        constexpr int SPLITS = 8;
        dim3 grid(1u, 132u);
        size_t lds_b = (size_t)SPLITS * BATCH_S1 * 1024;
        size_t lds_red = (size_t)SPLITS * 256 * sizeof(float);
        kern_hybrid_s1<BATCH_S1, SPLITS><<<grid, SPLITS*64, lds_b + lds_red>>>(
            (const uint32_t*)A.data_ptr(), (const uint8_t*)Bs.data_ptr(),
            (const uint8_t*)Bsc.data_ptr(), C_ptr,
            16, 2112, 7168, kt, pk);
    }
    else if (M==32 && N==4096 && K==512)  { LAUNCH_FUSED(32, 4096, 512); }
    else if (M==32 && N==2880 && K==512)  { LAUNCH_FUSED(32, 2880, 512); }
    // v64: GLL BATCH=12 for Shape 4 — best config
    // 12 GLL pre-loaded, burst 4 mid-loop, 8 MFMAs overlap for burst
    // LDS: A(24576) + B(8 waves × 12 slots × 1024 = 98304) = 122880 = 120KB
    else if (M==64 && N==7168 && K==2048) {
        // Dedicated kern_s4: all optimizations, hardcoded constants
        dim3 grid(4u, 56u);
        size_t lds_a = (size_t)1 * 16 * 384 * sizeof(uint32_t);
        size_t lds_b = (size_t)8 * 10 * 1024;
        kern_s4<<<grid, 512, lds_a + lds_b>>>(
            (const uint32_t*)A.data_ptr(), (const uint8_t*)Bs.data_ptr(),
            (const uint8_t*)Bsc.data_ptr(), C_ptr,
            64, 7168, 2048, kt, pk);
    }
    else if (M==256 && N==3072 && K==1536) {
        // Dedicated kern_s5: Ki=12, N_N=12, BATCH=6
        dim3 grid(16u, 16u);
        size_t lds_a = (size_t)1 * 12 * 384 * sizeof(uint32_t);
        size_t lds_b = (size_t)12 * 8 * 1024;
        kern_s5<<<grid, 768, lds_a + lds_b>>>(
            (const uint32_t*)A.data_ptr(), (const uint8_t*)Bs.data_ptr(),
            (const uint8_t*)Bsc.data_ptr(), C_ptr,
            256, 3072, 1536, kt, pk);
    }
    else if (K <= 1024) {
        // Fallback for unknown K<=1024 shapes: use original kern_fused
        dim3 grid((unsigned)mt, (unsigned)nt);
        kern_fused<<<grid, 64>>>(
            (const uint32_t*)A.data_ptr(), (const uint8_t*)Bs.data_ptr(),
            (const uint8_t*)Bsc.data_ptr(), C_ptr,
            (int)M, (int)N, (int)K, kt, pk);
    } else {
        int Ksc = (int)(K / 32);
        int Mpad = ((int)M + 255) / 256 * 256;
        int Kscpad = (Ksc + 7) / 8 * 8;

        int qtot = (int)M * Ksc * QPG, qthr = 256, qblk = (qtot+qthr-1)/qthr;
        kern_quant<<<qblk, qthr>>>(
            (const uint32_t*)A.data_ptr(),
            s_fp4_ptr, s_sc_ptr,
            (int)M, (int)K, Ksc, Mpad, Kscpad);

        if (tiles < NUM_CUS && Ki >= 16) {
            int ksplit = Ki / 4;
            if (ksplit < 2) ksplit = 2;
            if (ksplit > 4) ksplit = 4;

            hipMemsetAsync(s_temp_ptr, 0, (int)(M*N) * sizeof(float), 0);

            int total_blocks = mt * nt * ksplit;
            kern_splitk<<<total_blocks, 64>>>(
                (const uint8_t*)s_fp4_ptr, (const uint8_t*)s_sc_ptr,
                (const uint8_t*)Bs.data_ptr(), (const uint8_t*)Bsc.data_ptr(),
                s_temp_ptr,
                (int)M, (int)N, (int)K, kt, pk, pk, nt, ksplit);

            int MN = (int)(M * N);
            kern_convert<<<(MN+255)/256, 256>>>(s_temp_ptr, C_ptr, MN);
        } else {
            dim3 grid((unsigned)mt, (unsigned)nt);
            kern_gemm<<<grid, 64>>>(
                (const uint8_t*)s_fp4_ptr, (const uint8_t*)s_sc_ptr,
                (const uint8_t*)Bs.data_ptr(), (const uint8_t*)Bsc.data_ptr(),
                C_ptr,
                (int)M, (int)N, (int)K, kt, pk, pk);
        }
    }
    return s_C_tensor;
}

// Zero-Python-overhead dispatch: takes the data tuple directly in C++
// Avoids: Python tuple unpack, .shape access, .contiguous(), pybind11 6-arg conversion
// data = (A, B, B_q, B_shuffle, B_scale_sh)
torch::Tensor dispatch_fast(py::tuple data) {
    // Direct PyTuple access — no vector construction, no per-element type check
    torch::Tensor A   = py::cast<torch::Tensor>(data[0]);
    torch::Tensor Bs  = py::cast<torch::Tensor>(data[3]);  // B_shuffle
    torch::Tensor Bsc = py::cast<torch::Tensor>(data[4]);  // B_scale_sh
    int64_t M = A.size(0);
    int64_t K = A.size(1);
    int64_t N = py::cast<torch::Tensor>(data[1]).size(0);  // B.shape[0]
    return dispatch(A, Bs, Bsc, M, N, K);
}
"""

CPP_SRC = """
namespace py = pybind11;
void preallocate(int64_t max_M, int64_t max_N, int64_t max_K);
torch::Tensor dispatch(torch::Tensor A, torch::Tensor Bs, torch::Tensor Bsc,
    int64_t M, int64_t N, int64_t K);
torch::Tensor dispatch_fast(py::tuple data);
"""

print("[v102] Compiling all-C++ dispatch kernel...")
hip_module = load_inline(
    name='fused_mxfp4_v102',
    cpp_sources=[CPP_SRC],
    cuda_sources=[HIP_SRC],
    functions=['preallocate', 'dispatch', 'dispatch_fast'],
    verbose=False,
    extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3",
                       "-mllvm", "-amdgpu-kernarg-preload-count=16"],
)

# Pre-allocate ALL buffers for all benchmark shapes (max M=256, N=7168, K=7168)
hip_module.preallocate(256, 7168, 7168)

print("[v102] Done.")


def custom_kernel(data: input_t) -> output_t:
    return hip_module.dispatch_fast(data)
scrolls · 1862 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