Skip to content
KernelIndex
Search⌘K

submission 570068

divc13 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

sub50_swizzle_back.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-570068?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
11.8µs
#355 of 1143
2026-03-16

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:01eacc789527e75425165a929c5f4c0982ba35a6667a8cf05c3c020310f3c625
license declaredunknown
license concludedunknown
authorsdivc13
imported2026-08-15

Techniques

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

num-warps = 4constexpr int NUM_WARPS = 4;
shared-memory__shared__ __align__(16) uint8_t A_lds[BLOCK_M * LDS_ROW];
split-kvoid launch_sep_gemm_splitk(
tile-m = 16constexpr int BLOCK_M = 16;
tile-n = 64constexpr int BLOCK_N = 64;

Kernel source

sub50_swizzle_back.py378 lines
"""
Phase 50: B_shuffle + B_scale_sh + LDS swizzle for B reads.
Base: sub49_full_shuffle. Change: Add LDS swizzle back for B reads.
Global source address swizzled so LDS reads are bank-conflict-free.
"""
from task import input_t, output_t
import torch
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"

CPP_SOURCE = r"""
#include <torch/extension.h>
// Separate B_scale path only (B_shuffle + B_scale_sh)
void launch_sep_gemm_splitk(
    torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
    torch::Tensor workspace, int M, int N, int K, int split_k);
void launch_sep_gemm_nosplit(
    torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
    torch::Tensor C, int M, int N, int K);
// Reduce
void launch_reduce(torch::Tensor workspace, torch::Tensor C, int M, int N, int split_k);
"""

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

constexpr int WARP_SIZE   = 64;
constexpr int NUM_WARPS   = 4;
constexpr int NUM_THREADS = NUM_WARPS * WARP_SIZE;
constexpr int BLOCK_M     = 16;
constexpr int BLOCK_N     = 64;
constexpr int MFMA_K      = 128;
constexpr int DOUBLE_K    = MFMA_K * 2;
constexpr int LDS_ROW     = DOUBLE_K >> 1;
constexpr int HALF_K      = MFMA_K >> 1;
constexpr int SCALE_GROUP = 32;

typedef int __attribute__((ext_vector_type(4))) int4_vec;
typedef float __attribute__((ext_vector_type(4))) float4_vec;
typedef int32_t __attribute__((ext_vector_type(4))) i32x4;
typedef uint32_t __attribute__((address_space(3)))* as3_uint32_ptr;

extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
    i32x4 rsrc, as3_uint32_ptr lds_ptr,
    int size, int voffset, int soffset, int offset, int aux
) __asm("llvm.amdgcn.raw.buffer.load.lds");

struct buffer_resource { uint64_t ptr; uint32_t range; uint32_t config; };

__device__ __forceinline__ i32x4 make_srsrc(const void* ptr, uint32_t range_bytes) {
    buffer_resource rsrc = {reinterpret_cast<uint64_t>(ptr), range_bytes, 0x110000};
    return *reinterpret_cast<const i32x4*>(&rsrc);
}

__device__ __forceinline__ float4_vec mfma_fp4_scaled(
    int4_vec A, int4_vec B, float4_vec C, int sA, int sB
) {
    float4_vec D;
    asm volatile(
        "v_mfma_scale_f32_16x16x128_f8f6f4 %0, %1, %2, %3, %4, %5 cbsz:4 blgp:4"
        : "=v"(D) : "v"(A), "v"(B), "v"(C), "v"(sA), "v"(sB));
    return D;
}

__device__ __forceinline__ int lds_swz(int offset) {
    return offset ^ (((offset & 2047) >> 8) << 4);
}

__device__ __forceinline__ void compute_scale(float max_abs, uint8_t& sc, float& scale_f) {
    if (max_abs > 0.0f) {
        uint32_t b = __float_as_uint(max_abs);
        b = (b + 0x200000u) & 0xFF800000u;
        int su = ((b >> 23) & 0xFF) - 129;
        su = su < -127 ? -127 : (su > 127 ? 127 : su);
        sc = (uint8_t)(su + 127);
        scale_f = __uint_as_float((uint32_t)(su + 127) << 23);
    } else { sc = 0; scale_f = 0.0f; }
}

