Skip to content
KernelIndex
Search⌘K

submission 727205

vuxml · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v25_pf.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-727205?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, int32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD Instinct MI355X
46.3µs
#146 of 766
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:92473bd39f4fcead3c4cf408362e58dd20ac520389dbb09cf04826f126999a33
license declaredunknown
license concludedunknown
authorsvuxml
imported2026-08-15

Techniques

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

fp4f4, sc = kvd["mxfp4"]
shared-memory__shared__ uint8_t Ks8[BN][KPAD];
vector-width = uint4void qquant32(const bf16_t* Qh, uint4& out, int& e8){

Kernel source

submission_v25_pf.py504 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
v25: LEAN. v24 proved iglp_opt(1)/(0) = -3% consistent (256/8k
244->236.6, b64w4s1 VGPR 168->144). v24b timed out (8 compiles).

SCHED (all fp8+tr8, 4 variants only):
  1: iglp_opt(1) QK + iglp_opt(0) PV  [v24-PROVEN -3%]
  3: Cross-iter PF: bkN[NTW][5] primed in prologue. Main loop:
     QK+cvt on bkN (CURRENT), THEN issue LOADK(n+BN)->bkN (NEXT),
     THEN sched_barrier(0). Softmax+PV (~800cyc) follow -> next
     iter's bkN is warm by the time QK_CVT hits it. The SBAR0 pins
     LOADK above softmax (can't sink). +25*NTW VGPR loop-carried.
     iglp_opt(0) on PV region only (QK region has SBAR -> iglp noop).

b64w4 x {s1,s3} + b128w4 x {s1,s3} = 4. 4/8k unpinned (race it).
"""
import os
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")

import sys, math, time
import torch

_L = lambda *a: print("[v25]", *a, file=sys.stderr, flush=True)

NH, NKV, DQK, DV = 16, 1, 576, 512
SM_SCALE = 1.0 / math.sqrt(576.0)


_HIP_HEAD = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <cstdint>
typedef int    i32x2 __attribute__((ext_vector_type(2)));
typedef int    i32x8 __attribute__((ext_vector_type(8)));
typedef float  f32x2 __attribute__((ext_vector_type(2)));
typedef float  f32x4 __attribute__((ext_vector_type(4)));
typedef __hip_bfloat16 bf16_t;
typedef __attribute__((address_space(3))) i32x2* lds_i32x2p;
__device__ __forceinline__ float shx(float v,int m){return __shfl_xor(v,m,64);}
__device__ __forceinline__ float e8f(uint8_t e){
  union{uint32_t u;float f;}c;c.u=(uint32_t)e<<23;return c.f;}
__device__ __forceinline__ uint32_t f2u(float x){
  union{float f;uint32_t u;}c;c.f=x;return c.u;}
__device__ __forceinline__ void cvt_fp4_fp8_8(uint32_t w,float sc,int* o2){
  f32x2 a=__builtin_amdgcn_cvt_scalef32_pk_f32_fp4(w,sc,0);
  f32x2 b=__builtin_amdgcn_cvt_scalef32_pk_f32_fp4(w,sc,1);
  f32x2 c=__builtin_amdgcn_cvt_scalef32_pk_f32_fp4(w,sc,2);
  f32x2 d=__builtin_amdgcn_cvt_scalef32_pk_f32_fp4(w,sc,3);
  int r0=0,r1=0;
  r0=__builtin_amdgcn_cvt_pk_fp8_f32(a[0],a[1],r0,false);
  r0=__builtin_amdgcn_cvt_pk_fp8_f32(b[0],b[1],r0,true);
  r1=__builtin_amdgcn_cvt_pk_fp8_f32(c[0],c[1],r1,false);
  r1=__builtin_amdgcn_cvt_pk_fp8_f32(d[0],d[1],r1,true);
  o2[0]=r0;o2[1]=r1;}
#if defined(__HIP_DEVICE_COMPILE__)
#define LDTR8(p)  __builtin_amdgcn_ds_read_tr8_b64_v2i32 ((lds_i32x2p)(void*)(p))
#define SBAR0()   __builtin_amdgcn_sched_barrier(0)
#define IGLP(v)   __builtin_amdgcn_iglp_opt((v))
#else
#define LDTR8(p)  (*(i32x2*)(void*)(p))
#define SBAR0()   ((void)0)
#define IGLP(v)   ((void)0)
#endif
#define NH 16
#define DQK 576
#define DV 512
#define QCV(o,a,b,s,bs) __builtin_amdgcn_cvt_scalef32_pk_fp4_f32((o),(a),(b),(s),(bs))
#define P8SC 448.0f

__device__ __forceinline__
void qquant32(const bf16_t* Qh, uint4& out, int& e8){
  float amax=0.f;
  #pragma unroll
  for(int i=0;i<32;++i){float a=__builtin_fabsf((float)Qh[i]);amax=a>amax?a:amax;}
  uint32_t ab=(f2u(amax)+0x200000u)&0xFF800000u;
  int su=ab?(int)((ab>>23)&0xFFu)-129:-127;
  su=su<-127?-127:(su>127?127:su);e8=su+127;
  float bsc=e8f((uint8_t)e8);
  #define QF(i) ((float)Qh[i])
  int w0=0,w1=0,w2=0,w3=0;
  w0=QCV(w0,QF( 0),QF( 1),bsc,0);w0=QCV(w0,QF( 2),QF( 3),bsc,1);
  w0=QCV(w0,QF( 4),QF( 5),bsc,2);w0=QCV(w0,QF( 6),QF( 7),bsc,3);
  w1=QCV(w1,QF( 8),QF( 9),bsc,0);w1=QCV(w1,QF(10),QF(11),bsc,1);
  w1=QCV(w1,QF(12),QF(13),bsc,2);w1=QCV(w1,QF(14),QF(15),bsc,3);
  w2=QCV(w2,QF(16),QF(17),bsc,0);w2=QCV(w2,QF(18),QF(19),bsc,1);
  w2=QCV(w2,QF(20),QF(21),bsc,2);w2=QCV(w2,QF(22),QF(23),bsc,3);
  w3=QCV(w3,QF(24),QF(25),bsc,0);w3=QCV(w3,QF(26),QF(27),bsc,1);
  w3=QCV(w3,QF(28),QF(29),bsc,2);w3=QCV(w3,QF(30),QF(31),bsc,3);
  #undef QF
  out=make_uint4(w0,w1,w2,w3);
}
"""


def _hip_body(bn, nw, sched):
    return r"""
#define BN """ + str(bn) + r"""
#define NW """ + str(nw) + r"""
#define SCHED """ + str(sched) + r"""
#define KPAD 520
#define NT  (BN/16)
#define NTW (BN/16/NW)
#define VTW (DV/16/NW)
#define NKT (BN/32)
static_assert(NTW>=1 && DV%(16*NW)==0, "shape");

#define LOADK(tki, BK, SC) do { \
  long _gt=(long)(kv0+((tki)<kvn?(tki):kvn-1)); \
  const uint8_t* _Kf=KVF+_gt*(DQK/2); \
  const uint8_t* _Kc=KVS+_gt*scs; \
  (SC)[0]=*(const uint32_t*)(_Kc+0);(SC)[1]=*(const uint32_t*)(_Kc+4); \
  (SC)[2]=*(const uint32_t*)(_Kc+8);(SC)[3]=*(const uint32_t*)(_Kc+12); \
  (SC)[4]=*(const uint32_t*)(_Kc+16); \
  (BK)[0]=*(const uint4*)(_Kf+ 0*64+lg*16); \
  (BK)[1]=*(const uint4*)(_Kf+ 1*64+lg*16); \
  (BK)[2]=*(const uint4*)(_Kf+ 2*64+lg*16); \
  (BK)[3]=*(const uint4*)(_Kf+ 3*64+lg*16); \
  (BK)[4]=(lg<2)?*(const uint4*)(_Kf+4*64+lg*16):make_uint4(0,0,0,0); \
} while(0)
#define SCB(SC,bi) ((bi)<18?(int)(((SC)[(bi)>>2]>>(((bi)&3)*8))&0xFFu):127)

extern "C" __global__ __launch_bounds__(64*NW)
void mla_s1(
    const bf16_t*__restrict__ Q,const uint8_t*__restrict__ KVF,
    const uint8_t*__restrict__ KVS,const int32_t*__restrict__ KVI,
    float*__restrict__ PO,float*__restrict__ PM,float*__restrict__ PL,
    int bs,int ns,int scs,float sms)
{
  const int wg=blockIdx.x,b=wg/ns,sp=wg%ns;
  if(b>=bs)return;
  const int tid=threadIdx.x,w=tid>>6,l=tid&63,lr=l&15,lg=l>>4;
  const int kv0=KVI[b],kv1=KVI[b+1],kvn=kv1-kv0;
  const int ch=((kvn+ns-1)/ns+BN-1)/BN*BN,n0=sp*ch,n1=min(n0+ch,kvn);
  __shared__ uint8_t Ks8[BN][KPAD];
  __shared__ uint8_t Ps8[NH][BN];
  __shared__ uint8_t Qsf[64][88];
  __shared__ float Sm[NW][NH], Ss[NW][NH];

  if(w==0){
    const bf16_t* Qh0=Q+(long)b*NH*DQK+lr*DQK+lg*32;
    for(int j=0;j<5;++j){
      uint4 qf;int qs;
      if(j*128+lg*32<DQK) qquant32(Qh0+j*128,qf,qs);
      else{qf=make_uint4(0,0,0,0);qs=127;}
      *reinterpret_cast<uint4*>(&Qsf[l][j*16])=qf;
      Qsf[l][80+j]=(uint8_t)qs;
    }
  }
  __syncthreads();

  f32x4 oacc[VTW];
  #pragma unroll
  for(int j=0;j<VTW;++j)oacc[j]=(f32x4){0,0,0,0};
  float mrow[4]={-1e30f,-1e30f,-1e30f,-1e30f},lrow[4]={0,0,0,0};

  uint4 bkN[NTW][5]; uint32_t scN[NTW][5];
#if SCHED==3
  if(n0<n1){
    #pragma unroll
    for(int t=0;t<NTW;++t) LOADK(n0+(w*NTW+t)*16+lr, bkN[t], scN[t]);
  }
#endif

  for(int n=n0;n<n1;n+=BN){
    const int nv=min(BN,n1-n);
    f32x4 sc[NTW];
#if SCHED==1
    #pragma unroll
    for(int t=0;t<NTW;++t) LOADK(n+(w*NTW+t)*16+lr, bkN[t], scN[t]);
#endif
    #pragma unroll
    for(int t=0;t<NTW;++t){
      const int tw=w*NTW+t;
      sc[t]=(f32x4){0,0,0,0};
      #pragma unroll
      for(int j=0;j<5;++j){
        i32x8 a={0,0,0,0,0,0,0,0},bk={0,0,0,0,0,0,0,0};
        *reinterpret_cast<uint4*>(&a)=
          *reinterpret_cast<const uint4*>(&Qsf[l][j*16]);
        int sca=(int)Qsf[l][80+j];
        int scb=SCB(scN[t], j*4+lg);
        *reinterpret_cast<uint4*>(&bk)=bkN[t][j];
        sc[t]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            a,bk,sc[t],4,4,0,sca,0,scb);
      }
      #pragma unroll
      for(int j=0;j<4;++j){
        int bi=j*4+lg;
        float bsc=e8f((uint8_t)SCB(scN[t], bi));
        int o2[8];
        cvt_fp4_fp8_8(bkN[t][j].x,bsc,o2+0);cvt_fp4_fp8_8(bkN[t][j].y,bsc,o2+2);
        cvt_fp4_fp8_8(bkN[t][j].z,bsc,o2+4);cvt_fp4_fp8_8(bkN[t][j].w,bsc,o2+6);
        int* dst=reinterpret_cast<int*>(&Ks8[tw*16+lr][bi*32]);
        #pragma unroll
        for(int k=0;k<8;++k)dst[k]=o2[k];
      }
    }
#if SCHED==3
    #pragma unroll
    for(int t=0;t<NTW;++t) LOADK((n+BN)+(w*NTW+t)*16+lr, bkN[t], scN[t]);
    SBAR0();
#else
    IGLP(1);
#endif

    #pragma unroll
    for(int i=0;i<4;++i){
      float vmax=-1e30f;
      #pragma unroll
      for(int t=0;t<NTW;++t){
        float v=sc[t][i]*sms;
        if((w*NTW+t)*16+lr>=nv)v=-1e30f;
        sc[t][i]=v;
        v=fmaxf(v,shx(v,1));v=fmaxf(v,shx(v,2));
        v=fmaxf(v,shx(v,4));v=fmaxf(v,shx(v,8));
        vmax=fmaxf(vmax,v);
      }
      if(lr==0) Sm[w][lg*4+i]=vmax;
    }
    __syncthreads();

    float al[4],mn_[4];
    #pragma unroll
    for(int i=0;i<4;++i){
      float gm=Sm[0][lg*4+i];
      #pragma unroll
      for(int k=1;k<NW;++k) gm=fmaxf(gm,Sm[k][lg*4+i]);
      mn_[i]=fmaxf(mrow[i],gm);al[i]=__expf(mrow[i]-mn_[i]);
      float s=0.f;
      #pragma unroll
      for(int t=0;t<NTW;++t){
        float p=__expf(sc[t][i]-mn_[i]);
        int pk=0;pk=__builtin_amdgcn_cvt_pk_fp8_f32(p*P8SC,0.f,pk,false);
        Ps8[lg*4+i][(w*NTW+t)*16+lr]=((uint8_t*)&pk)[0];
        s+=p;
      }
      s+=shx(s,1);s+=shx(s,2);s+=shx(s,4);s+=shx(s,8);
      if(lr==0) Ss[w][lg*4+i]=s;
    }
    #pragma unroll
    for(int j=0;j<VTW;++j)
      #pragma unroll
      for(int i=0;i<4;++i)oacc[j][i]*=al[i];
    __syncthreads();

    #pragma unroll
    for(int i=0;i<4;++i){
      float ts=0.f;
      #pragma unroll
      for(int k=0;k<NW;++k) ts+=Ss[k][lg*4+i];
      mrow[i]=mn_[i]; lrow[i]=lrow[i]*al[i]+ts;
    }

    long aPt8[NKT];
    #pragma unroll
    for(int kt=0;kt<NKT;++kt)
      aPt8[kt]=*reinterpret_cast<const long*>(&Ps8[lr][kt*32+lg*8]);
    #pragma unroll
    for(int j=0;j<VTW;++j){
      const int vc0=(w*VTW+j)*16;
      #pragma unroll
      for(int kt=0;kt<NKT;++kt){
        i32x2 bv2=LDTR8(&Ks8[kt*32+lg*8+(lr>>1)][vc0+(lr&1)*8]);
        long bV=*reinterpret_cast<long*>(&bv2);
        oacc[j]=__builtin_amdgcn_mfma_f32_16x16x32_fp8_fp8(
            aPt8[kt],bV,oacc[j],0,0,0);
      }
    }
    IGLP(0);
    __syncthreads();
  }

  const float ip=1.0f/P8SC;
  const long ob=((long)sp*bs+b)*NH;float* POb=PO+ob*DV;
  #pragma unroll
  for(int j=0;j<VTW;++j){
    const int vc=(w*VTW+j)*16+lr;
    #pragma unroll
    for(int i=0;i<4;++i)POb[(lg*4+i)*DV+vc]=oacc[j][i]*ip;
  }
  if(w==0&&lr==0){
    #pragma unroll
    for(int i=0;i<4;++i){PM[ob+lg*4+i]=mrow[i];PL[ob+lg*4+i]=lrow[i];}
  }
}
"""


_HIP_TAIL = r"""
extern "C" __global__ __launch_bounds__(128)
void mla_rd(const float*__restrict__ PO,const float*__restrict__ PM,
    const float*__restrict__ PL,bf16_t*__restrict__ O,int bs,int ns){
  const int bh=blockIdx.x,l=threadIdx.x,d0=l*4;float mm=-1e30f;
  for(int s=0;s<ns;++s)mm=fmaxf(mm,PM[(long)s*bs*NH+bh]);
  f32x4 a={0,0,0,0};float ls=0;
  for(int s=0;s<ns;++s){const long bb=(long)s*bs*NH+bh;
    float wt=__expf(PM[bb]-mm);ls+=wt*PL[bb];
    f32x4 p=*reinterpret_cast<const f32x4*>(PO+bb*DV+d0);
    a[0]+=wt*p[0];a[1]+=wt*p[1];a[2]+=wt*p[2];a[3]+=wt*p[3];}
  float iv=1.f/ls;bf16_t* ob=O+(long)bh*DV+d0;
  ob[0]=(bf16_t)(a[0]*iv);ob[1]=(bf16_t)(a[1]*iv);
  ob[2]=(bf16_t)(a[2]*iv);ob[3]=(bf16_t)(a[3]*iv);
}
#include <torch/extension.h>
void run_s1(torch::Tensor Q,torch::Tensor F,torch::Tensor S,torch::Tensor I,
    torch::Tensor PO,torch::Tensor PM,torch::Tensor PL,
    int64_t bs,int64_t ns,int64_t scs,double sm){
  dim3 g(bs*ns),bk(64*NW);
  mla_s1<<<g,bk,0,0>>>(reinterpret_cast<const bf16_t*>(Q.data_ptr()),
    F.data_ptr<uint8_t>(),S.data_ptr<uint8_t>(),I.data_ptr<int32_t>(),
    PO.data_ptr<float>(),PM.data_ptr<float>(),PL.data_ptr<float>(),
    (int)bs,(int)ns,(int)scs,(float)sm);}
void run_rd(torch::Tensor PO,torch::Tensor PM,torch::Tensor PL,
    torch::Tensor O,int64_t bs,int64_t ns){
  dim3 g(bs*NH),bk(128);
  mla_rd<<<g,bk,0,0>>>(PO.data_ptr<float>(),PM.data_ptr<float>(),
    PL.data_ptr<float>(),reinterpret_cast<bf16_t*>(O.data_ptr()),
    (int)bs,(int)ns);}
void probe(){hipFuncAttributes a;
  if(hipFuncGetAttributes(&a,reinterpret_cast<const void*>(mla_s1))==0)
    printf("[v25 %s] VGPR=%d shmem=%zu spill=%zu\n",
      KTAG,a.numRegs,a.sharedSizeBytes,a.localSizeBytes);}
"""


_CPP = r"""
#include <torch/extension.h>
void run_s1(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
    torch::Tensor,torch::Tensor,torch::Tensor,int64_t,int64_t,int64_t,double);
void run_rd(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,int64_t,int64_t);
void probe();
"""


def _hip_src(bn, nw, sched):
    tag = f"b{bn}w{nw}s{sched}"
    return _HIP_HEAD + f'\n#define KTAG "{tag}"\n' + _hip_body(bn, nw, sched) + _HIP_TAIL


def _compile(name, src):
    from torch.utils.cpp_extension import load_inline
    try:
        return load_inline(name=name, cpp_sources=_CPP, cuda_sources=src,
            functions=["run_s1", "run_rd", "probe"], with_cuda=True,
            extra_cuda_cflags=["-O3", "--offload-arch=gfx950", "-ffast-math"],
            verbose=False)
    except Exception as e:
        _L(f"{name} COMPILE FAIL: ...{str(e)[-400:]}")
        return None


_MODS = {}
for bn, nw, sched in ((64, 4, 1), (64, 4, 3), (128, 4, 1), (128, 4, 3)):
    _t0 = time.time()
    m = _compile(f"mla_v25_{bn}_{nw}_{sched}", _hip_src(bn, nw, sched))
    if m is not None:
        _MODS[(bn, nw, sched)] = m
        _L(f"b{bn}w{nw}s{sched} compiled {time.time()-_t0:.1f}s")
        m.probe()
_L(f"hip modules: {sorted(_MODS.keys())}")


_L2B = None
def _cold(fn, n=5):
    global _L2B
    if _L2B is None:
        _L2B = torch.empty(int(2 * 1024**3), dtype=torch.int8, device="cuda")
    fn(); torch.cuda.synchronize()
    ts = []
    for _ in range(n):
        _L2B.zero_(); torch.cuda.synchronize()
        e0, e1 = torch.cuda.Event(True), torch.cuda.Event(True)
        e0.record(); fn(); e1.record(); torch.cuda.synchronize()
        ts.append(e0.elapsed_time(e1) * 1000)
    ts.sort(); return sum(ts[1:-1]) / max(1, n - 2)


def _make_ai(bs, kvlen, q0, qo_i, kv_i, dev):
    import aiter
    from aiter import dtypes as adt, get_mla_metadata_info_v1, get_mla_metadata_v1
    from aiter.ops.quant import static_per_tensor_quant as _sq
    FP8 = adt.fp8
    _s1 = next(getattr(aiter, n) for n in dir(aiter)
               if n.startswith("mla_decode_stage1_") and n.endswith("_fwd"))
    _rd = aiter.mla_reduce_v1
    tkv = bs * kvlen
    ki = torch.arange(tkv, dtype=torch.int32, device=dev)
    kl = torch.full((bs,), kvlen, dtype=torch.int32, device=dev)
    info = get_mla_metadata_info_v1(bs, 1, NH, FP8, FP8, is_sparse=False,
        fast_mode=False, num_kv_splits=32, intra_batch_mode=True)
    wm, wi, wis, ri, rfm, rpm = (torch.empty(s, dtype=t, device=dev) for s, t in info)
    get_mla_metadata_v1(qo_i, kv_i, kl, NH, NKV, True, wm, wis, wi, ri, rfm, rpm,
        page_size=1, kv_granularity=16, max_seqlen_qo=1, uni_seqlen_qo=1,
        fast_mode=False, max_split_per_batch=32, intra_batch_mode=True,
        dtype_q=FP8, dtype_kv=FP8)
    pn = rpm.size(0)
    lg = torch.empty((pn, 1, NH, DV), dtype=torch.float32, device=dev)
    ls = torch.empty((pn, 1, NH, 1), dtype=torch.float32, device=dev)
    out = torch.empty((bs, NH, DV), dtype=torch.bfloat16, device=dev)
    qs = torch.tensor([(q0.abs().amax().item() * 1.25) / float(torch.finfo(FP8).max)],
        dtype=torch.float32, device=dev)
    qf8 = torch.empty((bs, NH, DQK), dtype=FP8, device=dev)
    _sq(qf8, q0, qs)
    shp = (tkv, 1, NKV, DQK)
    def hot(q, kf8, ks8, qo, kv):
        _sq(qf8, q, qs)
        _s1(qf8, kf8.view(*shp), qo, kv, ki, kl, None, wm, wi, wis, 1, 1, NKV,
            SM_SCALE, lg, ls, out, qs, ks8)
        _rd(lg, ls, ri, rfm, rpm, 1, out, None)
        return out
    return hot


_cache = {}
def _build(bs, kvlen, data, dev):
    q, kvd, qo_i, kv_i, cfg = data
    kf8, ks8 = kvd["fp8"]
    ref = None; ah = None
    try:
        ah = _make_ai(bs, kvlen, q, qo_i, kv_i, dev)
        ref = ah(q, kf8, ks8, qo_i, kv_i).float()
    except Exception as e:
        _L(f"ai FAIL:{e}")
    t0 = time.time(); BUDGET = 20.0
    t_ai = _cold(lambda: ah(q, kf8, ks8, qo_i, kv_i)) if ah else 1e9
    _L(f"[{bs},{kvlen}] ai cold={t_ai:.1f}")

    ns_mem = max(1, (3 * 1024**3) // (bs * NH * DV * 4))
    ns_hi = min(max(1, 2048 // max(bs, 1)), ns_mem, max(1, kvlen // 32))
    cands = sorted({x for x in (2, 4, 8, 16, 32) if x <= ns_hi}) or [1]
    O = torch.empty((bs, NH, DV), dtype=torch.bfloat16, device=dev)
    f4, sc = kvd["mxfp4"]
    f4u = f4.view(torch.uint8).reshape(-1, DQK // 2).contiguous()
    scu = sc.view(torch.uint8).contiguous(); scs = scu.stride(0)
    PO, PM, PL = {}, {}, {}
    for ns in cands:
        PO[ns] = torch.empty((ns, bs, NH, DV), dtype=torch.float32, device=dev)
        PM[ns] = torch.empty((ns, bs, NH), dtype=torch.float32, device=dev)
        PL[ns] = torch.empty((ns, bs, NH), dtype=torch.float32, device=dev)

    best_hip = (1e9, None)
    for key, mod in sorted(_MODS.items()):
        bn, nw, sched = key
        for ns in cands:
            if time.time() - t0 > BUDGET: break
            try:
                def _h(ns=ns, mod=mod):
                    mod.run_s1(q, f4u, scu, kv_i, PO[ns], PM[ns], PL[ns],
                               bs, ns, scs, SM_SCALE)
                    mod.run_rd(PO[ns], PM[ns], PL[ns], O, bs, ns)
                _h(); torch.cuda.synchronize()
                mm = me = 0
                if ref is not None:
                    mm = (~torch.isclose(O.float(), ref, rtol=1e-1,
                          atol=1e-1)).float().mean().item()
                    me = (O.float() - ref).abs().max().item()
                if mm > 0.04:
                    _L(f"[{bs},{kvlen}] b{bn}w{nw}s{sched}/sp{ns} WRONG "
                       f"mm={mm:.2%} me={me:.3f}")
                    continue
                tc = _cold(_h)
                _L(f"[{bs},{kvlen}] b{bn}w{nw}s{sched}/sp{ns} cold={tc:.1f} "
                   f"mm={mm:.3%} me={me:.3f}")
                if tc < best_hip[0]:
                    best_hip = (tc, (bn, nw, sched, ns, mod))
            except Exception as e:
                _L(f"b{bn}w{nw}s{sched}/sp{ns} FAIL:{str(e)[:150]}")

    use_hip = best_hip[1] is not None and best_hip[0] < t_ai * 0.97
    _L(f"[{bs},{kvlen}] ai={t_ai:.1f} hip={best_hip[0]:.1f}"
       f"@{best_hip[1][:4] if best_hip[1] else None} -> "
       f"{'HIP' if use_hip else 'ai'} [{time.time()-t0:.1f}s]")

    if use_hip:
        bn, nw, sched, ns, mod = best_hip[1]
        po, pm, pl = PO[ns], PM[ns], PL[ns]; scs0 = [None]
        def hh(q, kvd, qo, kv):
            f4, sc = kvd["mxfp4"]
            f4u = f4.view(torch.uint8).reshape(-1, DQK // 2)
            scu = sc.view(torch.uint8)
            if scs0[0] is None: scs0[0] = scu.stride(0)
            mod.run_s1(q, f4u, scu, kv, po, pm, pl, bs, ns, scs0[0], SM_SCALE)
            mod.run_rd(po, pm, pl, O, bs, ns)
            return O
        return hh
    if ah is None:
        ah = _make_ai(bs, kvlen, q, qo_i, kv_i, dev)
    def ha(q, kvd, qo, kv):
        kf8, ks8 = kvd["fp8"]
        return ah(q, kf8, ks8, qo, kv)
    return ha


def custom_kernel(data):
    q, kvd, qo_i, kv_i, cfg = data
    bs, kvlen = cfg["batch_size"], cfg["kv_seq_len"]
    k = (bs, kvlen)
    h = _cache.get(k)
    if h is None:
        h = _build(bs, kvlen, data, q.device)
        _cache[k] = h
    return h(q, kvd, qo_i, kv_i)
scrolls · 504 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