Skip to content
KernelIndex
Search⌘K

submission 696489

bill_97933 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
9.31µs
#156 of 1143
2026-04-02

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e14f231602ee88c75c4ad5914294b030df001c67b582310b113cc3a64fa81846
license declaredunknown
license concludedunknown
authorsbill_97933
imported2026-08-15

Techniques

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

fp4k_scale = K_packed // 16 # K=1024 fp4 → 32 scale groups
num-warps = 4num_warps=4, num_stages=2, waves_per_eu=4, matrix_instr_nonkdim=16,
shared-memory__shared__ float lds[4][8][256];
split-k- S2 (M=16, K=7168): ext_splitk ks=16 → 12.60µs
stages = 2num_warps=4, num_stages=2, waves_per_eu=4, matrix_instr_nonkdim=16,
tile-k = 1024v594_triton_bn32_s5: v590 but S5 Triton uses AITer's tuned config (BN=32, BK=1024)
tile-m = 32SPLITK_BLOCK_SIZE=2*K_packed, BLOCK_SIZE_M=32, BLOCK_SIZE_N=32,
tile-n = 32v594_triton_bn32_s5: v590 but S5 Triton uses AITer's tuned config (BN=32, BK=1024)
vector-width = uint4uint4 a0,a1,a2,a3;

Kernel source

v594_triton_bn32_s5.py539 lines
"""
v594_triton_bn32_s5: v590 but S5 Triton uses AITer's tuned config (BN=32, BK=1024)
- S5 Triton: BN=16→32, BK=512→1024, warps=2→4, GROUP_SIZE_M=4→2
  448 CTAs (1.75 waves) vs v590's 896 CTAs (3.5 waves)
  Matches gfx950-GEMM-AFP4WFP4-N=7168-K=2048.json M_LEQ_64 config
- Hypothesis: fewer k-iters (2 vs 4) + larger BN reduces overhead per CTA

per-shape dispatch (expected):
- S1 (M=4):  gemm_fused_16x128 → 6.12µs
- S2 (M=16, K=7168): ext_splitk ks=16 → 12.60µs
- S3 (M=32, N=4096, K=512): gemm_fused_16x16 → 6.73µs
- S4 (M=32, N=2880, K=512): gemm_fused_16x16 → 6.65µs
- S5 (M=64): quant_raw_v2 + Triton BN=32 → ???µs (was 12.9µs with BN=16)
- S6 (M=256): quant_a + ASM GEMM 32x128 → 12.5µs
"""
import os
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
os.environ.setdefault("CXX", "clang++")

import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
import math

SCALE_GROUP_SIZE = 32

_HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_ext_ocp.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <torch/extension.h>
#include <cstdint>

typedef int __attribute__((ext_vector_type(8))) i32x8;
typedef float __attribute__((ext_vector_type(4))) f32x4;
typedef __bf16 bf16v2_t __attribute__((ext_vector_type(2)));

__device__ __forceinline__ int pack8_hw(const uint16_t* src, float hs) {
    unsigned int d = 0;
    d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(d, *reinterpret_cast<const bf16v2_t*>(&src[0]), hs, 0);
    d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(d, *reinterpret_cast<const bf16v2_t*>(&src[2]), hs, 1);
    d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(d, *reinterpret_cast<const bf16v2_t*>(&src[4]), hs, 2);
    d = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(d, *reinterpret_cast<const bf16v2_t*>(&src[6]), hs, 3);
    return (int)d;
}
__device__ __forceinline__ void quant_32(const uint16_t* ap, int out[4], int32_t& spk) {
    uint16_t mx = 0;
    #pragma unroll
    for (int j=0;j<32;j++) mx=max(mx,(uint16_t)(ap[j]&0x7FFF));
    uint32_t au=(((uint32_t)mx<<16)+0x200000u)&0xFF800000u;
    int ef=(au>>23u)&0xFFu;
    int su=(au==0u)?-127:max(-127,min(127,ef-127-2));
    float hs=(su>=-126)?__uint_as_float((uint32_t)(su+127)<<23):0.0f;
    #pragma unroll
    for (int j=0;j<4;j++) out[j]=pack8_hw(&ap[j*8],hs);
    spk=(int32_t)(uint8_t)(su+127);
}
__device__ __forceinline__ int32_t load_b_scale_shuf(const uint8_t* Bsc,int ng,int bb,int SNG,int N) {
    if (ng>=N) return 127;
    int ifl=(ng/32)*(SNG*256)+(bb/8)*256+(bb%4)*64+(ng%16)*4+((bb%8)/4)*2+(ng%32)/16;
    return (int32_t)Bsc[ifl];
}

