Skip to content
KernelIndex
Search⌘K

submission 740115

LunNova · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c00ca907b60877ff80e87721a29bdbf7c916c8e465cf22214c9008197d1b16a0
license declaredunknown
license concludedunknown
authorsLunNova
imported2026-08-15

Techniques

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

shared-memory__shared__ char lds_bytes[3136];
split-kconstexpr int kSplitK=7, kKPerSplit=1024, kNumWaves=4, kItersPerWave=2;
tile-n = 64constexpr int kN=2112, kK=7168, kKHalf=3584, kTileN=64, kMfmaCols=4;

Kernel source

submission_standalone_best_v3b.py913 lines
"""Standalone best-of GEMM: 6 shapes, load_inline C++ dispatch."""
from __future__ import annotations
import os, sys, tempfile
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t


# ── M=64 hand-unrolled iteration generator ──────────────────────
def _gen_m64_iters():
    table = [
        ("a_buf0", "a_buf2", True,  True),
        ("a_buf1", "a_buf0", True,  True),
        ("a_buf2", "a_buf1", True,  True),
        ("a_buf0", "a_buf2", True,  True),
        ("a_buf1", "a_buf0", True,  True),
        ("a_buf2", "a_buf1", True,  True),
        ("a_buf0", "a_buf2", False, True),
        ("a_buf1", "a_buf0", False, False),
    ]
    lines = []
    for k, (consume, far, load_a, load_b) in enumerate(table):
        lines.append(f"    {{ // iter {k}")
        lines.append(f"        const int kg_ = (ki_base + {k}) * 2 + half;")
        lines.append( "        int bsc0_ = (int)lds_bscale[g::soff<64>(b_col0, kg_) - bsc_base];")
        lines.append( "        int bsc1_ = (int)lds_bscale[g::soff<64>(b_col1, kg_) - bsc_base];")
        lines.append(f"        g::v4i32 ar_; int as_; g::quant_a({consume}, ar_, as_);")
        if load_b:
            lines.append(f"        {{ const uint32_t* bn0_ = reinterpret_cast<const uint32_t*>(bb0 + (ki_base+{k+1})*512 + half*256);")
            lines.append( "          b_nxt0 = {(int)bn0_[0],(int)bn0_[1],(int)bn0_[2],(int)bn0_[3]};")
            lines.append(f"          const uint32_t* bn1_ = reinterpret_cast<const uint32_t*>(bb1 + (ki_base+{k+1})*512 + half*256);")
            lines.append( "          b_nxt1 = {(int)bn1_[0],(int)bn1_[1],(int)bn1_[2],(int)bn1_[3]}; }")
        lines.append( "        g::mfma_pair(ar_, b_cur0, b_cur1, acc0, acc1, as_, bsc0_, bsc1_);")
        if load_a:
            lines.append(f"        g::load_a(a_row + ((ki_base+{k+2})*2 + half)*16, {far});")
        if load_b:
            lines.append( "        b_cur0 = b_nxt0; b_cur1 = b_nxt1;")
        lines.append( "    }")
    return "\n".join(lines)


# ═════════════════════════════════════════════════════════════════
# HIP source sections
# ═════════════════════════════════════════════════════════════════