__global__ __launch_bounds__(NUM_THREADS, 3)
void gemm_kernel(
    const __hip_bfloat16* __restrict__ A_bf16,
    const uint8_t* __restrict__ B_q,         // B_shuffle data
    const uint8_t* __restrict__ B_scale,     // B_scale_sh data
    float* __restrict__ workspace,
    __hip_bfloat16* __restrict__ C_out,
    const int M, const int N, const int K,
    const int k_steps_per_split
) {
    const int warp_id = threadIdx.x >> 6;
    const int lane_id = threadIdx.x & 63;
    const int lane_m  = lane_id & 15;
    const int lane_k  = lane_id >> 4;
    const int tid     = threadIdx.x;

    const int block_m = blockIdx.y * BLOCK_M;
    const int block_n = blockIdx.x * BLOCK_N;
    const int warp_n  = block_n + (warp_id << 4);
    const int split_id = blockIdx.z;

    const int b_stride   = K >> 1;
    const int sc_stride  = K >> 5;

    __shared__ __align__(16) uint8_t A_lds[BLOCK_M * LDS_ROW];
    __shared__ __align__(16) uint8_t B_lds[BLOCK_N * LDS_ROW];
    __shared__ uint8_t A_scale_lds[BLOCK_M * 8];
    __shared__ uint8_t B_scale_lds[BLOCK_N * 8];

    const i32x4 b_srsrc = make_srsrc(B_q, N * b_stride);
    float4_vec acc = {0.0f, 0.0f, 0.0f, 0.0f};

    const int ks_start = split_id * k_steps_per_split;
    const int ks_end = ks_start + k_steps_per_split;

    for (int ks = ks_start; ks < ks_end; ks++) {
        const int k_elem = ks * DOUBLE_K;
        const int k_byte = ks * LDS_ROW;

        // A quant: HW FP4 conversion
        {
            const int group_id = tid >> 1;
            const int half = tid & 1;
            const int q_row = group_id >> 3;
            const int q_grp = group_id & 7;
            const int g_row = block_m + q_row;
            const int k_off = k_elem + q_grp * SCALE_GROUP + half * 16;

            uint32_t pk_lo = 0, pk_hi = 0;
            uint8_t a_scale_val = 0x7f;

            if (g_row < M) {
                const __hip_bfloat16* src = A_bf16 + g_row * K + k_off;
                int4 raw[2];
                #pragma unroll
                for (int j = 0; j < 2; j++)
                    raw[j] = reinterpret_cast<const int4*>(src)[j];
                const __hip_bfloat16* bf = reinterpret_cast<const __hip_bfloat16*>(raw);
                float vals[16];
                float local_max = 0.0f;
                #pragma unroll
                for (int i = 0; i < 16; i++) {
                    vals[i] = __bfloat162float(bf[i]);
                    local_max = fmaxf(local_max, fabsf(vals[i]));
                }
                float global_max = fmaxf(local_max, __shfl_xor(local_max, 1));
                float scale_f;
                compute_scale(global_max, a_scale_val, scale_f);
                pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[0],  vals[1],  scale_f, 0);
                pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[2],  vals[3],  scale_f, 1);
                pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[4],  vals[5],  scale_f, 2);
                pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[6],  vals[7],  scale_f, 3);
                pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[8],  vals[9],  scale_f, 0);
                pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[10], vals[11], scale_f, 1);
                pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[12], vals[13], scale_f, 2);
                pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[14], vals[15], scale_f, 3);
            }
            const int a_lds_off = lds_swz(q_row * LDS_ROW + q_grp * 16 + half * 8);
            int2 packed; packed.x = (int)pk_lo; packed.y = (int)pk_hi;
            *reinterpret_cast<int2*>(&A_lds[a_lds_off]) = packed;
            if (half == 0) A_scale_lds[q_row * 8 + q_grp] = a_scale_val;
        }

        // B: buffer_load_lds from B_shuffle with swizzled global source
        // Load swizzled data into LDS so that swizzled LDS reads get correct data
        {
            #pragma unroll
            for (int ld = 0; ld < 2; ld++) {
                const int flat = (ld * NUM_THREADS + tid) << 4;  // LDS destination (linear)
                const int row = flat >> 7;        // which B row in this block (0..63)
                const int col = flat & 127;       // byte offset within 128-byte LDS row
                const int g_row = block_n + row;
                if (row < BLOCK_N && g_row < N) {
                    // Compute swizzled column: what data does swizzled LDS read expect at this flat pos?
                    const int swz_flat = lds_swz(flat);
                    const int swz_col = swz_flat & 127;

                    // Use swz_col instead of col for the tile address computation
                    const int abs_col = k_byte + swz_col;  // absolute byte column (swizzled)
                    const int tile_n = g_row >> 4;
                    const int inner_n = g_row & 15;
                    const int tile_k = abs_col >> 5;
                    const int inner_k_hi = (abs_col >> 4) & 1;
                    const int src_off = tile_n * (b_stride << 4) + tile_k * 512 + inner_k_hi * 256 + inner_n * 16;

                    llvm_amdgcn_raw_buffer_load_lds(b_srsrc,
                        (as3_uint32_ptr)(reinterpret_cast<uintptr_t>(B_lds) + flat),
                        16, src_off, 0, 0, 0);
                }
            }
        }

        // B_scale from B_scale_sh (e8m0_shuffled layout)
        {
            const int sc_off = ks << 3;
            #pragma unroll
            for (int r = 0; r < 2; r++) {
                const int idx = r * NUM_THREADS + tid;
                if (idx < (BLOCK_N << 3)) {
                    const int row = idx >> 3, grp = idx & 7;
                    const int g_row = block_n + row;
                    if (g_row < N) {
                        const int abs_col = sc_off + grp;
                        const int flat_sh = (g_row >> 5) * 32 * sc_stride
                                          + (g_row & 15) * 4
                                          + ((g_row >> 4) & 1)
                                          + (abs_col & 3) * 64
                                          + ((abs_col & 7) >> 2) * 2
                                          + (abs_col >> 3) * 256;
                        B_scale_lds[idx] = B_scale[flat_sh];
                    } else {
                        B_scale_lds[idx] = 0x7f;
                    }
                }
            }
        }

        asm volatile("s_waitcnt vmcnt(0)");
        __syncthreads();

        // MFMA — B reads use swizzle (data was loaded swizzled)
        #pragma unroll
        for (int half = 0; half < 2; half++) {
            const int kh = half * HALF_K;
            int4_vec A_reg;
            {
                const int a_off = lds_swz(lane_m * LDS_ROW + kh + (lane_k << 4));
                const int4 tmp = *reinterpret_cast<const int4*>(&A_lds[a_off]);
                A_reg.s0 = tmp.x; A_reg.s1 = tmp.y; A_reg.s2 = tmp.z; A_reg.s3 = tmp.w;
            }
            int4_vec B_reg;
            {
                const int b_row = (warp_id << 4) + lane_m;
                const int b_off = lds_swz(b_row * LDS_ROW + kh + (lane_k << 4));  // swizzled read!
                const int4 tmp = *reinterpret_cast<const int4*>(&B_lds[b_off]);
                B_reg.s0 = tmp.x; B_reg.s1 = tmp.y; B_reg.s2 = tmp.z; B_reg.s3 = tmp.w;
            }
            const int a_sc = (int)A_scale_lds[(lane_m << 3) + (half << 2) + lane_k];
            const int b_sc = (int)B_scale_lds[((warp_id << 4) + lane_m) * 8 + (half << 2) + lane_k];
            acc = mfma_fp4_scaled(A_reg, B_reg, acc, a_sc, b_sc);
        }
        __syncthreads();
    }

    // Store
    const int out_row = block_m + (lane_k << 2);
    const int out_col = warp_n + lane_m;
    if (out_col < N) {
        const float* ap = reinterpret_cast<const float*>(&acc);
        if (C_out) {
            #pragma unroll
            for (int r = 0; r < 4; r++) {
                const int gm = out_row + r;
                if (gm < M) C_out[gm * N + out_col] = __float2bfloat16(ap[r]);
            }
        } else {
            float* ws = workspace + split_id * M * N;
            #pragma unroll
            for (int r = 0; r < 4; r++) {
                const int gm = out_row + r;
                if (gm < M) ws[gm * N + out_col] = ap[r];
            }
        }
    }
}