// ══════ 16×128 fused kernel ══════
__global__ void __launch_bounds__(256,1)
gemm_fused_16x128_kernel(
    const uint16_t* __restrict__ A, const uint8_t* __restrict__ Bq,
    const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
    int M, int N, int K, int strA, int strBq, int BscSN)
{
    const int mt=blockIdx.x, nt=blockIdx.y;
    const int wid=threadIdx.x>>6, lid=threadIdx.x&63;
    const int t_row=lid&15, kpart=lid>>4;
    const int SNG=BscSN>>3;
    const int m_row=mt*16+t_row, a_row_off=m_row*strA;
    const int n_base=nt*128;
    int total_steps=K>>7, steps_per_warp=(total_steps+3)>>2;
    int k_start=wid*steps_per_warp*128, k_end=min(k_start+steps_per_warp*128,K);
    f32x4 c_acc[8];
    #pragma unroll
    for (int ni=0;ni<8;ni++) c_acc[ni]={0,0,0,0};
    for (int kb=k_start;kb<k_end;kb+=128) {
        int k_off=kb+kpart*32, bk=(kb>>1)+kpart*16, b_scale_k=(kb>>5)+kpart;
        uint4 a0,a1,a2,a3;
        if (m_row<M&&k_off+31<K) {
            const uint4*src=reinterpret_cast<const uint4*>(&A[a_row_off+k_off]);
            a0=src[0];a1=src[1];a2=src[2];a3=src[3];
        } else {a0={0,0,0,0};a1={0,0,0,0};a2={0,0,0,0};a3={0,0,0,0};}
        uint16_t a_local[32]; uint4*adst=reinterpret_cast<uint4*>(a_local);
        adst[0]=a0;adst[1]=a1;adst[2]=a2;adst[3]=a3;
        int a_i32[4]; int32_t a_spk;
        quant_32(a_local,a_i32,a_spk);
        i32x8 am={a_i32[0],a_i32[1],a_i32[2],a_i32[3],0,0,0,0};
        #pragma unroll
        for (int ni=0;ni<8;ni++) {
            int b_ng=n_base+ni*16+t_row;
            int b_i32[4];
            if (b_ng<N&&bk+15<(K>>1)) {
                uint4 bd=*reinterpret_cast<const uint4*>(&Bq[b_ng*strBq+bk]);
                b_i32[0]=((int*)&bd)[0];b_i32[1]=((int*)&bd)[1];
                b_i32[2]=((int*)&bd)[2];b_i32[3]=((int*)&bd)[3];
            } else {b_i32[0]=0;b_i32[1]=0;b_i32[2]=0;b_i32[3]=0;}
            int32_t b_spk=load_b_scale_shuf(Bsc,b_ng,b_scale_k,SNG,N);
            i32x8 bm={b_i32[0],b_i32[1],b_i32[2],b_i32[3],0,0,0,0};
            c_acc[ni]=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(am,bm,c_acc[ni],4,4,0,a_spk,0,b_spk);
        }
    }
    __shared__ float lds[4][8][256];
    #pragma unroll
    for (int ni=0;ni<8;ni++)
        #pragma unroll
        for (int j=0;j<4;j++) lds[wid][ni][lid*4+j]=c_acc[ni][j];
    __syncthreads();
    #pragma unroll
    for (int ni=0;ni<8;ni++) {
        #pragma unroll
        for (int j=0;j<4;j++) {
            float sum=lds[0][ni][lid*4+j]+lds[1][ni][lid*4+j]+lds[2][ni][lid*4+j]+lds[3][ni][lid*4+j];
            int mo=mt*16+kpart*4+j, no=n_base+ni*16+t_row;
            if (mo<M&&no<N) {
                uint32_t fp=__float_as_uint(sum); fp+=0x7FFFu+((fp>>16)&1u);
                C[mo*N+no]=(uint16_t)(fp>>16u);
            }
        }
    }
}

// ══════ 16×16 fused kernel ══════
__global__ void __launch_bounds__(256,3)
gemm_fused_16x16_kernel(
    const uint16_t* __restrict__ A, const uint8_t* __restrict__ Bq,
    const uint8_t* __restrict__ Bsc, uint16_t* __restrict__ C,
    int M, int N, int K, int strA, int strBq, int BscSN)
{
    const int mt=blockIdx.x,nt=blockIdx.y;
    const int wid=threadIdx.x>>6,lid=threadIdx.x&63;
    const int t_row=lid&15,kpart=lid>>4,SNG=BscSN>>3;
    const int m_row=mt*16+t_row,a_row_off=m_row*strA,b_ng=nt*16+t_row;
    int total_steps=K>>7,steps_per_warp=(total_steps+3)>>2;
    int k_start=wid*steps_per_warp*128,k_end=min(k_start+steps_per_warp*128,K);
    f32x4 c_acc={0,0,0,0};
    for (int kb=k_start;kb<k_end;kb+=128) {
        int k_off=kb+kpart*32,bk=(kb>>1)+kpart*16;
        uint4 a0,a1,a2,a3;
        if (m_row<M&&k_off+31<K){const uint4*s=reinterpret_cast<const uint4*>(&A[a_row_off+k_off]);a0=s[0];a1=s[1];a2=s[2];a3=s[3];}
        else{a0={0,0,0,0};a1={0,0,0,0};a2={0,0,0,0};a3={0,0,0,0};}
        int b_i32[4];
        if (b_ng<N&&bk+15<(K>>1)){uint4 bd=*reinterpret_cast<const uint4*>(&Bq[b_ng*strBq+bk]);b_i32[0]=((int*)&bd)[0];b_i32[1]=((int*)&bd)[1];b_i32[2]=((int*)&bd)[2];b_i32[3]=((int*)&bd)[3];}
        else{b_i32[0]=0;b_i32[1]=0;b_i32[2]=0;b_i32[3]=0;}
        int32_t b_spk=load_b_scale_shuf(Bsc,b_ng,(kb>>5)+kpart,SNG,N);
        uint16_t a_local[32];uint4*d=reinterpret_cast<uint4*>(a_local);d[0]=a0;d[1]=a1;d[2]=a2;d[3]=a3;
        int a_i32[4];int32_t a_spk;quant_32(a_local,a_i32,a_spk);
        i32x8 am={a_i32[0],a_i32[1],a_i32[2],a_i32[3],0,0,0,0};
        i32x8 bm={b_i32[0],b_i32[1],b_i32[2],b_i32[3],0,0,0,0};
        c_acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(am,bm,c_acc,4,4,0,a_spk,0,b_spk);
    }
    __shared__ float r16[4][64*4];
    #pragma unroll
    for (int i=0;i<4;i++) r16[wid][lid*4+i]=c_acc[i];
    __syncthreads();
    if (wid==0) {
        for (int j=0;j<4;j++) {
            float sum=r16[0][lid*4+j]+r16[1][lid*4+j]+r16[2][lid*4+j]+r16[3][lid*4+j];
            int mo=mt*16+kpart*4+j,no=nt*16+t_row;
            if (mo<M&&no<N){uint32_t fp=__float_as_uint(sum);fp+=0x7FFFu+((fp>>16)&1u);C[mo*N+no]=(uint16_t)(fp>>16u);}
        }
    }
}

