Skip to content
KernelIndex
Search⌘K

submission 596367

rjvkr2021 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-596367?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
9.29µs
#153 of 1143
2026-03-20

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:1fbdfb50df4b5af57cc0b2ef8b57625b81eff0858d119b04805c9b195183208a
license declaredunknown
license concludedunknown
authorsrjvkr2021
imported2026-08-15

Techniques

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

fp4FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
shared-memory__shared__ __align__(16) uint8_t Alds[BLOCK_M*LDS_ROW];__shared__ uint8_t Asclds[BLOCK_M*8];
tile-m = 16constexpr int MFMA_K=128, DOUBLE_K=256, LDS_ROW=128, HALF_K=64, SCALE_GROUP=32, BLOCK_M=16, BLOCK_M_LARGE=32;
tile-n = 64constexpr int BLOCK_N=64, BNT=128;
vector-width = int4const int4 r0 = reinterpret_cast<const int4*>(s)[0];

Kernel source

v3.py946 lines
"""
FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).
v3: Double-buffered K-loop for ksps>=4 to overlap memory loads with MFMA compute.
"""
import os

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t


os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")


CPP_SOURCE = r"""
#include <torch/extension.h>
torch::Tensor dispatch_gemm(torch::Tensor A, torch::Tensor Bq, torch::Tensor Bsc, int M, int N, int K);
"""


HIP_SOURCE = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <torch/extension.h>
constexpr int MFMA_K=128, DOUBLE_K=256, LDS_ROW=128, HALF_K=64, SCALE_GROUP=32, BLOCK_M=16, BLOCK_M_LARGE=32;
typedef int __attribute__((ext_vector_type(4))) int4_vec;
typedef float __attribute__((ext_vector_type(4))) float4_vec;
typedef int32_t __attribute__((ext_vector_type(4))) i32x4;
typedef __bf16 __attribute__((ext_vector_type(2))) bf16x2;
typedef uint32_t __attribute__((address_space(3)))* as3_uint32_ptr;
extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(i32x4 rsrc,as3_uint32_ptr lds_ptr,int size,int voffset,int soffset,int offset,int aux) __asm("llvm.amdgcn.raw.buffer.load.lds");
struct buffer_resource{uint64_t ptr;uint32_t range;uint32_t config;};
__device__ __forceinline__ i32x4 make_srsrc(const void*p,uint32_t r){buffer_resource s={reinterpret_cast<uint64_t>(p),r,0x110000};return *reinterpret_cast<const i32x4*>(&s);}
__device__ __forceinline__ float4_vec mfma_fp4(int4_vec A,int4_vec B,float4_vec C,int sA,int sB){float4_vec D;asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%3,%4,%5 cbsz:4 blgp:4":"=v"(D):"v"(A),"v"(B),"v"(C),"v"(sA),"v"(sB));return D;}
__device__ __forceinline__ int lds_swz(int o){return o^(((o&2047)>>8)<<4);}
__device__ __forceinline__ void compute_scale(float mx,uint8_t&sc,float&sf){if(mx>0.f){uint32_t b=__float_as_uint(mx);b=(b+0x200000u)&0xFF800000u;int su=((b>>23)&0xFF)-129;su=su<-127?-127:(su>127?127:su);sc=(uint8_t)(su+127);sf=__uint_as_float((uint32_t)(su+127)<<23);}else{sc=0;sf=0.f;}}
__device__ __forceinline__ bf16x2 as_bf16x2(uint32_t w){return __builtin_bit_cast(bf16x2,w);}
__device__ __forceinline__ void update_max_word(uint32_t w,float&mx){mx=fmaxf(mx,__uint_as_float((w&0x7fffu)<<16));mx=fmaxf(mx,__uint_as_float(w&0x7fff0000u));}
__device__ __forceinline__ uint32_t pack4_words(uint32_t w0,uint32_t w1,uint32_t w2,uint32_t w3,float sf){
    uint32_t p=0;
    p=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p,as_bf16x2(w0),sf,0);
    p=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p,as_bf16x2(w1),sf,1);
    p=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p,as_bf16x2(w2),sf,2);
    p=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(p,as_bf16x2(w3),sf,3);
    return p;
}

// A-quant helper: loads bf16 from global, quantizes to FP4, stores to LDS
__device__ __forceinline__ void do_a_quant(
    const __hip_bfloat16*__restrict__ A, int M, int K,
    int tid, int bm, int ke,
    uint8_t* Alds_buf, uint8_t* Asclds_buf)
{
    const int qr = tid >> 3, qg = tid & 7;
    const int gr = bm + qr, ko = ke + qg * 32;
    uint32_t p0=0, p1=0, p2=0, p3=0; uint8_t asc = 0x7f;
    if (gr < M) {
        const __hip_bfloat16* s = A + gr * K + ko;
        const int4 r0 = reinterpret_cast<const int4*>(s)[0];
        const int4 r1 = reinterpret_cast<const int4*>(s)[1];
        const int4 r2 = reinterpret_cast<const int4*>(s)[2];
        const int4 r3 = reinterpret_cast<const int4*>(s)[3];
        const uint32_t w0 = (uint32_t)r0.x, w1 = (uint32_t)r0.y, w2 = (uint32_t)r0.z, w3 = (uint32_t)r0.w;
        const uint32_t w4 = (uint32_t)r1.x, w5 = (uint32_t)r1.y, w6 = (uint32_t)r1.z, w7 = (uint32_t)r1.w;
        const uint32_t w8 = (uint32_t)r2.x, w9 = (uint32_t)r2.y, w10 = (uint32_t)r2.z, w11 = (uint32_t)r2.w;
        const uint32_t w12 = (uint32_t)r3.x, w13 = (uint32_t)r3.y, w14 = (uint32_t)r3.z, w15 = (uint32_t)r3.w;
        float l = 0;
        update_max_word(w0,l); update_max_word(w1,l); update_max_word(w2,l); update_max_word(w3,l);
        update_max_word(w4,l); update_max_word(w5,l); update_max_word(w6,l); update_max_word(w7,l);
        update_max_word(w8,l); update_max_word(w9,l); update_max_word(w10,l); update_max_word(w11,l);
        update_max_word(w12,l); update_max_word(w13,l); update_max_word(w14,l); update_max_word(w15,l);
        float sf; compute_scale(l, asc, sf);
        p0=pack4_words(w0,w1,w2,w3,sf);
        p1=pack4_words(w4,w5,w6,w7,sf);
        p2=pack4_words(w8,w9,w10,w11,sf);
        p3=pack4_words(w12,w13,w14,w15,sf);
    }
    int4 wr; wr.x=(int)p0; wr.y=(int)p1; wr.z=(int)p2; wr.w=(int)p3;
    *reinterpret_cast<int4*>(&Alds_buf[lds_swz(qr*LDS_ROW+qg*16)]) = wr;
    Asclds_buf[qr*8+qg] = asc;
}

// B-load helper: issues buffer_load_lds for B tiles
__device__ __forceinline__ void do_b_load(
    i32x4 srsrc, int btid, int bn, int N, int kb, int bstride,
    uint8_t* Blds_buf)
{
    constexpr int BLOCK_N=64, BNT=128;
    constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
    for (int ld = 0; ld < BL; ld++) {
        const int f = (ld * BNT + btid) << 4;
        const int r = f >> 7; const int gn = bn + r;
        if (r < BLOCK_N && gn < N) {
            const int sf = lds_swz(f), sc = sf & 127, ac = kb + sc;
            llvm_amdgcn_raw_buffer_load_lds(srsrc,
                (as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds_buf)+f),
                16,
                (gn>>4)*(bstride<<4)+(ac>>5)*512+((ac>>4)&1)*256+(gn&15)*16,
                0, 0, 0);
        }
    }
}