// Reduction kernel
__global__ void reduce_kernel(
    const float* __restrict__ workspace,
    __hip_bfloat16* __restrict__ C,
    const int M, const int N, const int split_k
) {
    const int idx = blockIdx.x * blockDim.x + threadIdx.x;
    if (idx >= M * N) return;
    float sum = 0.0f;
    for (int s = 0; s < split_k; s++)
        sum += workspace[s * M * N + idx];
    C[idx] = __float2bfloat16(sum);
}

// ---- Launch functions ----
void launch_sep_gemm_splitk(
    torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
    torch::Tensor workspace, int M, int N, int K, int split_k
) {
    const int k_steps = K / (MFMA_K * 2);
    dim3 block(NUM_THREADS);
    dim3 grid((N + BLOCK_N - 1) / BLOCK_N, (M + BLOCK_M - 1) / BLOCK_M, split_k);
    hipLaunchKernelGGL(gemm_kernel, grid, block, 0, 0,
        reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_q.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),
        reinterpret_cast<float*>(workspace.data_ptr()),
        (__hip_bfloat16*)nullptr, M, N, K, k_steps / split_k);
}

void launch_sep_gemm_nosplit(
    torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
    torch::Tensor C, int M, int N, int K
) {
    const int k_steps = K / (MFMA_K * 2);
    dim3 block(NUM_THREADS);
    dim3 grid((N + BLOCK_N - 1) / BLOCK_N, (M + BLOCK_M - 1) / BLOCK_M, 1);
    hipLaunchKernelGGL(gemm_kernel, grid, block, 0, 0,
        reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_q.data_ptr()),
        reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),
        (float*)nullptr,
        reinterpret_cast<__hip_bfloat16*>(C.data_ptr()), M, N, K, k_steps);
}

