Skip to content
KernelIndex
Search⌘K

submission 107168

XoTic · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-107168?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
225.3µs
#579 of 678
2025-11-26

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:32fad00e4950d0bc4f00fb0df952a9a24fde3f2359aadc7e1fcff702772c26bb
license declaredunknown
license concludedunknown
authorsXoTic
imported2026-08-26

Techniques

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

async-copy"cp.async.cg.shared.global [%0], [%1], 16;\n"
shared-memoryextern __shared__ half b_shared[];
vector-width = uint4const uint4 a_vec = *reinterpret_cast<const uint4*>(&row_ptr[byte_off]);

Kernel source

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

cuda_source = r"""
#include <cuda_runtime.h>
#include <cuda_fp16.h>

__device__ __constant__ float fp4_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__ __forceinline__ size_t row_scale_pointer(int m, int rest_m, int rest_k, int L) {
    const int block = m >> 7;
    const int idx32 = m & 31;
    const int idx4 = (m >> 5) & 0x3;
    const size_t base = static_cast<size_t>(((idx32 * 4 + idx4) * rest_m + block) * 4);
    return base * rest_k * L;
}

__device__ __forceinline__ int get_scale_idx_b(int kb, int rest_k, int L, int l) {
    const int kk = kb >> 2;
    const int kk4 = kb & 0x3;
    return (kk4 * rest_k + kk) * L + l;
}

__global__ void decode_b_kernel(
    const uint8_t* __restrict__ b_packed,
    const half* __restrict__ sfb_scales,
    half* __restrict__ b_decoded,
    int K,
    int L,
    int rest_k
) {
    const int K_half = K >> 1;
    const int total_bytes = K_half * L;
    for (int idx = blockIdx.x * blockDim.x + threadIdx.x; idx < total_bytes; idx += blockDim.x * gridDim.x) {
        const int byte_in_k = idx / L;
        const int l = idx - byte_in_k * L;
        const uint8_t packed = b_packed[idx];
        const int k = byte_in_k << 1;
        const int kb = k >> 4;
        const float scale = __half2float(__ldg(&sfb_scales[get_scale_idx_b(kb, rest_k, L, l)]));
        const float v0 = fp4_lut[packed & 0xF] * scale;
        const float v1 = fp4_lut[packed >> 4] * scale;
        const size_t base = static_cast<size_t>(l) * K + k;
        b_decoded[base] = __float2half(v0);
        b_decoded[base + 1] = __float2half(v1);
    }
}

// Double-buffered L=1 kernel with cp.async
// CHUNK_SIZE is the number of half elements per chunk (must be multiple of 8 for cp.async.cg alignment)
constexpr int L1_CHUNK_SIZE = 2048;

__global__ void __launch_bounds__(256) gemv_L1_kernel(
    const uint8_t* __restrict__ a_packed,
    const half* __restrict__ b_decoded,
    const half* __restrict__ sfa_scales,
    half* __restrict__ c,
    int M, int K,
    int rest_m, int rest_k
) {
    // Double buffer: two chunks of L1_CHUNK_SIZE halfs each
    extern __shared__ half b_shared[];
    half* buf0 = b_shared;
    half* buf1 = b_shared + L1_CHUNK_SIZE;

    const int tid = threadIdx.x;
    const int warp_id = tid >> 5;
    const int lane_id = tid & 31;
    constexpr int ROWS_PER_BLOCK = 8;
    const int m = blockIdx.x * ROWS_PER_BLOCK + warp_id;

    const int K_half = K >> 1;
    const int num_chunks = (K + L1_CHUNK_SIZE - 1) / L1_CHUNK_SIZE;

    // Prefetch first chunk into buf0
    const int elems_chunk0 = (L1_CHUNK_SIZE < K) ? L1_CHUNK_SIZE : K;
    // Use cp.async with 16-byte (8 half) granularity
    for (int idx = tid * 8; idx < elems_chunk0; idx += blockDim.x * 8) {
        if (idx + 8 <= elems_chunk0) {
            asm volatile(
                "cp.async.cg.shared.global [%0], [%1], 16;\n"
                :
                : "r"(static_cast<unsigned>(__cvta_generic_to_shared(&buf0[idx]))),
                  "l"(&b_decoded[idx])
            );
        }
    }
    asm volatile("cp.async.commit_group;\n");

    if (m >= M) {
        // Still need to participate in async completion
        for (int chunk = 0; chunk < num_chunks; ++chunk) {
            asm volatile("cp.async.wait_group 0;\n");
            __syncthreads();
            if (chunk + 1 < num_chunks) {
                asm volatile("cp.async.commit_group;\n");
            }
            __syncthreads();
        }
        return;
    }

    const uint8_t* row_ptr = a_packed + static_cast<size_t>(m) * K_half;
    const half* row_scales = sfa_scales + row_scale_pointer(m, rest_m, rest_k, 1);

    float sum = 0.0f;

    for (int chunk = 0; chunk < num_chunks; ++chunk) {
        const int k_chunk_start = chunk * L1_CHUNK_SIZE;
        const int k_chunk_end = ((k_chunk_start + L1_CHUNK_SIZE) < K) ? (k_chunk_start + L1_CHUNK_SIZE) : K;
        const int chunk_elems = k_chunk_end - k_chunk_start;

        // Prefetch next chunk into alternate buffer
        if (chunk + 1 < num_chunks) {
            const int next_k_start = (chunk + 1) * L1_CHUNK_SIZE;
            const int next_k_end = ((next_k_start + L1_CHUNK_SIZE) < K) ? (next_k_start + L1_CHUNK_SIZE) : K;
            const int next_elems = next_k_end - next_k_start;
            half* next_buf = (chunk & 1) ? buf0 : buf1;

            for (int idx = tid * 8; idx < next_elems; idx += blockDim.x * 8) {
                if (idx + 8 <= next_elems) {
                    asm volatile(
                        "cp.async.cg.shared.global [%0], [%1], 16;\n"
                        :
                        : "r"(static_cast<unsigned>(__cvta_generic_to_shared(&next_buf[idx]))),
                          "l"(&b_decoded[next_k_start + idx])
                    );
                }
            }
            asm volatile("cp.async.commit_group;\n");
        }

        // Wait for current chunk
        asm volatile("cp.async.wait_group 1;\n");
        __syncthreads();

        // Compute on current chunk
        half* cur_buf = (chunk & 1) ? buf1 : buf0;
        const int byte_start = k_chunk_start >> 1;
        const int byte_end = k_chunk_end >> 1;

        for (int byte_off = byte_start + (lane_id << 4); byte_off < byte_end; byte_off += 512) {
            const int bytes_to_read = ((byte_off + 16) <= byte_end) ? 16 : (byte_end - byte_off);
            if (bytes_to_read < 16) break;  // Skip partial reads for simplicity

            const uint4 a_vec = *reinterpret_cast<const uint4*>(&row_ptr[byte_off]);
            const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a_vec);

            const int k_start = byte_off << 1;
            const int kb0 = k_start >> 4;
            const int kb1 = kb0 + 1;

            const int kk0 = kb0 >> 2;
            const int kk4_0 = kb0 & 0x3;
            const int kk1 = kb1 >> 2;
            const int kk4_1 = kb1 & 0x3;

            const float sa0 = __half2float(__ldg(&row_scales[kk4_0 * rest_k + kk0]));
            const float sa1 = __half2float(__ldg(&row_scales[kk4_1 * rest_k + kk1]));

            #pragma unroll
            for (int i = 0; i < 8; ++i) {
                const uint8_t ab = a_bytes[i];
                const int k = k_start + (i << 1);
                const int k_local = k - k_chunk_start;
                const float a0 = fp4_lut[ab & 0xF] * sa0;
                const float a1 = fp4_lut[ab >> 4] * sa0;
                const float b0 = __half2float(cur_buf[k_local]);
                const float b1 = __half2float(cur_buf[k_local + 1]);
                sum = __fmaf_rn(a0, b0, sum);
                sum = __fmaf_rn(a1, b1, sum);
            }

            #pragma unroll
            for (int i = 8; i < 16; ++i) {
                const uint8_t ab = a_bytes[i];
                const int k = k_start + (i << 1);
                const int k_local = k - k_chunk_start;
                const float a0 = fp4_lut[ab & 0xF] * sa1;
                const float a1 = fp4_lut[ab >> 4] * sa1;
                const float b0 = __half2float(cur_buf[k_local]);
                const float b1 = __half2float(cur_buf[k_local + 1]);
                sum = __fmaf_rn(a0, b0, sum);
                sum = __fmaf_rn(a1, b1, sum);
            }
        }

        __syncthreads();
    }

    // Final wait for any remaining async ops
    asm volatile("cp.async.wait_group 0;\n");

    #pragma unroll
    for (int offset = 16; offset > 0; offset >>= 1) {
        sum += __shfl_down_sync(0xffffffff, sum, offset);
    }

    if (lane_id == 0) {
        c[m] = __float2half(sum);
    }
}

template<int L, int ROWS_PER_BLOCK, int L_PER_WARP>
__global__ void __launch_bounds__(256) gemv_multi_kernel(
    const uint8_t* __restrict__ a_packed,
    const half* __restrict__ b_decoded,
    const half* __restrict__ sfa_scales,
    half* __restrict__ c,
    int M, int K,
    int rest_m, int rest_k
) {
    constexpr int CHUNKS_PER_ROW = L / L_PER_WARP;
    constexpr int KHALF_PER_VEC = 16 / L;

    const int tid = threadIdx.x;
    const int warp_id = tid >> 5;
    const int lane_id = tid & 31;

    const int row_in_block = warp_id / CHUNKS_PER_ROW;
    if (row_in_block >= ROWS_PER_BLOCK) return;
    const int chunk_id = warp_id % CHUNKS_PER_ROW;
    const int l_base = chunk_id * L_PER_WARP;

    const int m = blockIdx.x * ROWS_PER_BLOCK + row_in_block;
    if (m >= M) return;

    const int K_half = K >> 1;
    const int iter_bound = K_half / KHALF_PER_VEC;
    const size_t a_row_base = static_cast<size_t>(m) * K_half * L;
    const uint8_t* row_ptr = a_packed + a_row_base;
    const half* row_scales = sfa_scales + row_scale_pointer(m, rest_m, rest_k, L);

    const half* b_ptrs[L_PER_WARP];
    #pragma unroll
    for (int lp = 0; lp < L_PER_WARP; ++lp) {
        const int l = l_base + lp;
        b_ptrs[lp] = (l < L) ? (b_decoded + static_cast<size_t>(l) * K) : nullptr;
    }

    float sum[L_PER_WARP] = {0.0f};

    for (int vec_idx = lane_id; vec_idx < iter_bound; vec_idx += 32) {
        const int k_half_base = vec_idx * KHALF_PER_VEC;
        const uint4 a_vec = *reinterpret_cast<const uint4*>(&row_ptr[k_half_base * L]);
        const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a_vec);

        const int kb = (k_half_base << 1) >> 4;
        const int kk = kb >> 2;
        const int kk4 = kb & 0x3;

        float sa_vals[L_PER_WARP];
        #pragma unroll
        for (int lp = 0; lp < L_PER_WARP; ++lp) {
            const int l = l_base + lp;
            if (l < L) {
                const int scale_idx = (kk4 * rest_k + kk) * L + l;
                sa_vals[lp] = __half2float(__ldg(&row_scales[scale_idx]));
            } else {
                sa_vals[lp] = 0.0f;
            }
        }

        #pragma unroll
        for (int i = 0; i < KHALF_PER_VEC; ++i) {
            const int byte_off = i * L + l_base;
            const int k_pair = ((k_half_base + i) << 1);
            #pragma unroll
            for (int lp = 0; lp < L_PER_WARP; ++lp) {
                const int l = l_base + lp;
                if (l >= L) continue;
                const uint8_t ab = a_bytes[byte_off + lp];
                const float a0 = fp4_lut[ab & 0xF] * sa_vals[lp];
                const float a1 = fp4_lut[ab >> 4] * sa_vals[lp];
                const half* b_vec = b_ptrs[lp];
                const float b0 = __half2float(b_vec[k_pair]);
                const float b1 = __half2float(b_vec[k_pair + 1]);
                sum[lp] = __fmaf_rn(a0, b0, sum[lp]);
                sum[lp] = __fmaf_rn(a1, b1, sum[lp]);
            }
        }
    }

    #pragma unroll
    for (int lp = 0; lp < L_PER_WARP; ++lp) {
        const int l = l_base + lp;
        if (l >= L) continue;
        float val = sum[lp];
        #pragma unroll
        for (int offset = 16; offset > 0; offset >>= 1) {
            val += __shfl_down_sync(0xffffffff, val, offset);
        }
        if (lane_id == 0) {
            c[m * L + l] = __float2half(val);
        }
    }
}

extern "C" void decode_nvfp4_vector(
    const uint8_t* b_packed,
    const half* sfb_scales,
    half* b_decoded,
    int K, int L,
    int rest_k
) {
    const int threads = 256;
    const int total_pairs = (K >> 1) * L;
    int blocks = (total_pairs + threads - 1) / threads;
    blocks = blocks == 0 ? 1 : blocks;
    blocks = blocks > 65535 ? 65535 : blocks;
    decode_b_kernel<<<blocks, threads>>>(b_packed, sfb_scales, b_decoded, K, L, rest_k);
}

extern "C" void launch_nvfp4_kernel(
    const uint8_t* a_packed,
    const half* b_decoded,
    const half* sfa_scales,
    half* c,
    int M, int K, int L,
    int rest_m, int rest_k
) {
    const int threads = 256;
    if (L == 1) {
        constexpr int ROWS_PER_BLOCK = 8;
        constexpr int L1_CHUNK_SIZE = 2048;
        const int blocks = (M + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
        // Double buffer: 2 chunks of L1_CHUNK_SIZE halfs
        const int smem = 2 * L1_CHUNK_SIZE * sizeof(half);
        gemv_L1_kernel<<<blocks, threads, smem>>>(
            a_packed, b_decoded, sfa_scales, c, M, K, rest_m, rest_k
        );
    } else if (L == 2) {
        constexpr int ROWS_PER_BLOCK = 4;
        const int blocks = (M + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
        gemv_multi_kernel<2, ROWS_PER_BLOCK, 2><<<blocks, threads>>>(
            a_packed, b_decoded, sfa_scales, c, M, K, rest_m, rest_k
        );
    } else if (L == 4) {
        constexpr int ROWS_PER_BLOCK = 4;
        const int blocks = (M + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
        gemv_multi_kernel<4, ROWS_PER_BLOCK, 2><<<blocks, threads>>>(
            a_packed, b_decoded, sfa_scales, c, M, K, rest_m, rest_k
        );
    } else if (L == 8) {
        constexpr int ROWS_PER_BLOCK = 2;
        const int blocks = (M + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
        gemv_multi_kernel<8, ROWS_PER_BLOCK, 2><<<blocks, threads>>>(
            a_packed, b_decoded, sfa_scales, c, M, K, rest_m, rest_k
        );
    }
}
"""