// B-scale load helper
__device__ __forceinline__ void do_b_scale_load(
    const uint8_t*__restrict__ Bsc, int stid, int bn, int N, int scs, int so,
    uint8_t* Bslds_buf)
{
    constexpr int BLOCK_N=64;
    const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
    const int gn = bn + row;
    if (gn < N) {
        const int rb=(gn>>5)*32*scs+(gn&15)*4+((gn>>4)&1);
        for(int g=0;g<4;g++){
            const int grp=gb+g,ac=so+grp;
            Bslds_buf[row*8+grp]=Bsc[rb+(ac&3)*64+((ac&7)>>2)*2+(ac>>3)*256];
        }
    } else {
        for(int g=0;g<4;g++) Bslds_buf[row*8+gb+g]=0x7f;
    }
}

// Original gemm_bn64 kernel (unchanged from baseline)
__global__ __launch_bounds__(256, 4)
void gemm_bn64(const __hip_bfloat16*__restrict__ A,const uint8_t*__restrict__ Bq,const uint8_t*__restrict__ Bsc,float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N,int K,int ksps){
    constexpr int BLOCK_N=64, NT=256;
    const int wid=threadIdx.x>>6,lm=threadIdx.x&15,lk=(threadIdx.x&63)>>4,tid=threadIdx.x;
    const int bn=blockIdx.x*BLOCK_N,bm=blockIdx.y*BLOCK_M,wn=bn+(wid<<4),sid=blockIdx.z;const int bstride=K>>1,scs=K>>5;
    __shared__ __align__(16) uint8_t Alds[BLOCK_M*LDS_ROW];__shared__ uint8_t Asclds[BLOCK_M*8];
    __shared__ __align__(16) uint8_t Blds[BLOCK_N*LDS_ROW];__shared__ uint8_t Bslds[BLOCK_N*8];
    const i32x4 srsrc=make_srsrc(Bq,N*bstride);float4_vec acc={0,0,0,0};
    for(int ks=sid*ksps;ks<sid*ksps+ksps;ks++){const int ke=ks*DOUBLE_K,kb=ks*LDS_ROW;
        if (tid < 128) {
            const int qr = tid >> 3, qg = tid & 7;
            const int gr = bm + qr, ko = ke + qg * 32;
            uint32_t p0=0, p1=0, p2=0, p3=0; uint8_t asc = 0x7f;
            if (gr < M) {
                const __hip_bfloat16* s = A + gr * K + ko;
                const int4 r0 = reinterpret_cast<const int4*>(s)[0];
                const int4 r1 = reinterpret_cast<const int4*>(s)[1];
                const int4 r2 = reinterpret_cast<const int4*>(s)[2];
                const int4 r3 = reinterpret_cast<const int4*>(s)[3];
                const uint32_t w0 = (uint32_t)r0.x, w1 = (uint32_t)r0.y, w2 = (uint32_t)r0.z, w3 = (uint32_t)r0.w;
                const uint32_t w4 = (uint32_t)r1.x, w5 = (uint32_t)r1.y, w6 = (uint32_t)r1.z, w7 = (uint32_t)r1.w;
                const uint32_t w8 = (uint32_t)r2.x, w9 = (uint32_t)r2.y, w10 = (uint32_t)r2.z, w11 = (uint32_t)r2.w;
                const uint32_t w12 = (uint32_t)r3.x, w13 = (uint32_t)r3.y, w14 = (uint32_t)r3.z, w15 = (uint32_t)r3.w;
                float l = 0;
                update_max_word(w0,l); update_max_word(w1,l); update_max_word(w2,l); update_max_word(w3,l);
                update_max_word(w4,l); update_max_word(w5,l); update_max_word(w6,l); update_max_word(w7,l);
                update_max_word(w8,l); update_max_word(w9,l); update_max_word(w10,l); update_max_word(w11,l);
                update_max_word(w12,l); update_max_word(w13,l); update_max_word(w14,l); update_max_word(w15,l);
                float sf; compute_scale(l, asc, sf);
                p0=pack4_words(w0,w1,w2,w3,sf);
                p1=pack4_words(w4,w5,w6,w7,sf);
                p2=pack4_words(w8,w9,w10,w11,sf);
                p3=pack4_words(w12,w13,w14,w15,sf);
            }
            int4 wr; wr.x=(int)p0; wr.y=(int)p1; wr.z=(int)p2; wr.w=(int)p3;
            *reinterpret_cast<int4*>(&Alds[lds_swz(qr*LDS_ROW+qg*16)]) = wr;
            Asclds[qr*8+qg] = asc;
        }

        if (tid >= 128) {
            const int btid = tid - 128;
            constexpr int BNT = 128;
            constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
            for (int ld = 0; ld < BL; ld++) {
                const int f = (ld * BNT + btid) << 4;
                const int r = f >> 7; const int gn = bn + r;
                if (r < BLOCK_N && gn < N) {
                    const int sf = lds_swz(f), sc = sf & 127, ac = kb + sc;
                    llvm_amdgcn_raw_buffer_load_lds(srsrc,(as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds)+f),16,(gn>>4)*(bstride<<4)+(ac>>5)*512+((ac>>4)&1)*256+(gn&15)*16,0,0,0);
                }
            }
        }

        {const int so=ks<<3;
        if (tid >= 128) {
            const int stid = tid - 128;
            const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
            const int gn = bn + row;
            if (gn < N) { const int rb=(gn>>5)*32*scs+(gn&15)*4+((gn>>4)&1);
                for(int g=0;g<4;g++){const int grp=gb+g,ac=so+grp;Bslds[row*8+grp]=Bsc[rb+(ac&3)*64+((ac&7)>>2)*2+(ac>>3)*256];}
            } else { for(int g=0;g<4;g++)Bslds[row*8+gb+g]=0x7f; }
        }}

        asm volatile("s_waitcnt vmcnt(0)");__syncthreads();

        int4_vec A0,A1;int as0,as1;
        {const int o=lds_swz(lm*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A0={t.x,t.y,t.z,t.w};}
        as0=(int)Asclds[(lm<<3)+lk];
        {const int o=lds_swz(lm*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A1={t.x,t.y,t.z,t.w};}
        as1=(int)Asclds[(lm<<3)+4+lk];
        const int br=(wid<<4)+lm;
        int4_vec B0,B1;int bs0,bs1;
        {const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B0={t.x,t.y,t.z,t.w};}
        bs0=(int)Bslds[br*8+lk];
        {const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B1={t.x,t.y,t.z,t.w};}
        bs1=(int)Bslds[br*8+4+lk];

        acc=mfma_fp4(A0,B0,acc,as0,bs0);
        acc=mfma_fp4(A1,B1,acc,as1,bs1);
    }

    const int or_=bm+(lk<<2),oc=wn+lm;
    if(oc<N){const float*ap=reinterpret_cast<const float*>(&acc);
    if(C){for(int r=0;r<4;r++){int g=or_+r;if(g<M)C[g*N+oc]=__float2bfloat16(ap[r]);}}
    else{float*w=ws+sid*M*N;for(int r=0;r<4;r++){int g=or_+r;if(g<M)w[g*N+oc]=ap[r];}}}
}

template<int KSPS>
__global__ __launch_bounds__(256, 4)
void gemm_bn64_k(const __hip_bfloat16*__restrict__ A,const uint8_t*__restrict__ Bq,const uint8_t*__restrict__ Bsc,float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N,int K){
    constexpr int BLOCK_N=64, NT=256;
    const int wid=threadIdx.x>>6,lm=threadIdx.x&15,lk=(threadIdx.x&63)>>4,tid=threadIdx.x;
    const int bn=blockIdx.x*BLOCK_N,bm=blockIdx.y*BLOCK_M,wn=bn+(wid<<4),sid=blockIdx.z;const int bstride=K>>1,scs=K>>5;
    __shared__ __align__(16) uint8_t Alds[BLOCK_M*LDS_ROW];__shared__ uint8_t Asclds[BLOCK_M*8];
    __shared__ __align__(16) uint8_t Blds[BLOCK_N*LDS_ROW];__shared__ uint8_t Bslds[BLOCK_N*8];
    const i32x4 srsrc=make_srsrc(Bq,N*bstride);float4_vec acc={0,0,0,0};
    #pragma unroll
    for(int kstep=0;kstep<KSPS;kstep++){const int ks=sid*KSPS+kstep,ke=ks*DOUBLE_K,kb=ks*LDS_ROW;
        if (tid < 128) {
            const int qr = tid >> 3, qg = tid & 7;
            const int gr = bm + qr, ko = ke + qg * 32;
            uint32_t p0=0, p1=0, p2=0, p3=0; uint8_t asc = 0x7f;
            if (gr < M) {
                const __hip_bfloat16* s = A + gr * K + ko;
                const int4 r0 = reinterpret_cast<const int4*>(s)[0];
                const int4 r1 = reinterpret_cast<const int4*>(s)[1];
                const int4 r2 = reinterpret_cast<const int4*>(s)[2];
                const int4 r3 = reinterpret_cast<const int4*>(s)[3];
                const uint32_t w0 = (uint32_t)r0.x, w1 = (uint32_t)r0.y, w2 = (uint32_t)r0.z, w3 = (uint32_t)r0.w;
                const uint32_t w4 = (uint32_t)r1.x, w5 = (uint32_t)r1.y, w6 = (uint32_t)r1.z, w7 = (uint32_t)r1.w;
                const uint32_t w8 = (uint32_t)r2.x, w9 = (uint32_t)r2.y, w10 = (uint32_t)r2.z, w11 = (uint32_t)r2.w;
                const uint32_t w12 = (uint32_t)r3.x, w13 = (uint32_t)r3.y, w14 = (uint32_t)r3.z, w15 = (uint32_t)r3.w;
                float l = 0;
                update_max_word(w0,l); update_max_word(w1,l); update_max_word(w2,l); update_max_word(w3,l);
                update_max_word(w4,l); update_max_word(w5,l); update_max_word(w6,l); update_max_word(w7,l);
                update_max_word(w8,l); update_max_word(w9,l); update_max_word(w10,l); update_max_word(w11,l);
                update_max_word(w12,l); update_max_word(w13,l); update_max_word(w14,l); update_max_word(w15,l);
                float sf; compute_scale(l, asc, sf);
                p0=pack4_words(w0,w1,w2,w3,sf);
                p1=pack4_words(w4,w5,w6,w7,sf);
                p2=pack4_words(w8,w9,w10,w11,sf);
                p3=pack4_words(w12,w13,w14,w15,sf);
            }
            int4 wr; wr.x=(int)p0; wr.y=(int)p1; wr.z=(int)p2; wr.w=(int)p3;
            *reinterpret_cast<int4*>(&Alds[lds_swz(qr*LDS_ROW+qg*16)]) = wr;
            Asclds[qr*8+qg] = asc;
        }

        if (tid >= 128) {
            const int btid = tid - 128;
            constexpr int BNT = 128;
            constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
            for (int ld = 0; ld < BL; ld++) {
                const int f = (ld * BNT + btid) << 4;
                const int r = f >> 7; const int gn = bn + r;
                if (r < BLOCK_N && gn < N) {
                    const int sf = lds_swz(f), sc = sf & 127, ac = kb + sc;
                    llvm_amdgcn_raw_buffer_load_lds(srsrc,(as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds)+f),16,(gn>>4)*(bstride<<4)+(ac>>5)*512+((ac>>4)&1)*256+(gn&15)*16,0,0,0);
                }
            }
        }

        {const int so=ks<<3;
        if (tid >= 128) {
            const int stid = tid - 128;
            const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
            const int gn = bn + row;
            if (gn < N) { const int rb=(gn>>5)*32*scs+(gn&15)*4+((gn>>4)&1);
                for(int g=0;g<4;g++){const int grp=gb+g,ac=so+grp;Bslds[row*8+grp]=Bsc[rb+(ac&3)*64+((ac&7)>>2)*2+(ac>>3)*256];}
            } else { for(int g=0;g<4;g++)Bslds[row*8+gb+g]=0x7f; }
        }}

        asm volatile("s_waitcnt vmcnt(0)");__syncthreads();

        int4_vec A0,A1;int as0,as1;
        {const int o=lds_swz(lm*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A0={t.x,t.y,t.z,t.w};}
        as0=(int)Asclds[(lm<<3)+lk];
        {const int o=lds_swz(lm*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A1={t.x,t.y,t.z,t.w};}
        as1=(int)Asclds[(lm<<3)+4+lk];
        const int br=(wid<<4)+lm;
        int4_vec B0,B1;int bs0,bs1;
        {const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B0={t.x,t.y,t.z,t.w};}
        bs0=(int)Bslds[br*8+lk];
        {const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B1={t.x,t.y,t.z,t.w};}
        bs1=(int)Bslds[br*8+4+lk];

        acc=mfma_fp4(A0,B0,acc,as0,bs0);
        acc=mfma_fp4(A1,B1,acc,as1,bs1);
    }

    const int or_=bm+(lk<<2),oc=wn+lm;
    if(oc<N){const float*ap=reinterpret_cast<const float*>(&acc);
    if(C){for(int r=0;r<4;r++){int g=or_+r;if(g<M)C[g*N+oc]=__float2bfloat16(ap[r]);}}
    else{float*w=ws+sid*M*N;for(int r=0;r<4;r++){int g=or_+r;if(g<M)w[g*N+oc]=ap[r];}}}
}

// Double-buffered version for KSPS >= 4
// Overlaps memory loads for step N+1 with MFMA compute for step N.
template<int KSPS>
__global__ __launch_bounds__(256, 2)
void gemm_bn64_db(const __hip_bfloat16*__restrict__ A,const uint8_t*__restrict__ Bq,const uint8_t*__restrict__ Bsc,float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N,int K){
    constexpr int BLOCK_N=64;
    const int wid=threadIdx.x>>6,lm=threadIdx.x&15,lk=(threadIdx.x&63)>>4,tid=threadIdx.x;
    const int bn=blockIdx.x*BLOCK_N,bm=blockIdx.y*BLOCK_M,wn=bn+(wid<<4),sid=blockIdx.z;
    const int bstride=K>>1,scs=K>>5;

    // Double-buffered LDS
    __shared__ __align__(16) uint8_t Alds[2][BLOCK_M*LDS_ROW];
    __shared__ uint8_t Asclds[2][BLOCK_M*8];
    __shared__ __align__(16) uint8_t Blds[2][BLOCK_N*LDS_ROW];
    __shared__ uint8_t Bslds[2][BLOCK_N*8];

    const i32x4 srsrc=make_srsrc(Bq,N*bstride);
    float4_vec acc={0,0,0,0};

    const int ks_base = sid * KSPS;

    // === PROLOGUE: load step 0 into buffer 0 ===
    {
        const int ks = ks_base;
        const int ke = ks * DOUBLE_K, kb = ks * LDS_ROW;
        if (tid < 128) {
            do_a_quant(A, M, K, tid, bm, ke, Alds[0], Asclds[0]);
        }
        if (tid >= 128) {
            do_b_load(srsrc, tid - 128, bn, N, kb, bstride, Blds[0]);
        }
        {
            const int so = ks << 3;
            if (tid >= 128) {
                do_b_scale_load(Bsc, tid - 128, bn, N, scs, so, Bslds[0]);
            }
        }
        asm volatile("s_waitcnt vmcnt(0)");
        __syncthreads();
    }

    // === MAIN LOOP: for steps 1..KSPS-1, load next step while computing current ===
    #pragma unroll
    for (int kstep = 1; kstep < KSPS; kstep++) {
        const int cur = (kstep - 1) & 1;  // buffer with ready data
        const int nxt = kstep & 1;         // buffer to load into

        const int ks_next = ks_base + kstep;
        const int ke_next = ks_next * DOUBLE_K, kb_next = ks_next * LDS_ROW;

        // Read current step's data from LDS into registers FIRST (before issuing new loads)
        int4_vec A0,A1; int as0,as1;
        {const int o=lds_swz(lm*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A0={t.x,t.y,t.z,t.w};}
        as0=(int)Asclds[cur][(lm<<3)+lk];
        {const int o=lds_swz(lm*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A1={t.x,t.y,t.z,t.w};}
        as1=(int)Asclds[cur][(lm<<3)+4+lk];

        const int br=(wid<<4)+lm;
        int4_vec B0,B1; int bs0,bs1;
        {const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B0={t.x,t.y,t.z,t.w};}
        bs0=(int)Bslds[cur][br*8+lk];
        {const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B1={t.x,t.y,t.z,t.w};}
        bs1=(int)Bslds[cur][br*8+4+lk];

        // Issue loads for next step into nxt buffer
        // These overlap with the MFMA below since they target different LDS regions
        if (tid < 128) {
            do_a_quant(A, M, K, tid, bm, ke_next, Alds[nxt], Asclds[nxt]);
        }
        if (tid >= 128) {
            do_b_load(srsrc, tid - 128, bn, N, kb_next, bstride, Blds[nxt]);
        }
        {
            const int so = ks_next << 3;
            if (tid >= 128) {
                do_b_scale_load(Bsc, tid - 128, bn, N, scs, so, Bslds[nxt]);
            }
        }

        // MFMA on current step's data (from registers)
        acc=mfma_fp4(A0,B0,acc,as0,bs0);
        acc=mfma_fp4(A1,B1,acc,as1,bs1);

        // Wait for next step's loads to complete
        asm volatile("s_waitcnt vmcnt(0)");
        __syncthreads();
    }

    // === EPILOGUE: compute last step from the last loaded buffer ===
    {
        const int cur = (KSPS - 1) & 1;

        int4_vec A0,A1; int as0,as1;
        {const int o=lds_swz(lm*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A0={t.x,t.y,t.z,t.w};}
        as0=(int)Asclds[cur][(lm<<3)+lk];
        {const int o=lds_swz(lm*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A1={t.x,t.y,t.z,t.w};}
        as1=(int)Asclds[cur][(lm<<3)+4+lk];

        const int br=(wid<<4)+lm;
        int4_vec B0,B1; int bs0,bs1;
        {const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B0={t.x,t.y,t.z,t.w};}
        bs0=(int)Bslds[cur][br*8+lk];
        {const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B1={t.x,t.y,t.z,t.w};}
        bs1=(int)Bslds[cur][br*8+4+lk];

        acc=mfma_fp4(A0,B0,acc,as0,bs0);
        acc=mfma_fp4(A1,B1,acc,as1,bs1);
    }

    const int or_=bm+(lk<<2),oc=wn+lm;
    if(oc<N){const float*ap=reinterpret_cast<const float*>(&acc);
    if(C){for(int r=0;r<4;r++){int g=or_+r;if(g<M)C[g*N+oc]=__float2bfloat16(ap[r]);}}
    else{float*w=ws+sid*M*N;for(int r=0;r<4;r++){int g=or_+r;if(g<M)w[g*N+oc]=ap[r];}}}
}

// Double-buffered version for large-M (BLOCK_M=32, 512 threads)
template<int KSPS>
__global__ __launch_bounds__(512, 1)
void gemm_bn64_m32_db(const __hip_bfloat16*__restrict__ A,const uint8_t*__restrict__ Bq,const uint8_t*__restrict__ Bsc,float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N,int K){
    constexpr int BLOCK_N=64;
    const int wid=threadIdx.x>>6,lm=threadIdx.x&15,lk=(threadIdx.x&63)>>4,tid=threadIdx.x;
    const int wg=wid>>2,wc=wid&3;
    const int bn=blockIdx.x*BLOCK_N,bm=blockIdx.y*BLOCK_M_LARGE,col=bn+(wc<<4),sid=blockIdx.z;
    const int bstride=K>>1,scs=K>>5;

    __shared__ __align__(16) uint8_t Alds[2][BLOCK_M_LARGE*LDS_ROW];
    __shared__ uint8_t Asclds[2][BLOCK_M_LARGE*8];
    __shared__ __align__(16) uint8_t Blds[2][BLOCK_N*LDS_ROW];
    __shared__ uint8_t Bslds[2][BLOCK_N*8];

    const i32x4 srsrc=make_srsrc(Bq,N*bstride);
    float4_vec acc={0,0,0,0};

    const int ks_base = sid * KSPS;

    // === PROLOGUE: load step 0 into buffer 0 ===
    {
        const int ks = ks_base;
        const int ke = ks * DOUBLE_K, kb = ks * LDS_ROW;

        if (tid < 256) {
            const int qr = tid >> 3, qg = tid & 7;
            const int gr = bm + qr, ko = ke + qg * 32;
            uint32_t p0=0, p1=0, p2=0, p3=0; uint8_t asc = 0x7f;
            if (gr < M) {
                const __hip_bfloat16* s = A + gr * K + ko;
                const int4 r0 = reinterpret_cast<const int4*>(s)[0];
                const int4 r1 = reinterpret_cast<const int4*>(s)[1];
                const int4 r2 = reinterpret_cast<const int4*>(s)[2];
                const int4 r3 = reinterpret_cast<const int4*>(s)[3];
                const uint32_t w0 = (uint32_t)r0.x, w1 = (uint32_t)r0.y, w2 = (uint32_t)r0.z, w3 = (uint32_t)r0.w;
                const uint32_t w4 = (uint32_t)r1.x, w5 = (uint32_t)r1.y, w6 = (uint32_t)r1.z, w7 = (uint32_t)r1.w;
                const uint32_t w8 = (uint32_t)r2.x, w9 = (uint32_t)r2.y, w10 = (uint32_t)r2.z, w11 = (uint32_t)r2.w;
                const uint32_t w12 = (uint32_t)r3.x, w13 = (uint32_t)r3.y, w14 = (uint32_t)r3.z, w15 = (uint32_t)r3.w;
                float l = 0;
                update_max_word(w0,l); update_max_word(w1,l); update_max_word(w2,l); update_max_word(w3,l);
                update_max_word(w4,l); update_max_word(w5,l); update_max_word(w6,l); update_max_word(w7,l);
                update_max_word(w8,l); update_max_word(w9,l); update_max_word(w10,l); update_max_word(w11,l);
                update_max_word(w12,l); update_max_word(w13,l); update_max_word(w14,l); update_max_word(w15,l);
                float sf; compute_scale(l, asc, sf);
                p0=pack4_words(w0,w1,w2,w3,sf);
                p1=pack4_words(w4,w5,w6,w7,sf);
                p2=pack4_words(w8,w9,w10,w11,sf);
                p3=pack4_words(w12,w13,w14,w15,sf);
            }
            int4 wr; wr.x=(int)p0; wr.y=(int)p1; wr.z=(int)p2; wr.w=(int)p3;
            *reinterpret_cast<int4*>(&Alds[0][lds_swz(qr*LDS_ROW+qg*16)]) = wr;
            Asclds[0][qr*8+qg] = asc;
        }

        if (tid >= 256) {
            const int btid = tid - 256;
            constexpr int BNT = 256;
            constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
            for (int ld = 0; ld < BL; ld++) {
                const int f = (ld * BNT + btid) << 4;
                const int r = f >> 7; const int gn = bn + r;
                if (r < BLOCK_N && gn < N) {
                    const int sf = lds_swz(f), sc = sf & 127, ac = kb + sc;
                    llvm_amdgcn_raw_buffer_load_lds(srsrc,(as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds[0])+f),16,(gn>>4)*(bstride<<4)+(ac>>5)*512+((ac>>4)&1)*256+(gn&15)*16,0,0,0);
                }
            }
        }

        {
            const int so = ks << 3;
            if (tid >= 256 && tid < 384) {
                const int stid = tid - 256;
                const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
                const int gn = bn + row;
                if (gn < N) { const int rb=(gn>>5)*32*scs+(gn&15)*4+((gn>>4)&1);
                    for(int g=0;g<4;g++){const int grp=gb+g,ac=so+grp;Bslds[0][row*8+grp]=Bsc[rb+(ac&3)*64+((ac&7)>>2)*2+(ac>>3)*256];}
                } else { for(int g=0;g<4;g++)Bslds[0][row*8+gb+g]=0x7f; }
            }
        }

        asm volatile("s_waitcnt vmcnt(0)");
        __syncthreads();
    }

    // === MAIN LOOP ===
    #pragma unroll
    for (int kstep = 1; kstep < KSPS; kstep++) {
        const int cur = (kstep - 1) & 1;
        const int nxt = kstep & 1;
        const int ks_next = ks_base + kstep;
        const int ke_next = ks_next * DOUBLE_K, kb_next = ks_next * LDS_ROW;

        // Read current data from LDS
        const int ar=(wg<<4)+lm;
        int4_vec A0,A1; int as0,as1;
        {const int o=lds_swz(ar*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A0={t.x,t.y,t.z,t.w};}
        as0=(int)Asclds[cur][ar*8+lk];
        {const int o=lds_swz(ar*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A1={t.x,t.y,t.z,t.w};}
        as1=(int)Asclds[cur][ar*8+4+lk];
        const int br=(wc<<4)+lm;
        int4_vec B0,B1; int bs0,bs1;
        {const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B0={t.x,t.y,t.z,t.w};}
        bs0=(int)Bslds[cur][br*8+lk];
        {const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B1={t.x,t.y,t.z,t.w};}
        bs1=(int)Bslds[cur][br*8+4+lk];

        // Issue loads for next step
        if (tid < 256) {
            const int qr = tid >> 3, qg = tid & 7;
            const int gr = bm + qr, ko = ke_next + qg * 32;
            uint32_t p0=0, p1=0, p2=0, p3=0; uint8_t asc = 0x7f;
            if (gr < M) {
                const __hip_bfloat16* s = A + gr * K + ko;
                const int4 r0 = reinterpret_cast<const int4*>(s)[0];
                const int4 r1 = reinterpret_cast<const int4*>(s)[1];
                const int4 r2 = reinterpret_cast<const int4*>(s)[2];
                const int4 r3 = reinterpret_cast<const int4*>(s)[3];
                const uint32_t ww0 = (uint32_t)r0.x, ww1 = (uint32_t)r0.y, ww2 = (uint32_t)r0.z, ww3 = (uint32_t)r0.w;
                const uint32_t ww4 = (uint32_t)r1.x, ww5 = (uint32_t)r1.y, ww6 = (uint32_t)r1.z, ww7 = (uint32_t)r1.w;
                const uint32_t ww8 = (uint32_t)r2.x, ww9 = (uint32_t)r2.y, ww10 = (uint32_t)r2.z, ww11 = (uint32_t)r2.w;
                const uint32_t ww12 = (uint32_t)r3.x, ww13 = (uint32_t)r3.y, ww14 = (uint32_t)r3.z, ww15 = (uint32_t)r3.w;
                float l = 0;
                update_max_word(ww0,l); update_max_word(ww1,l); update_max_word(ww2,l); update_max_word(ww3,l);
                update_max_word(ww4,l); update_max_word(ww5,l); update_max_word(ww6,l); update_max_word(ww7,l);
                update_max_word(ww8,l); update_max_word(ww9,l); update_max_word(ww10,l); update_max_word(ww11,l);
                update_max_word(ww12,l); update_max_word(ww13,l); update_max_word(ww14,l); update_max_word(ww15,l);
                float sf; compute_scale(l, asc, sf);
                p0=pack4_words(ww0,ww1,ww2,ww3,sf);
                p1=pack4_words(ww4,ww5,ww6,ww7,sf);
                p2=pack4_words(ww8,ww9,ww10,ww11,sf);
                p3=pack4_words(ww12,ww13,ww14,ww15,sf);
            }
            int4 wr; wr.x=(int)p0; wr.y=(int)p1; wr.z=(int)p2; wr.w=(int)p3;
            *reinterpret_cast<int4*>(&Alds[nxt][lds_swz(qr*LDS_ROW+qg*16)]) = wr;
            Asclds[nxt][qr*8+qg] = asc;
        }

        if (tid >= 256) {
            const int btid = tid - 256;
            constexpr int BNT = 256;
            constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
            for (int ld = 0; ld < BL; ld++) {
                const int f = (ld * BNT + btid) << 4;
                const int r = f >> 7; const int gn = bn + r;
                if (r < BLOCK_N && gn < N) {
                    const int sf = lds_swz(f), sc = sf & 127, ac = kb_next + sc;
                    llvm_amdgcn_raw_buffer_load_lds(srsrc,(as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds[nxt])+f),16,(gn>>4)*(bstride<<4)+(ac>>5)*512+((ac>>4)&1)*256+(gn&15)*16,0,0,0);
                }
            }
        }

        {
            const int so = ks_next << 3;
            if (tid >= 256 && tid < 384) {
                const int stid = tid - 256;
                const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
                const int gn = bn + row;
                if (gn < N) { const int rb=(gn>>5)*32*scs+(gn&15)*4+((gn>>4)&1);
                    for(int g=0;g<4;g++){const int grp=gb+g,ac=so+grp;Bslds[nxt][row*8+grp]=Bsc[rb+(ac&3)*64+((ac&7)>>2)*2+(ac>>3)*256];}
                } else { for(int g=0;g<4;g++)Bslds[nxt][row*8+gb+g]=0x7f; }
            }
        }

        // MFMA on current step's data
        acc=mfma_fp4(A0,B0,acc,as0,bs0);
        acc=mfma_fp4(A1,B1,acc,as1,bs1);

        asm volatile("s_waitcnt vmcnt(0)");
        __syncthreads();
    }

    // === EPILOGUE ===
    {
        const int cur = (KSPS - 1) & 1;
        const int ar=(wg<<4)+lm;
        int4_vec A0,A1; int as0,as1;
        {const int o=lds_swz(ar*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A0={t.x,t.y,t.z,t.w};}
        as0=(int)Asclds[cur][ar*8+lk];
        {const int o=lds_swz(ar*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[cur][o]);A1={t.x,t.y,t.z,t.w};}
        as1=(int)Asclds[cur][ar*8+4+lk];
        const int br=(wc<<4)+lm;
        int4_vec B0,B1; int bs0,bs1;
        {const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B0={t.x,t.y,t.z,t.w};}
        bs0=(int)Bslds[cur][br*8+lk];
        {const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[cur][o]);B1={t.x,t.y,t.z,t.w};}
        bs1=(int)Bslds[cur][br*8+4+lk];

        acc=mfma_fp4(A0,B0,acc,as0,bs0);
        acc=mfma_fp4(A1,B1,acc,as1,bs1);
    }

    const int or_=bm+(wg<<4)+(lk<<2),oc=col+lm;
    if(oc<N){const float*ap=reinterpret_cast<const float*>(&acc);
    if(C){for(int r=0;r<4;r++){int g=or_+r;if(g<M)C[g*N+oc]=__float2bfloat16(ap[r]);}}
    else{float*w=ws+sid*M*N;for(int r=0;r<4;r++){int g=or_+r;if(g<M)w[g*N+oc]=ap[r];}}}
}

// Original gemm_bn64_m32 kernel (unchanged)
__global__ void gemm_bn64_m32(const __hip_bfloat16*__restrict__ A,const uint8_t*__restrict__ Bq,const uint8_t*__restrict__ Bsc,float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N,int K,int ksps){
    constexpr int BLOCK_N=64;
    const int wid=threadIdx.x>>6,lm=threadIdx.x&15,lk=(threadIdx.x&63)>>4,tid=threadIdx.x;
    const int wg=wid>>2,wc=wid&3;
    const int bn=blockIdx.x*BLOCK_N,bm=blockIdx.y*BLOCK_M_LARGE,col=bn+(wc<<4),sid=blockIdx.z;const int bstride=K>>1,scs=K>>5;
    __shared__ __align__(16) uint8_t Alds[BLOCK_M_LARGE*LDS_ROW];__shared__ uint8_t Asclds[BLOCK_M_LARGE*8];
    __shared__ __align__(16) uint8_t Blds[BLOCK_N*LDS_ROW];__shared__ uint8_t Bslds[BLOCK_N*8];
    const i32x4 srsrc=make_srsrc(Bq,N*bstride);float4_vec acc={0,0,0,0};
    for(int ks=sid*ksps;ks<sid*ksps+ksps;ks++){const int ke=ks*DOUBLE_K,kb=ks*LDS_ROW;
        if (tid < 256) {
            const int qr = tid >> 3, qg = tid & 7;
            const int gr = bm + qr, ko = ke + qg * 32;
            uint32_t p0=0, p1=0, p2=0, p3=0; uint8_t asc = 0x7f;
            if (gr < M) {
                const __hip_bfloat16* s = A + gr * K + ko;
                const int4 r0 = reinterpret_cast<const int4*>(s)[0];
                const int4 r1 = reinterpret_cast<const int4*>(s)[1];
                const int4 r2 = reinterpret_cast<const int4*>(s)[2];
                const int4 r3 = reinterpret_cast<const int4*>(s)[3];
                const uint32_t w0 = (uint32_t)r0.x, w1 = (uint32_t)r0.y, w2 = (uint32_t)r0.z, w3 = (uint32_t)r0.w;
                const uint32_t w4 = (uint32_t)r1.x, w5 = (uint32_t)r1.y, w6 = (uint32_t)r1.z, w7 = (uint32_t)r1.w;
                const uint32_t w8 = (uint32_t)r2.x, w9 = (uint32_t)r2.y, w10 = (uint32_t)r2.z, w11 = (uint32_t)r2.w;
                const uint32_t w12 = (uint32_t)r3.x, w13 = (uint32_t)r3.y, w14 = (uint32_t)r3.z, w15 = (uint32_t)r3.w;
                float l = 0;
                update_max_word(w0,l); update_max_word(w1,l); update_max_word(w2,l); update_max_word(w3,l);
                update_max_word(w4,l); update_max_word(w5,l); update_max_word(w6,l); update_max_word(w7,l);
                update_max_word(w8,l); update_max_word(w9,l); update_max_word(w10,l); update_max_word(w11,l);
                update_max_word(w12,l); update_max_word(w13,l); update_max_word(w14,l); update_max_word(w15,l);
                float sf; compute_scale(l, asc, sf);
                p0=pack4_words(w0,w1,w2,w3,sf);
                p1=pack4_words(w4,w5,w6,w7,sf);
                p2=pack4_words(w8,w9,w10,w11,sf);
                p3=pack4_words(w12,w13,w14,w15,sf);
            }
            int4 wr; wr.x=(int)p0; wr.y=(int)p1; wr.z=(int)p2; wr.w=(int)p3;
            *reinterpret_cast<int4*>(&Alds[lds_swz(qr*LDS_ROW+qg*16)]) = wr;
            Asclds[qr*8+qg] = asc;
        }

        if (tid >= 256) {
            const int btid = tid - 256;
            constexpr int BNT = 256;
            constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
            for (int ld = 0; ld < BL; ld++) {
                const int f = (ld * BNT + btid) << 4;
                const int r = f >> 7; const int gn = bn + r;
                if (r < BLOCK_N && gn < N) {
                    const int sf = lds_swz(f), sc = sf & 127, ac = kb + sc;
                    llvm_amdgcn_raw_buffer_load_lds(srsrc,(as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds)+f),16,(gn>>4)*(bstride<<4)+(ac>>5)*512+((ac>>4)&1)*256+(gn&15)*16,0,0,0);
                }
            }
        }

        {const int so=ks<<3;
        if (tid >= 256 && tid < 384) {
            const int stid = tid - 256;
            const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
            const int gn = bn + row;
            if (gn < N) { const int rb=(gn>>5)*32*scs+(gn&15)*4+((gn>>4)&1);
                for(int g=0;g<4;g++){const int grp=gb+g,ac=so+grp;Bslds[row*8+grp]=Bsc[rb+(ac&3)*64+((ac&7)>>2)*2+(ac>>3)*256];}
            } else { for(int g=0;g<4;g++)Bslds[row*8+gb+g]=0x7f; }
        }}

        asm volatile("s_waitcnt vmcnt(0)");__syncthreads();

        const int ar=(wg<<4)+lm;
        int4_vec A0,A1;int as0,as1;
        {const int o=lds_swz(ar*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A0={t.x,t.y,t.z,t.w};}
        as0=(int)Asclds[ar*8+lk];
        {const int o=lds_swz(ar*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A1={t.x,t.y,t.z,t.w};}
        as1=(int)Asclds[ar*8+4+lk];
        const int br=(wc<<4)+lm;
        int4_vec B0,B1;int bs0,bs1;
        {const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B0={t.x,t.y,t.z,t.w};}
        bs0=(int)Bslds[br*8+lk];
        {const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B1={t.x,t.y,t.z,t.w};}
        bs1=(int)Bslds[br*8+4+lk];

        acc=mfma_fp4(A0,B0,acc,as0,bs0);
        acc=mfma_fp4(A1,B1,acc,as1,bs1);
    }

    const int or_=bm+(wg<<4)+(lk<<2),oc=col+lm;
    if(oc<N){const float*ap=reinterpret_cast<const float*>(&acc);
    if(C){for(int r=0;r<4;r++){int g=or_+r;if(g<M)C[g*N+oc]=__float2bfloat16(ap[r]);}}
    else{float*w=ws+sid*M*N;for(int r=0;r<4;r++){int g=or_+r;if(g<M)w[g*N+oc]=ap[r];}}}
}

template<int KSPS>
__global__ void gemm_bn64_m32_k(const __hip_bfloat16*__restrict__ A,const uint8_t*__restrict__ Bq,const uint8_t*__restrict__ Bsc,float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N,int K){
    constexpr int BLOCK_N=64;
    const int wid=threadIdx.x>>6,lm=threadIdx.x&15,lk=(threadIdx.x&63)>>4,tid=threadIdx.x;
    const int wg=wid>>2,wc=wid&3;
    const int bn=blockIdx.x*BLOCK_N,bm=blockIdx.y*BLOCK_M_LARGE,col=bn+(wc<<4),sid=blockIdx.z;const int bstride=K>>1,scs=K>>5;
    __shared__ __align__(16) uint8_t Alds[BLOCK_M_LARGE*LDS_ROW];__shared__ uint8_t Asclds[BLOCK_M_LARGE*8];
    __shared__ __align__(16) uint8_t Blds[BLOCK_N*LDS_ROW];__shared__ uint8_t Bslds[BLOCK_N*8];
    const i32x4 srsrc=make_srsrc(Bq,N*bstride);float4_vec acc={0,0,0,0};
    #pragma unroll
    for(int kstep=0;kstep<KSPS;kstep++){const int ks=sid*KSPS+kstep,ke=ks*DOUBLE_K,kb=ks*LDS_ROW;
        if (tid < 256) {
            const int qr = tid >> 3, qg = tid & 7;
            const int gr = bm + qr, ko = ke + qg * 32;
            uint32_t p0=0, p1=0, p2=0, p3=0; uint8_t asc = 0x7f;
            if (gr < M) {
                const __hip_bfloat16* s = A + gr * K + ko;
                const int4 r0 = reinterpret_cast<const int4*>(s)[0];
                const int4 r1 = reinterpret_cast<const int4*>(s)[1];
                const int4 r2 = reinterpret_cast<const int4*>(s)[2];
                const int4 r3 = reinterpret_cast<const int4*>(s)[3];
                const uint32_t w0 = (uint32_t)r0.x, w1 = (uint32_t)r0.y, w2 = (uint32_t)r0.z, w3 = (uint32_t)r0.w;
                const uint32_t w4 = (uint32_t)r1.x, w5 = (uint32_t)r1.y, w6 = (uint32_t)r1.z, w7 = (uint32_t)r1.w;
                const uint32_t w8 = (uint32_t)r2.x, w9 = (uint32_t)r2.y, w10 = (uint32_t)r2.z, w11 = (uint32_t)r2.w;
                const uint32_t w12 = (uint32_t)r3.x, w13 = (uint32_t)r3.y, w14 = (uint32_t)r3.z, w15 = (uint32_t)r3.w;
                float l = 0;
                update_max_word(w0,l); update_max_word(w1,l); update_max_word(w2,l); update_max_word(w3,l);
                update_max_word(w4,l); update_max_word(w5,l); update_max_word(w6,l); update_max_word(w7,l);
                update_max_word(w8,l); update_max_word(w9,l); update_max_word(w10,l); update_max_word(w11,l);
                update_max_word(w12,l); update_max_word(w13,l); update_max_word(w14,l); update_max_word(w15,l);
                float sf; compute_scale(l, asc, sf);
                p0=pack4_words(w0,w1,w2,w3,sf);
                p1=pack4_words(w4,w5,w6,w7,sf);
                p2=pack4_words(w8,w9,w10,w11,sf);
                p3=pack4_words(w12,w13,w14,w15,sf);
            }
            int4 wr; wr.x=(int)p0; wr.y=(int)p1; wr.z=(int)p2; wr.w=(int)p3;
            *reinterpret_cast<int4*>(&Alds[lds_swz(qr*LDS_ROW+qg*16)]) = wr;
            Asclds[qr*8+qg] = asc;
        }

        if (tid >= 256) {
            const int btid = tid - 256;
            constexpr int BNT = 256;
            constexpr int BL = (BLOCK_N * LDS_ROW / 16 + BNT - 1) / BNT;
            for (int ld = 0; ld < BL; ld++) {
                const int f = (ld * BNT + btid) << 4;
                const int r = f >> 7; const int gn = bn + r;
                if (r < BLOCK_N && gn < N) {
                    const int sf = lds_swz(f), sc = sf & 127, ac = kb + sc;
                    llvm_amdgcn_raw_buffer_load_lds(srsrc,(as3_uint32_ptr)(reinterpret_cast<uintptr_t>(Blds)+f),16,(gn>>4)*(bstride<<4)+(ac>>5)*512+((ac>>4)&1)*256+(gn&15)*16,0,0,0);
                }
            }
        }

        {const int so=ks<<3;
        if (tid >= 256 && tid < 384) {
            const int stid = tid - 256;
            const int row = stid & (BLOCK_N - 1), gb = (stid >> 6) << 2;
            const int gn = bn + row;
            if (gn < N) { const int rb=(gn>>5)*32*scs+(gn&15)*4+((gn>>4)&1);
                for(int g=0;g<4;g++){const int grp=gb+g,ac=so+grp;Bslds[row*8+grp]=Bsc[rb+(ac&3)*64+((ac&7)>>2)*2+(ac>>3)*256];}
            } else { for(int g=0;g<4;g++)Bslds[row*8+gb+g]=0x7f; }
        }}

        asm volatile("s_waitcnt vmcnt(0)");__syncthreads();

        const int ar=(wg<<4)+lm;
        int4_vec A0,A1;int as0,as1;
        {const int o=lds_swz(ar*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A0={t.x,t.y,t.z,t.w};}
        as0=(int)Asclds[ar*8+lk];
        {const int o=lds_swz(ar*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Alds[o]);A1={t.x,t.y,t.z,t.w};}
        as1=(int)Asclds[ar*8+4+lk];
        const int br=(wc<<4)+lm;
        int4_vec B0,B1;int bs0,bs1;
        {const int o=lds_swz(br*LDS_ROW+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B0={t.x,t.y,t.z,t.w};}
        bs0=(int)Bslds[br*8+lk];
        {const int o=lds_swz(br*LDS_ROW+HALF_K+(lk<<4));const int4 t=*reinterpret_cast<const int4*>(&Blds[o]);B1={t.x,t.y,t.z,t.w};}
        bs1=(int)Bslds[br*8+4+lk];

        acc=mfma_fp4(A0,B0,acc,as0,bs0);
        acc=mfma_fp4(A1,B1,acc,as1,bs1);
    }

    const int or_=bm+(wg<<4)+(lk<<2),oc=col+lm;
    if(oc<N){const float*ap=reinterpret_cast<const float*>(&acc);
    if(C){for(int r=0;r<4;r++){int g=or_+r;if(g<M)C[g*N+oc]=__float2bfloat16(ap[r]);}}
    else{float*w=ws+sid*M*N;for(int r=0;r<4;r++){int g=or_+r;if(g<M)w[g*N+oc]=ap[r];}}}
}

template<int SK>
__global__ void reduce_k_t(const float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N){
    const int i=blockIdx.x*blockDim.x+threadIdx.x;if(i>=M*N)return;
    const int stride=M*N;
    float vals[SK];
    #pragma unroll
    for(int k=0;k<SK;k++) vals[k]=ws[k*stride+i];
    float s=0;
    #pragma unroll
    for(int k=0;k<SK;k++) s+=vals[k];
    C[i]=__float2bfloat16(s);
}

__global__ void reduce_k_gen(const float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N,int sk){
    const int i=blockIdx.x*blockDim.x+threadIdx.x;if(i>=M*N)return;
    const int stride=M*N;const float*p=ws+i;
    float s0=0,s1=0,s2=0,s3=0;
    int k=0;
    for(;k+3<sk;k+=4){s0+=p[k*stride];s1+=p[(k+1)*stride];s2+=p[(k+2)*stride];s3+=p[(k+3)*stride];}
    for(;k<sk;k++)s0+=p[k*stride];
    C[i]=__float2bfloat16(s0+s1+s2+s3);}

__global__ void reduce_k_14(const float*__restrict__ ws,__hip_bfloat16*__restrict__ C,int M,int N){
    const int i=blockIdx.x*blockDim.x+threadIdx.x;if(i>=M*N)return;
    const int stride=M*N;const float*p=ws+i;
    float s0=p[0];
    s0+=p[4*stride];
    s0+=p[8*stride];
    s0+=p[12*stride];
    float s1=p[1*stride];
    s1+=p[5*stride];
    s1+=p[9*stride];
    s1+=p[13*stride];
    float s2=p[2*stride];
    s2+=p[6*stride];
    s2+=p[10*stride];
    float s3=p[3*stride];
    s3+=p[7*stride];
    s3+=p[11*stride];
    C[i]=__float2bfloat16(s0+s1+s2+s3);
}

torch::Tensor dispatch_gemm(torch::Tensor A, torch::Tensor Bq, torch::Tensor Bsc, int M, int N, int K) {
    auto C = torch::empty({M, N}, torch::dtype(torch::kBFloat16).device(torch::kCUDA));
    const auto* a_ptr = reinterpret_cast<const __hip_bfloat16*>(A.data_ptr());
    const auto* bq_ptr = reinterpret_cast<const uint8_t*>(Bq.data_ptr());
    const auto* bsc_ptr = reinterpret_cast<const uint8_t*>(Bsc.data_ptr());
    auto* c_ptr = reinterpret_cast<__hip_bfloat16*>(C.data_ptr());

    int ks = K / 256;
    const int bn = 64;
    const bool large_m = M >= 128;
    const int block_m = large_m ? BLOCK_M_LARGE : BLOCK_M;
    int bmn = ((N + bn - 1) / bn) * ((M + block_m - 1) / block_m);

    int sk = 1;
    if (!large_m && bmn < 304) {
        if (ks > 8) {
            int target = (462 + bmn - 1) / bmn;
            int best = 1;
            for (int s = 1; s <= ks; s++) {
                if (ks % s == 0 && s <= target) best = s;
            }
            sk = best;
        } else {
            int target = (304 + bmn - 1) / bmn;
            if (target < 1) target = 1;
            int best = 1;
            for (int s = 1; s <= ks; s++) {
                if (ks % s == 0 && s <= target) best = s;
            }
            while (best > 1 && ks / best < 2) best /= 2;
            sk = best > 0 ? best : 1;
        }
    }

    if (sk == 1) {
        if (large_m) {
            dim3 b(512), g((N + bn - 1) / bn, (M + BLOCK_M_LARGE - 1) / BLOCK_M_LARGE, 1);
            // Use double-buffered kernel for ks >= 4 (significant K-loop to overlap)
            switch(ks) {
                case 6: hipLaunchKernelGGL(gemm_bn64_m32_db<6>, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, (float*)nullptr, c_ptr, M, N, K); break;
                default: hipLaunchKernelGGL(gemm_bn64_m32, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, (float*)nullptr, c_ptr, M, N, K, ks); break;
            }
        } else {
            dim3 b(256), g((N + bn - 1) / bn, (M + BLOCK_M - 1) / BLOCK_M, 1);
            // Use double-buffered kernel for ks >= 4
            switch(ks) {
                case 2: hipLaunchKernelGGL(gemm_bn64_k<2>, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, (float*)nullptr, c_ptr, M, N, K); break;
                case 4: hipLaunchKernelGGL(gemm_bn64_db<4>, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, (float*)nullptr, c_ptr, M, N, K); break;
                case 6: hipLaunchKernelGGL(gemm_bn64_db<6>, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, (float*)nullptr, c_ptr, M, N, K); break;
                case 8: hipLaunchKernelGGL(gemm_bn64_db<8>, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, (float*)nullptr, c_ptr, M, N, K); break;
                default: hipLaunchKernelGGL(gemm_bn64, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, (float*)nullptr, c_ptr, M, N, K, ks); break;
            }
        }
    } else {
        auto ws = torch::empty({sk, M, N}, torch::dtype(torch::kFloat32).device(torch::kCUDA));
        auto* ws_ptr = reinterpret_cast<float*>(ws.data_ptr());
        int ksps = ks / sk;
        if (large_m) {
            dim3 b(512), g((N + bn - 1) / bn, (M + BLOCK_M_LARGE - 1) / BLOCK_M_LARGE, sk);
            if (ksps == 6) {
                hipLaunchKernelGGL(gemm_bn64_m32_k<6>, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, ws_ptr, (__hip_bfloat16*)nullptr, M, N, K);
            } else {
                hipLaunchKernelGGL(gemm_bn64_m32, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, ws_ptr, (__hip_bfloat16*)nullptr, M, N, K, ksps);
            }
        } else {
            dim3 b(256), g((N + bn - 1) / bn, (M + BLOCK_M - 1) / BLOCK_M, sk);
            if (ksps == 2) {
                hipLaunchKernelGGL(gemm_bn64_k<2>, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, ws_ptr, (__hip_bfloat16*)nullptr, M, N, K);
            } else {
                hipLaunchKernelGGL(gemm_bn64, g, b, 0, 0, a_ptr, bq_ptr, bsc_ptr, ws_ptr, (__hip_bfloat16*)nullptr, M, N, K, ksps);
            }
        }
        int n = M * N;
        dim3 rb(256), rg((n+255)/256);
        switch(sk) {
            case 2: hipLaunchKernelGGL(reduce_k_t<2>, rg, rb, 0, 0, ws_ptr, c_ptr, M, N); break;
            case 4: hipLaunchKernelGGL(reduce_k_t<4>, rg, rb, 0, 0, ws_ptr, c_ptr, M, N); break;
            case 7: hipLaunchKernelGGL(reduce_k_t<7>, rg, rb, 0, 0, ws_ptr, c_ptr, M, N); break;
            case 14: hipLaunchKernelGGL(reduce_k_14, rg, rb, 0, 0, ws_ptr, c_ptr, M, N); break;
            default: hipLaunchKernelGGL(reduce_k_gen, rg, rb, 0, 0, ws_ptr, c_ptr, M, N, sk); break;
        }
    }
    return C;
}
"""


_module = None


def _get_module():
    global _module
    if _module is None:
        _module = load_inline(
            name="gemm_v3",
            cpp_sources=CPP_SOURCE,
            cuda_sources=HIP_SOURCE,
            functions=["dispatch_gemm"],
            verbose=False,
            extra_cuda_cflags=["-O3", "-fno-gpu-rdc", "-ffp-contract=fast"],
        )
    return _module


@torch.inference_mode()
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]
    return _get_module().dispatch_gemm(A, B_shuffle.view(torch.uint8), B_scale_sh.view(torch.uint8), m, n, k)
scrolls · 946 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