// ══════ ext splitk for S2 ══════
__global__ void __launch_bounds__(64,8)
gemm_ext_splitk_16x16_bsh_kernel(
    const uint16_t* __restrict__ A,const uint8_t* __restrict__ Bsh,
    const uint8_t* __restrict__ Bsc,float* __restrict__ Cfp32,
    int M,int N,int K,int strA,int BscSN,int kper)
{
    const int mt=blockIdx.x,nt=blockIdx.y,ks=blockIdx.z;
    const int lid=threadIdx.x,t_row=lid&15,kpart=lid>>4,SNG=BscSN>>3;
    const int m_row=mt*16+t_row,a_row_off=m_row*strA,b_ng=nt*16+t_row;
    int kst=(ks*kper/128)*128,ken=min(((ks+1)*kper+127)/128*128,K);
    if (kst>=ken) return;
    const int nkt=K/32;
    const int bsh_base=(b_ng<N)?((b_ng/16)*nkt*256+(b_ng%16)*16):0;
    f32x4 c_acc={0,0,0,0};
    for (int kb=kst;kb<ken;kb+=128) {
        int k_off=kb+kpart*32;
        uint16_t a_local[32];
        if (m_row<M&&k_off+31<K){const uint4*s=reinterpret_cast<const uint4*>(&A[a_row_off+k_off]);uint4*d=reinterpret_cast<uint4*>(a_local);d[0]=s[0];d[1]=s[1];d[2]=s[2];d[3]=s[3];}
        else{
            #pragma unroll
            for(int j=0;j<32;j++) a_local[j]=(m_row<M&&k_off+j<K)?A[a_row_off+k_off+j]:0;
        }
        int a_i32[4];int32_t a_spk;quant_32(a_local,a_i32,a_spk);
        int bk=(kb>>1)+kpart*16;
        int b_i32[4];
        if (b_ng<N&&bk+15<K/2){uint4 bd=*reinterpret_cast<const uint4*>(&Bsh[bsh_base+(bk>>4)*256]);b_i32[0]=((int*)&bd)[0];b_i32[1]=((int*)&bd)[1];b_i32[2]=((int*)&bd)[2];b_i32[3]=((int*)&bd)[3];}
        else{b_i32[0]=0;b_i32[1]=0;b_i32[2]=0;b_i32[3]=0;}
        int32_t b_spk=load_b_scale_shuf(Bsc,b_ng,(kb>>5)+kpart,SNG,N);
        i32x8 am={a_i32[0],a_i32[1],a_i32[2],a_i32[3],0,0,0,0};
        i32x8 bm={b_i32[0],b_i32[1],b_i32[2],b_i32[3],0,0,0,0};
        c_acc=__builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(am,bm,c_acc,4,4,0,a_spk,0,b_spk);
    }
    #pragma unroll
    for (int j=0;j<4;j++){int mo=mt*16+kpart*4+j,no=nt*16+t_row;if(mo<M&&no<N)atomicAdd(&Cfp32[mo*N+no],c_acc[j]);}
}

__global__ void __launch_bounds__(1024)
fp32_to_bf16(float* __restrict__ src,uint16_t* __restrict__ dst,int n){
    int i=blockIdx.x*1024+threadIdx.x;
    if(i<n){float val=src[i];src[i]=0.0f;uint32_t fp=__float_as_uint(val);fp+=0x7FFFu+((fp>>16)&1u);dst[i]=(uint16_t)(fp>>16u);}
}