_HIP_HEADER = r"""
#include <torch/extension.h>
#include <torch/csrc/autograd/python_variable.h>
#include <pybind11/pybind11.h>
namespace py = pybind11;
#include <hip/hip_runtime.h>
#include <cstdint>
#include <cstdio>
#include <string>

#define CVT_PK_FP4_BF16_B0(dst, src, scale) \
    dst = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(dst, src, scale, 0)
#define CVT_PK_FP4_BF16_B1(dst, src, scale) \
    dst = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(dst, src, scale, 1)
#define CVT_PK_FP4_BF16_B2(dst, src, scale) \
    dst = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(dst, src, scale, 2)
#define CVT_PK_FP4_BF16_B3(dst, src, scale) \
    dst = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(dst, src, scale, 3)

// ═══════════════════════════════════════════════════════
// Shared utilities
// ═══════════════════════════════════════════════════════
namespace g {
using v4i32  = int   __attribute__((ext_vector_type(4)));
using v4f32  = float __attribute__((ext_vector_type(4)));
using v16f32 = float __attribute__((ext_vector_type(16)));
typedef __bf16 v2bf16 __attribute__((ext_vector_type(2)));

__device__ __forceinline__ uint32_t bcu32(float v) { union{float f;uint32_t u;}x; x.f=v; return x.u; }
__device__ __forceinline__ float bcf32(uint32_t v) { union{uint32_t u;float f;}x; x.u=v; return x.f; }
__device__ __forceinline__ uint16_t f2bf16(float v) {
    uint32_t b=bcu32(v); b+=((b>>16)&1u)+0x7FFFu; return (uint16_t)(b>>16);
}
__device__ __forceinline__ uint8_t e8m0sc(uint16_t m) {
    uint16_t r=(uint16_t)(m+0x20u); int e=(int)((r>>7)&0xFFu);
    return (e<=2)?0u:(e>=255)?254u:(uint8_t)(e-2);
}
template<int PG>
__device__ __forceinline__ int soff(int row, int group) {
    return (row>>5)*(32*PG)+(group>>3)*256+(group&3)*64+(row&15)*4+((group>>2)&1)*2+((row>>4)&1);
}
__device__ __forceinline__ void load_a(const uint32_t* __restrict__ p, uint32_t* d) {
    *reinterpret_cast<uint4*>(&d[0])  = *(reinterpret_cast<const uint4*>(p)+0);
    *reinterpret_cast<uint4*>(&d[4])  = *(reinterpret_cast<const uint4*>(p)+1);
    *reinterpret_cast<uint4*>(&d[8])  = *(reinterpret_cast<const uint4*>(p)+2);
    *reinterpret_cast<uint4*>(&d[12]) = *(reinterpret_cast<const uint4*>(p)+3);
}
__device__ __forceinline__ void quant_a(const uint32_t* dw, v4i32& out, int& sc) {
    uint32_t mx=0u;
    #pragma unroll
    for(int i=0;i<16;++i){uint32_t w=dw[i],hi=(w>>16)&0x7FFFu,lo=w&0x7FFFu;mx=(hi>mx)?hi:mx;mx=(lo>mx)?lo:mx;}
    uint8_t s=e8m0sc((uint16_t)mx); sc=(int)s;
    float fwd=(s==0u)?bcf32(0x00400000u):bcf32((uint32_t)s<<23);
    const v2bf16*p=reinterpret_cast<const v2bf16*>(dw);
    uint32_t pk[4]={0,0,0,0};
    CVT_PK_FP4_BF16_B0(pk[0],p[0],fwd); CVT_PK_FP4_BF16_B0(pk[1],p[4],fwd);
    CVT_PK_FP4_BF16_B0(pk[2],p[8],fwd); CVT_PK_FP4_BF16_B0(pk[3],p[12],fwd);
    CVT_PK_FP4_BF16_B1(pk[0],p[1],fwd); CVT_PK_FP4_BF16_B1(pk[1],p[5],fwd);
    CVT_PK_FP4_BF16_B1(pk[2],p[9],fwd); CVT_PK_FP4_BF16_B1(pk[3],p[13],fwd);
    CVT_PK_FP4_BF16_B2(pk[0],p[2],fwd); CVT_PK_FP4_BF16_B2(pk[1],p[6],fwd);
    CVT_PK_FP4_BF16_B2(pk[2],p[10],fwd); CVT_PK_FP4_BF16_B2(pk[3],p[14],fwd);
    CVT_PK_FP4_BF16_B3(pk[0],p[3],fwd); CVT_PK_FP4_BF16_B3(pk[1],p[7],fwd);
    CVT_PK_FP4_BF16_B3(pk[2],p[11],fwd); CVT_PK_FP4_BF16_B3(pk[3],p[15],fwd);
    out = {(int)pk[0],(int)pk[1],(int)pk[2],(int)pk[3]};
}
// v_pk_max_u16 tree reduction — fewer instructions than scalar loop
__device__ __forceinline__ void quant_a_pk(const uint32_t* dw, v4i32& out, int& sc) {
    const uint32_t mask = 0x7FFF7FFFu;
    uint32_t p0=dw[0]&mask,p1=dw[1]&mask,p2=dw[2]&mask,p3=dw[3]&mask;
    uint32_t p4=dw[4]&mask,p5=dw[5]&mask,p6=dw[6]&mask,p7=dw[7]&mask;
    uint32_t p8=dw[8]&mask,p9=dw[9]&mask,pA=dw[10]&mask,pB=dw[11]&mask;
    uint32_t pC=dw[12]&mask,pD=dw[13]&mask,pE=dw[14]&mask,pF=dw[15]&mask;
    uint32_t m01,m23,m45,m67,m89,mAB,mCD,mEF;
    asm("v_pk_max_u16 %0,%1,%2":"=v"(m01):"v"(p0),"v"(p1));
    asm("v_pk_max_u16 %0,%1,%2":"=v"(m23):"v"(p2),"v"(p3));
    asm("v_pk_max_u16 %0,%1,%2":"=v"(m45):"v"(p4),"v"(p5));
    asm("v_pk_max_u16 %0,%1,%2":"=v"(m67):"v"(p6),"v"(p7));
    asm("v_pk_max_u16 %0,%1,%2":"=v"(m89):"v"(p8),"v"(p9));
    asm("v_pk_max_u16 %0,%1,%2":"=v"(mAB):"v"(pA),"v"(pB));
    asm("v_pk_max_u16 %0,%1,%2":"=v"(mCD):"v"(pC),"v"(pD));
    asm("v_pk_max_u16 %0,%1,%2":"=v"(mEF):"v"(pE),"v"(pF));
    uint32_t t0,t1,t2,t3;
    asm("v_pk_max_u16 %0,%1,%2":"=v"(t0):"v"(m01),"v"(m23));
    asm("v_pk_max_u16 %0,%1,%2":"=v"(t1):"v"(m45),"v"(m67));
    asm("v_pk_max_u16 %0,%1,%2":"=v"(t2):"v"(m89),"v"(mAB));
    asm("v_pk_max_u16 %0,%1,%2":"=v"(t3):"v"(mCD),"v"(mEF));
    uint32_t u0,u1;
    asm("v_pk_max_u16 %0,%1,%2":"=v"(u0):"v"(t0),"v"(t1));
    asm("v_pk_max_u16 %0,%1,%2":"=v"(u1):"v"(t2),"v"(t3));
    uint32_t fpk;
    asm("v_pk_max_u16 %0,%1,%2":"=v"(fpk):"v"(u0),"v"(u1));
    uint32_t mx = (fpk>>16) > (fpk&0xFFFFu) ? (fpk>>16) : (fpk&0xFFFFu);

    uint8_t s=e8m0sc((uint16_t)mx); sc=(int)s;
    float fwd=(s==0u)?bcf32(0x00400000u):bcf32((uint32_t)s<<23);
    const v2bf16*p=reinterpret_cast<const v2bf16*>(dw);
    uint32_t pk[4]={0,0,0,0};
    CVT_PK_FP4_BF16_B0(pk[0],p[0],fwd); CVT_PK_FP4_BF16_B0(pk[1],p[4],fwd);
    CVT_PK_FP4_BF16_B0(pk[2],p[8],fwd); CVT_PK_FP4_BF16_B0(pk[3],p[12],fwd);
    CVT_PK_FP4_BF16_B1(pk[0],p[1],fwd); CVT_PK_FP4_BF16_B1(pk[1],p[5],fwd);
    CVT_PK_FP4_BF16_B1(pk[2],p[9],fwd); CVT_PK_FP4_BF16_B1(pk[3],p[13],fwd);
    CVT_PK_FP4_BF16_B2(pk[0],p[2],fwd); CVT_PK_FP4_BF16_B2(pk[1],p[6],fwd);
    CVT_PK_FP4_BF16_B2(pk[2],p[10],fwd); CVT_PK_FP4_BF16_B2(pk[3],p[14],fwd);
    CVT_PK_FP4_BF16_B3(pk[0],p[3],fwd); CVT_PK_FP4_BF16_B3(pk[1],p[7],fwd);
    CVT_PK_FP4_BF16_B3(pk[2],p[11],fwd); CVT_PK_FP4_BF16_B3(pk[3],p[15],fwd);
    out = {(int)pk[0],(int)pk[1],(int)pk[2],(int)pk[3]};
}
__device__ __forceinline__ v4f32 mfma16v(v4i32 a, v4i32 b, v4f32 c, int as, int bs) {
    asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
        :"+v"(c):"v"(a),"v"(b),"v"(as),"v"(bs)); return c;
}
__device__ __forceinline__ v4f32 mfma16a(v4i32 a, v4i32 b, v4f32 c, int as, int bs) {
    asm volatile("v_mfma_scale_f32_16x16x128_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
        :"+a"(c):"v"(a),"v"(b),"v"(as),"v"(bs)); return c;
}
__device__ __forceinline__ void mfma_pair(v4i32 a, v4i32 b0, v4i32 b1,
    v16f32& c0, v16f32& c1, int as, int bs0, int bs1) {
    __builtin_amdgcn_sched_barrier(0); __builtin_amdgcn_s_setprio(1);
    asm volatile(
        "v_mfma_scale_f32_32x32x64_f8f6f4 %0,%2,%3,%0,%4,%5 cbsz:4 blgp:4\n"
        "v_mfma_scale_f32_32x32x64_f8f6f4 %1,%2,%6,%1,%4,%7 cbsz:4 blgp:4\n"
        :"+a"(c0),"+a"(c1):"v"(a),"v"(b0),"v"(as),"v"(bs0),"v"(b1),"v"(bs1));
    __builtin_amdgcn_s_setprio(0);
}
} // namespace g
"""

