Skip to content
KernelIndex
Search⌘K

submission 100978

s.am._ · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

bang.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-100978?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
18.7µs
#5 of 678
2025-11-24

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:dd2f4db72c9d44d8b2d0a40304daf11330975cf641d72f0075dff6da0dffdd4f
license declaredunknown
license concludedunknown
authorss.am._
imported2026-08-15

Techniques

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

fp4if (params.k <= 256) { // <= 512 FP4 values
fp8const __nv_fp8_e4m3* rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;
shared-memory__shared__ float sdata[THREADS_PER_ROW];

Kernel source

bang.py593 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

# ---- C++ stub: declare the function so load_inline can bind it ----
gemv_cpp = r"""
#include <torch/extension.h>

// Forward declaration so PyTorch can bind it (definition is in the CUDA source).
torch::Tensor cuda_nvfp4_gemv(torch::Tensor A,
                            torch::Tensor B,
                            torch::Tensor C,
                            torch::Tensor SFA,
                            torch::Tensor SFB);
"""

# ---- CUDA source: struct, kernel, launcher, and Python-facing wrapper ----
gemv_cuda = r"""
#include <assert.h>
#include <cuda.h>
#include <stdio.h>
#include <cuda_runtime.h>

#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>

#include <cuda_fp4.h>
#include <cuda_bf16.h>
#include <cuda_fp8.h>

// ---- gemv.h ----
struct Gemv_params {
    using index_t = uint64_t;

    int b, m, k, real_k;

    void *__restrict__ a_ptr;
    void *__restrict__ b_ptr;
    void *__restrict__ sfa_ptr;
    void *__restrict__ sfb_ptr;
    void *__restrict__ o_ptr;

    index_t a_batch_stride;
    index_t b_batch_stride;
    index_t sfa_batch_stride;
    index_t sfb_batch_stride;
    index_t o_batch_stride;

    index_t a_row_stride;
    index_t b_row_stride;
    index_t sfa_row_stride;
    index_t sfb_row_stride;
    index_t o_row_stride;
};

static constexpr int BLOCK_SIZE = 128; // 128

// ---- load helpers ----
__device__ __forceinline__ void load_block_16x2fp4_generic(
    const __nv_fp4x2_e2m1* rowA,
    const __nv_fp4x2_e2m1* vecB,
    const uint16_t*        rowS_u16,
    const uint16_t*        vecS_u16,
    int                    elem_base,
    int                    block_base,
    uint64_t (&a_regs)[2],
    uint64_t (&b_regs)[2],
    uint16_t &sfa_regs,
    uint16_t &sfb_regs)
{
    uint64_t rowA_addr = reinterpret_cast<uint64_t>(rowA + elem_base);
    uint64_t vecB_addr = reinterpret_cast<uint64_t>(vecB + elem_base);
    uint64_t rowS_addr = reinterpret_cast<uint64_t>(rowS_u16 + block_base);
    uint64_t vecS_addr = reinterpret_cast<uint64_t>(vecS_u16 + block_base);

    asm volatile(
        "ld.global.u64.v2 {%0, %1}, [%4];\n\t"
        "ld.global.u64.v2 {%2, %3}, [%5];\n\t"
        : "=l"(a_regs[0]), "=l"(a_regs[1]),
          "=l"(b_regs[0]), "=l"(b_regs[1])
        : "l"(rowA_addr), "l"(vecB_addr)
    );

    asm volatile(
        "ld.global.u16 %0, [%2];\n\t"
        "ld.global.u16 %1, [%3];\n\t"
        : "=h"(sfa_regs), "=h"(sfb_regs)
        : "l"(rowS_addr), "l"(vecS_addr)
    );
}

// k = 3584
__device__ __forceinline__ void load_block_16x2fp4_k3584(
    const __nv_fp4x2_e2m1* rowA,
    const __nv_fp4x2_e2m1* vecB,
    const uint16_t*        rowS_u16,
    const uint16_t*        vecS_u16,
    int                    elem_base,
    int                    block_base,
    uint64_t (&a_regs)[2],
    uint64_t (&b_regs)[2],
    uint16_t &sfa_regs,
    uint16_t &sfb_regs)
{
    uint64_t rowA_addr = reinterpret_cast<uint64_t>(rowA + elem_base);
    uint64_t vecB_addr = reinterpret_cast<uint64_t>(vecB + elem_base);
    uint64_t rowS_addr = reinterpret_cast<uint64_t>(rowS_u16 + block_base);
    uint64_t vecS_addr = reinterpret_cast<uint64_t>(vecS_u16 + block_base);

    asm volatile(
        "ld.global.cs.u64.v2 {%0, %1}, [%4];\n\t"
        "ld.global.L2::128B.u64.v2 {%2, %3}, [%5];\n\t"
        : "=l"(a_regs[0]), "=l"(a_regs[1]),
          "=l"(b_regs[0]), "=l"(b_regs[1])
        : "l"(rowA_addr), "l"(vecB_addr)
    );

    asm volatile(
        "ld.global.cs.u16 %0, [%2];\n\t"
        "ld.global.L2::128B.u16 %1, [%3];\n\t"
        : "=h"(sfa_regs), "=h"(sfb_regs)
        : "l"(rowS_addr), "l"(vecS_addr)
    );
}

// k = 8192
__device__ __forceinline__ void load_block_16x2fp4_k8192(
    const __nv_fp4x2_e2m1* rowA,
    const __nv_fp4x2_e2m1* vecB,
    const uint16_t*        rowS_u16,
    const uint16_t*        vecS_u16,
    int                    elem_base,
    int                    block_base,
    uint64_t (&a_regs)[2],
    uint64_t (&b_regs)[2],
    uint16_t &sfa_regs,
    uint16_t &sfb_regs)
{
    uint64_t rowA_addr = reinterpret_cast<uint64_t>(rowA + elem_base);
    uint64_t vecB_addr = reinterpret_cast<uint64_t>(vecB + elem_base);
    uint64_t rowS_addr = reinterpret_cast<uint64_t>(rowS_u16 + block_base);
    uint64_t vecS_addr = reinterpret_cast<uint64_t>(vecS_u16 + block_base);

    asm volatile(
        "ld.global.cs.u64.v2 {%0, %1}, [%4];\n\t"
        "ld.global.u64.v2 {%2, %3}, [%5];\n\t"
        : "=l"(a_regs[0]), "=l"(a_regs[1]),
          "=l"(b_regs[0]), "=l"(b_regs[1])
        : "l"(rowA_addr), "l"(vecB_addr)
    );

    asm volatile(
        "ld.global.lu.u16 %0, [%2];\n\t"
        "ld.global.u16 %1, [%3];\n\t"
        : "=h"(sfa_regs), "=h"(sfb_regs)
        : "l"(rowS_addr), "l"(vecS_addr)
    );
}

// k = 1024
__device__ __forceinline__ void load_block_16x2fp4_k1024(
    const __nv_fp4x2_e2m1* rowA,
    const __nv_fp4x2_e2m1* vecB,
    const uint16_t*        rowS_u16,
    const uint16_t*        vecS_u16,
    int                    elem_base,
    int                    block_base,
    uint64_t (&a_regs)[2],
    uint64_t (&b_regs)[2],
    uint16_t &sfa_regs,
    uint16_t &sfb_regs)
{
    uint64_t rowA_addr = reinterpret_cast<uint64_t>(rowA + elem_base);
    uint64_t vecB_addr = reinterpret_cast<uint64_t>(vecB + elem_base);
    uint64_t rowS_addr = reinterpret_cast<uint64_t>(rowS_u16 + block_base);
    uint64_t vecS_addr = reinterpret_cast<uint64_t>(vecS_u16 + block_base);

    asm volatile(
        "ld.global.cs.u64.v2 {%0, %1}, [%4];\n\t"
        "ld.global.u64.v2 {%2, %3}, [%5];\n\t"
        : "=l"(a_regs[0]), "=l"(a_regs[1]),
          "=l"(b_regs[0]), "=l"(b_regs[1])
        : "l"(rowA_addr), "l"(vecB_addr)
    );

    asm volatile(
        "ld.global.cs.u16 %0, [%2];\n\t"
        "ld.global.u16 %1, [%3];\n\t"
        : "=h"(sfa_regs), "=h"(sfb_regs)
        : "l"(rowS_addr), "l"(vecS_addr)
    );
}

// Compile-time dispatcher
template<int K>
__device__ __forceinline__ void load_block_16x2fp4(
    const __nv_fp4x2_e2m1* rowA,
    const __nv_fp4x2_e2m1* vecB,
    const uint16_t*        rowS_u16,
    const uint16_t*        vecS_u16,
    int                    elem_base,
    int                    block_base,
    uint64_t (&a_regs)[2],
    uint64_t (&b_regs)[2],
    uint16_t &sfa_regs,
    uint16_t &sfb_regs)
{
    if constexpr (K == 3584) {
        load_block_16x2fp4_k3584(
            rowA, vecB, rowS_u16, vecS_u16,
            elem_base, block_base, a_regs, b_regs, sfa_regs, sfb_regs);
    } else if constexpr (K == 8192) {
        load_block_16x2fp4_k8192(
            rowA, vecB, rowS_u16, vecS_u16,
            elem_base, block_base, a_regs, b_regs, sfa_regs, sfb_regs);
    } else if constexpr (K == 1024) {
        load_block_16x2fp4_k1024(
            rowA, vecB, rowS_u16, vecS_u16,
            elem_base, block_base, a_regs, b_regs, sfa_regs, sfb_regs);
    } else {
        // generic / fallback
        load_block_16x2fp4_generic(
            rowA, vecB, rowS_u16, vecS_u16,
            elem_base, block_base, a_regs, b_regs, sfa_regs, sfb_regs);
    }
}



__device__ __forceinline__ float block_scaled_fma_16x2fp4(
    const uint64_t (&a_regs)[2],
    const uint64_t (&b_regs)[2],
    uint16_t       sfa_regs,
    uint16_t       sfb_regs)
{
    uint32_t const* a_regs_packed = reinterpret_cast<uint32_t const*>(&a_regs);
    uint32_t const* b_regs_packed = reinterpret_cast<uint32_t const*>(&b_regs);

    float out_f32;

    asm volatile(
        "{\n"
        // 8 bytes of A and B at a time (reused for upper half)
        ".reg .b8 a0_0, a0_1, a0_2, a0_3;\n"
        ".reg .b8 a0_4, a0_5, a0_6, a0_7;\n"
        ".reg .b8 b0_0, b0_1, b0_2, b0_3;\n"
        ".reg .b8 b0_4, b0_5, b0_6, b0_7;\n"

        // scales and accumulators
        ".reg .f16x2 sfa_f16x2, sfb_f16x2, sf_f16x2;\n"
        ".reg .f16x2 scale0_f16x2, scale1_f16x2;\n"
        ".reg .f16x2 accum_total, accum_group;\n"

        // converted fp4 -> f16x2 (only 8 per vector kept live)
        ".reg .f16x2 cvt_0_0, cvt_0_1, cvt_0_2, cvt_0_3;\n"
        ".reg .f16x2 cvt_0_4, cvt_0_5, cvt_0_6, cvt_0_7;\n"
        ".reg .f16x2 cvt_1_0, cvt_1_1, cvt_1_2, cvt_1_3;\n"
        ".reg .f16x2 cvt_1_4, cvt_1_5, cvt_1_6, cvt_1_7;\n"

        ".reg .f16 lane0, lane1, result_f16;\n"
        ".reg .f32 result_f32;\n"

        // scales
        "cvt.rn.f16x2.e4m3x2 sfa_f16x2, %5;\n"
        "cvt.rn.f16x2.e4m3x2 sfb_f16x2, %6;\n"
        "mul.rn.f16x2 sf_f16x2, sfa_f16x2, sfb_f16x2;\n"
        "mov.b32 {lane0, lane1}, sf_f16x2;\n"
        "mov.b32 scale0_f16x2, {lane0, lane0};\n"
        "mov.b32 scale1_f16x2, {lane1, lane1};\n"

        "mov.b32 accum_total, 0;\n"

        //----------------------------------------------------------------------
        // First 8×(2×FP4) -> uses scale0
        //----------------------------------------------------------------------
        "mov.b32 {a0_0, a0_1, a0_2, a0_3}, %1;\n"
        "mov.b32 {a0_4, a0_5, a0_6, a0_7}, %2;\n"
        "mov.b32 {b0_0, b0_1, b0_2, b0_3}, %3;\n"
        "mov.b32 {b0_4, b0_5, b0_6, b0_7}, %4;\n"

        // all conversions first
        "cvt.rn.f16x2.e2m1x2 cvt_0_0, a0_0;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_0, b0_0;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_1, a0_1;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_1, b0_1;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_2, a0_2;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_2, b0_2;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_3, a0_3;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_3, b0_3;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_4, a0_4;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_4, b0_4;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_5, a0_5;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_5, b0_5;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_6, a0_6;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_6, b0_6;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_7, a0_7;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_7, b0_7;\n"

        // then all FMAs into one group accumulator
        "mov.b32 accum_group, 0;\n"
        "fma.rn.f16x2 accum_group, cvt_0_0, cvt_1_0, accum_group;\n"
        "fma.rn.f16x2 accum_group, cvt_0_1, cvt_1_1, accum_group;\n"
        "fma.rn.f16x2 accum_group, cvt_0_2, cvt_1_2, accum_group;\n"
        "fma.rn.f16x2 accum_group, cvt_0_3, cvt_1_3, accum_group;\n"
        "fma.rn.f16x2 accum_group, cvt_0_4, cvt_1_4, accum_group;\n"
        "fma.rn.f16x2 accum_group, cvt_0_5, cvt_1_5, accum_group;\n"
        "fma.rn.f16x2 accum_group, cvt_0_6, cvt_1_6, accum_group;\n"
        "fma.rn.f16x2 accum_group, cvt_0_7, cvt_1_7, accum_group;\n"

        "mul.rn.f16x2 accum_group, scale0_f16x2, accum_group;\n"
        "add.rn.f16x2 accum_total, accum_total, accum_group;\n"

        //----------------------------------------------------------------------
        // Second 8×(2×FP4) -> uses scale1, reusing all the same regs
        //----------------------------------------------------------------------
        "mov.b32 {a0_0, a0_1, a0_2, a0_3}, %7;\n"
        "mov.b32 {a0_4, a0_5, a0_6, a0_7}, %8;\n"
        "mov.b32 {b0_0, b0_1, b0_2, b0_3}, %9;\n"
        "mov.b32 {b0_4, b0_5, b0_6, b0_7}, %10;\n"

        // conversions for upper half
        "cvt.rn.f16x2.e2m1x2 cvt_0_0, a0_0;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_0, b0_0;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_1, a0_1;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_1, b0_1;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_2, a0_2;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_2, b0_2;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_3, a0_3;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_3, b0_3;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_4, a0_4;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_4, b0_4;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_5, a0_5;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_5, b0_5;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_6, a0_6;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_6, b0_6;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_7, a0_7;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_7, b0_7;\n"

        "mov.b32 accum_group, 0;\n"
        "fma.rn.f16x2 accum_group, cvt_0_0, cvt_1_0, accum_group;\n"
        "fma.rn.f16x2 accum_group, cvt_0_1, cvt_1_1, accum_group;\n"
        "fma.rn.f16x2 accum_group, cvt_0_2, cvt_1_2, accum_group;\n"
        "fma.rn.f16x2 accum_group, cvt_0_3, cvt_1_3, accum_group;\n"
        "fma.rn.f16x2 accum_group, cvt_0_4, cvt_1_4, accum_group;\n"
        "fma.rn.f16x2 accum_group, cvt_0_5, cvt_1_5, accum_group;\n"
        "fma.rn.f16x2 accum_group, cvt_0_6, cvt_1_6, accum_group;\n"
        "fma.rn.f16x2 accum_group, cvt_0_7, cvt_1_7, accum_group;\n"

        "mul.rn.f16x2 accum_group, scale1_f16x2, accum_group;\n"
        "add.rn.f16x2 accum_total, accum_total, accum_group;\n"

        // final reduction to scalar f16 -> then upconvert to f32
        "mov.b32 {lane0, lane1}, accum_total;\n"
        "add.rn.f16 result_f16, lane0, lane1;\n"
        "cvt.f32.f16 result_f32, result_f16;\n"
        "mov.b32 %0, result_f32;\n"

        "}\n"
        : "=f"(out_f32)
        : "r"(a_regs_packed[0]), "r"(a_regs_packed[1]),
          "r"(b_regs_packed[0]), "r"(b_regs_packed[1]),
          "h"(sfa_regs), "h"(sfb_regs),
          "r"(a_regs_packed[2]), "r"(a_regs_packed[3]),
          "r"(b_regs_packed[2]), "r"(b_regs_packed[3])
        : "memory"
    );

    return out_f32;
}

template <int ROWS_PER_BLOCK, int THREADS_PER_ROW, int ITERS, int K_SPECIAL>
__global__ void __launch_bounds__(ROWS_PER_BLOCK*THREADS_PER_ROW, 8)
gemv_kernel_shared(const __grid_constant__ Gemv_params params)
{
    const int tid   = threadIdx.x;
    const int rib   = tid / THREADS_PER_ROW;
    const int lane  = tid % THREADS_PER_ROW;
    const int batch = blockIdx.z;
    const int row   = blockIdx.x * ROWS_PER_BLOCK + rib;

    const size_t A_batch_base   = static_cast<size_t>(batch) * params.a_batch_stride;
    const size_t SFA_batch_base = static_cast<size_t>(batch) * params.sfa_batch_stride;
    const size_t B_batch_base   = static_cast<size_t>(batch) * params.b_batch_stride;
    const size_t SFB_batch_base = static_cast<size_t>(batch) * params.sfb_batch_stride;
    const size_t C_batch_base   = static_cast<size_t>(batch) * params.o_batch_stride;

    const __nv_fp4x2_e2m1* rowA = static_cast<const __nv_fp4x2_e2m1*>(params.a_ptr) + A_batch_base   + row * params.a_row_stride;
    const __nv_fp8_e4m3*   rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;
    const __nv_fp4x2_e2m1* vecB = static_cast<const __nv_fp4x2_e2m1*>(params.b_ptr) + B_batch_base;
    const __nv_fp8_e4m3*   vecS = static_cast<const __nv_fp8_e4m3*>(params.sfb_ptr) + SFB_batch_base;

    const uint16_t* rowS_u16 = reinterpret_cast<const uint16_t*>(rowS);
    const uint16_t* vecS_u16 = reinterpret_cast<const uint16_t*>(vecS);

    float sum = 0.f;
    #pragma unroll
    for (int idx = 0; idx < ITERS; ++idx) {
        int block_base = idx * THREADS_PER_ROW + lane;
        int elem_base  = block_base * 16;

        uint64_t a_regs[2], b_regs[2];
        uint16_t sfa_regs, sfb_regs;

        load_block_16x2fp4<K_SPECIAL>(
            rowA, vecB,
            rowS_u16, vecS_u16,
            elem_base, block_base,
            a_regs, b_regs,
            sfa_regs, sfb_regs);
        sum += block_scaled_fma_16x2fp4(a_regs, b_regs, sfa_regs, sfb_regs);
    }

    __shared__ float sdata[THREADS_PER_ROW];
    sdata[lane] = sum;
    __syncthreads();

    if (tid < 64) sdata[lane] += sdata[lane+64];
    __syncthreads();

    if (lane < 32) {
        float val = sdata[lane] + sdata[lane + 32];
        val += __shfl_down_sync(0xffffffff, val, 16);
        val += __shfl_down_sync(0xffffffff, val, 8);
        val += __shfl_down_sync(0xffffffff, val, 4);
        val += __shfl_down_sync(0xffffffff, val, 2);
        val += __shfl_down_sync(0xffffffff, val, 1);

        if (lane == 0) {
            __half* out = (__half*)params.o_ptr + C_batch_base + row;
            out[0] = __float2half(val);
        }
    }
}

template <int ROWS_PER_BLOCK, int THREADS_PER_ROW, int ITERS, int K_SPECIAL>
__global__ void __launch_bounds__(ROWS_PER_BLOCK*THREADS_PER_ROW, 8)
gemv_kernel(const __grid_constant__ Gemv_params params)
{
    const int tid   = threadIdx.x;
    const int rib   = tid / THREADS_PER_ROW;
    const int lane  = tid % THREADS_PER_ROW;
    const int batch = blockIdx.z;
    const int row   = blockIdx.x * ROWS_PER_BLOCK + rib;

    const size_t A_batch_base   = static_cast<size_t>(batch) * params.a_batch_stride;
    const size_t SFA_batch_base = static_cast<size_t>(batch) * params.sfa_batch_stride;
    const size_t B_batch_base   = static_cast<size_t>(batch) * params.b_batch_stride;
    const size_t SFB_batch_base = static_cast<size_t>(batch) * params.sfb_batch_stride;
    const size_t C_batch_base   = static_cast<size_t>(batch) * params.o_batch_stride;

    const __nv_fp4x2_e2m1* rowA = static_cast<const __nv_fp4x2_e2m1*>(params.a_ptr) + A_batch_base   + row * params.a_row_stride;
    const __nv_fp8_e4m3*   rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;
    const __nv_fp4x2_e2m1* vecB = static_cast<const __nv_fp4x2_e2m1*>(params.b_ptr) + B_batch_base;
    const __nv_fp8_e4m3*   vecS = static_cast<const __nv_fp8_e4m3*>(params.sfb_ptr) + SFB_batch_base;

    const uint16_t* rowS_u16 = reinterpret_cast<const uint16_t*>(rowS);
    const uint16_t* vecS_u16 = reinterpret_cast<const uint16_t*>(vecS);

    float sum = 0.f;

    auto body = [&](int idx) {
        int block_base = idx * THREADS_PER_ROW + lane;
        int elem_base  = block_base * 16;

        uint64_t a_regs[2], b_regs[2];
        uint16_t sfa_regs, sfb_regs;

        load_block_16x2fp4<K_SPECIAL>(
            rowA, vecB,
            rowS_u16, vecS_u16,
            elem_base, block_base,
            a_regs, b_regs,
            sfa_regs, sfb_regs);
        sum += block_scaled_fma_16x2fp4(a_regs, b_regs, sfa_regs, sfb_regs);
    };

    if constexpr (ITERS > 0) {
        #pragma unroll
        for (int idx = 0; idx < ITERS; ++idx) {
            body(idx);
        }
    } else {
        int iters = params.k / (THREADS_PER_ROW * 16);
        for (int idx = 0; idx < iters; ++idx) {
            body(idx);
        }
    }

    #pragma unroll
    for (int offset = THREADS_PER_ROW / 2; offset > 0; offset /= 2) {
        sum += __shfl_down_sync(0xffffffffu, sum, offset, THREADS_PER_ROW);
    }

    if (lane == 0) {
        __half* out = (__half*)params.o_ptr + C_batch_base + row;
        out[0] = __float2half(sum);
    }
}


torch::Tensor cuda_nvfp4_gemv(torch::Tensor A,
                            torch::Tensor B,
                            torch::Tensor C,
                            torch::Tensor SFA,
                            torch::Tensor SFB)
{

    const auto sizes = A.sizes();
    const int M = sizes[0];
    const int K = sizes[1];
    const int L = sizes[2];

    Gemv_params params{};
    params.b = L;
    params.m = M;
    params.k = K;

    params.a_ptr  = A.data_ptr();
    params.b_ptr  = B.data_ptr();
    params.sfa_ptr= SFA.data_ptr();
    params.sfb_ptr= SFB.data_ptr();
    params.o_ptr  = C.data_ptr();

    params.a_batch_stride  = A.stride(2);
    params.b_batch_stride  = B.stride(2);
    params.sfa_batch_stride= SFA.stride(2);
    params.sfb_batch_stride= SFB.stride(2);
    params.o_batch_stride  = C.stride(2);

    params.a_row_stride  = A.stride(0);
    params.b_row_stride  = B.stride(0);
    params.sfa_row_stride= SFA.stride(0);
    params.sfb_row_stride= SFB.stride(0);
    params.o_row_stride  = C.stride(0);

    auto stream = at::cuda::getCurrentCUDAStream().stream();
    if (params.k <= 256) {  // <= 512 FP4 values
        dim3 grid(params.m / 16, 1, params.b);
        dim3 block(128, 1, 1);
        // generic loader (K_SPECIAL = 0)
        gemv_kernel<16, 8, 0, 0><<<grid, block, 0, stream>>>(params);
    } else if (params.k == 3584) {
        dim3 block(128, 1, 1);
        dim3 grid(params.m / 4, 1, params.b);
        // K_SPECIAL = 3584 -> uses k3584 loader
        gemv_kernel<4, 32, 7, 3584><<<grid, block, 0, stream>>>(params);
    } else if (params.k == 8192) {
        dim3 block(128, 1, 1);
        dim3 grid(params.m, 1, params.b);
        // K_SPECIAL = 8192 -> uses k8192 loader
        gemv_kernel_shared<1, 128, 4, 8192><<<grid, block, 0, stream>>>(params);
    } else if (params.k == 1024) {
        dim3 block(128, 1, 1);
        dim3 grid(params.m / 8, 1, params.b);
        // K_SPECIAL = 1024 -> uses k1024 loader
        gemv_kernel<8, 16, 4, 1024><<<grid, block, 0, stream>>>(params);
    } else {
        dim3 block(128, 1, 1);
        dim3 grid(params.m / 8, 1, params.b);
        // generic loader again
        gemv_kernel<8, 16, 0, 0><<<grid, block, 0, stream>>>(params);
    }


    return C;
}
"""

# ---- build the module ----
nvfp4_module = load_inline(
    name="nvfp4_gemv",
    cpp_sources=[gemv_cpp],
    cuda_sources=[gemv_cuda],
    functions=["cuda_nvfp4_gemv"],  # this exposes the function to Python
    extra_cuda_cflags=[
        "-std=c++17",
        "-maxrregcount=48",
        "-gencode=arch=compute_100a,code=sm_100a",
        "--ptxas-options=--gpu-name=sm_100a",
        "-O3",
        "-w",
        "--use_fast_math",
        "-allow-unsupported-compiler",
    ],
    extra_ldflags=["-lcuda", "-lcublas"],
    verbose=True,
)


def custom_kernel(data: input_t) -> output_t:
    return nvfp4_module.cuda_nvfp4_gemv(data[0], data[1], data[6], data[2], data[3])
scrolls · 593 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 95654.

⋯ 109 unchanged lines
asm volatile(
"ld.global.cs.u64.v2 {%0, %1}, [%4];\n\t"
- "ld.global.u64.v2 {%2, %3}, [%5];\n\t"
+ "ld.global.L2::128B.u64.v2 {%2, %3}, [%5];\n\t"
: "=l"(a_regs[0]), "=l"(a_regs[1]),
"=l"(b_regs[0]), "=l"(b_regs[1])
: "l"(rowA_addr), "l"(vecB_addr)
⋯ 1 unchanged lines
asm volatile(
"ld.global.cs.u16 %0, [%2];\n\t"
- "ld.global.u16 %1, [%3];\n\t"
+ "ld.global.L2::128B.u16 %1, [%3];\n\t"
: "=h"(sfa_regs), "=h"(sfb_regs)
: "l"(rowS_addr), "l"(vecS_addr)
);
⋯ 450 unchanged lines
functions=["cuda_nvfp4_gemv"], # this exposes the function to Python
extra_cuda_cflags=[
"-std=c++17",
+ "-maxrregcount=48",
"-gencode=arch=compute_100a,code=sm_100a",
"--ptxas-options=--gpu-name=sm_100a",
"-O3",
scrolls · 26 diff lines total

Best evidence level for this revision: reported

JSON