// ══════ quant kernels ══════
__global__ void __launch_bounds__(64,16)
quant_a_kernel(const uint16_t* __restrict__ A,uint8_t* __restrict__ Aq,
               uint8_t* __restrict__ Asc,int M,int K,int strA,int sm){
    int idx=blockIdx.x*blockDim.x+threadIdx.x;
    int n_scales=K/32,total_blocks=sm*n_scales;
    if(idx>=total_blocks)return;
    int row=idx/n_scales,blk=idx%n_scales,k_off=blk*32;
    uint16_t vals[32];
    if(row<M){const uint4*s=reinterpret_cast<const uint4*>(&A[row*strA+k_off]);uint4*d=reinterpret_cast<uint4*>(vals);d[0]=s[0];d[1]=s[1];d[2]=s[2];d[3]=s[3];}
    else{
        #pragma unroll
        for(int j=0;j<32;j++)vals[j]=0;
    }
    uint16_t mx=0;
    #pragma unroll
    for(int j=0;j<32;j++)mx=max(mx,(uint16_t)(vals[j]&0x7FFF));
    uint32_t au=(((uint32_t)mx<<16)+0x200000u)&0xFF800000u;
    int ef=(au>>23u)&0xFFu,su=(au==0u)?-127:max(-127,min(127,ef-127-2));
    float hs=(su>=-126)?__uint_as_float((uint32_t)(su+127)<<23):0.0f;
    uint8_t scale_val=(uint8_t)(su+127);
    unsigned int packed[4];
    packed[0]=0;packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[0]),hs,0);
    packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[2]),hs,1);
    packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[4]),hs,2);
    packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[6]),hs,3);
    packed[1]=0;packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[8]),hs,0);
    packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[10]),hs,1);
    packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[12]),hs,2);
    packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[14]),hs,3);
    packed[2]=0;packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[16]),hs,0);
    packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[18]),hs,1);
    packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[20]),hs,2);
    packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[22]),hs,3);
    packed[3]=0;packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[24]),hs,0);
    packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[26]),hs,1);
    packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[28]),hs,2);
    packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[30]),hs,3);
    if(row<M)*reinterpret_cast<uint4*>(&Aq[row*(K/2)+blk*16])=*reinterpret_cast<uint4*>(packed);
    int sn=K/32,i0=row/32,i1=(row%32)/16,i2=row%16,i3=blk/8,i4=(blk%8)/4,i5=blk%4;
    Asc[i0*(32*sn)+i3*256+i5*64+i2*4+i4*2+i1]=scale_val;
}

__global__ void __launch_bounds__(64,16)
quant_a_raw_v2_kernel(const uint16_t* __restrict__ A,uint8_t* __restrict__ Aq,
                      uint8_t* __restrict__ Asc,int M,int K,int strA){
    int idx=blockIdx.x*blockDim.x+threadIdx.x;
    int n_scales=K/32,total=M*n_scales;
    if(idx>=total)return;
    int row=idx/n_scales,blk=idx%n_scales,k_off=blk*32;
    uint16_t vals[32];
    if(row<M){const uint4*s=reinterpret_cast<const uint4*>(&A[row*strA+k_off]);uint4*d=reinterpret_cast<uint4*>(vals);d[0]=s[0];d[1]=s[1];d[2]=s[2];d[3]=s[3];}
    else{
        #pragma unroll
        for(int j=0;j<32;j++)vals[j]=0;
    }
    float fmax_val=0.0f;
    const __hip_bfloat16*bf16v=reinterpret_cast<const __hip_bfloat16*>(vals);
    #pragma unroll
    for(int j=0;j<32;j++){float fv=__bfloat162float(bf16v[j]);fmax_val=fmaxf(fmax_val,fabsf(fv));}
    int scale_exp=-127;
    if(fmax_val>0.0f){uint32_t ab;memcpy(&ab,&fmax_val,4);ab=(ab+0x00200000u)&0xFF800000u;float r;memcpy(&r,&ab,4);scale_exp=max(-127,min(127,(int)floorf(log2f(r))-2));}
    uint8_t scale_byte=(uint8_t)(scale_exp+127);
    float hs=(scale_exp>=-126)?__uint_as_float((uint32_t)(scale_exp+127)<<23):0.0f;
    unsigned int packed[4];
    packed[0]=0;packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[0]),hs,0);
    packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[2]),hs,1);
    packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[4]),hs,2);
    packed[0]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[0],*reinterpret_cast<bf16v2_t*>(&vals[6]),hs,3);
    packed[1]=0;packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[8]),hs,0);
    packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[10]),hs,1);
    packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[12]),hs,2);
    packed[1]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[1],*reinterpret_cast<bf16v2_t*>(&vals[14]),hs,3);
    packed[2]=0;packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[16]),hs,0);
    packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[18]),hs,1);
    packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[20]),hs,2);
    packed[2]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[2],*reinterpret_cast<bf16v2_t*>(&vals[22]),hs,3);
    packed[3]=0;packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[24]),hs,0);
    packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[26]),hs,1);
    packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[28]),hs,2);
    packed[3]=__builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(packed[3],*reinterpret_cast<bf16v2_t*>(&vals[30]),hs,3);
    if(row<M)*reinterpret_cast<uint4*>(&Aq[row*(K/2)+blk*16])=*reinterpret_cast<uint4*>(packed);
    Asc[row*n_scales+blk]=scale_byte;
}

