Skip to content
KernelIndex
Search⌘K

submission 79939

kimhyeonho_xel · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-79939?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVFP4 GEMVsuite of 3 cases
NVIDIA B200
248.1µs
#582 of 678
2025-11-16

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a514c7daffb92e24a26bc414924e0b955cefb9dfd3d333c45074d7ad9f04c11e
license declaredunknown
license concludedunknown
authorskimhyeonho_xel
imported2026-08-26

Techniques

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

fp4"NVFP4 batched GEMV optimized (4-row, warp-reduction, shared-b/sfb) kernel");
shared-memory__shared__ uint8_t s_b[MAX_K2];

Kernel source

submission.py389 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

# Scaling factor vector size (문제 스펙과 맞춤)
sf_vec_size = 16

_ext = None


def _load_ext():
    global _ext
    if _ext is not None:
        return _ext

    cpp_src = r"""
#include <torch/extension.h>

// FFMA 기반 4-row-per-block 커널 런처
at::Tensor batched_gemv_launcher(
    at::Tensor a_u8,
    at::Tensor b_u8,
    at::Tensor sfa_h,
    at::Tensor sfb_h,
    at::Tensor c_h
);

// (옵션) WGMMA/Tensor Core 기반 타일 커널 스켈레톤 런처
// 현재는 구현되지 않고, 필요하면 나중에 구현 후 Python에서 따로 호출 가능.
at::Tensor batched_gemv_wgmma_skeleton(
    at::Tensor a_u8,
    at::Tensor b_u8,
    at::Tensor sfa_h,
    at::Tensor sfb_h,
    at::Tensor c_h
);

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("batched_gemv", &batched_gemv_launcher,
          "NVFP4 batched GEMV optimized (4-row, warp-reduction, shared-b/sfb) kernel");

    m.def("batched_gemv_wgmma", &batched_gemv_wgmma_skeleton,
          "NVFP4 batched GEMV WGMMA-skeleton kernel (currently falls back to FFMA kernel)");
}
"""

    cuda_src = r"""
#include <torch/extension.h>
#include <cuda_fp16.h>
#include <math.h>

using at::Tensor;

// 문제 스펙 상:
//  - K <= 16384 → K2 = K/2 <= 8192
//  - SF_K = K/16 <= 1024
#define MAX_K2   8192
#define MAX_SF_K 1024
#define M_TILE   4      // block 당 처리하는 row 수

// ---- NVFP4(E2M1) nibble → float LUT (16개) ----
// 상위 비트(3)가 sign, 하위 3비트가 magnitude index.
// mag_table: [0, 0.5, 1, 1.5, 2, 3, 4, 6]
// 0..7  : +mag_table[idx]
// 8..15 : -mag_table[idx]
__device__ __constant__ float NVFP4_LUT[16] = {
    0.0f,  0.5f,  1.0f,  1.5f,  2.0f,  3.0f,  4.0f,  6.0f,
    0.0f, -0.5f, -1.0f, -1.5f, -2.0f, -3.0f, -4.0f, -6.0f
};

__device__ inline float nvfp4_nibble_to_float(uint8_t code) {
    return NVFP4_LUT[code & 0xF];
}

/* ============================================================
 * 1. 메인 커널
 *    - 4-row per block + shared b/sfb + warp-level reduction
 *    - nibble→float LUT
 *    - b FP4 디코드를 row 간 공유 (한 번 디코드해서 4 row 재사용)
 * ============================================================ */

__global__ void nvfp4_batched_gemv_4row_kernel(
    const uint8_t* __restrict__ a,    // [M, K2, L]
    const uint8_t* __restrict__ b,    // [1, K2, L]
    const __half*  __restrict__ sfa,  // [M, SF_K, L]
    const __half*  __restrict__ sfb,  // [1, SF_K, L]
    __half*        __restrict__ c,    // [M, 1, L]
    int M, int K2, int L, int SF_K
) {
    int tile_m = blockIdx.x;         // tile index along M
    int l      = blockIdx.y;         // batch index
    int tid    = threadIdx.x;

    if (l >= L) return;

    int m_base = tile_m * M_TILE;

    const int warpSize_ = 32;
    int lane_id   = tid & (warpSize_ - 1);   // 0..31
    int warp_id   = tid >> 5;                // tid / 32
    int num_warps = blockDim.x >> 5;         // blockDim.x / 32

    // shared: b[0,:,l], sfb[0,:,l]
    __shared__ uint8_t s_b[MAX_K2];
    __shared__ __half  s_sfb[MAX_SF_K];

    // warp reduction용 shared (동적 shared 메모리)
    extern __shared__ float s_warp_partial[];  // 크기 = M_TILE * num_warps

    if (tid == 0) {
        if (K2 > MAX_K2 || SF_K > MAX_SF_K) {
            // 스펙 밖이면 아무 것도 하지 않고 리턴
        }
    }
    __syncthreads();

    // ---- 1) b, sfb를 shared memory에 로드 ----
    // b: [1, K2, L] → index = k2 * L + l
    for (int k2 = tid; k2 < K2; k2 += blockDim.x) {
        int b_idx = k2 * L + l;
        s_b[k2] = b[b_idx];
    }

    // sfb: [1, SF_K, L] → index = blk * L + l
    for (int blk = tid; blk < SF_K; blk += blockDim.x) {
        int sfb_idx = blk * L + l;
        s_sfb[blk] = sfb[sfb_idx];
    }

    __syncthreads();

    // ---- 2) 각 thread가 M_TILE개의 row에 대한 partial sum(float)을 갖는다 ----
    float acc[M_TILE];
    #pragma unroll
    for (int r = 0; r < M_TILE; ++r) {
        acc[r] = 0.0f;
    }

    // a: [M, K2, L]   → ((m * K2) + k2) * L + l
    // sfa: [M, SF_K,L]→ ((m * SF_K) + blk) * L + l
    for (int blk = tid; blk < SF_K; blk += blockDim.x) {
        __half sfb_h = s_sfb[blk];
        float sb = __half2float(sfb_h);

        int base_k2 = blk * 8;  // 8 bytes per block → 16 FP4 값

        // row별 scale factor 곱(sa*sb)을 미리 계산
        float s_row[M_TILE];
        #pragma unroll
        for (int r = 0; r < M_TILE; ++r) {
            int m = m_base + r;
            if (m < M) {
                int sfa_idx = ((m * SF_K) + blk) * L + l;
                float sa = __half2float(sfa[sfa_idx]);
                s_row[r] = sa * sb;
            } else {
                s_row[r] = 0.0f;
            }
        }

        // 같은 blk에 대해 k2 바이트를 돈다.
        #pragma unroll
        for (int i_byte = 0; i_byte < 8; ++i_byte) {
            int k2 = base_k2 + i_byte;
            if (k2 >= K2) break;

            uint8_t b_byte = s_b[k2];

            // ✅ b nibble 디코드는 row loop 밖에서 한 번만
            float fb_low  = nvfp4_nibble_to_float(b_byte & 0xF);
            float fb_high = nvfp4_nibble_to_float((b_byte >> 4) & 0xF);

            // 각 row마다 a nibble만 디코드
            #pragma unroll
            for (int r = 0; r < M_TILE; ++r) {
                int m = m_base + r;
                if (m >= M) continue;

                int a_idx = ((m * K2) + k2) * L + l;
                uint8_t a_byte = a[a_idx];

                float fa_low  = nvfp4_nibble_to_float(a_byte & 0xF);
                float fa_high = nvfp4_nibble_to_float((a_byte >> 4) & 0xF);

                float s = s_row[r];

                acc[r] = __fmaf_rn(s * fa_low,  fb_low,  acc[r]);
                acc[r] = __fmaf_rn(s * fa_high, fb_high, acc[r]);
            }
        }
    }

    // ---- 3) row별 warp-level reduction → warp별 partial ----
    unsigned int full_mask = 0xffffffffu;

    // s_warp_partial는 [M_TILE, num_warps] 레이아웃으로 사용
    // index = row * num_warps + warp_id
    #pragma unroll
    for (int r = 0; r < M_TILE; ++r) {
        float warp_sum = acc[r];
        for (int offset = warpSize_ / 2; offset > 0; offset >>= 1) {
            warp_sum += __shfl_down_sync(full_mask, warp_sum, offset);
        }

        if (lane_id == 0) {
            int idx = r * num_warps + warp_id;
            s_warp_partial[idx] = warp_sum;
        }
    }

    __syncthreads();

    // ---- 4) warp 0이 row별로 warp-partial을 다시 warp reduction ----
    if (warp_id == 0) {
        #pragma unroll
        for (int r = 0; r < M_TILE; ++r) {
            float block_sum = 0.0f;
            if (lane_id < num_warps) {
                int idx = r * num_warps + lane_id;
                block_sum = s_warp_partial[idx];
            }

            for (int offset = warpSize_ / 2; offset > 0; offset >>= 1) {
                block_sum += __shfl_down_sync(full_mask, block_sum, offset);
            }

            if (lane_id == 0) {
                int m = m_base + r;
                if (m < M) {
                    int c_idx = (m * L) + l; // [M,1,L] contiguous
                    c[c_idx] = __float2half(block_sum);
                }
            }
        }
    }
}

/* ============================================================
 * 2. WGMMA/Tensor Core용 타일 스켈레톤 커널 (현재 사용 안 함)
 * ============================================================ */

__global__ void nvfp4_batched_gemv_wgmma_tile_kernel(
    const uint8_t* __restrict__ a,    // [M, K2, L]
    const uint8_t* __restrict__ b,    // [1, K2, L]
    const __half*  __restrict__ sfa,  // [M, SF_K, L]
    const __half*  __restrict__ sfb,  // [1, SF_K, L]
    __half*        __restrict__ c,    // [M, 1, L]
    int M, int K2, int L, int SF_K
) {
    // 스켈레톤: 현재는 어디에서도 호출하지 않는다.
}

/* ============================================================
 * 3. C++ 런처
 * ============================================================ */

at::Tensor batched_gemv_launcher(
    at::Tensor a_u8,   // [M, K2, L], uint8
    at::Tensor b_u8,   // [1, K2, L], uint8
    at::Tensor sfa_h,  // [M, SF_K, L], fp16
    at::Tensor sfb_h,  // [1, SF_K, L], fp16
    at::Tensor c_h     // [M, 1, L], fp16 (output)
) {
    TORCH_CHECK(a_u8.is_cuda(), "a_u8 must be CUDA");
    TORCH_CHECK(b_u8.is_cuda(), "b_u8 must be CUDA");
    TORCH_CHECK(sfa_h.is_cuda(), "sfa_h must be CUDA");
    TORCH_CHECK(sfb_h.is_cuda(), "sfb_h must be CUDA");
    TORCH_CHECK(c_h.is_cuda(), "c_h must be CUDA");

    TORCH_CHECK(a_u8.dtype() == torch::kUInt8, "a_u8 must be uint8");
    TORCH_CHECK(b_u8.dtype() == torch::kUInt8, "b_u8 must be uint8");
    TORCH_CHECK(sfa_h.dtype() == torch::kFloat16, "sfa_h must be float16");
    TORCH_CHECK(sfb_h.dtype() == torch::kFloat16, "sfb_h must be float16");
    TORCH_CHECK(c_h.dtype() == torch::kFloat16, "c_h must be float16");

    TORCH_CHECK(a_u8.dim() == 3, "a_u8 must be [M,K2,L]");
    TORCH_CHECK(b_u8.dim() == 3, "b_u8 must be [1,K2,L]");
    TORCH_CHECK(sfa_h.dim() == 3, "sfa_h must be [M,SF_K,L]");
    TORCH_CHECK(sfb_h.dim() == 3, "sfb_h must be [1,SF_K,L]");
    TORCH_CHECK(c_h.dim() == 3, "c_h must be [M,1,L]");

    int64_t M    = a_u8.size(0);
    int64_t K2   = a_u8.size(1);
    int64_t L    = a_u8.size(2);
    int64_t SF_K = sfa_h.size(1);

    TORCH_CHECK(b_u8.size(1) == K2 && b_u8.size(2) == L,  "b_u8 shape mismatch");
    TORCH_CHECK(sfa_h.size(0) == M && sfa_h.size(2) == L, "sfa_h shape mismatch");
    TORCH_CHECK(sfb_h.size(1) == SF_K && sfb_h.size(2) == L, "sfb_h shape mismatch");
    TORCH_CHECK(c_h.size(0) == M && c_h.size(1) == 1 && c_h.size(2) == L,
                "c_h must be [M,1,L]");

    TORCH_CHECK(K2  <= MAX_K2,  "K2 exceeds MAX_K2");
    TORCH_CHECK(SF_K <= MAX_SF_K, "SF_K exceeds MAX_SF_K");

    // grid: tile_m × l
    int grid_x = (M + M_TILE - 1) / M_TILE;
    dim3 grid(grid_x, L, 1);

    // 256 threads/block (8 warps) – 가벼운 블록을 많이 띄우는 전략
    int block_size = 256;
    dim3 block(block_size, 1, 1);

    int num_warps = block_size / 32;
    size_t shmem_bytes = M_TILE * num_warps * sizeof(float);  // warp partial sums

    nvfp4_batched_gemv_4row_kernel<<<grid, block, shmem_bytes>>>(
        reinterpret_cast<const uint8_t*>(a_u8.data_ptr<uint8_t>()),
        reinterpret_cast<const uint8_t*>(b_u8.data_ptr<uint8_t>()),
        reinterpret_cast<const __half*>(sfa_h.data_ptr<at::Half>()),
        reinterpret_cast<const __half*>(sfb_h.data_ptr<at::Half>()),
        reinterpret_cast<__half*>(c_h.data_ptr<at::Half>()),
        static_cast<int>(M),
        static_cast<int>(K2),
        static_cast<int>(L),
        static_cast<int>(SF_K)
    );

    return c_h;
}

// WGMMA 스켈레톤 런처: 현재는 그냥 FFMA 커널을 재사용.
// (Python 쪽에서 안 쓰므로 성능 영향 없음)
at::Tensor batched_gemv_wgmma_skeleton(
    at::Tensor a_u8,
    at::Tensor b_u8,
    at::Tensor sfa_h,
    at::Tensor sfb_h,
    at::Tensor c_h
) {
    return batched_gemv_launcher(a_u8, b_u8, sfa_h, sfb_h, c_h);
}
"""

    _ext = load_inline(
        name="nvfp4_batched_gemv_ext",
        cpp_sources=[cpp_src],
        cuda_sources=[cuda_src],
        extra_cuda_cflags=[
            "-gencode=arch=compute_100,code=sm_100",
            "-gencode=arch=compute_100,code=compute_100",
            "-O3",
        ],
        verbose=False,
    )
    return _ext


