Skip to content
KernelIndex
Search⌘K

submission 687670

Niall · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_k2_half_lds.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-687670?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, fp32, fp8_e8m0, int32, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 MoEsuite of 7 cases
AMD Instinct MI355X
179.4µs
#448 of 782
2026-04-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:cf928adbc73164abe0197ad61f83d77866ad8bf019d6883a9726c7a81534dcf9
license declaredunknown
license concludedunknown
authorsNiall
imported2026-08-26

Techniques

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

shared-memory__shared__ uint32_t s_data[TOTAL];
vector-width = uint4__device__ __forceinline__ float unpack_f16x8_maxabs(const uint4 &packed, float* v) {

Kernel source

submission_k2_half_lds.py870 lines
import torch, sys, os, subprocess, importlib.util, sysconfig
from task import input_t, output_t
def _log(msg): print(msg, file=sys.stderr, flush=True)

HIP_SOURCE = r"""
#include <hip/hip_runtime.h>
#include <torch/extension.h>
#include <stdint.h>
#include <stdio.h>

/* ══════════════════════════════════════════════════════════════
 * Constants
 * ══════════════════════════════════════════════════════════════ */

#define K_DIM          7168
#define THREADS        512
#define N_BLOCKS       256
#define COLS_PER_WF    8
#define PACK_PAIR(eid, tok) (((uint32_t)(eid) << 16) | ((uint32_t)(tok) & 0xFFFF))

/* Activation LDS layout */
#define ACT_ROW_DATA   7168
#define ACT_ROW_PAD    16
#define ACT_ROW_STRIDE 7184
#define ACT_STEP       128
#define ACT_BUF_SIZE   (16 * ACT_ROW_STRIDE)   /* 114,944 */
#define ACT_BUF_OFF    0

/* Weight register map */
#define WT_AREG        4
#define WT_AGPR_STEPS  28
#define N_STEPS        56

/* Benchmark parameters */
#define N_ITERS        100

typedef int v8i32 __attribute__((ext_vector_type(8)));
typedef int v4i32 __attribute__((ext_vector_type(4)));


/* ══════════════════════════════════════════════════════════════
 * K1 — Real sort + quantize kernel
 * ══════════════════════════════════════════════════════════════ */

__device__ __forceinline__ void unpack_bf16x2(uint32_t packed, float &lo, float &hi) {
    uint32_t lo_bits = (packed & 0xFFFF) << 16;
    uint32_t hi_bits = packed & 0xFFFF0000;
    __builtin_memcpy(&lo, &lo_bits, 4);
    __builtin_memcpy(&hi, &hi_bits, 4);
}

template<int BASE>
__device__ __forceinline__ float unpack_f16x8_maxabs(const uint4 &packed, float* v) {
    float lm = 0;
    unpack_bf16x2(packed.x, v[BASE+0], v[BASE+1]);
    unpack_bf16x2(packed.y, v[BASE+2], v[BASE+3]);
    unpack_bf16x2(packed.z, v[BASE+4], v[BASE+5]);
    unpack_bf16x2(packed.w, v[BASE+6], v[BASE+7]);
    lm = fmaxf(lm, fmaxf(fabsf(v[BASE+0]), fabsf(v[BASE+1])));
    lm = fmaxf(lm, fmaxf(fabsf(v[BASE+2]), fabsf(v[BASE+3])));
    lm = fmaxf(lm, fmaxf(fabsf(v[BASE+4]), fabsf(v[BASE+5])));
    lm = fmaxf(lm, fmaxf(fabsf(v[BASE+6]), fabsf(v[BASE+7])));
    return lm;
}

template<int BASE>
__device__ __forceinline__ void convert_fp8x8(float* v, float inv, uint2* dst, int off, int tid) {
    for (int i = BASE; i < BASE + 8; i++) v[i] *= inv;
    uint32_t a = __builtin_amdgcn_cvt_pk_fp8_f32(v[BASE],   v[BASE+1], 0u,  false);
    a = __builtin_amdgcn_cvt_pk_fp8_f32(v[BASE+2], v[BASE+3], a,   true);
    uint32_t b = __builtin_amdgcn_cvt_pk_fp8_f32(v[BASE+4], v[BASE+5], 0u,  false);
    b = __builtin_amdgcn_cvt_pk_fp8_f32(v[BASE+6], v[BASE+7], b,   true);
    dst[off + tid] = make_uint2(a, b);
}

template <int TM, int TE>
__device__ void k1_sort(const int* __restrict__ topk_ids, uint32_t* __restrict__ sorted_output,
    uint2* __restrict__ slot_pairs, int n_groups)
{
    constexpr int TOP_K = 9, TOTAL = TM * TOP_K, N_ROUTED = TE - 1, TPT = (TM + 255) >> 8;
    __shared__ uint32_t s_data[TOTAL];
    __shared__ uint32_t s_counts[N_ROUTED];
    int tid = threadIdx.x;
    int r[8 * TPT];

    if (tid < N_ROUTED) s_counts[tid] = 0;
    for (int i = tid; i < TOTAL; i += 256) s_data[i] = (uint32_t)topk_ids[i];
    __syncthreads();

    for (int t = 0; t < TPT; t++) {
        int tok = tid + t * 256;
        if (tok < TM) for (int i = 0; i < 8; i++) r[t*8+i] = (int)s_data[tok*TOP_K+i];
    }
    __syncthreads();

    for (int t = 0; t < TPT; t++) {
        int tok = tid + t * 256;
        if (tok < TM) for (int i = 0; i < 8; i++) atomicAdd(&s_counts[r[t*8+i]], 1);
    }
    __syncthreads();

    if constexpr (TE == 33) {
        if (tid < 32) {
            uint32_t val = s_counts[tid], orig = val;
            for (int d = 1; d < 32; d *= 2) { uint32_t n = __shfl_up(val, d); if (tid >= d) val += n; }
            s_counts[tid] = val - orig;
        }
    } else if constexpr (TE == 257) {
        if (tid < 64) {
            uint4 ld = ((const uint4*)s_counts)[tid];
            uint32_t c0=ld.x, c1=ld.y, c2=ld.z, c3=ld.w;
            uint32_t s0=0, s1=c0, s2=c0+c1, s3=c0+c1+c2, tt=s3+c3, val=tt;
            for (int d = 1; d < 64; d *= 2) { uint32_t n = __shfl_up(val, d); if (tid >= d) val += n; }
            uint32_t to = val - tt;
            ((uint4*)s_counts)[tid] = make_uint4(s0+to, s1+to, s2+to, s3+to);
        }
    }
    __syncthreads();

    for (int t = 0; t < TPT; t++) {
        int tok = tid + t * 256;
        if (tok < TM) for (int i = 0; i < 8; i++) {
            int eid = r[t*8+i];
            int pos = atomicAdd(&s_counts[eid], 1);
            sorted_output[pos] = PACK_PAIR(eid, tok);
        }
    }
    constexpr int SB = TM * 8;
    for (int tok = tid; tok < TM; tok += 256)
        sorted_output[SB + tok] = PACK_PAIR(TE-1, tok);

    if (tid < n_groups) {
        if (tid < n_groups-1) slot_pairs[tid] = make_uint2(tid*2, tid*2+2);
        else                  slot_pairs[tid] = make_uint2(tid*2+1, TOTAL);
    }
}

template <int TM>
__device__ void k1_convert(const uint16_t* __restrict__ hidden, uint8_t* __restrict__ hidden_fp8,
    float* __restrict__ scales, int d_hidden)
{
    __shared__ float s_reduce[4];
    __shared__ float s_scale;
    int tid = threadIdx.x;
    int token = blockIdx.x - 1;
    const uint4* src = (const uint4*)(hidden + token * d_hidden);
    float v[32], lm = 0;

    uint4 c0 = src[tid], c1 = src[256+tid], c2 = src[512+tid];
    uint4 c3; if (tid < 128) c3 = src[768+tid];

    lm = fmaxf(lm, unpack_f16x8_maxabs<0>(c0, v));
    lm = fmaxf(lm, unpack_f16x8_maxabs<8>(c1, v));
    lm = fmaxf(lm, unpack_f16x8_maxabs<16>(c2, v));
    if (tid < 128) lm = fmaxf(lm, unpack_f16x8_maxabs<24>(c3, v));

    for (int off = 32; off > 0; off >>= 1) lm = fmaxf(lm, __shfl_xor(lm, off));
    int warp_id = tid >> 6, lane_id = tid & 63;
    if (lane_id == 0) s_reduce[warp_id] = lm;
    __syncthreads();

    if (tid == 0) {
        float bm = 0; for (int i = 0; i < 4; i++) bm = fmaxf(bm, s_reduce[i]);
        s_scale = bm > 0 ? bm * (1.0f / 448.0f) : 1.0f;
    }
    __syncthreads();

    float inv = 1.0f / s_scale;
    uint2* dst = (uint2*)(hidden_fp8 + token * d_hidden);
    convert_fp8x8<0>(v, inv, dst, 0, tid);
    convert_fp8x8<8>(v, inv, dst, 256, tid);
    convert_fp8x8<16>(v, inv, dst, 512, tid);
    if (tid < 128) convert_fp8x8<24>(v, inv, dst, 768, tid);
    if (tid == 0) scales[token] = s_scale;
}

template <int TM>
__device__ void k1_zerofill(uint8_t* __restrict__ hidden_fp8, float* __restrict__ scales, int d_hidden) {
    int tid = threadIdx.x;
    uint4* zd = (uint4*)(hidden_fp8 + TM * d_hidden);
    uint4 z4 = {0,0,0,0};
    zd[tid] = z4; if (tid < 192) zd[256+tid] = z4;
    if (tid == 0) scales[TM] = 1.0f;
}

template <int TM, int TE>
__global__ __launch_bounds__(768)
void moe_k1(const uint16_t* __restrict__ hidden, uint8_t* __restrict__ hidden_fp8,
    float* __restrict__ scales, const int* __restrict__ topk_ids,
    uint32_t* __restrict__ sorted_output, uint2* __restrict__ slot_pairs,
    int d_hidden, int n_groups)
{
    if (blockIdx.x == 0)         k1_sort<TM, TE>(topk_ids, sorted_output, slot_pairs, n_groups);
    else if ((int)blockIdx.x <= TM) k1_convert<TM>(hidden, hidden_fp8, scales, d_hidden);
    else                         k1_zerofill<TM>(hidden_fp8, scales, d_hidden);
}

template __global__ void moe_k1<16,33>(const uint16_t*,uint8_t*,float*,const int*,uint32_t*,uint2*,int,int);
template __global__ void moe_k1<16,257>(const uint16_t*,uint8_t*,float*,const int*,uint32_t*,uint2*,int,int);
template __global__ void moe_k1<128,33>(const uint16_t*,uint8_t*,float*,const int*,uint32_t*,uint2*,int,int);
template __global__ void moe_k1<128,257>(const uint16_t*,uint8_t*,float*,const int*,uint32_t*,uint2*,int,int);
template __global__ void moe_k1<512,33>(const uint16_t*,uint8_t*,float*,const int*,uint32_t*,uint2*,int,int);
template __global__ void moe_k1<512,257>(const uint16_t*,uint8_t*,float*,const int*,uint32_t*,uint2*,int,int);


/* ══════════════════════════════════════════════════════════════
 * K2 GEMM Benchmark — ASM Primitives
 * ══════════════════════════════════════════════════════════════ */

template<int BSTART>
__device__ __forceinline__ void mfma_vaa_zero(v8i32 a, uint32_t sa, uint32_t sb) {
    asm volatile(
        "v_mfma_scale_f32_16x16x128_f8f6f4 a[0:3], %0, a[%1:%2], 0, %3, %4 op_sel_hi:[0,0,0] blgp:4"
        : : "v"(a), "n"(BSTART), "n"(BSTART+3), "v"(sa), "v"(sb)
        : "a0", "a1", "a2", "a3");
}

template<int BSTART, int IDX_SEL>
__device__ __forceinline__ void mfma_vaa(v8i32 a, uint32_t sa, uint32_t sb) {
    if constexpr (IDX_SEL == 0) {
        asm volatile(
            "v_mfma_scale_f32_16x16x128_f8f6f4 a[0:3], %0, a[%1:%2], a[0:3], %3, %4 op_sel_hi:[0,0,0] blgp:4"
            : : "v"(a), "n"(BSTART), "n"(BSTART+3), "v"(sa), "v"(sb)
            : "a0", "a1", "a2", "a3");
    } else if constexpr (IDX_SEL == 1) {
        asm volatile(
            "v_mfma_scale_f32_16x16x128_f8f6f4 a[0:3], %0, a[%1:%2], a[0:3], %3, %4 op_sel:[0,1,0] op_sel_hi:[0,0,0] blgp:4"
            : : "v"(a), "n"(BSTART), "n"(BSTART+3), "v"(sa), "v"(sb)
            : "a0", "a1", "a2", "a3");
    } else if constexpr (IDX_SEL == 2) {
        asm volatile(
            "v_mfma_scale_f32_16x16x128_f8f6f4 a[0:3], %0, a[%1:%2], a[0:3], %3, %4 op_sel_hi:[0,1,0] blgp:4"
            : : "v"(a), "n"(BSTART), "n"(BSTART+3), "v"(sa), "v"(sb)
            : "a0", "a1", "a2", "a3");
    } else if constexpr (IDX_SEL == 3) {
        asm volatile(
            "v_mfma_scale_f32_16x16x128_f8f6f4 a[0:3], %0, a[%1:%2], a[0:3], %3, %4 op_sel:[0,1,0] op_sel_hi:[0,1,0] blgp:4"
            : : "v"(a), "n"(BSTART), "n"(BSTART+3), "v"(sa), "v"(sb)
            : "a0", "a1", "a2", "a3");
    }
}

template<int IDX_SEL>
__device__ __forceinline__ void mfma_vva(v8i32 a, v4i32 b, uint32_t sa, uint32_t sb) {
    if constexpr (IDX_SEL == 0) {
        asm volatile(
            "v_mfma_scale_f32_16x16x128_f8f6f4 a[0:3], %0, %1, a[0:3], %2, %3 op_sel_hi:[0,0,0] blgp:4"
            : : "v"(a), "v"(b), "v"(sa), "v"(sb)
            : "a0", "a1", "a2", "a3");
    } else if constexpr (IDX_SEL == 1) {
        asm volatile(
            "v_mfma_scale_f32_16x16x128_f8f6f4 a[0:3], %0, %1, a[0:3], %2, %3 op_sel:[0,1,0] op_sel_hi:[0,0,0] blgp:4"
            : : "v"(a), "v"(b), "v"(sa), "v"(sb)
            : "a0", "a1", "a2", "a3");
    } else if constexpr (IDX_SEL == 2) {
        asm volatile(
            "v_mfma_scale_f32_16x16x128_f8f6f4 a[0:3], %0, %1, a[0:3], %2, %3 op_sel_hi:[0,1,0] blgp:4"
            : : "v"(a), "v"(b), "v"(sa), "v"(sb)
            : "a0", "a1", "a2", "a3");
    } else if constexpr (IDX_SEL == 3) {
        asm volatile(
            "v_mfma_scale_f32_16x16x128_f8f6f4 a[0:3], %0, %1, a[0:3], %2, %3 op_sel:[0,1,0] op_sel_hi:[0,1,0] blgp:4"
            : : "v"(a), "v"(b), "v"(sa), "v"(sb)
            : "a0", "a1", "a2", "a3");
    }
}


/* ══════════════════════════════════════════════════════════════
 * Load A operand from LDS — matches k2_engine.hip addressing
 *
 * base = ACT_BUF_OFF + m_lane * ACT_ROW_STRIDE + kg_lane * 32
 * Each thread reads two ds_read_b128: lo (16B) and hi (16B at +16)
 * giving 32 bytes per thread, 128 bytes per step across 4 kg_lanes.
 * ══════════════════════════════════════════════════════════════ */

template<int STEP>
__device__ __forceinline__ v8i32 load_a(const uint8_t* lds, int base) {
    constexpr int OFF = STEP * ACT_STEP;
    const uint4* lo = (const uint4*)(lds + base + OFF);
    const uint4* hi = (const uint4*)(lds + base + OFF + 16);
    uint4 lo_val = *lo;
    uint4 hi_val = *hi;
    v8i32 a;
    a[0] = (int)lo_val.x; a[1] = (int)lo_val.y; a[2] = (int)lo_val.z; a[3] = (int)lo_val.w;
    a[4] = (int)hi_val.x; a[5] = (int)hi_val.y; a[6] = (int)hi_val.z; a[7] = (int)hi_val.w;
    return a;
}


/* ══════════════════════════════════════════════════════════════
 * Triple-Buffered Runners — BASELINE (no yield)
 * ══════════════════════════════════════════════════════════════ */

template<int S, int END>
struct VaaRunner {
    static __device__ __forceinline__ void run(
        uint32_t* sb, v8i32 a0, v8i32 a1, v8i32 a2,
        const uint8_t* lds, int base, float& ys)
    {
        if constexpr (S == 0) mfma_vaa_zero<WT_AREG + S*4>(a0, 0x7F7F7F7F, sb[S>>2]);
        else                  mfma_vaa<WT_AREG + S*4, S&3>(a0, 0x7F7F7F7F, sb[S>>2]);
        if constexpr (S+3 < END) a0 = load_a<S+3>(lds, base);

        constexpr int S1 = S+1;
        if constexpr (S1 < END) {
            mfma_vaa<WT_AREG+S1*4, S1&3>(a1, 0x7F7F7F7F, sb[S1>>2]);
            if constexpr (S1+3 < END) a1 = load_a<S1+3>(lds, base);
        }
        constexpr int S2 = S+2;
        if constexpr (S2 < END) {
            mfma_vaa<WT_AREG+S2*4, S2&3>(a2, 0x7F7F7F7F, sb[S2>>2]);
            if constexpr (S2+3 < END) a2 = load_a<S2+3>(lds, base);
        }
        if constexpr (S+3 < END) VaaRunner<S+3, END>::run(sb, a0, a1, a2, lds, base, ys);
    }
};
template<int N> struct VaaRunner<N,N> {
    static __device__ __forceinline__ void run(uint32_t*, v8i32, v8i32, v8i32, const uint8_t*, int, float&) {}
};

template<int S, int END>
struct VvaRunner {
    static __device__ __forceinline__ void run(
        v4i32* w_v, uint32_t* sb, v8i32 a0, v8i32 a1, v8i32 a2,
        const uint8_t* lds, int base, float& ys)
    {
        mfma_vva<S&3>(a0, w_v[S-WT_AGPR_STEPS], 0x7F7F7F7F, sb[S>>2]);
        if constexpr (S+3 < END) a0 = load_a<S+3>(lds, base);

        constexpr int S1 = S+1;
        if constexpr (S1 < END) {
            mfma_vva<S1&3>(a1, w_v[S1-WT_AGPR_STEPS], 0x7F7F7F7F, sb[S1>>2]);
            if constexpr (S1+3 < END) a1 = load_a<S1+3>(lds, base);
        }
        constexpr int S2 = S+2;
        if constexpr (S2 < END) {
            mfma_vva<S2&3>(a2, w_v[S2-WT_AGPR_STEPS], 0x7F7F7F7F, sb[S2>>2]);
            if constexpr (S2+3 < END) a2 = load_a<S2+3>(lds, base);
        }
        if constexpr (S+3 < END) VvaRunner<S+3, END>::run(w_v, sb, a0, a1, a2, lds, base, ys);
    }
};
template<int N> struct VvaRunner<N,N> {
    static __device__ __forceinline__ void run(v4i32*, uint32_t*, v8i32, v8i32, v8i32, const uint8_t*, int, float&) {}
};


/* ══════════════════════════════════════════════════════════════
 * Half-LDS load: only lo read from LDS, hi is constant
 * ══════════════════════════════════════════════════════════════ */

template<int STEP>
__device__ __forceinline__ v8i32 load_a_half(const uint8_t* lds, int base) {
    constexpr int OFF = STEP * ACT_STEP;
    const uint4* lo = (const uint4*)(lds + base + OFF);
    uint4 lo_val = *lo;
    v8i32 a;
    a[0] = (int)lo_val.x; a[1] = (int)lo_val.y; a[2] = (int)lo_val.z; a[3] = (int)lo_val.w;
    a[4] = 0x12345678; a[5] = 0x23456789; a[6] = (int)0x3456789A; a[7] = 0x456789AB;
    return a;
}

template<int S, int END>
struct VaaHalf {
    static __device__ __forceinline__ void run(
        uint32_t* sb, v8i32 a0, v8i32 a1, v8i32 a2,
        const uint8_t* lds, int base, float& ys)
    {
        if constexpr (S == 0) mfma_vaa_zero<WT_AREG + S*4>(a0, 0x7F7F7F7F, sb[S>>2]);
        else                  mfma_vaa<WT_AREG + S*4, S&3>(a0, 0x7F7F7F7F, sb[S>>2]);
        if constexpr (S+3 < END) a0 = load_a_half<S+3>(lds, base);

        constexpr int S1 = S+1;
        if constexpr (S1 < END) {
            mfma_vaa<WT_AREG+S1*4, S1&3>(a1, 0x7F7F7F7F, sb[S1>>2]);
            if constexpr (S1+3 < END) a1 = load_a_half<S1+3>(lds, base);
        }
        constexpr int S2 = S+2;
        if constexpr (S2 < END) {
            mfma_vaa<WT_AREG+S2*4, S2&3>(a2, 0x7F7F7F7F, sb[S2>>2]);
            if constexpr (S2+3 < END) a2 = load_a_half<S2+3>(lds, base);
        }
        if constexpr (S+3 < END) VaaHalf<S+3, END>::run(sb, a0, a1, a2, lds, base, ys);
    }
};
template<int N> struct VaaHalf<N,N> {
    static __device__ __forceinline__ void run(uint32_t*, v8i32, v8i32, v8i32, const uint8_t*, int, float&) {}
};

template<int S, int END>
struct VvaHalf {
    static __device__ __forceinline__ void run(
        v4i32* w_v, uint32_t* sb, v8i32 a0, v8i32 a1, v8i32 a2,
        const uint8_t* lds, int base, float& ys)
    {
        mfma_vva<S&3>(a0, w_v[S-WT_AGPR_STEPS], 0x7F7F7F7F, sb[S>>2]);
        if constexpr (S+3 < END) a0 = load_a_half<S+3>(lds, base);

        constexpr int S1 = S+1;
        if constexpr (S1 < END) {
            mfma_vva<S1&3>(a1, w_v[S1-WT_AGPR_STEPS], 0x7F7F7F7F, sb[S1>>2]);
            if constexpr (S1+3 < END) a1 = load_a_half<S1+3>(lds, base);
        }
        constexpr int S2 = S+2;
        if constexpr (S2 < END) {
            mfma_vva<S2&3>(a2, w_v[S2-WT_AGPR_STEPS], 0x7F7F7F7F, sb[S2>>2]);
            if constexpr (S2+3 < END) a2 = load_a_half<S2+3>(lds, base);
        }
        if constexpr (S+3 < END) VvaHalf<S+3, END>::run(w_v, sb, a0, a1, a2, lds, base, ys);
    }
};
template<int N> struct VvaHalf<N,N> {
    static __device__ __forceinline__ void run(v4i32*, uint32_t*, v8i32, v8i32, v8i32, const uint8_t*, int, float&) {}
};


/* ══════════════════════════════════════════════════════════════
 * Baseline kernel: MFMA + LDS + sync (no yield) — same as Test 1
 * ══════════════════════════════════════════════════════════════ */

extern "C"
__global__ __launch_bounds__(512)
void k2_baseline(
    volatile float* __restrict__ out,
    volatile uint64_t* __restrict__ cycles_out)
{
    __shared__ uint8_t lds[ACT_BUF_SIZE];

    int tid  = threadIdx.x;
    int lane = tid & 63;
    int m_lane  = lane & 15;
    int kg_lane = lane >> 4;

    for (int i = tid; i < ACT_BUF_SIZE / 4; i += THREADS)
        ((uint32_t*)lds)[i] = 0xDEADBEEF;
    __syncthreads();

    uint32_t fill = 0x12345678;
    #pragma unroll
    for (int i = WT_AREG; i < WT_AREG + 112; i++)
        asm volatile("v_accvgpr_write_b32 a[%0], %1" : : "n"(i), "v"(fill));

    v4i32 w_v[28];
    v4i32 nz = {0x12345678, 0x23456789, 0x3456789A, 0x456789AB};
    #pragma unroll
    for (int i = 0; i < 28; i++) w_v[i] = nz;

    uint32_t sb[14];
    #pragma unroll
    for (int i = 0; i < 14; i++) sb[i] = 0x7F7F7F7F;

    int base = ACT_BUF_OFF + m_lane * ACT_ROW_STRIDE + (kg_lane << 5);
    float dce0 = 0, dce1 = 0, dce2 = 0, dce3 = 0;
    float ys = 0;
    uint64_t elapsed = 0;

    #pragma unroll 1
    for (int pass = 0; pass < 3; pass++) {
        __syncthreads();
        uint64_t t0 = clock64();

        #pragma unroll 1
        for (int iter = 0; iter < N_ITERS; iter++) {
            {
                v8i32 a0 = load_a<0>(lds, base);
                v8i32 a1 = load_a<1>(lds, base);
                v8i32 a2 = load_a<2>(lds, base);
                VaaRunner<0, 28>::run(sb, a0, a1, a2, lds, base, ys);
            }
            __syncthreads();
            {
                v8i32 a0 = load_a<28>(lds, base);
                v8i32 a1 = load_a<29>(lds, base);
                v8i32 a2 = load_a<30>(lds, base);
                VvaRunner<28, 56>::run(w_v, sb, a0, a1, a2, lds, base, ys);
            }
            {
                float r0, r1, r2, r3;
                asm volatile("v_accvgpr_read_b32 %0, a[0]" : "=v"(r0));
                asm volatile("v_accvgpr_read_b32 %0, a[1]" : "=v"(r1));
                asm volatile("v_accvgpr_read_b32 %0, a[2]" : "=v"(r2));
                asm volatile("v_accvgpr_read_b32 %0, a[3]" : "=v"(r3));
                dce0 += r0; dce1 += r1; dce2 += r2; dce3 += r3;
            }
            __syncthreads();
        }

        uint64_t t1 = clock64();
        if (pass == 2) elapsed = t1 - t0;
    }

    int boff = blockIdx.x * THREADS * 4;
    out[boff + tid * 4 + 0] = dce0 + ys;
    out[boff + tid * 4 + 1] = dce1;
    out[boff + tid * 4 + 2] = dce2;
    out[boff + tid * 4 + 3] = dce3;
    if (tid == 0) cycles_out[blockIdx.x] = elapsed;
}


/* ══════════════════════════════════════════════════════════════
 * Half-LDS kernel: MFMA + 1 ds_read per step + sync (50% LDS BW)
 * ══════════════════════════════════════════════════════════════ */

extern "C"
__global__ __launch_bounds__(512)
void k2_half_lds(
    volatile float* __restrict__ out,
    volatile uint64_t* __restrict__ cycles_out)
{
    __shared__ uint8_t lds[ACT_BUF_SIZE];

    int tid  = threadIdx.x;
    int lane = tid & 63;
    int m_lane  = lane & 15;
    int kg_lane = lane >> 4;

    for (int i = tid; i < ACT_BUF_SIZE / 4; i += THREADS)
        ((uint32_t*)lds)[i] = 0xDEADBEEF;
    __syncthreads();

    uint32_t fill = 0x12345678;
    #pragma unroll
    for (int i = WT_AREG; i < WT_AREG + 112; i++)
        asm volatile("v_accvgpr_write_b32 a[%0], %1" : : "n"(i), "v"(fill));

    v4i32 w_v[28];
    v4i32 nz = {0x12345678, 0x23456789, 0x3456789A, 0x456789AB};
    #pragma unroll
    for (int i = 0; i < 28; i++) w_v[i] = nz;

    uint32_t sb[14];
    #pragma unroll
    for (int i = 0; i < 14; i++) sb[i] = 0x7F7F7F7F;

    int base = ACT_BUF_OFF + m_lane * ACT_ROW_STRIDE + (kg_lane << 5);
    float dce0 = 0, dce1 = 0, dce2 = 0, dce3 = 0;
    float ys = 0;
    uint64_t elapsed = 0;

    #pragma unroll 1
    for (int pass = 0; pass < 3; pass++) {
        __syncthreads();
        uint64_t t0 = clock64();

        #pragma unroll 1
        for (int iter = 0; iter < N_ITERS; iter++) {
            {
                v8i32 a0 = load_a_half<0>(lds, base);
                v8i32 a1 = load_a_half<1>(lds, base);
                v8i32 a2 = load_a_half<2>(lds, base);
                VaaHalf<0, 28>::run(sb, a0, a1, a2, lds, base, ys);
            }
            __syncthreads();
            {
                v8i32 a0 = load_a_half<28>(lds, base);
                v8i32 a1 = load_a_half<29>(lds, base);
                v8i32 a2 = load_a_half<30>(lds, base);
                VvaHalf<28, 56>::run(w_v, sb, a0, a1, a2, lds, base, ys);
            }
            {
                float r0, r1, r2, r3;
                asm volatile("v_accvgpr_read_b32 %0, a[0]" : "=v"(r0));
                asm volatile("v_accvgpr_read_b32 %0, a[1]" : "=v"(r1));
                asm volatile("v_accvgpr_read_b32 %0, a[2]" : "=v"(r2));
                asm volatile("v_accvgpr_read_b32 %0, a[3]" : "=v"(r3));
                dce0 += r0; dce1 += r1; dce2 += r2; dce3 += r3;
            }
            __syncthreads();
        }

        uint64_t t1 = clock64();
        if (pass == 2) elapsed = t1 - t0;
    }

    int boff = blockIdx.x * THREADS * 4;
    out[boff + tid * 4 + 0] = dce0 + ys;
    out[boff + tid * 4 + 1] = dce1;
    out[boff + tid * 4 + 2] = dce2;
    out[boff + tid * 4 + 3] = dce3;
    if (tid == 0) cycles_out[blockIdx.x] = elapsed;
}


/* ══════════════════════════════════════════════════════════════
 * Full-LDS NO SYNC
 * ══════════════════════════════════════════════════════════════ */

extern "C"
__global__ __launch_bounds__(512)
void k2_full_nosync(
    volatile float* __restrict__ out,
    volatile uint64_t* __restrict__ cycles_out)
{
    __shared__ uint8_t lds[ACT_BUF_SIZE];
    int tid = threadIdx.x, lane = tid & 63, m_lane = lane & 15, kg_lane = lane >> 4;
    for (int i = tid; i < ACT_BUF_SIZE / 4; i += THREADS) ((uint32_t*)lds)[i] = 0xDEADBEEF;
    __syncthreads();
    uint32_t fill = 0x12345678;
    #pragma unroll
    for (int i = WT_AREG; i < WT_AREG + 112; i++) asm volatile("v_accvgpr_write_b32 a[%0], %1" : : "n"(i), "v"(fill));
    v4i32 w_v[28]; v4i32 nz = {0x12345678, 0x23456789, 0x3456789A, 0x456789AB};
    #pragma unroll
    for (int i = 0; i < 28; i++) w_v[i] = nz;
    uint32_t sb[14];
    #pragma unroll
    for (int i = 0; i < 14; i++) sb[i] = 0x7F7F7F7F;
    int base = ACT_BUF_OFF + m_lane * ACT_ROW_STRIDE + (kg_lane << 5);
    float dce0=0, dce1=0, dce2=0, dce3=0, ys=0;
    uint64_t elapsed = 0;
    #pragma unroll 1
    for (int pass = 0; pass < 3; pass++) {
        uint64_t t0 = clock64();
        __syncthreads();
        #pragma unroll 1
        for (int iter = 0; iter < N_ITERS; iter++) {
            { v8i32 a0=load_a<0>(lds,base), a1=load_a<1>(lds,base), a2=load_a<2>(lds,base);
              VaaRunner<0,28>::run(sb, a0, a1, a2, lds, base, ys); }
            { v8i32 a0=load_a<28>(lds,base), a1=load_a<29>(lds,base), a2=load_a<30>(lds,base);
              VvaRunner<28,56>::run(w_v, sb, a0, a1, a2, lds, base, ys); }
            { float r0,r1,r2,r3;
              asm volatile("v_accvgpr_read_b32 %0, a[0]" : "=v"(r0));
              asm volatile("v_accvgpr_read_b32 %0, a[1]" : "=v"(r1));
              asm volatile("v_accvgpr_read_b32 %0, a[2]" : "=v"(r2));
              asm volatile("v_accvgpr_read_b32 %0, a[3]" : "=v"(r3));
              dce0+=r0; dce1+=r1; dce2+=r2; dce3+=r3; }
        }
        __syncthreads();
        uint64_t t1 = clock64();
        if (pass == 2) elapsed = t1 - t0;
    }
    int boff = blockIdx.x * THREADS * 4;
    out[boff+tid*4+0]=dce0+ys; out[boff+tid*4+1]=dce1; out[boff+tid*4+2]=dce2; out[boff+tid*4+3]=dce3;
    if (tid == 0) cycles_out[blockIdx.x] = elapsed;
}


/* ══════════════════════════════════════════════════════════════
 * Half-LDS NO SYNC
 * ══════════════════════════════════════════════════════════════ */

extern "C"
__global__ __launch_bounds__(512)
void k2_half_nosync(
    volatile float* __restrict__ out,
    volatile uint64_t* __restrict__ cycles_out)
{
    __shared__ uint8_t lds[ACT_BUF_SIZE];
    int tid = threadIdx.x, lane = tid & 63, m_lane = lane & 15, kg_lane = lane >> 4;
    for (int i = tid; i < ACT_BUF_SIZE / 4; i += THREADS) ((uint32_t*)lds)[i] = 0xDEADBEEF;
    __syncthreads();
    uint32_t fill = 0x12345678;
    #pragma unroll
    for (int i = WT_AREG; i < WT_AREG + 112; i++) asm volatile("v_accvgpr_write_b32 a[%0], %1" : : "n"(i), "v"(fill));
    v4i32 w_v[28]; v4i32 nz = {0x12345678, 0x23456789, 0x3456789A, 0x456789AB};
    #pragma unroll
    for (int i = 0; i < 28; i++) w_v[i] = nz;
    uint32_t sb[14];
    #pragma unroll
    for (int i = 0; i < 14; i++) sb[i] = 0x7F7F7F7F;
    int base = ACT_BUF_OFF + m_lane * ACT_ROW_STRIDE + (kg_lane << 5);
    float dce0=0, dce1=0, dce2=0, dce3=0, ys=0;
    uint64_t elapsed = 0;
    #pragma unroll 1
    for (int pass = 0; pass < 3; pass++) {
        uint64_t t0 = clock64();
        __syncthreads();
        #pragma unroll 1
        for (int iter = 0; iter < N_ITERS; iter++) {
            { v8i32 a0=load_a_half<0>(lds,base), a1=load_a_half<1>(lds,base), a2=load_a_half<2>(lds,base);
              VaaHalf<0,28>::run(sb, a0, a1, a2, lds, base, ys); }
            { v8i32 a0=load_a_half<28>(lds,base), a1=load_a_half<29>(lds,base), a2=load_a_half<30>(lds,base);
              VvaHalf<28,56>::run(w_v, sb, a0, a1, a2, lds, base, ys); }
            { float r0,r1,r2,r3;
              asm volatile("v_accvgpr_read_b32 %0, a[0]" : "=v"(r0));
              asm volatile("v_accvgpr_read_b32 %0, a[1]" : "=v"(r1));
              asm volatile("v_accvgpr_read_b32 %0, a[2]" : "=v"(r2));
              asm volatile("v_accvgpr_read_b32 %0, a[3]" : "=v"(r3));
              dce0+=r0; dce1+=r1; dce2+=r2; dce3+=r3; }
        }
        __syncthreads();
        uint64_t t1 = clock64();
        if (pass == 2) elapsed = t1 - t0;
    }
    int boff = blockIdx.x * THREADS * 4;
    out[boff+tid*4+0]=dce0+ys; out[boff+tid*4+1]=dce1; out[boff+tid*4+2]=dce2; out[boff+tid*4+3]=dce3;
    if (tid == 0) cycles_out[blockIdx.x] = elapsed;
}


/* ══════════════════════════════════════════════════════════════
 * Host
 * ══════════════════════════════════════════════════════════════ */

torch::Tensor run_bench(
    torch::Tensor hidden_states,
    torch::Tensor topk_ids,
    torch::Tensor gate_up_weight,
    torch::Tensor gate_up_scale,
    int M, int E, int d_expert, int d_hidden)
{
    int cu_tiles_per_expert = d_expert / 64;
    int n_groups = N_BLOCKS / cu_tiles_per_expert;

    fprintf(stderr, "\n=== K2 Half-LDS Test (1 vs 2 ds_read per step) ===\n");
    fprintf(stderr, "Config: M=%d, E=%d, d_expert=%d\n", M, E, d_expert);

    /* K1 outputs */
    uint8_t *d_fp8; float *d_scales; uint32_t *d_sorted; uint2 *d_slots;
    hipMalloc(&d_fp8, (M+1)*d_hidden);
    hipMalloc(&d_scales, (M+1)*sizeof(float));
    hipMalloc(&d_sorted, M*9*sizeof(uint32_t));
    hipMalloc(&d_slots, n_groups*sizeof(uint2));

    #define LAUNCH_K1 do { \
        if(M==16&&E<=33)       moe_k1<16,33> <<<M+2,256>>>((const uint16_t*)hidden_states.data_ptr(),d_fp8,d_scales,(const int*)topk_ids.data_ptr(),d_sorted,d_slots,d_hidden,n_groups); \
        else if(M==128&&E<=33) moe_k1<128,33><<<M+2,256>>>((const uint16_t*)hidden_states.data_ptr(),d_fp8,d_scales,(const int*)topk_ids.data_ptr(),d_sorted,d_slots,d_hidden,n_groups); \
        else if(M==512&&E<=33) moe_k1<512,33><<<M+2,256>>>((const uint16_t*)hidden_states.data_ptr(),d_fp8,d_scales,(const int*)topk_ids.data_ptr(),d_sorted,d_slots,d_hidden,n_groups); \
        else if(M==16)         moe_k1<16,257><<<M+2,256>>>((const uint16_t*)hidden_states.data_ptr(),d_fp8,d_scales,(const int*)topk_ids.data_ptr(),d_sorted,d_slots,d_hidden,n_groups); \
        else if(M==128)        moe_k1<128,257><<<M+2,256>>>((const uint16_t*)hidden_states.data_ptr(),d_fp8,d_scales,(const int*)topk_ids.data_ptr(),d_sorted,d_slots,d_hidden,n_groups); \
        else                   moe_k1<512,257><<<M+2,256>>>((const uint16_t*)hidden_states.data_ptr(),d_fp8,d_scales,(const int*)topk_ids.data_ptr(),d_sorted,d_slots,d_hidden,n_groups); \
    } while(0)

    /* K1 timing */
    for (int w = 0; w < 3; w++) { LAUNCH_K1; hipDeviceSynchronize(); }
    hipEvent_t ev0, ev1;
    hipEventCreate(&ev0); hipEventCreate(&ev1);
    float k1_total = 0;
    for (int run = 0; run < 10; run++) {
        hipEventRecord(ev0); LAUNCH_K1; hipEventRecord(ev1); hipDeviceSynchronize();
        float ms; hipEventElapsedTime(&ms, ev0, ev1); k1_total += ms;
    }
    fprintf(stderr, "K1 avg: %.1f us\n\n", k1_total / 10 * 1000.0f);

    /* Bench outputs */
    float *d_out;
    uint64_t *d_cycles;
    hipMalloc(&d_out, THREADS * 4 * sizeof(float));
    hipMalloc(&d_cycles, sizeof(uint64_t));

    auto run_variant = [&](const char* name, auto kernel_fn) {
        hipMemset(d_cycles, 0, sizeof(uint64_t));
        kernel_fn<<<1, 512>>>(d_out, d_cycles);
        hipDeviceSynchronize();
        uint64_t h_cyc;
        hipMemcpy(&h_cyc, d_cycles, sizeof(uint64_t), hipMemcpyDeviceToHost);
        float cyc_iter = (float)h_cyc / N_ITERS;
        fprintf(stderr, "  %-30s %7lu total, %7.1f/iter, %5.1f/step, %.2f us/iter\n",
                name, (unsigned long)h_cyc, cyc_iter, cyc_iter/56.0f, cyc_iter/2400.0f);
    };

    fprintf(stderr, "Half-LDS test (1 CU, %d iters, nosync=no inner sync, bookended):\n", N_ITERS);
    run_variant("full LDS + sync:", k2_baseline);
    run_variant("full LDS nosync (bookend):", k2_full_nosync);
    run_variant("half LDS + sync:", k2_half_lds);
    run_variant("half LDS nosync (bookend):", k2_half_nosync);

    fprintf(stderr, "\n  Reference:\n");
    fprintf(stderr, "    MFMA+LDS sync (Test 1):    7344/iter  (100%% LDS BW)\n");
    fprintf(stderr, "    MFMA+LDS nosync (Test 3):  4156/iter\n");
    fprintf(stderr, "    MFMA only (Test 5):        3644/iter  (0%% LDS BW)\n");
    fprintf(stderr, "=== END ===\n\n");

    hipEventDestroy(ev0); hipEventDestroy(ev1);
    hipFree(d_fp8); hipFree(d_scales); hipFree(d_sorted); hipFree(d_slots);
    hipFree(d_out); hipFree(d_cycles);
    return hidden_states;
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("run_bench", &run_bench, "K1 + K2 GEMM-only benchmark");
}
"""

_module = None; _build_failed = False
def _get_module():
    global _module, _build_failed
    if _build_failed: return None
    if _module is not None: return _module
    try:
        build_dir = os.path.expanduser('~/.cache/k2_half_lds_v2')
        os.makedirs(build_dir, exist_ok=True)
        hip_file = os.path.join(build_dir, 'k2_half_lds.hip')
        so_file  = os.path.join(build_dir, 'k2_half_lds.so')
        with open(hip_file, 'w') as f: f.write(HIP_SOURCE)

        torch_dir  = os.path.dirname(torch.__file__)
        torch_inc  = os.path.join(torch_dir, 'include')
        torch_api  = os.path.join(torch_inc, 'torch', 'csrc', 'api', 'include')
        torch_lib  = os.path.join(torch_dir, 'lib')
        python_inc = sysconfig.get_path('include')
        try:    abi = int(torch._C._GLIBCXX_USE_CXX11_ABI)
        except: abi = 0

        cmd = [
            'hipcc', '-shared', '-fPIC', '--offload-arch=gfx950', '-O2', '-ffast-math',
            '-D__HIP_PLATFORM_AMD__', f'-D_GLIBCXX_USE_CXX11_ABI={abi}',
            '-DTORCH_API_INCLUDE_EXTENSION_H', '-DTORCH_EXTENSION_NAME=k2_half_lds',
            f'-I{torch_inc}', f'-I{torch_api}', f'-I{python_inc}',
            f'-L{torch_lib}', '-ltorch', '-ltorch_cpu', '-lc10',
            '-ltorch_hip', '-ltorch_python',
            f'-Wl,-rpath,{torch_lib}',
            hip_file, '-o', so_file,
        ]
        _log("Building K2 GEMM bench module...")
        r = subprocess.run(cmd, capture_output=True, text=True, timeout=600)
        _log(f"hipcc exit: {r.returncode}")
        if r.stderr:
            err_lines = [l for l in r.stderr.split('\n') if 'error' in l.lower() and 'warning' not in l.lower()]
            if err_lines:
                _log("ERRORS:")
                for l in err_lines[:20]: _log(f"  {l}")
            _log(f"stderr tail: {r.stderr[-500:]}")
        if r.returncode != 0:
            _build_failed = True; return None
        _log(f"Build OK! .so size: {os.path.getsize(so_file)}")

        import ctypes
        for lib in ['libc10.so','libtorch_cpu.so','libtorch.so',
                     'libtorch_hip.so','libtorch_python.so']:
            p = os.path.join(torch_lib, lib)
            if os.path.exists(p):
                try: ctypes.CDLL(p, mode=ctypes.RTLD_GLOBAL)
                except: pass

        spec = importlib.util.spec_from_file_location("k2_half_lds", so_file)
        mod  = importlib.util.module_from_spec(spec)
        spec.loader.exec_module(mod)
        _module = mod
        _log("K2 GEMM bench module loaded OK!")
        return _module
    except Exception as e:
        import traceback
        _log(f"BUILD FAILED: {traceback.format_exc()[-2000:]}")
        _build_failed = True
        return None


_probed = False
def custom_kernel(data: input_t) -> output_t:
    global _probed
    if not _probed:
        _probed = True
        mod = _get_module()
        if mod:
            (hs, guw, dw, gus, ds, guws, dws, guss, dss, tw, ti, cfg) = data
            M = cfg["bs"]
            E = cfg["n_routed_experts"] + cfg["n_shared_experts"]
            d_expert = cfg["d_expert_pad"]
            dh = cfg["d_hidden"]
            _log(f"Running: M={M}, E={E}, d_expert={d_expert}, dh={dh}")
            mod.run_bench(hs, ti, guw, gus, M, E, d_expert, dh)

    from aiter import ActivationType, QuantType
    from aiter.fused_moe import fused_moe as aiter_fused_moe

    (hs, guw, dw, gus, ds, guws, dws, guss, dss, tw, ti, cfg) = data
    return aiter_fused_moe(
        hs, guws, dws, tw, ti,
        expert_mask=None,
        activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32,
        doweight_stage1=False,
        w1_scale=guss, w2_scale=dss,
        a1_scale=None, a2_scale=None,
        hidden_pad=cfg["d_hidden_pad"] - cfg["d_hidden"],
        intermediate_pad=cfg["d_expert_pad"] - cfg["d_expert"],
    )
scrolls · 870 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