static hipModule_t _asm_mod=nullptr;
static hipFunction_t _asm_fn=nullptr;
struct __attribute__((packed)) AsmArgs{
    void*D;char _0[8];void*C;char _1[8];void*A;char _2[8];void*B;char _3[8];
    float alpha;char _4[12];float beta;char _5[12];
    unsigned int sD0;char _6[12];unsigned int sD1;char _7[12];
    unsigned int sC0;char _8[12];unsigned int sC1;char _9[12];
    unsigned int sA0;char _10[12];unsigned int sA1;char _11[12];
    unsigned int sB0;char _12[12];unsigned int sB1;char _13[12];
    unsigned int M;char _14[12];unsigned int N;char _15[12];unsigned int K;char _16[12];
    void*SA;char _17[8];void*SB;char _18[8];
    unsigned int sSA0;char _19[12];unsigned int sSA1;char _20[12];
    unsigned int sSB0;char _21[12];unsigned int sSB1;char _22[12];
    int log2ks;
};
torch::Tensor asm_gemm_a4w4(torch::Tensor Aq,torch::Tensor Bsh,torch::Tensor Asc,torch::Tensor Bsc,int64_t m,int64_t n,int64_t k,torch::Tensor out){
    if(!_asm_mod){hipModuleLoad(&_asm_mod,"/home/runner/aiter/hsa//gfx950/f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co");hipModuleGetFunction(&_asm_fn,_asm_mod,"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E");}
    int mp=((int)m+31)/32*32;
    AsmArgs a={};a.D=out.data_ptr();a.C=out.data_ptr();a.A=Aq.data_ptr();a.B=Bsh.data_ptr();
    a.alpha=1.0f;a.beta=0.0f;a.sC0=(unsigned)out.stride(0);a.sC1=1;
    a.sA0=(unsigned)(Aq.stride(0)*2);a.sA1=1;a.sB0=(unsigned)(Bsh.stride(0)*2);a.sB1=1;
    a.M=(unsigned)m;a.N=(unsigned)n;a.K=(unsigned)k;a.SA=Asc.data_ptr();a.SB=Bsc.data_ptr();
    a.sSA0=(unsigned)Asc.stride(0);a.sSA1=1;a.sSB0=(unsigned)Bsc.stride(0);a.sSB1=1;a.log2ks=0;
    size_t asz=sizeof(a);
    void*cfg[]={HIP_LAUNCH_PARAM_BUFFER_POINTER,&a,HIP_LAUNCH_PARAM_BUFFER_SIZE,&asz,HIP_LAUNCH_PARAM_END};
    hipModuleLaunchKernel(_asm_fn,((unsigned)n+127)/128,(mp+31)/32,1,256,1,1,0,0,nullptr,(void**)cfg);
    if(mp>(int)m)return out.slice(0,0,(int)m);return out;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME,m){
    m.def("gemm_fused_16x128",[](torch::Tensor A,torch::Tensor Bq,torch::Tensor Bsc,int64_t N,torch::Tensor C)->torch::Tensor{
        int M=(int)A.size(0),K=(int)A.size(1);
        dim3 grid((M+15)/16,((int)N+127)/128);
        gemm_fused_16x128_kernel<<<grid,256>>>(reinterpret_cast<const uint16_t*>(A.data_ptr()),reinterpret_cast<const uint8_t*>(Bq.data_ptr()),reinterpret_cast<const uint8_t*>(Bsc.data_ptr()),reinterpret_cast<uint16_t*>(C.data_ptr()),M,(int)N,K,(int)A.stride(0),(int)(Bq.stride(0)*Bq.element_size()),(int)Bsc.size(1));
        return C;
    });
    m.def("gemm_fused_16x16",[](torch::Tensor A,torch::Tensor Bq,torch::Tensor Bsc,int64_t N,torch::Tensor C)->torch::Tensor{
        int M=(int)A.size(0),K=(int)A.size(1);
        dim3 grid((M+15)/16,(N+15)/16);
        gemm_fused_16x16_kernel<<<grid,256>>>(reinterpret_cast<const uint16_t*>(A.data_ptr()),reinterpret_cast<const uint8_t*>(Bq.data_ptr()),reinterpret_cast<const uint8_t*>(Bsc.data_ptr()),reinterpret_cast<uint16_t*>(C.data_ptr()),M,(int)N,K,(int)A.stride(0),(int)(Bq.stride(0)*Bq.element_size()),(int)Bsc.size(1));
        return C;
    });
    m.def("gemm_ext_splitk_16x16_bsh_cached",[](torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,int64_t N,int64_t ks,torch::Tensor Cfp,torch::Tensor Cbf)->torch::Tensor{
        int M=(int)A.size(0),K=(int)A.size(1);
        int kper=((K/128+(int)ks-1)/(int)ks)*128,aks=(K+kper-1)/kper;
        dim3 grid((M+15)/16,((int)N+15)/16,aks);
        gemm_ext_splitk_16x16_bsh_kernel<<<grid,64>>>(reinterpret_cast<const uint16_t*>(A.data_ptr()),reinterpret_cast<const uint8_t*>(Bsh.data_ptr()),reinterpret_cast<const uint8_t*>(Bsc.data_ptr()),Cfp.data_ptr<float>(),M,(int)N,K,(int)A.stride(0),(int)Bsc.size(1),kper);
        int total=M*(int)N;
        fp32_to_bf16<<<(total+1023)/1024,1024>>>(Cfp.data_ptr<float>(),reinterpret_cast<uint16_t*>(Cbf.data_ptr()),total);
        return Cbf;
    });
    m.def("asm_gemm_a4w4",&asm_gemm_a4w4);
    m.def("warmup_asm",[](){if(!_asm_mod){hipModuleLoad(&_asm_mod,"/home/runner/aiter/hsa//gfx950/f4gemm/f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co");hipModuleGetFunction(&_asm_fn,_asm_mod,"_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E");}});
    m.def("quant_and_asm_gemm",[](torch::Tensor A,torch::Tensor Bsh,torch::Tensor Bsc,int64_t N,torch::Tensor Aq,torch::Tensor Asc,torch::Tensor out)->torch::Tensor{
        int M=(int)A.size(0),K=(int)A.size(1),sm=(int)Asc.size(0);
        quant_a_kernel<<<(sm*(K/32)+63)/64,64>>>(reinterpret_cast<const uint16_t*>(A.data_ptr()),reinterpret_cast<uint8_t*>(Aq.data_ptr()),reinterpret_cast<uint8_t*>(Asc.data_ptr()),M,K,(int)A.stride(0),sm);
        return asm_gemm_a4w4(Aq,Bsh,Asc,Bsc,M,(int)N,K,out);
    });
    m.def("quant_raw_v2",[](torch::Tensor A,torch::Tensor Aq,torch::Tensor Asc){
        int M=(int)A.size(0),K=(int)A.size(1);
        quant_a_raw_v2_kernel<<<(M*(K/32)+63)/64,64>>>(reinterpret_cast<const uint16_t*>(A.data_ptr()),Aq.data_ptr<uint8_t>(),Asc.data_ptr<uint8_t>(),M,K,(int)A.stride(0));
    });
}
"""

_EXT = None
def _get_ext():
    global _EXT
    if _EXT is None:
        _EXT = load_inline(
            name="fused_mxfp4_v594",
            cpp_sources=[""],
            cuda_sources=[_HIP_SRC],
            extra_cuda_cflags=["-O3", "-std=c++17", "--offload-arch=gfx950"],
            verbose=False,
        )
    return _EXT


@triton.jit
def remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
    pids_per_xcd=(GRID_MN+NUM_XCDS-1)//NUM_XCDS
    tall_xcds=GRID_MN%NUM_XCDS; tall_xcds=NUM_XCDS if tall_xcds==0 else tall_xcds
    xcd=pid%NUM_XCDS; local_pid=pid//NUM_XCDS
    if xcd<tall_xcds: pid=xcd*pids_per_xcd+local_pid
    else: pid=(tall_xcds*pids_per_xcd+(xcd-tall_xcds)*(pids_per_xcd-1)+local_pid)
    return pid

@triton.jit
def pid_grid(pid,num_pid_m,num_pid_n,GROUP_SIZE_M:tl.constexpr=1):
    if GROUP_SIZE_M==1: pid_m=pid//num_pid_n; pid_n=pid%num_pid_n
    else:
        num_pid_in_group=GROUP_SIZE_M*num_pid_n; group_id=pid//num_pid_in_group
        first_pid_m=group_id*GROUP_SIZE_M; group_size_m=min(num_pid_m-first_pid_m,GROUP_SIZE_M)
        tl.assume(group_size_m>=0)
        pid_m=first_pid_m+(pid%group_size_m); pid_n=(pid%num_pid_in_group)//group_size_m
    return pid_m,pid_n

@triton.heuristics({"EVEN_K":lambda args:(args["K"]%(args["BLOCK_SIZE_K"]//2)==0)and(args["SPLITK_BLOCK_SIZE"]%args["BLOCK_SIZE_K"]==0)and(args["K"]%(args["SPLITK_BLOCK_SIZE"]//2)==0)})
@triton.jit
def _gemm_fp4_kernel(
    a_ptr,b_ptr,c_ptr,a_scales_ptr,b_scales_ptr,M,N,K,
    stride_am,stride_ak,stride_bk,stride_bn,stride_ck,stride_cm,stride_cn,
    stride_asm,stride_ask,stride_bsn,stride_bsk,B_SCALE_SN_PAD,
    BLOCK_SIZE_M:tl.constexpr,BLOCK_SIZE_N:tl.constexpr,BLOCK_SIZE_K:tl.constexpr,
    GROUP_SIZE_M:tl.constexpr,NUM_KSPLIT:tl.constexpr,SPLITK_BLOCK_SIZE:tl.constexpr,
    EVEN_K:tl.constexpr,num_warps:tl.constexpr,num_stages:tl.constexpr,
    waves_per_eu:tl.constexpr,matrix_instr_nonkdim:tl.constexpr,
):
    tl.assume(stride_am>0);tl.assume(stride_ak>0);tl.assume(stride_bk>0);tl.assume(stride_bn>0)
    tl.assume(stride_cm>0);tl.assume(stride_cn>0);tl.assume(stride_asm>0);tl.assume(stride_ask>0)
    tl.assume(stride_bsk>0);tl.assume(stride_bsn>0)
    GRID_MN=tl.cdiv(M,BLOCK_SIZE_M)*tl.cdiv(N,BLOCK_SIZE_N)
    pid_unified=tl.program_id(axis=0); pid_unified=remap_xcd(pid_unified,GRID_MN*NUM_KSPLIT,NUM_XCDS=8)
    pid_k=pid_unified%NUM_KSPLIT; pid=pid_unified//NUM_KSPLIT
    num_pid_m=tl.cdiv(M,BLOCK_SIZE_M); num_pid_n=tl.cdiv(N,BLOCK_SIZE_N)
    if NUM_KSPLIT==1: pid_m,pid_n=pid_grid(pid,num_pid_m,num_pid_n,GROUP_SIZE_M=GROUP_SIZE_M)
    else: pid_m=pid//num_pid_n; pid_n=pid%num_pid_n
    tl.assume(pid_m>=0); tl.assume(pid_n>=0)
    SCALE_GROUP_SIZE:tl.constexpr=32
    if (pid_k*SPLITK_BLOCK_SIZE//2)<K:
        num_k_iter=tl.cdiv(SPLITK_BLOCK_SIZE//2,BLOCK_SIZE_K//2)
        offs_k=tl.arange(0,BLOCK_SIZE_K//2); offs_k_split=pid_k*(SPLITK_BLOCK_SIZE//2)+offs_k
        offs_am=(pid_m*BLOCK_SIZE_M+tl.arange(0,BLOCK_SIZE_M))%M
        offs_bn=(pid_n*BLOCK_SIZE_N+tl.arange(0,BLOCK_SIZE_N))%N
        a_ptrs=a_ptr+(offs_am[:,None]*stride_am+offs_k_split[None,:]*stride_ak)
        b_ptrs=b_ptr+(offs_k_split[:,None]*stride_bk+offs_bn[None,:]*stride_bn)
        offs_ks=(pid_k*(SPLITK_BLOCK_SIZE//SCALE_GROUP_SIZE))+tl.arange(0,BLOCK_SIZE_K//SCALE_GROUP_SIZE)
        a_scale_ptrs=a_scales_ptr+offs_am[:,None]*stride_asm+offs_ks[None,:]*stride_ask
        bn=offs_bn[:,None];bk=offs_ks[None,:]
        b_i0=bn//32;b_i1=(bn%32)//16;b_i2=bn%16;b_i3=bk//8;b_i4=(bk%8)//4;b_i5=bk%4
        b_sh_idx=b_i0*(32*B_SCALE_SN_PAD)+b_i3*256+b_i5*64+b_i2*4+b_i4*2+b_i1
        accumulator=tl.zeros((BLOCK_SIZE_M,BLOCK_SIZE_N),dtype=tl.float32)
        for k in range(pid_k*num_k_iter,(pid_k+1)*num_k_iter):
            a_scales=tl.load(a_scale_ptrs); b_scales=tl.load(b_scales_ptr+b_sh_idx,cache_modifier=".cg")
            if EVEN_K: a=tl.load(a_ptrs);b=tl.load(b_ptrs,cache_modifier=".cg")
            else:
                a=tl.load(a_ptrs,mask=offs_k[None,:]<K-k*(BLOCK_SIZE_K//2),other=0)
                b=tl.load(b_ptrs,mask=offs_k[:,None]<K-k*(BLOCK_SIZE_K//2),other=0,cache_modifier=".cg")
            accumulator=tl.dot_scaled(a,a_scales,"e2m1",b,b_scales,"e2m1",accumulator)
            a_ptrs+=(BLOCK_SIZE_K//2)*stride_ak;b_ptrs+=(BLOCK_SIZE_K//2)*stride_bk
            a_scale_ptrs+=(BLOCK_SIZE_K//SCALE_GROUP_SIZE)*stride_ask
            bk=bk+(BLOCK_SIZE_K//SCALE_GROUP_SIZE);b_i3=bk//8;b_i4=(bk%8)//4;b_i5=bk%4
            b_sh_idx=b_i0*(32*B_SCALE_SN_PAD)+b_i3*256+b_i5*64+b_i2*4+b_i4*2+b_i1
        c=accumulator.to(c_ptr.type.element_ty)
        offs_cm=(pid_m*BLOCK_SIZE_M+tl.arange(0,BLOCK_SIZE_M)).to(tl.int64)
        offs_cn=(pid_n*BLOCK_SIZE_N+tl.arange(0,BLOCK_SIZE_N)).to(tl.int64)
        c_ptrs=c_ptr+stride_cm*offs_cm[:,None]+stride_cn*offs_cn[None,:]+pid_k*stride_ck
        tl.store(c_ptrs,c,mask=(offs_cm[:,None]<M)&(offs_cn[None,:]<N))


_ws_cache: dict = {}
_warmed = False

def _prewarm_triton():
    """Pre-compile Triton kernel with BN=32, BK=1024, warps=4 config."""
    try:
        import torch
        dev = torch.device('cuda', 0)
        # Use minimal sizes matching BN=32, BK=1024
        # K_packed=512 (K=1024 fp4 elems), BLOCK_SIZE_K=1024 → 1 k-iter
        M, N, K_packed = 32, 32, 512
        k_scale = K_packed // 16  # K=1024 fp4 → 32 scale groups
        aq = torch.zeros(M, K_packed, dtype=torch.uint8, device=dev)
        bt = torch.zeros(K_packed, N, dtype=torch.uint8, device=dev)
        y  = torch.zeros(M, N, dtype=torch.bfloat16, device=dev)
        asc = torch.zeros(M, k_scale, dtype=torch.uint8, device=dev)
        bsc = torch.zeros(N, k_scale, dtype=torch.uint8, device=dev)
        grid = lambda META: (1,)
        _gemm_fp4_kernel[grid](
            aq, bt, y, asc, bsc, M, N, K_packed,
            aq.stride(0), aq.stride(1), bt.stride(0), bt.stride(1),
            0, y.stride(0), y.stride(1), asc.stride(0), asc.stride(1), 0, 0, k_scale,
            SPLITK_BLOCK_SIZE=2*K_packed, BLOCK_SIZE_M=32, BLOCK_SIZE_N=32,
            BLOCK_SIZE_K=1024, GROUP_SIZE_M=2, NUM_KSPLIT=1,
            num_warps=4, num_stages=2, waves_per_eu=4, matrix_instr_nonkdim=16,
        )
        torch.cuda.synchronize()
    except Exception:
        pass

_prewarm_triton()

def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    global _warmed
    ext = _get_ext()
    if not _warmed:
        ext.warmup_asm(); _warmed = True

    M, K, N = A.shape[0], A.shape[1], B.shape[0]
    key = (M, K, N); dev = A.device
    if key not in _ws_cache:
        mp=((M+31)//32)*32; sm=((mp+255)//256)*256
        _ws_cache[key]=(
            torch.zeros(M,N,dtype=torch.float32,device=dev),
            torch.empty(M,N,dtype=torch.bfloat16,device=dev),
            torch.empty(M,K//2,dtype=torch.uint8,device=dev),
            torch.empty(sm,K//32,dtype=torch.uint8,device=dev),
            torch.empty(mp,N,dtype=torch.bfloat16,device=dev),
        )
    fp32_ws,bf16_out,aq_buf,asc_buf,asm_out=_ws_cache[key]

    # S5: M=64 — quant_raw_v2 + Triton BN=32, BK=1024, warps=4 (AITer config)
    if 64<=M<128:
        k_scale=K//SCALE_GROUP_SIZE
        asc_key=(M,K,'raw')
        if asc_key not in _ws_cache:
            _ws_cache[asc_key]=torch.empty(M,k_scale,dtype=torch.uint8,device=dev)
        asc_raw=_ws_cache[asc_key]
        ext.quant_raw_v2(A,aq_buf,asc_raw)
        B_q_u8=B_q.view(torch.uint8) if B_q.dtype!=torch.uint8 else B_q
        B_t=B_q_u8.T
        B_sc_u8=B_scale_sh.view(torch.uint8) if B_scale_sh.dtype!=torch.uint8 else B_scale_sh
        sn_pad=B_sc_u8.shape[1]; K_packed=K//2; y=bf16_out
        grid=lambda META:(META["NUM_KSPLIT"]*triton.cdiv(M,META["BLOCK_SIZE_M"])*triton.cdiv(N,META["BLOCK_SIZE_N"]),)
        _gemm_fp4_kernel[grid](
            aq_buf,B_t,y,asc_raw,B_sc_u8,M,N,K_packed,
            aq_buf.stride(0),aq_buf.stride(1),B_t.stride(0),B_t.stride(1),
            0,y.stride(0),y.stride(1),asc_raw.stride(0),asc_raw.stride(1),0,0,sn_pad,
            SPLITK_BLOCK_SIZE=2*K_packed,BLOCK_SIZE_M=32,BLOCK_SIZE_N=32,BLOCK_SIZE_K=1024,
            GROUP_SIZE_M=2,NUM_KSPLIT=1,num_warps=4,num_stages=2,waves_per_eu=4,matrix_instr_nonkdim=16,
        )
        return y

    # S6: M>=128 — quant_a + ASM GEMM (v517 path, best for large M)
    if M>=128:
        return ext.quant_and_asm_gemm(A,B_shuffle,B_scale_sh,N,aq_buf,asc_buf,asm_out)

    # S2: M<=16, large K — ext_splitk (v517 path)
    if M<=16 and K>1024:
        return ext.gemm_ext_splitk_16x16_bsh_cached(A,B_shuffle,B_scale_sh,N,16,fp32_ws,bf16_out)

    B_q_u8=B_q.view(torch.uint8) if B_q.dtype!=torch.uint8 else B_q
    if M<=16 and N % 128 == 0:
        return ext.gemm_fused_16x128(A,B_q_u8,B_scale_sh,N,bf16_out)
    else:
        return ext.gemm_fused_16x16(A,B_q_u8,B_scale_sh,N,bf16_out)
scrolls · 539 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