# ═══════════════════════════════════════════════════════
# K1: M=4, N=2880, K=512  (4-wave cooperative A-quant, 32x32x64)
# grid=90, block=256, LDS=3136
# From sub_f1b3a2ab493a/4x2880x512.hip
# ═══════════════════════════════════════════════════════
_HIP_M4 = r"""
namespace fused_m4_common {

using v4i32 = int __attribute__((ext_vector_type(4)));
using v16f32 = float __attribute__((ext_vector_type(16)));

__device__ __forceinline__ uint32_t m4_bcu32(float v) { union{float f;uint32_t u;}x; x.f=v; return x.u; }
__device__ __forceinline__ float m4_bcf32(uint32_t v) { union{uint32_t u;float f;}x; x.u=v; return x.f; }
__device__ __forceinline__ uint16_t m4_f2bf16(float v) {
    uint32_t bits=m4_bcu32(v); bits+=((bits>>16)&1u)+0x7FFFu; return (uint16_t)(bits>>16);
}
__device__ __forceinline__ uint8_t m4_e8m0sc(uint16_t max_abs) {
    uint16_t r=(uint16_t)(max_abs+0x20u); int e=(int)((r>>7)&0xFFu);
    return (e<=2)?0u:(e>=255)?254u:(uint8_t)(e-2);
}
__device__ __forceinline__ int m4_soff(int row, int group) {
    return (row>>5)*(32*16)+(group>>3)*256+(group&3)*64+(row&15)*4+((group>>2)&1)*2+((row>>4)&1);
}
__device__ __forceinline__ float m4_e8m0_inv(uint8_t biased) {
    if (biased==0u) return m4_bcf32(0x7F000000u);
    return m4_bcf32((uint32_t)(254u-biased)<<23);
}
__device__ __forceinline__ uint8_t m4_pack_fp4(float scaled) {
    uint32_t bits=m4_bcu32(scaled);
    uint32_t sign=bits&0x80000000u; bits^=sign;
    float mag=m4_bcf32(bits);
    uint8_t encoded=0x7u;
    if (mag<6.0f) {
        if (mag<1.0f) {
            constexpr uint32_t kDM=(uint32_t)(149)<<23;
            uint32_t d=m4_bcu32(mag+m4_bcf32(kDM)); d-=kDM;
            encoded=(uint8_t)d;
        } else {
            uint32_t normal=bits;
            uint32_t mo=(normal>>22)&0x1u;
            normal=(uint32_t)((int32_t)normal+(-126*(1<<23)+(1<<21)-1));
            normal+=mo; normal>>=22;
            encoded=(uint8_t)normal;
        }
    }
    return (uint8_t)((encoded|(uint8_t)(sign>>28))&0x0Fu);
}
__device__ __forceinline__ v16f32 m4_mfma(v4i32 a, v4i32 b, v16f32 acc, int as, int bs) {
    asm volatile("v_mfma_scale_f32_32x32x64_f8f6f4 %0,%1,%2,%0,%3,%4 cbsz:4 blgp:4"
        :"+v"(acc):"v"(a),"v"(b),"v"(as),"v"(bs)); return acc;
}

} // namespace fused_m4_common

extern "C" __global__ __launch_bounds__(256)
void fused_m4_m4_n2880_k512(
    const uint16_t* __restrict__ A,
    const uint8_t*  __restrict__ B_shuffle,
    const uint8_t*  __restrict__ B_scale_sh,
    uint16_t*       __restrict__ C)
{
    using namespace fused_m4_common;
    constexpr int M=4, N=2880, K=512, K_HALF=256, K_tiles=8, padded_groups=16;

    const int col_base = blockIdx.x * 32;
    if (col_base >= N) return;

    const int tid=threadIdx.x;
    const int wave_id=tid>>6;
    const int tid_in_wave=tid&63;
    const int lane=tid&31;
    const int half=(tid>>5)&1;

    __shared__ char lds_bytes[3136];
    uint8_t* lds_a_fp4   = reinterpret_cast<uint8_t*>(lds_bytes+0);
    uint8_t* lds_a_scale  = reinterpret_cast<uint8_t*>(lds_bytes+1024);
    float*   lds_reduce   = reinterpret_cast<float*>(lds_bytes+1088);

    const int b_col = col_base + lane;
    const bool b_valid = (b_col < N);
    const int n_in_tile = lane & 15;
    const int n_tile = (col_base>>4) + (lane>>4);

    const int ki0 = wave_id*2;
    const int ki1 = ki0+1;

    v4i32 br0={0,0,0,0}, br1={0,0,0,0};
    int bs0=127, bs1=127;
    if (b_valid) {
        const int b_flat0 = ((n_tile*K_tiles+ki0)*2+half)*256 + n_in_tile*16;
        const uint32_t* s0 = reinterpret_cast<const uint32_t*>(B_shuffle+b_flat0);
        br0 = {(int)s0[0],(int)s0[1],(int)s0[2],(int)s0[3]};
        bs0 = (int)B_scale_sh[m4_soff(b_col, ki0*2+half)];

        const int b_flat1 = ((n_tile*K_tiles+ki1)*2+half)*256 + n_in_tile*16;
        const uint32_t* s1 = reinterpret_cast<const uint32_t*>(B_shuffle+b_flat1);
        br1 = {(int)s1[0],(int)s1[1],(int)s1[2],(int)s1[3]};
        bs1 = (int)B_scale_sh[m4_soff(b_col, ki1*2+half)];
    }

    // Cooperative A quantization into LDS
    {
        const int row = wave_id;
        const int group = tid_in_wave>>2;
        const int tib = tid_in_wave&3;
        const uint16_t* ap = A + row*K + group*32 + tib*8;

        uint16_t ar[8]; uint16_t mx=0u;
        #pragma unroll
        for (int j=0;j<8;++j) { ar[j]=ap[j]; uint16_t ab=ar[j]&0x7FFFu; mx=(ab>mx)?ab:mx; }
        mx=(uint16_t)max((int)mx, __shfl_xor((int)mx,1,64));
        mx=(uint16_t)max((int)mx, __shfl_xor((int)mx,2,64));

        uint8_t sb = m4_e8m0sc(mx);
        float is = m4_e8m0_inv(sb);

        uint32_t pd=0u;
        #pragma unroll
        for (int j=0;j<4;++j) {
            uint8_t lo = m4_pack_fp4(m4_bcf32((uint32_t)(ar[j*2])<<16)*is);
            uint8_t hi = m4_pack_fp4(m4_bcf32((uint32_t)(ar[j*2+1])<<16)*is);
            pd |= ((uint32_t)((lo&0xFu)|(hi<<4))<<(j*8));
        }
        *reinterpret_cast<uint32_t*>(lds_a_fp4 + row*K_HALF + group*16 + tib*4) = pd;
        if (tib==0) lds_a_scale[row*16+group] = sb;
    }
    __syncthreads();

    const bool real = (lane < M);
    v16f32 acc;
    #pragma unroll
    for (int i=0;i<16;++i) acc[i]=0.0f;

    { // MFMA iteration 0
        v4i32 ar={0,0,0,0}; int as=127;
        if (real) {
            const uint32_t* s=reinterpret_cast<const uint32_t*>(lds_a_fp4+lane*K_HALF+ki0*32+half*16);
            ar={(int)s[0],(int)s[1],(int)s[2],(int)s[3]};
            as=(int)lds_a_scale[lane*16+ki0*2+half];
        }
        acc = m4_mfma(ar, br0, acc, as, bs0);
    }
    { // MFMA iteration 1
        v4i32 ar={0,0,0,0}; int as=127;
        if (real) {
            const uint32_t* s=reinterpret_cast<const uint32_t*>(lds_a_fp4+lane*K_HALF+ki1*32+half*16);
            ar={(int)s[0],(int)s[1],(int)s[2],(int)s[3]};
            as=(int)lds_a_scale[lane*16+ki1*2+half];
        }
        acc = m4_mfma(ar, br1, acc, as, bs1);
    }

    // Reduction
    if (half==0) {
        #pragma unroll
        for (int r=0;r<M;++r) lds_reduce[wave_id*32*M+lane*M+r] = acc[r];
    }
    __syncthreads();

    if (wave_id==0 && half==0 && (col_base+lane)<N) {
        #pragma unroll
        for (int r=0;r<M;++r) {
            float sum=0.0f;
            #pragma unroll
            for (int w=0;w<4;++w) sum += lds_reduce[w*32*M+lane*M+r];
            C[r*N+col_base+lane] = m4_f2bf16(sum);
        }
    }
}
"""