def custom_kernel(data: input_t) -> output_t:
    """
    NVFP4 block-scaled GEMV
      – 4-row-per-block + shared b/sfb + warp-level reduction
      – nibble→float LUT
      – b FP4 디코드를 row 간 공유 (디코드 오버헤드 감소)
    """

    a_ref, b_ref, sfa_ref_cpu, sfb_ref_cpu, _, _, c_ref = data

    device = a_ref.device
    assert device.type == "cuda", "Inputs must be on CUDA"

    M, _, L = c_ref.shape
    K2 = a_ref.size(1)
    K = K2 * 2
    SF_K = (K + sf_vec_size - 1) // sf_vec_size  # 문제 조건상 K는 16의 배수

    # NVFP4 데이터: float4_e2m1fn_x2 -> uint8 view
    a_u8 = a_ref.view(torch.uint8).contiguous()          # [M, K2, L]
    b_u8_full = b_ref.view(torch.uint8).contiguous()     # [128, K2, L]
    b_u8 = b_u8_full[:1, :, :].contiguous()              # [1, K2, L]

    # 스케일 팩터: float8_e4m3fn -> fp16
    sfa_f8 = sfa_ref_cpu.to(device)
    sfb_f8_full = sfb_ref_cpu.to(device)

    sfa_h = sfa_f8.to(torch.float16).contiguous()            # [M, SF_K, L]
    sfb_h_full = sfb_f8_full.to(torch.float16).contiguous()  # [128, SF_K, L]
    sfb_h = sfb_h_full[:1, :, :].contiguous()                # [1, SF_K, L]

    # Output: [M,1,L] contiguous (커널에서 이 가정)
    c_out = torch.zeros((M, 1, L), dtype=torch.float16, device=device)

    ext = _load_ext()

    # 항상 FFMA 커널 경로 사용
    c_res = ext.batched_gemv(a_u8, b_u8, sfa_h, sfb_h, c_out)

    return c_res
scrolls · 389 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