cpp_source = """
#include <torch/extension.h>

extern "C" void decode_nvfp4_vector(
    const uint8_t* b_packed,
    const at::Half* sfb_scales,
    at::Half* b_decoded,
    int K, int L,
    int rest_k
);

extern "C" void launch_nvfp4_kernel(
    const uint8_t* a_packed,
    const at::Half* b_decoded,
    const at::Half* sfa_scales,
    at::Half* c,
    int M, int K, int L,
    int rest_m, int rest_k
);

torch::Tensor nvfp4_gemv(
    torch::Tensor a_packed,
    torch::Tensor b_packed,
    torch::Tensor sfa,
    torch::Tensor sfb,
    torch::Tensor c_out,
    int64_t M, int64_t K, int64_t L,
    int64_t rest_m, int64_t rest_k
) {
    auto options = torch::dtype(torch::kFloat16).device(a_packed.device());
    auto b_decoded = torch::empty({L, K}, options);

    decode_nvfp4_vector(
        b_packed.data_ptr<uint8_t>(),
        sfb.data_ptr<at::Half>(),
        b_decoded.data_ptr<at::Half>(),
        static_cast<int>(K),
        static_cast<int>(L),
        static_cast<int>(rest_k)
    );

    launch_nvfp4_kernel(
        a_packed.data_ptr<uint8_t>(),
        b_decoded.data_ptr<at::Half>(),
        sfa.data_ptr<at::Half>(),
        c_out.data_ptr<at::Half>(),
        static_cast<int>(M),
        static_cast<int>(K),
        static_cast<int>(L),
        static_cast<int>(rest_m),
        static_cast<int>(rest_k)
    );

    return c_out;
}
"""