# ═══════════════════════════════════════════════════════
# K2: M=16, N=2112, K=7168  (two-pass split-K=7, workspace + reduce)
# From gen_m16_varI.py — no atomics, no device-scope staging
# GEMM: grid=231, block=256, shared=16384 (extern)
# Reduce: grid=132, block=256
# ═══════════════════════════════════════════════════════
_HIP_M16 = r"""
// ── PASS 1: GEMM kernel — stores f32 partials to workspace[7][16][2112] ──
extern "C" __global__ __attribute__((amdgpu_flat_work_group_size(256,256)))
void m16_varI_gemm(
    const uint16_t* __restrict__ A,
    const uint8_t*  __restrict__ B_shuffle,
    const uint8_t*  __restrict__ B_scale_sh,
    float*          __restrict__ workspace)
{
    using namespace g;
    constexpr int kN=2112, kK=7168, kKHalf=3584, kTileN=64, kMfmaCols=4;
    constexpr int kSplitK=7, kKPerSplit=1024, kNumWaves=4, kItersPerWave=2;
    constexpr int kNumNTiles=33, kBPanelStride=kKHalf*16;
    constexpr int kLdsFloatsPerCol=16*kNumWaves*16; // 1024
    constexpr int kNXCD=8, kC=4, kBPC=32, kLimit=224;

    int xy = blockIdx.x;
    if (xy < kLimit) {
        int xcd=xy%kNXCD, local_=xy/kNXCD;
        int chunk=local_/kC, pos=local_%kC;
        xy = chunk*kBPC + xcd*kC + pos;
    }

    const int k_split = xy / kNumNTiles;
    const int n_tile  = xy % kNumNTiles;
    const int col_base = n_tile * kTileN;

    const int tid=threadIdx.x, wave_id=tid>>6, lane=tid&63;
    const int row=lane&15, k_quarter=lane>>4;

    extern __shared__ char lds_raw[];
    float* lds = reinterpret_cast<float*>(lds_raw);

    const uint32_t* a_row = reinterpret_cast<const uint32_t*>(A + row * kK);

    const uint8_t* b_bases[kMfmaCols];
    #pragma unroll
    for (int c=0;c<kMfmaCols;++c) {
        int b_col = col_base + c*16 + row;
        b_bases[c] = B_shuffle + (b_col>>4)*kBPanelStride + (b_col&15)*16;
    }

    const int split_iter_base = k_split * (kKPerSplit/128);
    const int ki_base = split_iter_base + wave_id * kItersPerWave;
    const int kg0 = ki_base*4+k_quarter;
    const int kg1 = (ki_base+1)*4+k_quarter;

    // All loads upfront
    asm volatile("" ::: "memory");
    uint32_t a_buf0[16], a_buf1[16];
    { const uint32_t* s = a_row + kg0*16;
      #pragma unroll
      for (int i=0;i<16;++i) a_buf0[i]=s[i]; }
    { const uint32_t* s = a_row + kg1*16;
      #pragma unroll
      for (int i=0;i<16;++i) a_buf1[i]=s[i]; }

    v4i32 b0[kMfmaCols], b1[kMfmaCols];
    int bs0[kMfmaCols], bs1[kMfmaCols];
    #pragma unroll
    for (int c=0;c<kMfmaCols;++c) {
        const uint32_t* s0=reinterpret_cast<const uint32_t*>(b_bases[c]+ki_base*1024+k_quarter*256);
        b0[c]={(int)s0[0],(int)s0[1],(int)s0[2],(int)s0[3]};
        const uint32_t* s1=reinterpret_cast<const uint32_t*>(b_bases[c]+(ki_base+1)*1024+k_quarter*256);
        b1[c]={(int)s1[0],(int)s1[1],(int)s1[2],(int)s1[3]};
        int b_col = col_base+c*16+row;
        bs0[c]=(int)B_scale_sh[soff<224>(b_col,kg0)];
        bs1[c]=(int)B_scale_sh[soff<224>(b_col,kg1)];
    }

    asm volatile("":"+v"(bs0[0]),"+v"(bs0[1]),"+v"(bs0[2]),"+v"(bs0[3]),
                     "+v"(bs1[0]),"+v"(bs1[1]),"+v"(bs1[2]),"+v"(bs1[3])::"memory");

    // Quant + MFMA
    v4f32 acc[kMfmaCols];
    #pragma unroll
    for (int c=0;c<kMfmaCols;++c) acc[c]={0,0,0,0};

    { v4i32 ar; int as; quant_a(a_buf0, ar, as);
      #pragma unroll
      for (int c=0;c<kMfmaCols;++c) acc[c]=mfma16v(ar, b0[c], acc[c], as, bs0[c]); }
    { v4i32 ar; int as; quant_a(a_buf1, ar, as);
      #pragma unroll
      for (int c=0;c<kMfmaCols;++c) acc[c]=mfma16v(ar, b1[c], acc[c], as, bs1[c]); }

    // LDS reduce 4 waves
    const int oc16=lane&15, rb4=(lane>>4)*4;
    #pragma unroll
    for (int c=0;c<kMfmaCols;++c)
        #pragma unroll
        for (int r=0;r<4;++r)
            lds[c*kLdsFloatsPerCol + (rb4+r)*kNumWaves*16 + wave_id*16 + oc16] = acc[c][r];
    __syncthreads();

    // Reduce and store to workspace
    const int ws_base = k_split * 16 * kN;
    { int c=wave_id;
      #pragma unroll
      for (int r=0;r<4;++r) {
        int out_row=rb4+r;
        if (out_row < 16) {
            int lb = c*kLdsFloatsPerCol + out_row*kNumWaves*16 + oc16;
            float sum=0.0f;
            #pragma unroll
            for (int w=0;w<kNumWaves;++w) sum += lds[lb+w*16];
            int gc = col_base + c*16 + oc16;
            workspace[ws_base + out_row*kN + gc] = sum;
        }
      }
    }
}

// ── PASS 2: Reduce kernel — sum 7 partials, convert to bf16 ──
extern "C" __global__ __attribute__((amdgpu_flat_work_group_size(256,256)))
void m16_varI_reduce(
    const float*    __restrict__ workspace,
    uint16_t*       __restrict__ C)
{
    using namespace g;
    constexpr int kM=16, kN=2112, kSplitK=7;
    const int idx = blockIdx.x * 256 + threadIdx.x;
    if (idx >= kM*kN) return;

    float sum=0.0f;
    #pragma unroll
    for (int s=0;s<kSplitK;++s) sum += workspace[s*kM*kN + idx];
    C[idx] = f2bf16(sum);
}
"""

