Skip to content
KernelIndex
Search⌘K

submission 753968

vuxml · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

fused-epilogueSTILL +0.95us vs 2L. POSTMORTEM_1L: epilogue floor = Ctr-serialize(~530ns)
num-warps = 8num_warps=8,num_stages=2,matrix_instr_nonkdim=16)
shared-memoryv32 BUG: __shared__ int _arr (4B static LDS) + hipFuncSetAttribute(160KB
stages = 2num_warps=8,num_stages=2,matrix_instr_nonkdim=16)
tile-m = 16BM=16,BN=32,BK=BK,EN=(n%32==0),
tile-n = 32BM=16,BN=32,BK=BK,EN=(n%32==0),
vector-width = int4void hw_quant32(const int4* __restrict__ a4,i32x4& o,int& e8){

Kernel source

submission_v34_1Lnt.py1283 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
v34: MR=2 m=256 RACE + askx1L hardened. See INTERROGATE.md.

v33b RESULT: askx1L CORRECT (chk=0.000%, zero-fence memory model proven) but
  STILL +0.95us vs 2L. POSTMORTEM_1L: epilogue floor = Ctr-serialize(~530ns)
  + winner-wait(~400ns) = 930ns > cast-launch(~560ns). Structural. MLA same
  verdict via different physics (fence tax). s_sleep dead too (removed).

INTERROGATE Q3 REOPENED MR=2:
  HANDOFF §5 "alds MR>1 dead" was from v13/v15 at m=64 (grid 112=44% CU).
  At m=256 MR=2: grid 384->192 = 75% CU. DIFFERENT REGIME.
  B-L2 re-reads: 15x (m_tiles=16) -> 7x (m_tiles=8). -18.9MB / 20TB/s =
  -945ns. +141ns BW loss (75% grid). Net MODEL -800ns on m=256.
  RISK: VGPR. acc[2]+ab[KU][2] could push past 85 (2-wave) or 128 (1-wave).
  RACE auto-detects: _build runs both, _tcold picks, prints VGPR.

v34 CHANGES:
  A. ⭐ MR=2 a3x instance for m=256 + _build race. If VGPR<85 & faster: WIN.
     Pure upside (race picks MR=1 if MR=2 loses). +1 kernel instance.
  B. askx1L: NT-load+NT-store (pipelined, no atomicExch dep-chain). Info
     probe only -- HARD stays aldskx (proven 10.8). Auto-switch if 1L wins
     by >0.3us (unlikely per POSTMORTEM but test is free).
  C. askx1L: OOB waves branch-skip K-loop (if(vn){...}). ~249ns pipe-time
     on tail ntg=16. INTERROGATE Q2. 2L already has this (early return).
  D. s_sleep removed. a3x back to v29-identical.

ONE BENCH: MR=2 VGPR + MR=1vsMR=2 head-to-head + 1Lvs2L + all 6/6 chk.

v32 RESULT: bench GM 7.89 (best ever). askx1L hipErrorInvalidValue -> fallback
  picked aldskx(8,7,4) -> m=16=10.8 (=v31). Physics probe: a3x m=64 cold=10.09
  warm=7.00 dHBM=+3.09us -- PHYSICS_a3x model EXACT. a3x is pure-HBM above 7us
  floor; stop optimizing compute. a3xu1 mixed VGPR -> dead end.

v32 BUG: __shared__ int _arr (4B static LDS) + hipFuncSetAttribute(160KB
  dynamic) > device max -> hipErrorInvalidValue, sticky -> sync raised on 1st
  call. 2nd call skips attr -> runs but static/dynamic LDS alias -> 35.7us.

v33b ZERO-FENCE (see BRAINSTORM_v33.md):
  MLA v41b postmortem: __threadfence = buffer_wbinvl1_vol = 0.12-0.21us PER WG.
  At our 136 WGs x 2 fences = ~25us tax. THAT was v32's 35.7us, not (only) LDS.

  BUT MM DOESN'T NEED FENCES. Cf-writes are atomicAdd = GLC = L1-bypass =
  L2-direct, complete-on-return. Program-ordered before Ctr-atomicAdd (also
  GLC). Winner observes Ctr=SK-1 at L2 -> all siblings' Ctr++ done -> their
  Cf-adds done (each ordered-before each Ctr++). Winner's Cf-load: never in
  L1 (GLC bypassed), misses to L2 -> sees all 8 partials. Correct w/o fence.

  Only L2-commit needed: Cf/Ctr zero via atomicExch (GLC). Next call's atomic
  GLC-reads L2 directly -> sees the zero. ~750ns epilogue vs ~1520ns 2nd-launch.

  (a) NO __threadfence (MLA's trap, our no-op)
  (b) Cf/Ctr zero via atomicExch (GLC, L2-committed before next call)
  (c) _arr in _sh[0] (LDS reuse after K-loop, no static __shared__)
  (d) NO hipFuncSetAttribute (LDS 7616B < 64KB default)
  (e) hipGetLastError() pre-launch clears sticky
  (f) s_sleep(2) at a3x entry (IDEA D: let HBM write-buffer drain post-flush)
  (g) LEAN: stripped a3xu1/twarm probes

Expected: m=16 ~10.0 (-0.8). GM ~7.97 neutral, ~7.90 good roll (= guojun rk5).
Fallback: askx1L fail -> aldskx (v32-proven 10.8 ranked) -> aldsk. LB-safe.
"""
import os, sys, time
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
import warnings; warnings.filterwarnings("ignore")
import torch
import triton
import triton.language as tl

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


_HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <cstdint>

typedef int   i32x4 __attribute__((ext_vector_type(4)));
typedef int   i32x8 __attribute__((ext_vector_type(8)));
typedef float f32x4 __attribute__((ext_vector_type(4)));
typedef __hip_bfloat16 bf16;

__device__ __forceinline__ uint32_t f2u(float x){
  union{float f;uint32_t u;}c;c.f=x;return c.u;}
__device__ __forceinline__ float e8f(uint8_t e){
  union{uint32_t u;float f;}c;c.u=(uint32_t)e<<23;return c.f;}
#define QCV(o,a,b,s,bs) __builtin_amdgcn_cvt_scalef32_pk_fp4_f32((o),(a),(b),(s),(bs))

__device__ __forceinline__ i32x8 w8(i32x4 x){
  i32x8 r={0,0,0,0,0,0,0,0};r[0]=x[0];r[1]=x[1];r[2]=x[2];r[3]=x[3];return r;}

__device__ __forceinline__
void hw_quant32(const int4* __restrict__ a4,i32x4& o,int& e8){
  bf16 ab[32] __attribute__((aligned(16)));
  *reinterpret_cast<int4*>(&ab[ 0])=a4[0];
  *reinterpret_cast<int4*>(&ab[ 8])=a4[1];
  *reinterpret_cast<int4*>(&ab[16])=a4[2];
  *reinterpret_cast<int4*>(&ab[24])=a4[3];
  float v[32];float amax=0.f;
  #pragma unroll
  for(int i=0;i<32;++i){v[i]=(float)ab[i];
    float t=__builtin_fabsf(v[i]);amax=t>amax?t:amax;}
  uint32_t au=(f2u(amax)+0x200000u)&0xFF800000u;
  int su=au?(int)((au>>23)&0xFFu)-129:-127;
  su=su<-127?-127:(su>127?127:su);e8=su+127;
  float bsc=e8f((uint8_t)e8);
  int w0=0,w1=0,w2=0,w3=0;
  w0=QCV(w0,v[ 0],v[ 1],bsc,0);w0=QCV(w0,v[ 2],v[ 3],bsc,1);
  w0=QCV(w0,v[ 4],v[ 5],bsc,2);w0=QCV(w0,v[ 6],v[ 7],bsc,3);
  w1=QCV(w1,v[ 8],v[ 9],bsc,0);w1=QCV(w1,v[10],v[11],bsc,1);
  w1=QCV(w1,v[12],v[13],bsc,2);w1=QCV(w1,v[14],v[15],bsc,3);
  w2=QCV(w2,v[16],v[17],bsc,0);w2=QCV(w2,v[18],v[19],bsc,1);
  w2=QCV(w2,v[20],v[21],bsc,2);w2=QCV(w2,v[22],v[23],bsc,3);
  w3=QCV(w3,v[24],v[25],bsc,0);w3=QCV(w3,v[26],v[27],bsc,1);
  w3=QCV(w3,v[28],v[29],bsc,2);w3=QCV(w3,v[30],v[31],bsc,3);
  o=(i32x4){w0,w1,w2,w3};
}

// ═══════ fgemm_alds3x: v29's EXACT-SHAPE branch-free (the only alds we need) ═══════
// When MX=NX=1 (bench m=64/m=256): zero bounds checks. Straight-line barrier->store.
// K = K128T*128 constexpr -> ptr strides are shifts. K-loop fully unrolls.
template<int WAVES,int M_REP,int KU,int KP,int BSC,int NTW,int K128T,
         int M_EXACT,int NT_EXACT>
__global__ __launch_bounds__(WAVES*64)
void fgemm_alds3x(
    const bf16* __restrict__ A,const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bsc,bf16* __restrict__ C,
    int M,int N,long sn8,int NT)
{
  constexpr int  K128 = K128T;
  constexpr int  K32  = K128T * 4;
  constexpr long K    = (long)K128T * 128;
  constexpr long Ald_stride = (long)K128T * 1024;
  constexpr long Asd_stride = (long)K128T * 64;
  constexpr long Asd_base   = (long)M_REP * Ald_stride;
  constexpr int  nthr       = WAVES * 64;
  constexpr int  ngrp       = M_REP * 16 * K32;

  const int tid=threadIdx.x,L=tid&63,w=tid>>6;
  const int m16=L&15,kg=L>>4;const int bid=blockIdx.x;
  const int m_tile=bid/NTW, ntg=bid%NTW;
  const int n_tile=ntg*WAVES+w;

  extern __shared__ uint8_t _sh[];
  uint8_t* Ald=_sh; uint8_t* Asd=_sh+Asd_base;
  {
    const int m_base=m_tile*M_REP*16;
    #pragma unroll
    for(int g=tid; g<ngrp; g+=nthr){
      const int r=g/K32; const int kb=g%K32;
      const int r_tile=r>>4; const int r16=r&15;
      const int k128s=kb>>2; const int kg4=kb&3;
      const int Lw=kg4*16+r16;
      const int m=m_base+r;
      i32x4 o; int e8;
      if constexpr(M_EXACT){
        const bf16* Ap=A+(long)m*K+(long)kb*32;
        int4 ai[4];
        ai[0]=*reinterpret_cast<const int4*>(Ap);
        ai[1]=*reinterpret_cast<const int4*>(Ap+8);
        ai[2]=*reinterpret_cast<const int4*>(Ap+16);
        ai[3]=*reinterpret_cast<const int4*>(Ap+24);
        hw_quant32(ai,o,e8);
      } else {
        o=(i32x4){0,0,0,0}; e8=0;
        if(m<M){
          const bf16* Ap=A+(long)m*K+(long)kb*32;
          int4 ai[4];
          ai[0]=*reinterpret_cast<const int4*>(Ap);
          ai[1]=*reinterpret_cast<const int4*>(Ap+8);
          ai[2]=*reinterpret_cast<const int4*>(Ap+16);
          ai[3]=*reinterpret_cast<const int4*>(Ap+24);
          hw_quant32(ai,o,e8);
        }
      }
      *reinterpret_cast<i32x4*>(Ald+(long)r_tile*Ald_stride+(long)k128s*1024+Lw*16)=o;
      Asd[(long)r_tile*Asd_stride+(long)k128s*64+Lw]=(uint8_t)e8;
    }
  }
  __syncthreads();

  if constexpr(!NT_EXACT){
    if(n_tile>=NT) return;
  }

  const long n_col=(long)n_tile*16+m16;
  const long bsc_row=(n_col>>5)*(sn8*256)+(n_col&15)*4+((n_col>>4)&1)+(long)kg*64;
  const uint8_t* Bsc_r=Bsc+bsc_row;
  const uint8_t* Bsh_L=Bsh+(long)n_tile*K*8+L*16;
  const uint8_t* Ald_L=Ald+L*16;
  const uint8_t* Asd_L=Asd+L;
  f32x4 acc[M_REP];
  #pragma unroll
  for(int r=0;r<M_REP;++r) acc[r]=(f32x4){0,0,0,0};

  int mmsk[M_REP];
  if constexpr(!M_EXACT){
    #pragma unroll
    for(int r=0;r<M_REP;++r){
      int m=(m_tile*M_REP+r)*16+m16;mmsk[r]=(m<M)?0xFF:0;}
  }

  int bsc_pk[KP>0?KP:1];
  if constexpr(BSC){
    #pragma unroll
    for(int p=0;p<KP;++p)bsc_pk[p]=*reinterpret_cast<const int*>(Bsc_r+(long)p*256);
  }

  #pragma unroll
  for(int ks=0;ks<K128;ks+=KU){
    i32x4 bb[KU];int bsv[KU];i32x4 ab[KU][M_REP];int asv[KU][M_REP];
    #pragma unroll
    for(int u=0;u<KU;++u){const int ksi=ks+u;
      bb[u]=*reinterpret_cast<const i32x4*>(Bsh_L+(long)ksi*1024);
      if constexpr(BSC){
        bsv[u]=(bsc_pk[ksi>>1]>>((ksi&1)*16))&0xFF;
      } else {
        bsv[u]=(int)Bsc_r[(long)(ksi>>1)*256+(ksi&1)*2];
      }
      #pragma unroll
      for(int r=0;r<M_REP;++r){
        ab[u][r]=*reinterpret_cast<const i32x4*>(Ald_L+(long)r*Ald_stride+(long)ksi*1024);
        if constexpr(M_EXACT){
          asv[u][r]=(int)Asd_L[(long)r*Asd_stride+(long)ksi*64];
        } else {
          asv[u][r]=(int)Asd_L[(long)r*Asd_stride+(long)ksi*64]&mmsk[r];
        }
      }}
    __builtin_amdgcn_sched_barrier(0);
    #pragma unroll
    for(int u=0;u<KU;++u){i32x8 b8=w8(bb[u]);
      #pragma unroll
      for(int r=0;r<M_REP;++r)
        acc[r]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            w8(ab[u][r]),b8,acc[r],4,4,0,asv[u][r],0,bsv[u]);}}

  #pragma unroll
  for(int r=0;r<M_REP;++r)
    #pragma unroll
    for(int i=0;i<4;++i){
      int mo=(m_tile*M_REP+r)*16+kg*4+i;
      if constexpr(M_EXACT){
        C[(long)mo*N+n_col]=(bf16)acc[r][i];
      } else {
        if(mo<M) C[(long)mo*N+n_col]=(bf16)acc[r][i];
      }
    }
}

// ═══ alds3: RUNTIME fallback for secret shapes without compiled (NTW,K128) pair ═══
// v18's original. Only 2 instances compiled: W=8,4 MR=1 KU=4 (covers anything).
template<int WAVES,int M_REP,int KU>
__global__ __launch_bounds__(WAVES*64)
void fgemm_alds3(
    const bf16* __restrict__ A,const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bsc,bf16* __restrict__ C,
    int M,int N,int K,long sn8,int NT)
{
  const int tid=threadIdx.x,L=tid&63,w=tid>>6;
  const int m16=L&15,kg=L>>4;const int bid=blockIdx.x;
  const int NTW=(NT+WAVES-1)/WAVES;
  const int m_tile=bid/NTW, ntg=bid%NTW;
  const int n_tile=ntg*WAVES+w;const bool vn=(n_tile<NT);
  const int K32=K>>5, K128=K>>7;
  const long Ald_stride=(long)K128*1024;
  const long Asd_stride=(long)K128*64;
  const long Asd_base=(long)M_REP*Ald_stride;
  extern __shared__ uint8_t _sh[];
  uint8_t* Ald=_sh; uint8_t* Asd=_sh+Asd_base;
  { const int nthr=WAVES*64;
    const int ngrp=M_REP*16*K32;
    const int m_base=m_tile*M_REP*16;
    for(int g=tid; g<ngrp; g+=nthr){
      const int r=g/K32; const int kb=g%K32;
      const int r_tile=r>>4; const int r16=r&15;
      const int k128s=kb>>2; const int kg4=kb&3;
      const int Lw=kg4*16+r16;
      const int m=m_base+r;
      i32x4 o={0,0,0,0}; int e8=0;
      if(m<M){
        const bf16* Ap=A+(long)m*K+(long)kb*32;
        int4 ai[4];
        ai[0]=*reinterpret_cast<const int4*>(Ap);
        ai[1]=*reinterpret_cast<const int4*>(Ap+8);
        ai[2]=*reinterpret_cast<const int4*>(Ap+16);
        ai[3]=*reinterpret_cast<const int4*>(Ap+24);
        hw_quant32(ai,o,e8);
      }
      *reinterpret_cast<i32x4*>(Ald+(long)r_tile*Ald_stride+(long)k128s*1024+Lw*16)=o;
      Asd[(long)r_tile*Asd_stride+(long)k128s*64+Lw]=(uint8_t)e8;
    } }
  __syncthreads();
  if(!vn) return;
  const long n_col=(long)n_tile*16+m16;
  const long bsc_row=(n_col>>5)*(sn8*256)+(n_col&15)*4+((n_col>>4)&1)+(long)kg*64;
  const uint8_t* Bsc_r=Bsc+bsc_row;
  const uint8_t* Bsh_L=Bsh+(long)n_tile*(long)K*8+L*16;
  const uint8_t* Ald_L=Ald+L*16;const uint8_t* Asd_L=Asd+L;
  f32x4 acc[M_REP];
  #pragma unroll
  for(int r=0;r<M_REP;++r) acc[r]=(f32x4){0,0,0,0};
  int mmsk[M_REP];
  #pragma unroll
  for(int r=0;r<M_REP;++r){
    int m=(m_tile*M_REP+r)*16+m16;mmsk[r]=(m<M)?0xFF:0;}
  for(int ks=0;ks<K128;ks+=KU){
    i32x4 bb[KU];int bsv[KU];i32x4 ab[KU][M_REP];int asv[KU][M_REP];
    #pragma unroll
    for(int u=0;u<KU;++u){const int ksi=ks+u;
      bb[u]=*reinterpret_cast<const i32x4*>(Bsh_L+(long)ksi*1024);
      bsv[u]=(int)Bsc_r[(long)(ksi>>1)*256+(ksi&1)*2];
      #pragma unroll
      for(int r=0;r<M_REP;++r){
        ab[u][r]=*reinterpret_cast<const i32x4*>(Ald_L+(long)r*Ald_stride+(long)ksi*1024);
        asv[u][r]=(int)Asd_L[(long)r*Asd_stride+(long)ksi*64]&mmsk[r];}}
    __builtin_amdgcn_sched_barrier(0);
    #pragma unroll
    for(int u=0;u<KU;++u){i32x8 b8=w8(bb[u]);
      #pragma unroll
      for(int r=0;r<M_REP;++r)
        acc[r]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            w8(ab[u][r]),b8,acc[r],4,4,0,asv[u][r],0,bsv[u]);}}
  #pragma unroll
  for(int r=0;r<M_REP;++r)
    #pragma unroll
    for(int i=0;i<4;++i){int mo=(m_tile*M_REP+r)*16+kg*4+i;
      if(mo<M)C[(long)mo*N+n_col]=(bf16)acc[r][i];}
}

// ═══ alds_skx_1L: 1-LAUNCH split-K with winner-casts epilogue (v32) ═══
// bid mapping FLIPPED: ntg=bid/SK, pk=bid%SK. pk-siblings are consecutive bids
// -> dispatched together -> finish together -> winner's Cf-reads are L2-hot.
// After atomicAdd partials to Cf: threadfence(release), atomicInc(Ctr[ntg]).
// Winner (old==SK-1): threadfence(acquire), cast Cf->C, zero Cf+Ctr in-place.
// Between calls there IS torch.cuda.synchronize() so Ctr/Cf reset is safe.
template<int WAVES,int KU,int NTW,int K128T,int SK,int M_EXACT>
__global__ __launch_bounds__(WAVES*64)
void fgemm_alds_skx_1L(
    const bf16* __restrict__ A,const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bsc,float* __restrict__ Cf,
    bf16* __restrict__ C,int* __restrict__ Ctr,
    int M,int N,long sn8,int NT)
{
  static_assert(K128T % SK == 0, "skx requires exact K split");
  constexpr int  K128   = K128T;
  constexpr long K      = (long)K128T * 128;
  constexpr int  Kps128 = K128T / SK;
  constexpr int  nsl    = Kps128;
  constexpr int  nthr   = WAVES * 64;
  constexpr int  nkg    = nsl * 4;
  constexpr int  ngrp   = 16 * nkg;
  constexpr long Asd_base = (long)nsl * 1024;

  const int tid=threadIdx.x,L=tid&63,w=tid>>6;
  const int m16=L&15,kg=L>>4;const int bid=blockIdx.x;
  // FLIPPED: pk-siblings consecutive. SK template -> bid%SK = bid&7 at SK=8.
  const int ntg=bid/SK, pk=bid%SK;
  const int n_tile=ntg*WAVES+w;
  const int ks_lo=pk*Kps128;
  extern __shared__ uint8_t _sh[];
  uint8_t* Ald=_sh;uint8_t* Asd=_sh+Asd_base;
  {
    const int kb_lo=ks_lo*4;
    #pragma unroll
    for(int g=tid;g<ngrp;g+=nthr){
      const int r=g/nkg;const int kbs=g%nkg;const int kb=kb_lo+kbs;
      const int k128s=kbs>>2;const int kg4=kbs&3;const int Lw=kg4*16+r;
      i32x4 o;int e8;
      if constexpr(M_EXACT){
        const bf16* Ap=A+(long)r*K+(long)kb*32;int4 ai[4];
        ai[0]=*reinterpret_cast<const int4*>(Ap);
        ai[1]=*reinterpret_cast<const int4*>(Ap+8);
        ai[2]=*reinterpret_cast<const int4*>(Ap+16);
        ai[3]=*reinterpret_cast<const int4*>(Ap+24);
        hw_quant32(ai,o,e8);
      } else {
        o=(i32x4){0,0,0,0};e8=0;
        if(r<M){
          const bf16* Ap=A+(long)r*K+(long)kb*32;int4 ai[4];
          ai[0]=*reinterpret_cast<const int4*>(Ap);
          ai[1]=*reinterpret_cast<const int4*>(Ap+8);
          ai[2]=*reinterpret_cast<const int4*>(Ap+16);
          ai[3]=*reinterpret_cast<const int4*>(Ap+24);
          hw_quant32(ai,o,e8);
        }
      }
      *reinterpret_cast<i32x4*>(Ald+(long)k128s*1024+Lw*16)=o;
      Asd[(long)k128s*64+Lw]=(uint8_t)e8;
    }
  }
  __syncthreads();
  const bool vn=(n_tile<NT);
  // OOB waves (vn=false, wave-uniform) skip K-loop entirely (s_cbranch, ~1cy).
  // Saves ~249ns pipe-time on ntg=16's tail (4 OOB waves × 7 mfma × 16cy).
  // 2L aldskx already has `if(!vn)return` post-barrier; 1L can't return
  // (Ctr++ needed) but CAN skip the work. (INTERROGATE Q2.)
  if(vn){
    const long n_col=(long)n_tile*16+m16;
    const long bsc_row=(n_col>>5)*(sn8*256)+(n_col&15)*4+((n_col>>4)&1)+(long)kg*64;
    const uint8_t* Bsc_r=Bsc+bsc_row;
    const uint8_t* Bsh_L=Bsh+(long)n_tile*K*8+L*16;
    const uint8_t* Ald_L=Ald+L*16;const uint8_t* Asd_L=Asd+L;
    f32x4 acc={0,0,0,0};
    int mmsk;if constexpr(M_EXACT){mmsk=0xFF;}else{mmsk=(m16<M)?0xFF:0;}
    #pragma unroll
    for(int ks_l=0;ks_l<nsl;ks_l+=KU){
      i32x4 bb[KU];int bsv[KU];i32x4 ab[KU];int asv[KU];
      #pragma unroll
      for(int u=0;u<KU;++u){const int ksi_l=ks_l+u,ksi_g=ks_lo+ksi_l;
        bb[u]=*reinterpret_cast<const i32x4*>(Bsh_L+(long)ksi_g*1024);
        bsv[u]=(int)Bsc_r[(long)(ksi_g>>1)*256+(ksi_g&1)*2];
        ab[u]=*reinterpret_cast<const i32x4*>(Ald_L+(long)ksi_l*1024);
        if constexpr(M_EXACT){asv[u]=(int)Asd_L[(long)ksi_l*64];}
        else{asv[u]=(int)Asd_L[(long)ksi_l*64]&mmsk;}
      }
      __builtin_amdgcn_sched_barrier(0);
      #pragma unroll
      for(int u=0;u<KU;++u)
        acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
            w8(ab[u]),w8(bb[u]),acc,4,4,0,asv[u],0,bsv[u]);}
    #pragma unroll
    for(int i=0;i<4;++i){int mo=kg*4+i;
      if constexpr(M_EXACT){atomicAdd(&Cf[(long)mo*N+n_col],acc[i]);}
      else{if(mo<M)atomicAdd(&Cf[(long)mo*N+n_col],acc[i]);}}
  }
  // ── winner-casts epilogue (v34: pipelined NT-load, trust-eviction zero) ──
  // v33b: correct (chk=0) but +0.95us vs 2L. Culprit: atomicExch w/ USED
  // return = 4 sequential ~250cy RMW-read-backs/lane = ~560ns critical path.
  //
  // v34: DECOUPLE read from zero.
  //   READ: __builtin_nontemporal_load -> global_load glc slc. L1-bypass,
  //     L2-read (sees siblings' GLC-committed atomicAdds). 4 loads issued
  //     back-to-back (no deps) -> pipeline, amortized ~150cy total not 4x.
  //   ZERO: __builtin_nontemporal_store -> global_store glc slc. L1-bypass
  //     write-through to L2. Next call's atomicAdd (GLC) reads L2 -> sees 0.
  //     Fire-and-forget, zero wait. (Harness sync guarantees completion.)
  //
  // Ordering: s_waitcnt(0) drains Cf-atomicAdds, Ctr-atomicAdd (return used)
  //   ordered after. Winner sees Ctr=SK-1 -> siblings done -> NT-loads see L2.
  //
  // Cost: s_waitcnt(~50cy/WG all) + Ctr serialize(~400cy/ntg) + winner L2-load
  //   pipelined(~150cy) + winner NT-store fire-forget(~5cy) + winner wait
  //   on last-sibling(~500cy for slowest of 8 consecutive-bid). ~0.55us vs
  //   2L's ~1.52us. Projected NET -0.97. m=16 target 10.8 -> ~9.8.
  __builtin_amdgcn_s_waitcnt(0);            // drain Cf-atomicAdds to L2
  int* _arr = reinterpret_cast<int*>(_sh);  // A-LDS dead post-K-loop, reuse
  if(tid==0) _arr[0] = atomicAdd(&Ctr[ntg],1);
  __syncthreads();
  if(_arr[0] != SK-1) return;
  if(vn){
    const long n_col=(long)n_tile*16+m16;
    long ix[4]; float v[4];
    #pragma unroll
    for(int i=0;i<4;++i) ix[i]=(long)(kg*4+i)*N+n_col;
    #pragma unroll
    for(int i=0;i<4;++i) v[i]=__builtin_nontemporal_load(&Cf[ix[i]]);  // glc slc, pipelined
    #pragma unroll
    for(int i=0;i<4;++i){
      if constexpr(M_EXACT){
        C[ix[i]]=(bf16)v[i];
        __builtin_nontemporal_store(0.0f, &Cf[ix[i]]);    // glc slc, fire-forget
      } else { int mo=kg*4+i; if(mo<M){
        C[ix[i]]=(bf16)v[i];
        __builtin_nontemporal_store(0.0f, &Cf[ix[i]]);
      }}
    }
  }
  if(tid==0) __builtin_nontemporal_store(0, &Ctr[ntg]);
}

// ═══ alds_skx: TEMPLATE-EVERYTHING split-K (v31). m=16 bench + m=8 secret. ═══
// When K128%SK==0 (56%8==0): every pk gets EXACTLY Kps128 slices, nsl constexpr.
// K-loop bound=Kps128=7 at KU=7 -> single fully-unrolled iteration, branch-free.
template<int WAVES,int KU,int NTW,int K128T,int SK,int M_EXACT>
__global__ __launch_bounds__(WAVES*64)
void fgemm_alds_skx(
    const bf16* __restrict__ A,const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bsc,float* __restrict__ Cf,
    int M,int N,long sn8,int NT)
{
  static_assert(K128T % SK == 0, "skx requires exact K split");
  constexpr int  K128   = K128T;
  constexpr long K      = (long)K128T * 128;
  constexpr int  Kps128 = K128T / SK;
  constexpr int  nsl    = Kps128;
  constexpr int  nthr   = WAVES * 64;
  constexpr int  nkg    = nsl * 4;
  constexpr int  ngrp   = 16 * nkg;
  constexpr long Asd_base = (long)nsl * 1024;

  const int tid=threadIdx.x,L=tid&63,w=tid>>6;
  const int m16=L&15,kg=L>>4;const int bid=blockIdx.x;
  const int pk=bid/NTW, ntg=bid%NTW;
  const int n_tile=ntg*WAVES+w;
  const int ks_lo=pk*Kps128;
  extern __shared__ uint8_t _sh[];
  uint8_t* Ald=_sh;uint8_t* Asd=_sh+Asd_base;
  {
    const int kb_lo=ks_lo*4;
    #pragma unroll
    for(int g=tid;g<ngrp;g+=nthr){
      const int r=g/nkg;const int kbs=g%nkg;const int kb=kb_lo+kbs;
      const int k128s=kbs>>2;const int kg4=kbs&3;const int Lw=kg4*16+r;
      i32x4 o;int e8;
      if constexpr(M_EXACT){
        const bf16* Ap=A+(long)r*K+(long)kb*32;int4 ai[4];
        ai[0]=*reinterpret_cast<const int4*>(Ap);
        ai[1]=*reinterpret_cast<const int4*>(Ap+8);
        ai[2]=*reinterpret_cast<const int4*>(Ap+16);
        ai[3]=*reinterpret_cast<const int4*>(Ap+24);
        hw_quant32(ai,o,e8);
      } else {
        o=(i32x4){0,0,0,0};e8=0;
        if(r<M){
          const bf16* Ap=A+(long)r*K+(long)kb*32;int4 ai[4];
          ai[0]=*reinterpret_cast<const int4*>(Ap);
          ai[1]=*reinterpret_cast<const int4*>(Ap+8);
          ai[2]=*reinterpret_cast<const int4*>(Ap+16);
          ai[3]=*reinterpret_cast<const int4*>(Ap+24);
          hw_quant32(ai,o,e8);
        }
      }
      *reinterpret_cast<i32x4*>(Ald+(long)k128s*1024+Lw*16)=o;
      Asd[(long)k128s*64+Lw]=(uint8_t)e8;
    }
  }
  __syncthreads();
  if(n_tile>=NT)return;
  const long n_col=(long)n_tile*16+m16;
  const long bsc_row=(n_col>>5)*(sn8*256)+(n_col&15)*4+((n_col>>4)&1)+(long)kg*64;
  const uint8_t* Bsc_r=Bsc+bsc_row;
  const uint8_t* Bsh_L=Bsh+(long)n_tile*K*8+L*16;
  const uint8_t* Ald_L=Ald+L*16;const uint8_t* Asd_L=Asd+L;
  f32x4 acc={0,0,0,0};
  int mmsk;if constexpr(M_EXACT){mmsk=0xFF;}else{mmsk=(m16<M)?0xFF:0;}
  #pragma unroll
  for(int ks_l=0;ks_l<nsl;ks_l+=KU){
    i32x4 bb[KU];int bsv[KU];i32x4 ab[KU];int asv[KU];
    #pragma unroll
    for(int u=0;u<KU;++u){const int ksi_l=ks_l+u,ksi_g=ks_lo+ksi_l;
      bb[u]=*reinterpret_cast<const i32x4*>(Bsh_L+(long)ksi_g*1024);
      bsv[u]=(int)Bsc_r[(long)(ksi_g>>1)*256+(ksi_g&1)*2];
      ab[u]=*reinterpret_cast<const i32x4*>(Ald_L+(long)ksi_l*1024);
      if constexpr(M_EXACT){asv[u]=(int)Asd_L[(long)ksi_l*64];}
      else                 {asv[u]=(int)Asd_L[(long)ksi_l*64]&mmsk;}
    }
    __builtin_amdgcn_sched_barrier(0);
    #pragma unroll
    for(int u=0;u<KU;++u)
      acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
          w8(ab[u]),w8(bb[u]),acc,4,4,0,asv[u],0,bsv[u]);}
  #pragma unroll
  for(int i=0;i<4;++i){int mo=kg*4+i;
    if constexpr(M_EXACT){atomicAdd(&Cf[(long)mo*N+n_col],acc[i]);}
    else{if(mo<M)atomicAdd(&Cf[(long)mo*N+n_col],acc[i]);}}
}

// ═══ alds_sk: m=16 (v18 proven). 2 instances only. ═══
template<int WAVES,int KU>
__global__ __launch_bounds__(WAVES*64)
void fgemm_alds_sk(
    const bf16* __restrict__ A,const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bsc,float* __restrict__ Cf,
    int M,int N,int K,long sn8,int NT,int SK)
{
  const int tid=threadIdx.x,L=tid&63,w=tid>>6;
  const int m16=L&15,kg=L>>4;const int bid=blockIdx.x;
  const int NTW=(NT+WAVES-1)/WAVES;
  const int pk=bid/NTW, ntg=bid%NTW;
  const int n_tile=ntg*WAVES+w;const bool vn=(n_tile<NT);
  const int K32=K>>5, K128=K>>7;
  const int Kps128=(K128+SK-1)/SK;
  const int ks_lo=pk*Kps128;const int ks_hi=min(ks_lo+Kps128,K128);
  const int nsl=ks_hi-ks_lo;if(nsl<=0)return;
  extern __shared__ uint8_t _sh[];
  uint8_t* Ald=_sh;uint8_t* Asd=_sh+(long)nsl*1024;
  { const int nthr=WAVES*64;const int nkg=nsl*4;
    const int ngrp=16*nkg;const int kb_lo=ks_lo*4;
    for(int g=tid;g<ngrp;g+=nthr){
      const int r=g/nkg;const int kbs=g%nkg;const int kb=kb_lo+kbs;
      const int k128s=kbs>>2;const int kg4=kbs&3;const int Lw=kg4*16+r;
      i32x4 o={0,0,0,0};int e8=0;
      if(r<M){const bf16* Ap=A+(long)r*K+(long)kb*32;int4 ai[4];
        ai[0]=*reinterpret_cast<const int4*>(Ap);
        ai[1]=*reinterpret_cast<const int4*>(Ap+8);
        ai[2]=*reinterpret_cast<const int4*>(Ap+16);
        ai[3]=*reinterpret_cast<const int4*>(Ap+24);
        hw_quant32(ai,o,e8);}
      *reinterpret_cast<i32x4*>(Ald+(long)k128s*1024+Lw*16)=o;
      Asd[(long)k128s*64+Lw]=(uint8_t)e8;}}
  __syncthreads();if(!vn)return;
  const long n_col=(long)n_tile*16+m16;
  const long bsc_row=(n_col>>5)*(sn8*256)+(n_col&15)*4+((n_col>>4)&1)+(long)kg*64;
  const uint8_t* Bsc_r=Bsc+bsc_row;
  const uint8_t* Bsh_L=Bsh+(long)n_tile*(long)K*8+L*16;
  const uint8_t* Ald_L=Ald+L*16;const uint8_t* Asd_L=Asd+L;
  f32x4 acc={0,0,0,0};const int mmsk=(m16<M)?0xFF:0;
  for(int ks_l=0;ks_l<nsl;ks_l+=KU){
    i32x4 bb[KU];int bsv[KU];i32x4 ab[KU];int asv[KU];
    const int lim=min(KU,nsl-ks_l);
    #pragma unroll
    for(int u=0;u<KU;++u){const int ksi_l=ks_l+u,ksi_g=ks_lo+ksi_l;
      if(u<lim){bb[u]=*reinterpret_cast<const i32x4*>(Bsh_L+(long)ksi_g*1024);
        bsv[u]=(int)Bsc_r[(long)(ksi_g>>1)*256+(ksi_g&1)*2];
        ab[u]=*reinterpret_cast<const i32x4*>(Ald_L+(long)ksi_l*1024);
        asv[u]=(int)Asd_L[(long)ksi_l*64]&mmsk;
      }else{bb[u]=(i32x4){0,0,0,0};bsv[u]=0;ab[u]=(i32x4){0,0,0,0};asv[u]=0;}}
    __builtin_amdgcn_sched_barrier(0);
    #pragma unroll
    for(int u=0;u<KU;++u)
      acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
          w8(ab[u]),w8(bb[u]),acc,4,4,0,asv[u],0,bsv[u]);}
  #pragma unroll
  for(int i=0;i<4;++i){int mo=kg*4+i;
    if(mo<M)atomicAdd(&Cf[(long)mo*N+n_col],acc[i]);}
}

__global__ __launch_bounds__(256)
void cast_f32_bf16_z(const float* __restrict__ S,bf16* __restrict__ C,
                     float* __restrict__ Sz,long N){
  long g=(long)blockIdx.x*256+threadIdx.x;
  if(g<N){C[g]=(bf16)S[g];Sz[g]=0.0f;}
}

// ═══ fqn / fq: m<=32 (v9 proven). 2 instances total. ═══
template<int WAVES,int N_REP>
__global__ __launch_bounds__(WAVES*64)
void fgemm_fqn(
    const bf16* __restrict__ A,const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bsc,bf16* __restrict__ C,
    int M,int N,int K,long sn8,int NT)
{
  const int tid=threadIdx.x,L=tid&63,w=tid>>6;
  const int m16=L&15,kg=L>>4;const int bid=blockIdx.x;
  const int NTG=(NT+N_REP-1)/N_REP;
  const int m_tile=bid/NTG,ntg=bid%NTG;
  long ksz=((K/128+WAVES-1)/WAVES)*128;
  long k_lo=(long)w*ksz,k_hi=min(k_lo+ksz,(long)K);
  long n_col[N_REP];int vnm[N_REP];const uint8_t* Bsh_t[N_REP];
  #pragma unroll
  for(int nr=0;nr<N_REP;++nr){int nt=ntg*N_REP+nr;int v=(nt<NT);
    vnm[nr]=v?0xFF:0;long ntr=v?nt:0;
    n_col[nr]=ntr*16+m16;Bsh_t[nr]=Bsh+ntr*(long)K*8;}
  const int m_row=m_tile*16+m16;const bool vm=m_row<M;
  const long mrow=vm?m_row:0;
  f32x4 acc[N_REP];
  #pragma unroll
  for(int nr=0;nr<N_REP;++nr)acc[nr]=(f32x4){0,0,0,0};
  for(long k=k_lo;k<k_hi;k+=128){
    long kb_=(k>>5)*256+L*16,ks=(k>>5)+kg;
    i32x4 bb[N_REP];int bsv[N_REP];
    #pragma unroll
    for(int nr=0;nr<N_REP;++nr){
      bb[nr]=*reinterpret_cast<const i32x4*>(Bsh_t[nr]+kb_);
      long c=ks;
      bsv[nr]=(int)Bsc[(n_col[nr]>>5)*(sn8*256)+(n_col[nr]&15)*4+((n_col[nr]>>4)&1)
                      +(c>>3)*256+(c&3)*64+((c>>2)&1)*2]&vnm[nr];}
    const long kba=k+(long)kg*32;const bf16* Ap=A+mrow*K+kba;int4 ai[4];
    ai[0]=*reinterpret_cast<const int4*>(Ap);
    ai[1]=*reinterpret_cast<const int4*>(Ap+8);
    ai[2]=*reinterpret_cast<const int4*>(Ap+16);
    ai[3]=*reinterpret_cast<const int4*>(Ap+24);
    i32x4 a4;int a_sc;hw_quant32(ai,a4,a_sc);if(!vm)a_sc=0;
    i32x8 a8=w8(a4);
    #pragma unroll
    for(int nr=0;nr<N_REP;++nr)
      acc[nr]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
          a8,w8(bb[nr]),acc[nr],4,4,0,a_sc,0,bsv[nr]);}
  extern __shared__ float red[];
  #pragma unroll
  for(int nr=0;nr<N_REP;++nr)
    #pragma unroll
    for(int i=0;i<4;++i)red[((long)w*N_REP+nr)*256+L*4+i]=acc[nr][i];
  __syncthreads();if(w!=0)return;
  #pragma unroll
  for(int nr=0;nr<N_REP;++nr)
    #pragma unroll
    for(int i=0;i<4;++i){float s=0;
      #pragma unroll
      for(int ww=0;ww<WAVES;++ww)s+=red[((long)ww*N_REP+nr)*256+L*4+i];
      acc[nr][i]=s;}
  #pragma unroll
  for(int i=0;i<4;++i){int mo=m_tile*16+kg*4+i;if(mo>=M)continue;
    #pragma unroll
    for(int nr=0;nr<N_REP;++nr)
      if(vnm[nr])C[(long)mo*N+n_col[nr]]=(bf16)acc[nr][i];}
}

template<int WAVES>
__global__ __launch_bounds__(WAVES*64)
void fgemm_fq(
    const bf16* __restrict__ A,const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bsc,bf16* __restrict__ C,
    int M,int N,int K,long sn8,int NT)
{
  const int tid=threadIdx.x,L=tid&63,w=tid>>6;
  const int m16=L&15,kg=L>>4;const int bid=blockIdx.x;
  int m_tile=bid/NT,n_tile=bid%NT;
  long ksz=((K/128+WAVES-1)/WAVES)*128;
  long k_lo=(long)w*ksz,k_hi=min(k_lo+ksz,(long)K);
  const long n_col=(long)n_tile*16+m16;
  const uint8_t* Bsh_t=Bsh+(long)n_tile*(long)K*8;
  f32x4 acc={0,0,0,0};const int m_row=m_tile*16+m16;const bool vm=m_row<M;
  for(long k=k_lo;k<k_hi;k+=128){
    i32x4 b4=*reinterpret_cast<const i32x4*>(Bsh_t+(k>>5)*256+L*16);
    long c=(k>>5)+kg;
    int b_sc=(int)Bsc[(n_col>>5)*(sn8*256)+(n_col&15)*4+((n_col>>4)&1)
                     +(c>>3)*256+(c&3)*64+((c>>2)&1)*2];
    const long kb=k+(long)kg*32;
    const bf16* Ap=A+(long)(vm?m_row:0)*K+kb;int4 ai[4];
    ai[0]=*reinterpret_cast<const int4*>(Ap);
    ai[1]=*reinterpret_cast<const int4*>(Ap+8);
    ai[2]=*reinterpret_cast<const int4*>(Ap+16);
    ai[3]=*reinterpret_cast<const int4*>(Ap+24);
    i32x4 a4;int a_sc;hw_quant32(ai,a4,a_sc);if(!vm)a_sc=0;
    acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
        w8(a4),w8(b4),acc,4,4,0,a_sc,0,b_sc);}
  extern __shared__ float red[];
  #pragma unroll
  for(int i=0;i<4;++i)red[(long)w*256+L*4+i]=acc[i];
  __syncthreads();if(w!=0)return;
  #pragma unroll
  for(int i=0;i<4;++i){float s=0;
    #pragma unroll
    for(int ww=0;ww<WAVES;++ww)s+=red[(long)ww*256+L*4+i];acc[i]=s;}
  #pragma unroll
  for(int i=0;i<4;++i){int mo=m_tile*16+kg*4+i;
    if(mo<M)C[(long)mo*N+n_col]=(bf16)acc[i];}
}

#include <torch/extension.h>

// ═══════ MINIMAL DISPATCH: only what's actually called ═══════
template<int W,int MR,int KU,int KP,int BSC,int NTW,int K128T,int MX,int NX>
static void _ga3x(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor C,int64_t M,int64_t N,int64_t sn8,int64_t NT){
  const int64_t MT=(M+16*MR-1)/(16*MR);
  const int64_t gx=MT*(int64_t)NTW;
  constexpr int64_t lds=(int64_t)MR*K128T*1088;
  static bool _s=false;if(!_s){(void)hipFuncSetAttribute(
    (const void*)fgemm_alds3x<W,MR,KU,KP,BSC,NTW,K128T,MX,NX>,
    hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}
  fgemm_alds3x<W,MR,KU,KP,BSC,NTW,K128T,MX,NX><<<dim3(gx),dim3(W*64),lds,0>>>(
    reinterpret_cast<const bf16*>(A.data_ptr()),
    Bsh.data_ptr<uint8_t>(),Bsc.data_ptr<uint8_t>(),
    reinterpret_cast<bf16*>(C.data_ptr()),(int)M,(int)N,sn8,(int)NT);
}

// a3x dispatch: TIGHT. 2 bench winners + secret coverage + 2 a3 fallback.
// Total alds kernels: ~12 instances (was ~100 in v29).
int64_t launch_alds3x(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,int64_t NT,
    int64_t W,int64_t MR,int64_t KU,int64_t BSC){
  int64_t K128=K>>7;
  int64_t NTW=(NT+W-1)/W;
  if((int64_t)MR*K128*1088>160*1024)return -2;
  if(K128%KU!=0)return -3;
  if(BSC&&(K128&1))return -4;
  if(K!=K128*128)return -7;
  int64_t MX=(M%(16*MR)==0)?1:0;
  int64_t NX=(NT%W==0)?1:0;
  #define DX(Ww,Rr,Uu,Bb,Nw,Kt,Mx,Nx) if(W==Ww&&MR==Rr&&KU==Uu&&BSC==Bb&& \
      NTW==Nw&&K128==Kt&&MX==Mx&&NX==Nx){ \
      _ga3x<Ww,Rr,Uu,(Bb)?(Kt)/2:1,Bb,Nw,Kt,Mx,Nx>(A,Bsh,Bsc,C,M,N,sn8,NT);return 0;}
  // BENCH WINNERS (from v29 LB run — exact configs)
  DX(8,1,4, 0, 56,16, 1,1);   // m=64 n=7168 k=2048 HARD
  DX(8,1,4, 1, 24,12, 1,1);   // m=256 n=3072 k=1536 HARD (bsc KP=6)
  // ⭐ MR=2 at m=256 (INTERROGATE Q3): halves B-L2 re-read (15x -> 7x). grid
  //   384->192 = 75% CU. v13/v15 tested MR>1 at m=64 (grid 44% -> crushed)
  //   but NEVER at m=256. Traffic model: -945ns B-L2. VGPR is the risk.
  //   M=256 % 32 == 0 -> M_EXACT=1. NT=192 % 8 == 0 -> NT_EXACT=1. LDS=2*13056.
  DX(8,2,4, 1, 24,12, 1,1);   // m=256 MR=2 -- RACED in _build
  // SECRET SHAPES (from v16fix test log, all passed max_err=0.0)
  //   m=64/m=16 n=3072 k=1536: same NTW=24 K128=12 as m=256 above. M exact, NT exact.
  DX(8,1,4, 0, 24,12, 1,1);   // no-bsc variant if bsc mismatch
  DX(4,1,12, 1, 48,12, 1,1);  // W=4 NTW=48 KU=12 (secret m=64 won with this in v16)
  DX(4,1,4, 1, 48,12, 1,1);   // W=4 KU=4 alt
  //   m=256 n=2880 k=512: NT=180 180%8=4 180%4=0. M=256 exact. W=4 NTW=45 NX=1.
  DX(4,1,4, 1, 45,4, 1,1);    // secret m=256
  DX(4,1,4, 0, 45,4, 1,1);
  //   m=8/m=16 n=3072 k=1536 where M not %16: MX=0 fallbacks
  DX(8,1,4, 1, 24,12, 0,1);DX(8,1,4, 0, 24,12, 0,1);
  // Generic m=32 coverage (a3x on small-M if NTW happens to match)
  DX(8,1,4, 1, 32,4, 1,1);DX(4,1,4, 1, 64,4, 1,1);
  #undef DX
  return -1;
}

// alds3 runtime fallback: 2 instances ONLY. Covers any secret shape a3x misses.
template<int W,int MR,int KU>
static void _ga3(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,int64_t NT){
  const int64_t K128=K>>7;const int64_t MT=(M+16*MR-1)/(16*MR);
  const int64_t NTW=(NT+W-1)/W;const int64_t gx=MT*NTW;
  const int64_t lds=(int64_t)MR*K128*1088;
  static bool _s=false;if(!_s){(void)hipFuncSetAttribute(
    (const void*)fgemm_alds3<W,MR,KU>,hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}
  fgemm_alds3<W,MR,KU><<<dim3(gx),dim3(W*64),lds,0>>>(
    reinterpret_cast<const bf16*>(A.data_ptr()),
    Bsh.data_ptr<uint8_t>(),Bsc.data_ptr<uint8_t>(),
    reinterpret_cast<bf16*>(C.data_ptr()),(int)M,(int)N,(int)K,sn8,(int)NT);
}
int64_t launch_alds3(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,int64_t NT,
    int64_t W,int64_t MR,int64_t KU){
  int64_t K128=K>>7;
  if((int64_t)MR*K128*1088>160*1024)return -2;if(K128%KU!=0)return -3;
  // ONLY 2 instances. KU=4 covers everything (K128 always mult of 4).
  if(W==8&&MR==1&&KU==4){_ga3<8,1,4>(A,Bsh,Bsc,C,M,N,K,sn8,NT);return 0;}
  if(W==4&&MR==1&&KU==4){_ga3<4,1,4>(A,Bsh,Bsc,C,M,N,K,sn8,NT);return 0;}
  return -1;
}

// aldskx_1L: 1-launch winner-casts. 2 instances (m=16 MX + m=8 non-MX).
// No hipFuncSetAttribute: Kps128*1088 = 7*1088 = 7616B << 64KB default.
// (v32 showed: setting 160KB with any static LDS -> hipErrorInvalidValue,
//  sticky -> next sync raises. Discarding return hides it. Just don't call.)
template<int W,int KU,int NTW,int K128T,int SK,int MX>
static void _gask1L(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor Cf,torch::Tensor C,torch::Tensor Ctr,
    int64_t M,int64_t N,int64_t sn8,int64_t NT){
  constexpr int64_t Kps128=K128T/SK;constexpr int64_t lds=Kps128*1088;
  static_assert(lds < 64*1024, "1L lds must fit default limit");
  const int64_t gx=(int64_t)NTW*SK;
  (void)hipGetLastError();
  fgemm_alds_skx_1L<W,KU,NTW,K128T,SK,MX><<<dim3(gx),dim3(W*64),lds,0>>>(
    reinterpret_cast<const bf16*>(A.data_ptr()),
    Bsh.data_ptr<uint8_t>(),Bsc.data_ptr<uint8_t>(),
    Cf.data_ptr<float>(),reinterpret_cast<bf16*>(C.data_ptr()),
    Ctr.data_ptr<int>(),(int)M,(int)N,sn8,(int)NT);
}
int64_t launch_alds_skx_1L(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor Cf,torch::Tensor C,torch::Tensor Ctr,
    int64_t M,int64_t N,int64_t K,int64_t sn8,
    int64_t NT,int64_t SK,int64_t W,int64_t KU){
  int64_t K128=K>>7;if(K!=K128*128)return -7;
  if(K128%SK!=0)return -5;
  int64_t Kps128=K128/SK;if(Kps128*1088>160*1024)return -2;
  if(Kps128%KU!=0)return -3;
  int64_t NTW=(NT+W-1)/W;
  int64_t MX=(M%16==0)?1:0;
  #define DL(Ww,Uu,Nw,Kt,Sk,Mx) if(W==Ww&&KU==Uu&&NTW==Nw&&K128==Kt&& \
      SK==Sk&&MX==Mx){ \
      _gask1L<Ww,Uu,Nw,Kt,Sk,Mx>(A,Bsh,Bsc,Cf,C,Ctr,M,N,sn8,NT);return 0;}
  // BENCH m=16: SK=8 won 2/3 on v32 LB (10.70/10.74 vs SK=7 10.86/10.83). SK=8 HARD.
  DL(8,7, 17,56,8, 1);
  // SECRET m=8 n=2112 k=7168
  DL(8,7, 17,56,8, 0);
  #undef DL
  return -1;
}

// aldskx: 2-3 templated instances (bench m=16 + secret m=8). Rest -> runtime aldsk.
template<int W,int KU,int NTW,int K128T,int SK,int MX>
static void _gaskx(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor Cf,int64_t M,int64_t N,int64_t sn8,int64_t NT){
  constexpr int64_t Kps128=K128T/SK;constexpr int64_t lds=Kps128*1088;
  const int64_t gx=(int64_t)NTW*SK;
  static bool _s=false;if(!_s){(void)hipFuncSetAttribute(
    (const void*)fgemm_alds_skx<W,KU,NTW,K128T,SK,MX>,
    hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}
  fgemm_alds_skx<W,KU,NTW,K128T,SK,MX><<<dim3(gx),dim3(W*64),lds,0>>>(
    reinterpret_cast<const bf16*>(A.data_ptr()),
    Bsh.data_ptr<uint8_t>(),Bsc.data_ptr<uint8_t>(),
    Cf.data_ptr<float>(),(int)M,(int)N,sn8,(int)NT);
}
int64_t launch_alds_skx(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor Cf,int64_t M,int64_t N,int64_t K,int64_t sn8,
    int64_t NT,int64_t SK,int64_t W,int64_t KU){
  int64_t K128=K>>7;if(K!=K128*128)return -7;
  if(K128%SK!=0)return -5;
  int64_t Kps128=K128/SK;if(Kps128*1088>160*1024)return -2;
  if(Kps128%KU!=0)return -3;
  int64_t NTW=(NT+W-1)/W;
  int64_t MX=(M%16==0)?1:0;
  #define DS(Ww,Uu,Nw,Kt,Sk,Mx) if(W==Ww&&KU==Uu&&NTW==Nw&&K128==Kt&& \
      SK==Sk&&MX==Mx){ \
      _gaskx<Ww,Uu,Nw,Kt,Sk,Mx>(A,Bsh,Bsc,Cf,M,N,sn8,NT);return 0;}
  // BENCH m=16: v32 LB PICK picked SK=8 KU=7 (10.70/10.74). HARD for fallback.
  DS(8,7, 17,56,8, 1);
  // SECRET m=8
  DS(8,7, 17,56,8, 0);
  #undef DS
  return -1;
}

// aldsk: 2 instances (W=8 KU=7 winner + W=4 KU=4 backup).
template<int W,int KU>
static void _galdsk(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor Cf,int64_t M,int64_t N,int64_t K,int64_t sn8,int64_t NT,int64_t SK){
  const int64_t K128=K>>7;const int64_t Kps128=(K128+SK-1)/SK;
  const int64_t NTW=(NT+W-1)/W;const int64_t gx=NTW*SK;
  const int64_t lds=Kps128*1088;
  static bool _s=false;if(!_s){(void)hipFuncSetAttribute(
    (const void*)fgemm_alds_sk<W,KU>,hipFuncAttributeMaxDynamicSharedMemorySize,160*1024);_s=true;}
  fgemm_alds_sk<W,KU><<<dim3(gx),dim3(W*64),lds,0>>>(
    reinterpret_cast<const bf16*>(A.data_ptr()),
    Bsh.data_ptr<uint8_t>(),Bsc.data_ptr<uint8_t>(),
    Cf.data_ptr<float>(),(int)M,(int)N,(int)K,sn8,(int)NT,(int)SK);
}
void go_cast(torch::Tensor Cf,torch::Tensor C,torch::Tensor Cfz,int64_t Ne){
  int64_t g=(Ne+255)/256;
  cast_f32_bf16_z<<<dim3(g),dim3(256),0,0>>>(
    Cf.data_ptr<float>(),reinterpret_cast<bf16*>(C.data_ptr()),
    Cfz.data_ptr<float>(),Ne);
}
int64_t launch_alds_sk(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor Cf,int64_t M,int64_t N,int64_t K,int64_t sn8,
    int64_t NT,int64_t SK,int64_t W,int64_t KU){
  int64_t K128=K>>7;int64_t Kps128=(K128+SK-1)/SK;
  if(Kps128*1088>160*1024)return -2;
  if(W==8&&KU==7){_galdsk<8,7>(A,Bsh,Bsc,Cf,M,N,K,sn8,NT,SK);return 0;}
  if(W==4&&KU==4){_galdsk<4,4>(A,Bsh,Bsc,Cf,M,N,K,sn8,NT,SK);return 0;}
  return -1;
}

// fqn/fq: wrappers (hipify mangles inlined chevrons after if-open-brace).
template<int W,int NR>
static void _gfqn(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,int64_t MT,int64_t NT){
  int64_t NTG=(NT+NR-1)/NR;int64_t gx=MT*NTG,lds=(int64_t)W*NR*256*4;
  fgemm_fqn<W,NR><<<dim3(gx),dim3(W*64),lds,0>>>(
    reinterpret_cast<const bf16*>(A.data_ptr()),
    Bsh.data_ptr<uint8_t>(),Bsc.data_ptr<uint8_t>(),
    reinterpret_cast<bf16*>(C.data_ptr()),(int)M,(int)N,(int)K,sn8,(int)NT);
}
template<int W>
static void _gfq(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,int64_t MT,int64_t NT){
  int64_t gx=MT*NT,lds=(int64_t)W*256*4;
  fgemm_fq<W><<<dim3(gx),dim3(W*64),lds,0>>>(
    reinterpret_cast<const bf16*>(A.data_ptr()),
    Bsh.data_ptr<uint8_t>(),Bsc.data_ptr<uint8_t>(),
    reinterpret_cast<bf16*>(C.data_ptr()),(int)M,(int)N,(int)K,sn8,(int)NT);
}
int64_t launch_fqn(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,
    int64_t MT,int64_t NT,int64_t W,int64_t NR){
  if(W==4&&NR==2){_gfqn<4,2>(A,Bsh,Bsc,C,M,N,K,sn8,MT,NT);return 0;}
  return -1;
}
int64_t launch_fq(torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,
    torch::Tensor C,int64_t M,int64_t N,int64_t K,int64_t sn8,
    int64_t MT,int64_t NT,int64_t W){
  if(W==4){_gfq<4>(A,Bsh,Bsc,C,M,N,K,sn8,MT,NT);return 0;}
  if(W==8){_gfq<8>(A,Bsh,Bsc,C,M,N,K,sn8,MT,NT);return 0;}
  return -1;
}

void probe(){
  hipFuncAttributes a;
  #define P(k,s) (void)hipFuncGetAttributes(&a,(const void*)k); \
    printf("[v34] %-42s VGPR=%3d spill=%zu bin=%zu\n", \
           s,a.numRegs,a.localSizeBytes,a.binaryVersion);
  P((fgemm_alds3x<8,1,4,1,0,56,16,1,1>), "a3x<N56,K16,XX>    m=64  MR=1");
  P((fgemm_alds3x<8,1,4,6,1,24,12,1,1>), "a3x<N24,K12,XX>    m=256 MR=1");
  P((fgemm_alds3x<8,2,4,6,1,24,12,1,1>), "a3x<N24,K12,XX,R2> m=256 MR=2");
  P((fgemm_alds3x<4,1,4,2,1,45,4,1,1>),  "a3x<N45,K4,XX>  secret");
  P((fgemm_alds3<8,1,4>),                "alds3<8,1,4>    fallback");
  P((fgemm_alds_skx_1L<8,7,17,56,8,1>),  "askx1L<N17,K56,S8,MX> m=16");
  P((fgemm_alds_skx_1L<8,7,17,56,8,0>),  "askx1L<N17,K56,S8,_ > m=8 sec");
  P((fgemm_alds_skx<8,7,17,56,8,1>),     "aldskx<N17,K56,S8,MX> fallback 2L");
  P((fgemm_alds_sk<8,7>),                "alds_sk<8,7>    fallback");
  P((fgemm_fqn<4,2>),                    "fqn<4,2>");
  P((fgemm_fq<4>),                       "fq<4>");
  #undef P
}
"""

_CPP = r"""
#include <torch/extension.h>
int64_t launch_alds3x(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
    int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t);
int64_t launch_alds3(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
    int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t);
int64_t launch_alds_skx_1L(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
    torch::Tensor,torch::Tensor,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,
    int64_t,int64_t);
int64_t launch_alds_skx(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
    int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t);
int64_t launch_alds_sk(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
    int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t);
void go_cast(torch::Tensor,torch::Tensor,torch::Tensor,int64_t);
int64_t launch_fqn(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
    int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t);
int64_t launch_fq(torch::Tensor,torch::Tensor,torch::Tensor,torch::Tensor,
    int64_t,int64_t,int64_t,int64_t,int64_t,int64_t,int64_t);
void probe();
"""

_hip = None
try:
    from torch.utils.cpp_extension import load_inline
    _t0 = time.time()
    _hip = load_inline(name="v34_1Lnt", cpp_sources=_CPP,
        cuda_sources=_HIP_SRC,
        functions=["launch_alds3x","launch_alds3",
                   "launch_alds_skx_1L","launch_alds_skx","launch_alds_sk",
                   "go_cast","launch_fqn","launch_fq","probe"],
        with_cuda=True,
        extra_cuda_cflags=["-O3","--offload-arch=gfx950","-ffast-math",
                           "-mllvm","-amdgpu-early-inline-all=true",
                           "-mllvm","-amdgpu-function-calls=false",
                           "-munsafe-fp-atomics"],
        verbose=False)
    _L(f"[v34] HIP compiled {time.time()-_t0:.1f}s"); _hip.probe()
except Exception as ex:
    import traceback
    _L(f"[v34] HIP FAIL: {type(ex).__name__}: {str(ex)[:2000]}")
    for ln in traceback.format_exc().splitlines()[-25:]:
        _L(f"   {ln[:200]}")


@triton.jit
def _sh_row(r,sn8):return (r//32)*(sn8*256)+(r%16)*4+(r//16)%2
@triton.jit
def _sh_col(c):return (c//8)*256+(c%4)*64+(c//4)%2*2
def _make_gemm_ref():
    from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
    @triton.jit
    def _k(A,Bq,Bsc,C,M,N,K,sAm,sBn,sn8,
           BM:tl.constexpr,BN:tl.constexpr,BK:tl.constexpr,EN:tl.constexpr):
        pid=tl.program_id(0);nn=tl.cdiv(N,BN);pm=pid//nn;pn=pid%nn
        om=pm*BM+tl.arange(0,BM);on=pn*BN+tl.arange(0,BN)
        o64=on.to(tl.int64);mm=om<M;mn=on<N
        rk=tl.arange(0,BK);r2=tl.arange(0,BK//2);r32=tl.arange(0,BK//32)
        bp=Bq+o64[:,None]*sBn+r2[None,:]
        br=_sh_row(o64,sn8);acc=tl.zeros((BM,BN),dtype=tl.float32)
        ap=A+om[:,None].to(tl.int64)*sAm+rk[None,:]
        for kk in tl.range(0,K,BK):
            ab=tl.load(ap,mask=mm[:,None],other=0.)
            af,asc=_mxfp4_quant_op(ab.to(tl.float32),BK,BM,32);ap+=BK
            if EN:
                bf=tl.load(bp);bs=tl.load(Bsc+br[:,None]+_sh_col(kk//32+r32)[None,:])
            else:
                bf=tl.load(bp,mask=mn[:,None],other=0)
                bs=tl.load(Bsc+br[:,None]+_sh_col(kk//32+r32)[None,:],mask=mn[:,None],other=0)
            acc=tl.dot_scaled(af,asc,"e2m1",tl.trans(bf),bs,"e2m1",acc);bp+=BK//2
        cm=mm[:,None]&mn[None,:]
        tl.store(C+om[:,None].to(tl.int64)*N+on[None,:],acc.to(tl.bfloat16),mask=cm)
    return _k
_gemm_ref=_make_gemm_ref()


# ═══════ HARDCODED (v34: aldskx proven; MR=1 vs MR=2 RACED on m=256) ═══════
# m=256 MR race (INTERROGATE Q3): halved m_tiles -> halved B-L2 re-read.
# Risk is VGPR (acc[2]+ab[KU][2]). Value is a LIST -> _build races them.
_HARD = {
    (4,   2880, 512 ): ("fqn", 4, 2),
    (16,  2112, 7168): ("aldskx", 8, 8, 7),
    (32,  4096, 512 ): ("fqn", 4, 2),
    (32,  2880, 512 ): ("fqn", 4, 2),
    (64,  7168, 2048): ("a3x", 8, 1, 4, 0),
    (256, 3072, 1536): [("a3x", 8, 1, 4, 1), ("a3x", 8, 2, 4, 1)],  # raced
}

def _pick_cands(m,n,k):
    NT=-(-n//16);MT16=-(-m//16);K128=k//128
    h=_HARD.get((m,n,k))
    # HARD config(s) go first; full cascade follows as fallback so PICK loop
    # has known-good options if the HARD cfg errors. List = race candidates.
    if isinstance(h,list): out=list(h)
    elif h: out=[h]
    else: out=[]
    if MT16>=2 and K128*1088<=160*1024 and K128%4==0:
        bsc_ok=(K128%2==0)and((K128//2)in(2,4,6,8))
        for W in(8,4):
            for KU in(4,12):
                if K128%KU!=0:continue
                if bsc_ok:out.append(("a3x",W,1,KU,1))
                out.append(("a3x",W,1,KU,0))
            out.append(("a3",W,1,4))   # runtime fallback
    # askx1L/aldskx/aldsk for small-M large-K
    if MT16<=2 and K128>=16:
        for SK in (8,7):
            if K128%SK==0:
                Kps=K128//SK
                for KU in (Kps,7,4):
                    if Kps%KU==0:
                        out.append(("askx1L",8,SK,KU))
                        out.append(("aldskx",8,SK,KU))
        out.append(("aldsk",8,8,7))
        out.append(("aldsk",4,8,4))
    # Ultimate fallbacks
    out.append(("fqn",4,2))
    out.append(("fq",min(8,max(4,K128))))
    # Drop dups keeping order; HARD cfg stays at index 0.
    seen=set(); uniq=[]
    for c in out:
        if c not in seen: seen.add(c); uniq.append(c)
    return uniq


_L2=torch.empty(512*1024*1024,dtype=torch.int8,device="cuda")
def _tcold(fn,n=7):
    for _ in range(2):fn()
    torch.cuda.synchronize()
    evs=[(torch.cuda.Event(True),torch.cuda.Event(True))for _ in range(n)]
    for e0,e1 in evs:_L2.zero_();e0.record();fn();e1.record()
    torch.cuda.synchronize()
    ts=sorted(e0.elapsed_time(e1)for e0,e1 in evs)
    return sum(ts[1:-1])*1000/(n-2)


_ST={}
def _build(data):
    A,B,Bq_,Bsh_,Bsc_=data
    m,k=A.shape;n=B.shape[0]
    sn=Bsc_.shape[1];sn8=sn//8;dev=A.device
    NT=-(-n//16);MT16=-(-m//16);K128=k//128;mn=m*n
    Bq=Bq_.view(torch.uint8);Bsh=Bsh_.view(torch.uint8);Bsc=Bsc_.view(torch.uint8)
    C=torch.empty((m,n),dtype=torch.bfloat16,device=dev)
    Cf=torch.zeros((m,n),dtype=torch.float32,device=dev)
    Cf2=torch.zeros((m,n),dtype=torch.float32,device=dev)
    _pp=[Cf,Cf2]
    NTW8=(NT+7)//8
    # askx1L uses its OWN Cf/Ctr (isolated from 2L ping-pong). Winner-reset
    # keeps them zero across calls; explicit .zero_() only on error recovery.
    Cf1L=torch.zeros((m,n),dtype=torch.float32,device=dev)
    Ctr=torch.zeros((max(NTW8,64),),dtype=torch.int32,device=dev)

    def _reset1L():
        Cf1L.zero_();Ctr.zero_();torch.cuda.synchronize()

    def _run(cfg,_A,_Bq,_Bsh,_Bsc):
        kind=cfg[0]
        if kind=="a3x":
            _,Wv,MR,KU,BSC=cfg
            rc=_hip.launch_alds3x(_A,_Bsh,_Bsc,C,m,n,k,sn8,NT,Wv,MR,KU,BSC)
            if rc!=0:raise RuntimeError(f"a3x rc={rc}")
            return C
        if kind=="a3":
            _,Wv,MR,KU=cfg
            rc=_hip.launch_alds3(_A,_Bsh,_Bsc,C,m,n,k,sn8,NT,Wv,MR,KU)
            if rc!=0:raise RuntimeError(f"a3 rc={rc}")
            return C
        if kind=="askx1L":
            _,Wv,SK,KU=cfg
            rc=_hip.launch_alds_skx_1L(_A,_Bsh,_Bsc,Cf1L,C,Ctr,m,n,k,sn8,NT,SK,Wv,KU)
            if rc!=0:raise RuntimeError(f"askx1L rc={rc}")
            return C
        if kind=="aldskx":
            _,Wv,SK,KU=cfg
            rc=_hip.launch_alds_skx(_A,_Bsh,_Bsc,_pp[0],m,n,k,sn8,NT,SK,Wv,KU)
            if rc!=0:raise RuntimeError(f"aldskx rc={rc}")
            _hip.go_cast(_pp[0],C,_pp[1],mn)
            _pp[0],_pp[1]=_pp[1],_pp[0]
            return C
        if kind=="aldsk":
            _,Wv,SK,KU=cfg
            rc=_hip.launch_alds_sk(_A,_Bsh,_Bsc,_pp[0],m,n,k,sn8,NT,SK,Wv,KU)
            if rc!=0:raise RuntimeError(f"aldsk rc={rc}")
            _hip.go_cast(_pp[0],C,_pp[1],mn)
            _pp[0],_pp[1]=_pp[1],_pp[0]
            return C
        if kind=="fqn":
            _,Wv,NR=cfg
            rc=_hip.launch_fqn(_A,_Bsh,_Bsc,C,m,n,k,sn8,MT16,NT,Wv,NR)
            if rc!=0:raise RuntimeError(f"fqn rc={rc}")
            return C
        if kind=="fq":
            _,Wv=cfg
            rc=_hip.launch_fq(_A,_Bsh,_Bsc,C,m,n,k,sn8,MT16,NT,Wv)
            if rc!=0:raise RuntimeError(f"fq rc={rc}")
            return C
        raise RuntimeError(f"?{cfg}")

    def _ref_tri(_A,_Bq,_Bsc):
        Cref=torch.empty_like(C)
        BK=min(512,k);gx=MT16*triton.cdiv(n,32)
        _gemm_ref[(gx,)](_A,_Bq,_Bsc,Cref,m,n,k,k,k//2,sn8,
            BM=16,BN=32,BK=BK,EN=(n%32==0),
            num_warps=8,num_stages=2,matrix_instr_nonkdim=16)
        return Cref

    rf=_ref_tri(A,Bq,Bsc).float();mag=rf.abs().mean().item()+1e-9

    cands=_pick_cands(m,n,k)
    _h=_HARD.get((m,n,k))
    is_hard=_h is not None
    is_race=isinstance(_h,list)
    MX=1 if m%16==0 else 0; NX8=1 if NT%8==0 else 0
    _L(f"\n[v34 m={m} n={n} k={k}] {'HARD' if is_hard else 'PICK'}: {len(cands)}c K128={K128} MX={MX} NX8={NX8}")

    if is_hard:
        cfg=cands[0]
        try:
            C.fill_(float('nan'))
            o=_run(cfg,A,Bq,Bsh,Bsc);torch.cuda.synchronize()
            err=((o.float()-rf).abs().mean()/mag).item()
            if err<5e-3:
                A2=torch.randn_like(A)
                rf2=_ref_tri(A2,Bq,Bsc).float()
                C.fill_(float('nan'))
                o2=_run(cfg,A2,Bq,Bsh,Bsc);torch.cuda.synchronize()
                e2=((o2.float()-rf2).abs().mean()/(rf2.abs().mean()+1e-9)).item()
                if e2<5e-3:
                    _L(f"  {cfg} chk={err:.3%},{e2:.3%} OK")
                    # HARD-list = race. Auto-detect (e.g. MR=2 VGPR regression).
                    if is_race:
                        try:
                            res=[(cfg, _tcold(lambda:_run(cfg,A,Bq,Bsh,Bsc),n=7))]
                            for rc in _h[1:]:
                                C.fill_(float('nan'))
                                o=_run(rc,A,Bq,Bsh,Bsc);torch.cuda.synchronize()
                                e=((o.float()-rf).abs().mean()/mag).item()
                                if e<5e-3:
                                    t=_tcold(lambda c=rc:_run(c,A,Bq,Bsh,Bsc),n=7)
                                    res.append((rc,t));_L(f"  [race] {rc}:{t:.2f}us e={e:.3%}")
                                else:_L(f"  [race] {rc}:ERR{e:.2%}")
                            res.sort(key=lambda x:x[1])
                            cfg=res[0][0]
                            _L(f"  [race] ->{cfg}@{res[0][1]:.2f}us (of {len(res)})")
                        except Exception as _e:_L(f"  [race] exc {_e}")
                    # m=16 aldskx HARD: also probe askx1L NT (info; switch if decisively wins)
                    if cfg[0]=="aldskx" and m<=16:
                        try:
                            _reset1L()
                            t1L=_tcold(lambda:_run(("askx1L",8,8,7),A,Bq,Bsh,Bsc),n=5)
                            t2L=_tcold(lambda:_run(cfg,A,Bq,Bsh,Bsc),n=5)
                            _L(f"  [1Lvs2L] askx1L={t1L:.2f} aldskx={t2L:.2f} d={t2L-t1L:+.2f}us")
                            if t1L<t2L-0.3:
                                C.fill_(float('nan'))
                                o=_run(("askx1L",8,8,7),A,Bq,Bsh,Bsc);torch.cuda.synchronize()
                                e=((o.float()-rf).abs().mean()/mag).item()
                                if e<5e-3:cfg=("askx1L",8,8,7);_L(f"  [1Lvs2L] SWITCH")
                        except Exception as _e:_L(f"  [1Lvs2L] exc {_e}");_reset1L()
                    return {"cfg":cfg,"run":_run,"C":C}
                _L(f"  RECHECK FAIL {cfg} e2={e2:.2%}")
            else:
                _L(f"  ERR {cfg} {err:.2%}")
        except Exception as e:
            _L(f"  EXC {cfg} {type(e).__name__}:{e}")
        _reset1L()

    best=None;bt=1e18;log=[];t0=time.time();nmiss=0
    for cfg in cands:
        if time.time()-t0>40:break
        try:
            C.fill_(float('nan'))
            o=_run(cfg,A,Bq,Bsh,Bsc);torch.cuda.synchronize()
            err=((o.float()-rf).abs().mean()/mag).item()
            if not(err<5e-3):
                _L(f"  {cfg}:ERR{err:.2%}")
                if cfg[0]=="askx1L":_reset1L()
                continue
            t=_tcold(lambda c=cfg:_run(c,A,Bq,Bsh,Bsc))
            log.append((cfg,t))
            if t<bt:bt,best=t,cfg;_L(f"  {cfg}:{t:.2f}us*")
        except RuntimeError as e:
            if "rc=-1" in str(e):nmiss+=1
            else:_L(f"  {cfg}:EXC{str(e)[:100]}")
            torch.cuda.synchronize()
        except Exception as e:
            _L(f"  {cfg}:EXC{type(e).__name__}:{str(e)[:100]}")
            torch.cuda.synchronize()
    if nmiss:_L(f"  (miss={nmiss})")
    if best is None:
        _L("  ->fq");best=("fq",min(8,max(4,K128)))
    try:
        A2=torch.randn_like(A)
        rf2=_ref_tri(A2,Bq,Bsc).float()
        C.fill_(float('nan'))
        o2=_run(best,A2,Bq,Bsh,Bsc);torch.cuda.synchronize()
        e2=((o2.float()-rf2).abs().mean()/(rf2.abs().mean()+1e-9)).item()
        if not(e2<5e-3):_L(f"  RECHECK FAIL {best} {e2:.2%}");best=("fq",min(8,max(4,K128)))
    except Exception as e:_L(f"  recheck exc {e}")
    log.sort(key=lambda x:x[1])
    for c,t in log[:8]:_L(f"  top{c}:{t:.2f}")
    _L(f"  ->best={best}@{bt:.2f}us")
    return {"cfg":best,"run":_run,"C":C}


def custom_kernel(data):
    A=data[0];m,k=A.shape;n=data[2].shape[0]
    S=_ST.get((m,n,k))
    if S is None:
        S=_build(data);_ST[(m,n,k)]=S
    return S["run"](S["cfg"],A,
        data[2].view(torch.uint8),
        data[3].view(torch.uint8),data[4].view(torch.uint8))
scrolls · 1283 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 751752.

⋯ diff truncated: revisions differ almost entirely

Best evidence level for this revision: reported

JSON