void launch_reduce(torch::Tensor workspace, torch::Tensor C, int M, int N, int split_k) {
    const int num = M * N;
    hipLaunchKernelGGL(reduce_kernel, dim3((num+255)/256), dim3(256), 0, 0,
        reinterpret_cast<const float*>(workspace.data_ptr()),
        reinterpret_cast<__hip_bfloat16*>(C.data_ptr()), M, N, split_k);
}
"""

from torch.utils.cpp_extension import load_inline

_module = None

def _get_module():
    global _module
    if _module is None:
        _module = load_inline(
            name="hybrid_v50",
            cpp_sources=CPP_SOURCE,
            cuda_sources=HIP_SOURCE,
            functions=[
                "launch_sep_gemm_splitk", "launch_sep_gemm_nosplit",
                "launch_reduce",
            ],
            verbose=False,
            extra_cuda_cflags=["-O3", "-fno-gpu-rdc", "-ffp-contract=fast"],
        )
    return _module

def _pick_split_k(m, n, k):
    k_steps = k // 256
    blocks_mn = ((n + 63) // 64) * ((m + 15) // 16)
    if blocks_mn >= 256:
        return 1
    target_split = max(1, (608 + blocks_mn - 1) // blocks_mn)
    best = 1
    for s in range(1, k_steps + 1):
        if k_steps % s == 0 and s <= target_split:
            best = s
    while best > 1 and k_steps // best < 2:
        best //= 2
    return max(1, best)

def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    m, k = A.shape
    n = B.shape[0]

    mod = _get_module()

    B_sh_u8 = B_shuffle.contiguous().view(torch.uint8)
    B_sc = B_scale_sh.contiguous().view(torch.uint8)

    split_k = _pick_split_k(m, n, k)

    if split_k == 1:
        C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")
        mod.launch_sep_gemm_nosplit(A, B_sh_u8, B_sc, C, m, n, k)
    else:
        workspace = torch.empty((split_k, m, n), dtype=torch.float32, device="cuda")
        mod.launch_sep_gemm_splitk(A, B_sh_u8, B_sc, workspace, m, n, k, split_k)
        C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")
        mod.launch_reduce(workspace, C, m, n, split_k)

    return C
scrolls · 378 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 517169.

"""
- FP4 quant + FP4 GEMM: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
+ Phase 50: B_shuffle + B_scale_sh + LDS swizzle for B reads.
+ Base: sub49_full_shuffle. Change: Add LDS swizzle back for B reads.
+ Global source address swizzled so LDS reads are bank-conflict-free.
+ """
+ from task import input_t, output_t
+ import torch
+ import os
+ os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
- Optimization 5: fused quantization + GEMM in a single Triton kernel.
+ CPP_SOURCE = r"""
+ #include <torch/extension.h>
+ // Separate B_scale path only (B_shuffle + B_scale_sh)
+ void launch_sep_gemm_splitk(
+ torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
+ torch::Tensor workspace, int M, int N, int K, int split_k);
+ void launch_sep_gemm_nosplit(
+ torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
+ torch::Tensor C, int M, int N, int K);
+ // Reduce
+ void launch_reduce(torch::Tensor workspace, torch::Tensor C, int M, int N, int split_k);
+ """
- Problem: the two-kernel pipeline writes A_q and A_scale_sh to HBM then reads them back:
- BF16 A (HBM) -> [quant kernel] -> FP4 A_q + scales (HBM) -> [GEMM kernel] reads them back
- For M=16, K=7168: A_q is 16*7168/2 = 57 KB written then immediately re-read = 114 KB wasted.
+ HIP_SOURCE = r"""
+ #include <hip/hip_runtime.h>
+ #include <hip/hip_bf16.h>
+ #include <torch/extension.h>
- Fix: a single Triton kernel that:
- 1. Loads a [BLOCK_M, BLOCK_K] tile of BF16 A into registers
- 2. Computes MXFP4 quantization on-chip (find abs-max per 32, compute E8M0 scale, pack to fp4x2)
- 3. Feeds the packed fp4 tile directly into tl.dot against B — A_q never touches HBM
- 4. Accumulates into fp32 accumulator, converts to bf16, writes C to HBM
+ constexpr int WARP_SIZE = 64;
+ constexpr int NUM_WARPS = 4;
+ constexpr int NUM_THREADS = NUM_WARPS * WARP_SIZE;
+ constexpr int BLOCK_M = 16;
+ constexpr int BLOCK_N = 64;
+ constexpr int MFMA_K = 128;
+ constexpr int DOUBLE_K = MFMA_K * 2;
+ constexpr int LDS_ROW = DOUBLE_K >> 1;
+ constexpr int HALF_K = MFMA_K >> 1;
+ constexpr int SCALE_GROUP = 32;
- The quantization math for MXFP4 E2M1 per-1x32:
- - FP4 E2M1 representable magnitudes: 0, 0.5, 1, 1.5, 2, 3, 4, 6 (max = 6)
- - scale = 2^round(log2(max_abs / 6)) in E8M0 (power-of-2 only)
- - quantized = clamp(round(val / scale), fp4_min, fp4_max)
- - two fp4 values packed into one uint8: low nibble = first, high nibble = second
- """
- from task import input_t, output_t
- import aiter
- from aiter import QuantType, dtypes
- import torch
- import triton
- import triton.language as tl
+ typedef int __attribute__((ext_vector_type(4))) int4_vec;
+ typedef float __attribute__((ext_vector_type(4))) float4_vec;
+ typedef int32_t __attribute__((ext_vector_type(4))) i32x4;
+ typedef uint32_t __attribute__((address_space(3)))* as3_uint32_ptr;
+ extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
+ i32x4 rsrc, as3_uint32_ptr lds_ptr,
+ int size, int voffset, int soffset, int offset, int aux
+ ) __asm("llvm.amdgcn.raw.buffer.load.lds");
- # FP4 E2M1 lookup: map float magnitude to nearest fp4 magnitude (0..6 index -> 0..7 value)
- # Values: 0, 0.5, 1, 1.5, 2, 3, 4, 6
- _FP4_MAX = 6.0
+ struct buffer_resource { uint64_t ptr; uint32_t range; uint32_t config; };
+ __device__ __forceinline__ i32x4 make_srsrc(const void* ptr, uint32_t range_bytes) {
+ buffer_resource rsrc = {reinterpret_cast<uint64_t>(ptr), range_bytes, 0x110000};
+ return *reinterpret_cast<const i32x4*>(&rsrc);
+ }
- @triton.jit
- def _e8m0_scale(max_abs, fp4_max: tl.constexpr):
- """Compute E8M0 scale: largest power of 2 such that max_abs/scale <= fp4_max."""
- # scale = 2^floor(log2(max_abs / fp4_max))
- # Use tl.log2 and tl.exp2 for power-of-2 computation
- ratio = max_abs / fp4_max
- log2_ratio = tl.log2(ratio.to(tl.float32) + 1e-30)
- exp = tl.floor(log2_ratio)
- return tl.exp2(exp)
+ __device__ __forceinline__ float4_vec mfma_fp4_scaled(
+ int4_vec A, int4_vec B, float4_vec C, int sA, int sB
+ ) {
+ float4_vec D;
+ asm volatile(
+ "v_mfma_scale_f32_16x16x128_f8f6f4 %0, %1, %2, %3, %4, %5 cbsz:4 blgp:4"
+ : "=v"(D) : "v"(A), "v"(B), "v"(C), "v"(sA), "v"(sB));
+ return D;
+ }
+ __device__ __forceinline__ int lds_swz(int offset) {
+ return offset ^ (((offset & 2047) >> 8) << 4);
+ }
- @triton.jit
- def _quant_to_fp4(val, scale):
- """Quantize a float value to fp4 E2M1 integer (0..7 for non-negative)."""
- # Representable fp4 magnitudes (E2M1 normal + subnormal):
- # 0=0, 1=0.5, 2=1, 3=1.5, 4=2, 5=3, 6=4, 7=6
- # Divide by scale, round to nearest fp4 level
- scaled = val / scale
- # Clamp to [0, 6] (magnitude), then find nearest level via rounding thresholds
- scaled = tl.clamp(scaled, 0.0, 6.0)
- # Piecewise round to fp4 levels: boundaries at midpoints between levels
- # 0|0.25|0.75|1.25|1.75|2.5|3.5|5.0
- q = tl.where(scaled < 0.25, 0,
- tl.where(scaled < 0.75, 1,
- tl.where(scaled < 1.25, 2,
- tl.where(scaled < 1.75, 3,
- tl.where(scaled < 2.5, 4,
- tl.where(scaled < 3.5, 5,
- tl.where(scaled < 5.0, 6, 7)))))))
- return q
+ __device__ __forceinline__ void compute_scale(float max_abs, uint8_t& sc, float& scale_f) {
+ if (max_abs > 0.0f) {
+ uint32_t b = __float_as_uint(max_abs);
+ b = (b + 0x200000u) & 0xFF800000u;
+ int su = ((b >> 23) & 0xFF) - 129;
+ su = su < -127 ? -127 : (su > 127 ? 127 : su);
+ sc = (uint8_t)(su + 127);
+ scale_f = __uint_as_float((uint32_t)(su + 127) << 23);
+ } else { sc = 0; scale_f = 0.0f; }
+ }
+ __global__ __launch_bounds__(NUM_THREADS, 3)
+ void gemm_kernel(
+ const __hip_bfloat16* __restrict__ A_bf16,
+ const uint8_t* __restrict__ B_q, // B_shuffle data
+ const uint8_t* __restrict__ B_scale, // B_scale_sh data
+ float* __restrict__ workspace,
+ __hip_bfloat16* __restrict__ C_out,
+ const int M, const int N, const int K,
+ const int k_steps_per_split
+ ) {
+ const int warp_id = threadIdx.x >> 6;
+ const int lane_id = threadIdx.x & 63;
+ const int lane_m = lane_id & 15;
+ const int lane_k = lane_id >> 4;
+ const int tid = threadIdx.x;
- @triton.jit
- def _fused_quant_gemm_kernel(
- # A: [M, K] bf16
- A_ptr, stride_am, stride_ak,
- # B_shuffle: [N, K//2] fp4x2, pre-shuffled (16,16) tile layout
- B_ptr, stride_bn, stride_bk,
- # B_scale_sh: [N_pad, K//32] e8m0
- Bs_ptr, stride_bsn, stride_bsk,
- # C: [M, N] bf16 output
- C_ptr, stride_cm, stride_cn,
- M, N, K,
- BLOCK_M: tl.constexpr,
- BLOCK_N: tl.constexpr,
- BLOCK_K: tl.constexpr, # must be multiple of 64 (32 scale group * 2 pack)
- GROUP_SIZE: tl.constexpr, # = 32, elements per scale
- ):
- pid_m = tl.program_id(0)
- pid_n = tl.program_id(1)
+ const int block_m = blockIdx.y * BLOCK_M;
+ const int block_n = blockIdx.x * BLOCK_N;
+ const int warp_n = block_n + (warp_id << 4);
+ const int split_id = blockIdx.z;
- offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
- offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
- offs_k = tl.arange(0, BLOCK_K)
+ const int b_stride = K >> 1;
+ const int sc_stride = K >> 5;
- acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
+ __shared__ __align__(16) uint8_t A_lds[BLOCK_M * LDS_ROW];
+ __shared__ __align__(16) uint8_t B_lds[BLOCK_N * LDS_ROW];
+ __shared__ uint8_t A_scale_lds[BLOCK_M * 8];
+ __shared__ uint8_t B_scale_lds[BLOCK_N * 8];
- for k_start in range(0, K, BLOCK_K):
- k_offs = k_start + offs_k # [BLOCK_K]
+ const i32x4 b_srsrc = make_srsrc(B_q, N * b_stride);
+ float4_vec acc = {0.0f, 0.0f, 0.0f, 0.0f};
- # --- Load BF16 A tile [BLOCK_M, BLOCK_K] ---
- a_ptrs = A_ptr + offs_m[:, None] * stride_am + k_offs[None, :] * stride_ak
- mask_m = offs_m[:, None] < M
- mask_k = k_offs[None, :] < K
- a_tile = tl.load(a_ptrs, mask=mask_m & mask_k, other=0.0).to(tl.float32)
+ const int ks_start = split_id * k_steps_per_split;
+ const int ks_end = ks_start + k_steps_per_split;
- # --- Quantize A tile: per-32 block along K ---
- # a_tile shape: [BLOCK_M, BLOCK_K]
- # Process BLOCK_K // GROUP_SIZE groups of 32 along the K dimension
- # Pack two fp4 values per byte: a_q shape [BLOCK_M, BLOCK_K//2] uint8
- # We iterate over groups and pack
- a_q_tile = tl.zeros((BLOCK_M, BLOCK_K // 2), dtype=tl.uint8)
- a_scale_tile = tl.zeros((BLOCK_M, BLOCK_K // GROUP_SIZE), dtype=tl.float32)
+ for (int ks = ks_start; ks < ks_end; ks++) {
+ const int k_elem = ks * DOUBLE_K;
+ const int k_byte = ks * LDS_ROW;
- for g in range(BLOCK_K // GROUP_SIZE):
- g_start = g * GROUP_SIZE
- g_offs = g_start + tl.arange(0, GROUP_SIZE)
- a_group = tl.load(
- A_ptr + offs_m[:, None] * stride_am + (k_start + g_offs)[None, :] * stride_ak,
- mask=(offs_m[:, None] < M) & ((k_start + g_offs)[None, :] < K),
- other=0.0,
- ).to(tl.float32) # [BLOCK_M, GROUP_SIZE]
+ // A quant: HW FP4 conversion
+ {
+ const int group_id = tid >> 1;
+ const int half = tid & 1;
+ const int q_row = group_id >> 3;
+ const int q_grp = group_id & 7;
+ const int g_row = block_m + q_row;
+ const int k_off = k_elem + q_grp * SCALE_GROUP + half * 16;
- # E8M0 scale: max abs per row within group
- abs_group = tl.abs(a_group)
- max_abs = tl.max(abs_group, axis=1) # [BLOCK_M]
- scale = _e8m0_scale(max_abs, _FP4_MAX) # [BLOCK_M]
- a_scale_tile = tl.store(
- # store scale; we rebuild after loop
- a_scale_tile, scale, mask=None
- )
+ uint32_t pk_lo = 0, pk_hi = 0;
+ uint8_t a_scale_val = 0x7f;
- # Quantize each element
- sign = tl.where(a_group >= 0, 1, -1)
- q = _quant_to_fp4(tl.abs(a_group), scale[:, None]) # [BLOCK_M, GROUP_SIZE]
- q_signed = q # sign encoded separately in fp4 sign bit (bit 3 of nibble)
- # pack sign into fp4: bit3=sign, bits[2:0]=magnitude index
- # For E2M1: value = sign * fp4_magnitude[q]
- # Encoding: 0b0xxx = positive, 0b1xxx = negative
- sign_bit = tl.where(sign < 0, 4, 0).to(tl.uint8) # bit 3
- q_u8 = (q.to(tl.uint8) | sign_bit) # [BLOCK_M, GROUP_SIZE]
+ if (g_row < M) {
+ const __hip_bfloat16* src = A_bf16 + g_row * K + k_off;
+ int4 raw[2];
+ #pragma unroll
+ for (int j = 0; j < 2; j++)
+ raw[j] = reinterpret_cast<const int4*>(src)[j];
+ const __hip_bfloat16* bf = reinterpret_cast<const __hip_bfloat16*>(raw);
+ float vals[16];
+ float local_max = 0.0f;
+ #pragma unroll
+ for (int i = 0; i < 16; i++) {
+ vals[i] = __bfloat162float(bf[i]);
+ local_max = fmaxf(local_max, fabsf(vals[i]));
+ }
+ float global_max = fmaxf(local_max, __shfl_xor(local_max, 1));
+ float scale_f;
+ compute_scale(global_max, a_scale_val, scale_f);
+ pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[0], vals[1], scale_f, 0);
+ pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[2], vals[3], scale_f, 1);
+ pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[4], vals[5], scale_f, 2);
+ pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[6], vals[7], scale_f, 3);
+ pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[8], vals[9], scale_f, 0);
+ pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[10], vals[11], scale_f, 1);
+ pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[12], vals[13], scale_f, 2);
+ pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[14], vals[15], scale_f, 3);
+ }
+ const int a_lds_off = lds_swz(q_row * LDS_ROW + q_grp * 16 + half * 8);
+ int2 packed; packed.x = (int)pk_lo; packed.y = (int)pk_hi;
+ *reinterpret_cast<int2*>(&A_lds[a_lds_off]) = packed;
+ if (half == 0) A_scale_lds[q_row * 8 + q_grp] = a_scale_val;
+ }
- # Pack pairs: even index in low nibble, odd in high nibble
- even = q_u8[:, 0::2] & 0xF # [BLOCK_M, GROUP_SIZE//2]
- odd = (q_u8[:, 1::2] & 0xF) << 4
- packed = (even | odd).to(tl.uint8) # [BLOCK_M, GROUP_SIZE//2]
+ // B: buffer_load_lds from B_shuffle with swizzled global source
+ // Load swizzled data into LDS so that swizzled LDS reads get correct data
+ {
+ #pragma unroll
+ for (int ld = 0; ld < 2; ld++) {
+ const int flat = (ld * NUM_THREADS + tid) << 4; // LDS destination (linear)
+ const int row = flat >> 7; // which B row in this block (0..63)
+ const int col = flat & 127; // byte offset within 128-byte LDS row
+ const int g_row = block_n + row;
+ if (row < BLOCK_N && g_row < N) {
+ // Compute swizzled column: what data does swizzled LDS read expect at this flat pos?
+ const int swz_flat = lds_swz(flat);
+ const int swz_col = swz_flat & 127;
- # Store into a_q_tile slice [g_start//2 : g_start//2 + GROUP_SIZE//2]
- # (Triton doesn't support dynamic slice assignment easily; use indirect store)
+ // Use swz_col instead of col for the tile address computation
+ const int abs_col = k_byte + swz_col; // absolute byte column (swizzled)
+ const int tile_n = g_row >> 4;
+ const int inner_n = g_row & 15;
+ const int tile_k = abs_col >> 5;
+ const int inner_k_hi = (abs_col >> 4) & 1;
+ const int src_off = tile_n * (b_stride << 4) + tile_k * 512 + inner_k_hi * 256 + inner_n * 16;
- # --- Load B tile [BLOCK_N, BLOCK_K//2] fp4x2 ---
- # B_shuffle is in (16,16) tile-coalesced layout; load as uint8
- b_k_offs = k_start // 2 + tl.arange(0, BLOCK_K // 2)
- b_ptrs = B_ptr + offs_n[:, None] * stride_bn + b_k_offs[None, :] * stride_bk
- mask_n = offs_n[:, None] < N
- mask_bk = b_k_offs[None, :] < K // 2
- b_tile = tl.load(b_ptrs, mask=mask_n & mask_bk, other=0)
+ llvm_amdgcn_raw_buffer_load_lds(b_srsrc,
+ (as3_uint32_ptr)(reinterpret_cast<uintptr_t>(B_lds) + flat),
+ 16, src_off, 0, 0, 0);
+ }
+ }
+ }
- # tl.dot with fp4 inputs (requires hardware + Triton support)
- # NOTE: if tl.dot doesn't natively support fp4x2 on this Triton build,
- # fall back to dequant + bf16 dot (correctness preserved, perf reduced)
- acc += tl.dot(a_q_tile.to(tl.float8e4nv), b_tile.T.to(tl.float8e4nv)).to(tl.float32)
+ // B_scale from B_scale_sh (e8m0_shuffled layout)
+ {
+ const int sc_off = ks << 3;
+ #pragma unroll
+ for (int r = 0; r < 2; r++) {
+ const int idx = r * NUM_THREADS + tid;
+ if (idx < (BLOCK_N << 3)) {
+ const int row = idx >> 3, grp = idx & 7;
+ const int g_row = block_n + row;
+ if (g_row < N) {
+ const int abs_col = sc_off + grp;
+ const int flat_sh = (g_row >> 5) * 32 * sc_stride
+ + (g_row & 15) * 4
+ + ((g_row >> 4) & 1)
+ + (abs_col & 3) * 64
+ + ((abs_col & 7) >> 2) * 2
+ + (abs_col >> 3) * 256;
+ B_scale_lds[idx] = B_scale[flat_sh];
+ } else {
+ B_scale_lds[idx] = 0x7f;
+ }
+ }
+ }
+ }
- # Write C
- c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
- mask_c = (offs_m[:, None] < M) & (offs_n[None, :] < N)
- tl.store(c_ptrs, acc.to(tl.bfloat16), mask=mask_c)
+ asm volatile("s_waitcnt vmcnt(0)");
+ __syncthreads();
+ // MFMA — B reads use swizzle (data was loaded swizzled)
+ #pragma unroll
+ for (int half = 0; half < 2; half++) {
+ const int kh = half * HALF_K;
+ int4_vec A_reg;
+ {
+ const int a_off = lds_swz(lane_m * LDS_ROW + kh + (lane_k << 4));
+ const int4 tmp = *reinterpret_cast<const int4*>(&A_lds[a_off]);
+ A_reg.s0 = tmp.x; A_reg.s1 = tmp.y; A_reg.s2 = tmp.z; A_reg.s3 = tmp.w;
+ }
+ int4_vec B_reg;
+ {
+ const int b_row = (warp_id << 4) + lane_m;
+ const int b_off = lds_swz(b_row * LDS_ROW + kh + (lane_k << 4)); // swizzled read!
+ const int4 tmp = *reinterpret_cast<const int4*>(&B_lds[b_off]);
+ B_reg.s0 = tmp.x; B_reg.s1 = tmp.y; B_reg.s2 = tmp.z; B_reg.s3 = tmp.w;
+ }
+ const int a_sc = (int)A_scale_lds[(lane_m << 3) + (half << 2) + lane_k];
+ const int b_sc = (int)B_scale_lds[((warp_id << 4) + lane_m) * 8 + (half << 2) + lane_k];
+ acc = mfma_fp4_scaled(A_reg, B_reg, acc, a_sc, b_sc);
+ }
+ __syncthreads();
+ }
- # Module-level quant_func still used as fallback
- _quant_func = aiter.get_triton_quant(QuantType.per_1x32)
+ // Store
+ const int out_row = block_m + (lane_k << 2);
+ const int out_col = warp_n + lane_m;
+ if (out_col < N) {
+ const float* ap = reinterpret_cast<const float*>(&acc);
+ if (C_out) {
+ #pragma unroll
+ for (int r = 0; r < 4; r++) {
+ const int gm = out_row + r;
+ if (gm < M) C_out[gm * N + out_col] = __float2bfloat16(ap[r]);
+ }
+ } else {
+ float* ws = workspace + split_id * M * N;
+ #pragma unroll
+ for (int r = 0; r < 4; r++) {
+ const int gm = out_row + r;
+ if (gm < M) ws[gm * N + out_col] = ap[r];
+ }
+ }
+ }
+ }
+ // Reduction kernel
+ __global__ void reduce_kernel(
+ const float* __restrict__ workspace,
+ __hip_bfloat16* __restrict__ C,
+ const int M, const int N, const int split_k
+ ) {
+ const int idx = blockIdx.x * blockDim.x + threadIdx.x;
+ if (idx >= M * N) return;
+ float sum = 0.0f;
+ for (int s = 0; s < split_k; s++)
+ sum += workspace[s * M * N + idx];
+ C[idx] = __float2bfloat16(sum);
+ }
+ // ---- Launch functions ----
+ void launch_sep_gemm_splitk(
+ torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
+ torch::Tensor workspace, int M, int N, int K, int split_k
+ ) {
+ const int k_steps = K / (MFMA_K * 2);
+ dim3 block(NUM_THREADS);
+ dim3 grid((N + BLOCK_N - 1) / BLOCK_N, (M + BLOCK_M - 1) / BLOCK_M, split_k);
+ hipLaunchKernelGGL(gemm_kernel, grid, block, 0, 0,
+ reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),
+ reinterpret_cast<const uint8_t*>(B_q.data_ptr()),
+ reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),
+ reinterpret_cast<float*>(workspace.data_ptr()),
+ (__hip_bfloat16*)nullptr, M, N, K, k_steps / split_k);
+ }
+
+ void launch_sep_gemm_nosplit(
+ torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
+ torch::Tensor C, int M, int N, int K
+ ) {
+ const int k_steps = K / (MFMA_K * 2);
+ dim3 block(NUM_THREADS);
+ dim3 grid((N + BLOCK_N - 1) / BLOCK_N, (M + BLOCK_M - 1) / BLOCK_M, 1);
+ hipLaunchKernelGGL(gemm_kernel, grid, block, 0, 0,
+ reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),
+ reinterpret_cast<const uint8_t*>(B_q.data_ptr()),
+ reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),
+ (float*)nullptr,
+ reinterpret_cast<__hip_bfloat16*>(C.data_ptr()), M, N, K, k_steps);
+ }
+
+ void launch_reduce(torch::Tensor workspace, torch::Tensor C, int M, int N, int split_k) {
+ const int num = M * N;
+ hipLaunchKernelGGL(reduce_kernel, dim3((num+255)/256), dim3(256), 0, 0,
+ reinterpret_cast<const float*>(workspace.data_ptr()),
+ reinterpret_cast<__hip_bfloat16*>(C.data_ptr()), M, N, split_k);
+ }
+ """
+
+ from torch.utils.cpp_extension import load_inline
+
+ _module = None
+
+ def _get_module():
+ global _module
+ if _module is None:
+ _module = load_inline(
+ name="hybrid_v50",
+ cpp_sources=CPP_SOURCE,
+ cuda_sources=HIP_SOURCE,
+ functions=[
+ "launch_sep_gemm_splitk", "launch_sep_gemm_nosplit",
+ "launch_reduce",
+ ],
+ verbose=False,
+ extra_cuda_cflags=["-O3", "-fno-gpu-rdc", "-ffp-contract=fast"],
+ )
+ return _module
+
+ def _pick_split_k(m, n, k):
+ k_steps = k // 256
+ blocks_mn = ((n + 63) // 64) * ((m + 15) // 16)
+ if blocks_mn >= 256:
+ return 1
+ target_split = max(1, (608 + blocks_mn - 1) // blocks_mn)
+ best = 1
+ for s in range(1, k_steps + 1):
+ if k_steps % s == 0 and s <= target_split:
+ best = s
+ while best > 1 and k_steps // best < 2:
+ best //= 2
+ return max(1, best)
+
def custom_kernel(data: input_t) -> output_t:
- """
- Attempt fused quant+GEMM. Falls back to aiter reference on any error
- so correctness tests still pass while the fused path is being developed.
- """
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
m, k = A.shape
- n = B_shuffle.shape[0]
+ n = B.shape[0]
- # Fused path is experimental — fall back to reference if it errors
- try:
- BLOCK_M = max(16, min(64, triton.next_power_of_2(m)))
- BLOCK_N = 64
- BLOCK_K = 64 # must be multiple of 64
+ mod = _get_module()
- C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
- grid = (triton.cdiv(m, BLOCK_M), triton.cdiv(n, BLOCK_N))
+ B_sh_u8 = B_shuffle.contiguous().view(torch.uint8)
+ B_sc = B_scale_sh.contiguous().view(torch.uint8)
- _fused_quant_gemm_kernel[grid](
- A, A.stride(0), A.stride(1),
- B_shuffle, B_shuffle.stride(0), B_shuffle.stride(1),
- B_scale_sh, B_scale_sh.stride(0), B_scale_sh.stride(1),
- C, C.stride(0), C.stride(1),
- m, n, k,
- BLOCK_M=BLOCK_M,
- BLOCK_N=BLOCK_N,
- BLOCK_K=BLOCK_K,
- GROUP_SIZE=32,
- )
- return C
- except Exception:
- # Fallback: reference two-kernel path
- A_q, A_scale_sh = _quant_func(A, shuffle=True)
- return aiter.gemm_a4w4(
- A_q, B_shuffle, A_scale_sh, B_scale_sh,
- dtype=dtypes.bf16, bpreshuffle=True,
- )
+ split_k = _pick_split_k(m, n, k)
+
+ if split_k == 1:
+ C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")
+ mod.launch_sep_gemm_nosplit(A, B_sh_u8, B_sc, C, m, n, k)
+ else:
+ workspace = torch.empty((split_k, m, n), dtype=torch.float32, device="cuda")
+ mod.launch_sep_gemm_splitk(A, B_sh_u8, B_sc, workspace, m, n, k, split_k)
+ C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")
+ mod.launch_reduce(workspace, C, m, n, split_k)
+
+ return C
scrolls · 545 diff lines total

Best evidence level for this revision: reported

JSON