# ═══════════════════════════════════════════════════════
# K3: M=32 kernels (template body, 2 shapes)
# 32x2880x512: grid=180, 32x4096x512: grid=256
# block=256, shared=8192 (static)
# ═══════════════════════════════════════════════════════
_HIP_M32 = r"""
template <int kN>
__device__ void m32_body(const uint16_t* __restrict__ A,
    const uint8_t* __restrict__ B_shuffle,
    const uint8_t* __restrict__ B_scale_sh,
    uint16_t* __restrict__ C)
{
    using namespace g;
    const int n_tile=blockIdx.x/2, m_tile=blockIdx.x%2;
    const int col_base=n_tile*32, row_base=m_tile*16;
    const int tid=threadIdx.x, wave_id=tid>>6, lane=tid&63;
    const int row16=lane&15, kq=lane>>4;
    __shared__ char lds_raw[8192];
    float* lds = reinterpret_cast<float*>(lds_raw);
    const int ki=wave_id, kg=ki*4+kq;
    const int a_row=row_base+row16;

    int bsc[2];
    #pragma unroll
    for (int c=0;c<2;++c) {
        int bc=col_base+c*16+row16;
        bsc[c]=(int)B_scale_sh[soff<16>(bc,kg)];
    }
    asm volatile("":"+v"(bsc[0]),"+v"(bsc[1])::"memory");

    v4i32 b_regs[2];
    #pragma unroll
    for (int c=0;c<2;++c) {
        int bc=col_base+c*16+row16;
        const uint8_t* bb=B_shuffle+(bc>>4)*(256*16)+(bc&15)*16;
        const uint32_t* bs=reinterpret_cast<const uint32_t*>(bb+ki*1024+kq*256);
        b_regs[c]={(int)bs[0],(int)bs[1],(int)bs[2],(int)bs[3]};
    }
    uint32_t a_buf[16];
    { const uint32_t* as=reinterpret_cast<const uint32_t*>(A+a_row*512+kg*32);
      #pragma unroll
      for(int i=0;i<16;++i) a_buf[i]=as[i]; }

    v4f32 acc0={0,0,0,0}, acc1={0,0,0,0};
    { v4i32 ar; int asc; quant_a(a_buf,ar,asc);
      acc0=mfma16a(ar,b_regs[0],acc0,asc,bsc[0]);
      acc1=mfma16a(ar,b_regs[1],acc1,asc,bsc[1]); }

    const int oc=lane&15, rb4=(lane>>4)*4;
    #pragma unroll
    for (int r=0;r<4;++r) {
        int or_=rb4+r;
        lds[or_*128+wave_id*32+oc]=acc0[r];
        lds[or_*128+wave_id*32+oc+16]=acc1[r];
    }
    __syncthreads();

    const int e0=tid, e1=tid+256;
    const int lr0=e0/32,lc0=e0%32, lr1=e1/32,lc1=e1%32;
    const int lb0=lr0*128+lc0, lb1=lr1*128+lc1;
    float s0=lds[lb0]+lds[lb0+32]+lds[lb0+64]+lds[lb0+96];
    float s1=lds[lb1]+lds[lb1+32]+lds[lb1+64]+lds[lb1+96];
    C[(row_base+lr0)*kN+col_base+lc0]=f2bf16(s0);
    C[(row_base+lr1)*kN+col_base+lc1]=f2bf16(s1);
}

extern "C" __global__ __attribute__((amdgpu_flat_work_group_size(256,256)))
void m32_mk4l_m32_n2880_k512(const uint16_t* __restrict__ A,
    const uint8_t* __restrict__ B_shuffle, const uint8_t* __restrict__ B_scale_sh,
    uint16_t* __restrict__ C) { m32_body<2880>(A,B_shuffle,B_scale_sh,C); }

extern "C" __global__ __attribute__((amdgpu_flat_work_group_size(256,256)))
void m32_mk4l_m32_n4096_k512(const uint16_t* __restrict__ A,
    const uint8_t* __restrict__ B_shuffle, const uint8_t* __restrict__ B_scale_sh,
    uint16_t* __restrict__ C) { m32_body<4096>(A,B_shuffle,B_scale_sh,C); }
"""