module = None

def get_module():
    global module
    if module is None:
        module = load_inline(
            name="nvfp4_gemv",
            cpp_sources=cpp_source,
            cuda_sources=cuda_source,
            functions=["nvfp4_gemv"],
            verbose=False,
            extra_cuda_cflags=["-O3", "--use_fast_math", "-std=c++17"],
        )
    return module

def custom_kernel(data: input_t) -> output_t:
    a_ref, b_ref, _, _, sfa_permuted, sfb_permuted, c_ref = data
    M, _, L = c_ref.shape
    K = a_ref.shape[1] * 2

    a_packed = a_ref.view(torch.uint8).contiguous()
    b_packed = b_ref.view(torch.uint8).contiguous()
    device = a_ref.device

    sfa_scales = sfa_permuted.to(dtype=torch.float16, device=device, non_blocking=True).contiguous()
    sfb_scales = sfb_permuted.to(dtype=torch.float16, device=device, non_blocking=True).contiguous()

    rest_m = sfa_scales.shape[2]
    rest_k = sfa_scales.shape[4]

    if not c_ref.is_contiguous():
        c_ref = c_ref.contiguous()
    c_matrix = c_ref.view(M, L)

    mod = get_module()
    mod.nvfp4_gemv(a_packed, b_packed, sfa_scales, sfb_scales, c_matrix, M, K, L, rest_m, rest_k)

    return c_ref
scrolls · 453 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