# ═══════════════════════════════════════════════════════
# K4: M=64, N=7168, K=2048  (32x32x64, fused MFMA pair)
# grid=224, block=256, shared=32768 (static)
# ═══════════════════════════════════════════════════════
_HIP_M64_PRE = r"""
extern "C" __global__ __attribute__((amdgpu_flat_work_group_size(256,256)))
void m64_mk3a_m64_n7168_k2048(
    const uint16_t* __restrict__ A,
    const uint8_t*  __restrict__ B_shuffle,
    const uint8_t*  __restrict__ B_scale_sh,
    uint16_t*       __restrict__ C)
{
    using namespace g;
    // XCD: W=2, C=14, BPC=112, LIMIT=224=TOTAL
    int xy = blockIdx.x;
    { int xcd=xy%8, loc=xy/8, chunk=loc/14, pos=loc%14;
      xy = chunk*112 + xcd*14 + pos; }
    const int l=xy%224, m_tile=(xy/224)*2+(l%2), n_tile=l/2;
    const int col_base=n_tile*64, row_base=m_tile*32;
    const int tid=threadIdx.x, wave_id=tid>>6, lane=tid&31, half=(tid>>5)&1;

    __shared__ char lds_bytes[32768];
    constexpr int lds_half = 4096;

    const int b_col0=col_base+lane, b_col1=col_base+32+lane;
    const uint8_t* bb0 = B_shuffle + (b_col0>>4)*(1024*16) + (b_col0&15)*16;
    const uint8_t* bb1 = B_shuffle + (b_col1>>4)*(1024*16) + (b_col1&15)*16;

    // Cooperative B scale preload (4096 bytes, 256 threads x 16 bytes)
    { int bsc_gbase = (col_base>>5)*(32*64);
      const uint8_t* bsc_src = B_scale_sh + bsc_gbase;
      int my_off = tid * 16;
      #pragma unroll
      for (int bi=0;bi<16;++bi)
          reinterpret_cast<uint8_t*>(lds_bytes)[my_off+bi] = bsc_src[my_off+bi];
    }
    __syncthreads();
    const uint8_t* lds_bscale = reinterpret_cast<const uint8_t*>(lds_bytes);
    const int bsc_base = (col_base>>5)*(32*64);

    const int my_row = row_base + lane;
    const uint32_t* a_row = reinterpret_cast<const uint32_t*>(A + my_row * 2048);
    const int ki_base = wave_id * 8;

    v16f32 acc0, acc1;
    #pragma unroll
    for (int i=0;i<16;++i) { acc0[i]=0.0f; acc1[i]=0.0f; }

    // Prologue B data
    v4i32 b_cur0, b_cur1;
    { const uint32_t* bs0=reinterpret_cast<const uint32_t*>(bb0+ki_base*512+half*256);
      b_cur0={(int)bs0[0],(int)bs0[1],(int)bs0[2],(int)bs0[3]};
      const uint32_t* bs1=reinterpret_cast<const uint32_t*>(bb1+ki_base*512+half*256);
      b_cur1={(int)bs1[0],(int)bs1[1],(int)bs1[2],(int)bs1[3]}; }

    uint32_t a_buf0[16], a_buf1[16], a_buf2[16];
    load_a(a_row + (ki_base*2+half)*16, a_buf0);
    load_a(a_row + ((ki_base+1)*2+half)*16, a_buf1);
    v4i32 b_nxt0, b_nxt1;

    // 8 hand-unrolled iterations
"""

_HIP_M64_POST = r"""
    // Barrier + LDS double reduction
    __syncthreads();
    float* lds_r = reinterpret_cast<float*>(lds_bytes);
    #pragma unroll
    for (int i=0;i<4;++i) {
        #pragma unroll
        for (int j=0;j<4;++j) {
            int row=half*4+j+i*8;
            int base=row*128+wave_id*32+lane;
            lds_r[base]=acc0[i*4+j];
            lds_r[base+lds_half]=acc1[i*4+j];
        }
    }
    __syncthreads();
    { int i=wave_id;
      #pragma unroll
      for (int j=0;j<4;++j) {
          int row=half*4+j+i*8;
          int gr=row_base+row;
          float s0=0.0f, s1=0.0f;
          #pragma unroll
          for (int w=0;w<4;++w) {
              s0 += lds_r[row*128+w*32+lane];
              s1 += lds_r[lds_half+row*128+w*32+lane];
          }
          C[gr*7168+col_base+lane]=f2bf16(s0);
          C[gr*7168+col_base+32+lane]=f2bf16(s1);
      }
    }
}
"""

# ═══════════════════════════════════════════════════════
# K5: M=256, N=3072, K=1536  (32x32x64, fused MFMA pair)
# grid=384, block=256, shared=32768 (static)
# ═══════════════════════════════════════════════════════
_HIP_M256 = r"""
extern "C" __global__ __attribute__((amdgpu_flat_work_group_size(256,256)))
void m256_v6a_kernel(
    const uint16_t* __restrict__ A,
    const uint8_t*  __restrict__ B_shuffle,
    const uint8_t*  __restrict__ B_scale_sh,
    uint16_t*       __restrict__ C)
{
    using namespace g;
    // XCD: W=8, C=6, BPC=48, LIMIT=384=TOTAL
    int xy = blockIdx.x;
    { int xcd=xy%8, loc=xy/8, chunk=loc/6, pos=loc%6;
      xy = chunk*48 + xcd*6 + pos; }
    const int l=xy%384, m_tile=(xy/384)*8+(l%8), n_tile=l/8;
    const int col_base=n_tile*64, row_base=m_tile*32;
    const int tid=threadIdx.x, wave_id=tid>>6, lane=tid&31, half=(tid>>5)&1;

    __shared__ char lds_bytes[32768];
    constexpr int lds_half = 4096;

    const int bc0=col_base+lane, bc1=col_base+32+lane;
    const uint8_t* bb0 = B_shuffle + (bc0>>4)*(768*16) + (bc0&15)*16;
    const uint8_t* bb1 = B_shuffle + (bc1>>4)*(768*16) + (bc1&15)*16;

    const int my_row = row_base + lane;
    const uint32_t* a_row = reinterpret_cast<const uint32_t*>(A + my_row * 1536);
    const int ki_base = wave_id * 6;

    v16f32 acc0, acc1;
    #pragma unroll
    for (int i=0;i<16;++i) { acc0[i]=0.0f; acc1[i]=0.0f; }

    // Prologue A/B loads (fly during LDS preload)
    uint32_t a_triple[3][16];
    load_a(a_row + (ki_base*2+half)*16, a_triple[0]);
    load_a(a_row + ((ki_base+1)*2+half)*16, a_triple[1]);
    v4i32 b_cur0, b_cur1;
    { const uint32_t* bs0=reinterpret_cast<const uint32_t*>(bb0+ki_base*512+half*256);
      b_cur0={(int)bs0[0],(int)bs0[1],(int)bs0[2],(int)bs0[3]};
      const uint32_t* bs1=reinterpret_cast<const uint32_t*>(bb1+ki_base*512+half*256);
      b_cur1={(int)bs1[0],(int)bs1[1],(int)bs1[2],(int)bs1[3]}; }

    // Anti-stagger
    if (blockIdx.x & 4) asm volatile("s_sleep 1" :::);
    if (blockIdx.x & 8) asm volatile("s_sleep 1" :::);

    // Cooperative B scale preload (3072 bytes, 12 bytes/thread)
    { int bsc_gbase = (col_base>>5)*(32*48);
      const uint8_t* bsc_src = B_scale_sh + bsc_gbase;
      constexpr int bpt = 12;
      int my_off = tid * bpt;
      #pragma unroll
      for (int bi=0;bi<bpt;++bi)
          reinterpret_cast<uint8_t*>(lds_bytes)[my_off+bi] = bsc_src[my_off+bi]; }
    __syncthreads();
    const uint8_t* lds_bscale = reinterpret_cast<const uint8_t*>(lds_bytes);
    const int bsc_base = (col_base>>5)*(32*48);

    // Main K loop: 6 iters
    #pragma unroll
    for (int k=0;k<6;++k) {
        const int ki=ki_base+k;
        int kg=ki*2+half;
        int bsc0_=(int)lds_bscale[soff<48>(bc0,kg)-bsc_base];
        int bsc1_=(int)lds_bscale[soff<48>(bc1,kg)-bsc_base];
        v4i32 ar; int asc;
        quant_a(a_triple[k%3], ar, asc);
        v4i32 b_nxt0={0,0,0,0}, b_nxt1={0,0,0,0};
        if (k+1<6) {
            const uint32_t* bs0=reinterpret_cast<const uint32_t*>(bb0+(ki+1)*512+half*256);
            b_nxt0={(int)bs0[0],(int)bs0[1],(int)bs0[2],(int)bs0[3]};
            const uint32_t* bs1=reinterpret_cast<const uint32_t*>(bb1+(ki+1)*512+half*256);
            b_nxt1={(int)bs1[0],(int)bs1[1],(int)bs1[2],(int)bs1[3]};
        }
        mfma_pair(ar, b_cur0, b_cur1, acc0, acc1, asc, bsc0_, bsc1_);
        if (k+2<6) load_a(a_row+((ki+2)*2+half)*16, a_triple[(k+2)%3]);
        b_cur0=b_nxt0; b_cur1=b_nxt1;
    }

    __syncthreads();
    float* lds_r = reinterpret_cast<float*>(lds_bytes);
    #pragma unroll
    for (int i=0;i<4;++i) {
        #pragma unroll
        for (int j=0;j<4;++j) {
            int row=half*4+j+i*8;
            int base=row*128+wave_id*32+lane;
            lds_r[base]=acc0[i*4+j];
            lds_r[base+lds_half]=acc1[i*4+j];
        }
    }
    __syncthreads();
    { int i=wave_id;
      #pragma unroll
      for (int j=0;j<4;++j) {
          int row=half*4+j+i*8;
          int gr=row_base+row;
          float s0=0.0f, s1=0.0f;
          #pragma unroll
          for (int w=0;w<4;++w) {
              s0 += lds_r[row*128+w*32+lane];
              s1 += lds_r[lds_half+row*128+w*32+lane];
          }
          C[gr*3072+col_base+lane]=f2bf16(s0);
          C[gr*3072+col_base+32+lane]=f2bf16(s1);
      }
    }
}
"""

# ═══════════════════════════════════════════════════════
# C++ Dispatch (hip_module_v2 style)
# ═══════════════════════════════════════════════════════
_HIP_DISPATCH = r"""
static py::function g_fallback_fn;
static bool g_has_fallback = false;

#define SK(m,n,k) ((uint64_t)(m)|((uint64_t)(n)<<16)|((uint64_t)(k)<<32))

torch::Tensor dispatch(py::tuple data) {
    const auto& A = THPVariable_Unpack(data[0].ptr());
    const auto& B = THPVariable_Unpack(data[1].ptr());
    const auto& B_shuffle = THPVariable_Unpack(data[3].ptr());
    const auto& B_scale_sh = THPVariable_Unpack(data[4].ptr());

    int64_t m=A.size(0), k=A.size(1), n=B.size(0);
    uint64_t sk = SK(m,n,k);

    auto opts = torch::TensorOptions().dtype(torch::kBFloat16).device(torch::kCUDA);
    const uint16_t* a  = (const uint16_t*)A.data_ptr();
    const uint8_t*  bs = (const uint8_t*)B_shuffle.data_ptr();
    const uint8_t*  bsc= (const uint8_t*)B_scale_sh.data_ptr();

    torch::Tensor output;
    switch (sk) {
    case SK(4,2880,512):
        output = torch::empty({m,n}, opts);
        hipLaunchKernelGGL(fused_m4_m4_n2880_k512, dim3(90),dim3(256),0,0,
            a,bs,bsc,(uint16_t*)output.data_ptr());
        break;
    case SK(16,2112,7168): {
        output = torch::empty({m,n}, opts);
        auto ws_opts = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
        auto workspace = torch::empty({7*16*2112}, ws_opts);
        hipLaunchKernelGGL(m16_varI_gemm, dim3(231),dim3(256),16384,0,
            a,bs,bsc,(float*)workspace.data_ptr());
        hipLaunchKernelGGL(m16_varI_reduce, dim3(132),dim3(256),0,0,
            (const float*)workspace.data_ptr(),(uint16_t*)output.data_ptr());
        break;
    }
    case SK(32,2880,512):
        output = torch::empty({m,n}, opts);
        hipLaunchKernelGGL(m32_mk4l_m32_n2880_k512, dim3(180),dim3(256),0,0,
            a,bs,bsc,(uint16_t*)output.data_ptr());
        break;
    case SK(32,4096,512):
        output = torch::empty({m,n}, opts);
        hipLaunchKernelGGL(m32_mk4l_m32_n4096_k512, dim3(256),dim3(256),0,0,
            a,bs,bsc,(uint16_t*)output.data_ptr());
        break;
    case SK(64,7168,2048):
        output = torch::empty({m,n}, opts);
        hipLaunchKernelGGL(m64_mk3a_m64_n7168_k2048, dim3(224),dim3(256),0,0,
            a,bs,bsc,(uint16_t*)output.data_ptr());
        break;
    case SK(256,3072,1536):
        output = torch::empty({m,n}, opts);
        hipLaunchKernelGGL(m256_v6a_kernel, dim3(384),dim3(256),0,0,
            a,bs,bsc,(uint16_t*)output.data_ptr());
        break;
    default:
        if (g_has_fallback)
            return g_fallback_fn(data).cast<torch::Tensor>();
        throw std::runtime_error("[dispatch] unsupported shape "
            + std::to_string(m)+"x"+std::to_string(n)+"x"+std::to_string(k));
    }
    #undef SK
    return output;
}

void set_fallback(py::function fn) {
    g_fallback_fn = std::move(fn);
    g_has_fallback = true;
}
"""

# ═══════════════════════════════════════════════════════
# Assemble full HIP source
# ═══════════════════════════════════════════════════════
_FULL_HIP = (
    _HIP_HEADER
    + _HIP_M4
    + _HIP_M16
    + _HIP_M32
    + _HIP_M64_PRE + "\n" + _gen_m64_iters() + "\n" + _HIP_M64_POST
    + _HIP_M256
    + _HIP_DISPATCH
)

_CPP_DECL = r"""
#include <torch/extension.h>
#include <pybind11/pybind11.h>
namespace py = pybind11;
torch::Tensor dispatch(py::tuple data);
void set_fallback(py::function fn);
"""


# ═══════════════════════════════════════════════════════
# Build
# ═══════════════════════════════════════════════════════
print("[standalone_best_v3b] Compiling...", file=sys.stderr)

_name = "gemm_standalone_best_v3b_ext"
_bdir = os.path.join(tempfile.gettempdir(), _name)
os.makedirs(_bdir, exist_ok=True)

os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
_mod = load_inline(
    name=_name,
    cpp_sources=_CPP_DECL,
    cuda_sources=_FULL_HIP,
    functions=["dispatch", "set_fallback"],
    with_cuda=True,
    extra_cflags=["-O3"],
    extra_cuda_cflags=[
        "-O3", "-std=c++17",
        "-ffast-math",
        "-fgpu-flush-denormals-to-zero",
        "-mllvm", "-amdgpu-kernarg-preload-count=10",
    ],
    build_directory=_bdir,
    verbose=bool(int(os.environ.get("INLINE_HIP_VERBOSE", "0"))),
)

# Register aiter fallback for unsupported shapes
def _aiter_fallback(data):
    from aiter import dtypes
    import aiter
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle
    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    A_fp4, A_scale = dynamic_mxfp4_quant(A)
    A_scale_sh = e8m0_shuffle(A_scale)
    A_q = A_fp4.view(dtypes.fp4x2)
    A_scale_sh = A_scale_sh.view(dtypes.fp8_e8m0)
    return aiter.gemm_a4w4(
        A_q, B_shuffle, A_scale_sh, B_scale_sh,
        dtype=dtypes.bf16, bpreshuffle=True,
    )

_mod.set_fallback(_aiter_fallback)
print("[standalone_best_v3b] Ready.", file=sys.stderr)

custom_kernel = _mod.dispatch
scrolls · 913 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