Skip to content
KernelIndex
Search⌘K

submission 77405

s.am._ · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

less_naive_v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-77405?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
23.7µs
#72 of 678
2025-11-14

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

async-copy__device__ __forceinline__ void cp_async_ca(uint32_t smem_addr, void const *global_ptr, bool pred_guard=true) {
fp8alignas(128) __nv_fp8_e4m3 SFA[BUFFER][STAGES][128 * 2];
shared-memory__device__ __forceinline__ void cp_async_ca(uint32_t smem_addr, void const *global_ptr, bool pred_guard=true) {
stages = 4static constexpr int STAGES = 4;

Kernel source

less_naive_v3.py1114 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>

// Single Python-visible entry point that dispatches between kernels
torch::Tensor nvfp4_gemv_dispatch(torch::Tensor A,
                                  torch::Tensor B,
                                  torch::Tensor C,
                                  torch::Tensor SFA,
                                  torch::Tensor SFB);
"""

# ---- CUDA source: params struct, both kernels, launchers, 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-ish: params struct ----
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 ROWS_PER_BLOCK  = 8;
static constexpr int THREADS_PER_ROW = 16;
static constexpr int BLOCK_SIZE = ROWS_PER_BLOCK * THREADS_PER_ROW; // 128
static constexpr int BUFFER = 2;
static constexpr int STAGES = 4;

template <int LoadBytes>
__device__ __forceinline__ void cp_async_ca(uint32_t smem_addr, void const *global_ptr, bool pred_guard=true) {
    asm volatile(
        "{\n"
        "   .reg .pred p;\n"
        "   setp.ne.b32 p, %0, 0;\n"
        "   @p cp.async.ca.shared.global.L2::128B [%1], [%2], %3;\n"
        "}\n" 
        :
        :"r"((int)pred_guard), "r"(smem_addr), "l"(global_ptr), "n"(LoadBytes)
    );
}

struct SharedStorage {
    alignas(128) __nv_fp4x2_e2m1 A[BUFFER][STAGES][128 * 16];
    alignas(128) __nv_fp4x2_e2m1 B[BUFFER][STAGES][THREADS_PER_ROW * 16];
    alignas(128) __nv_fp8_e4m3 SFA[BUFFER][STAGES][128 * 2];
    alignas(128) __nv_fp8_e4m3 SFB[BUFFER][STAGES][THREADS_PER_ROW * 2];
};


__device__ __forceinline__ void cpasync_load(
    int buffer_idx,
    int &global_k, // should be in 2xfp4
    int smem_offset_A,
    int smem_offset_B,
    int smem_sf_write_offset,
    int lane,
    bool is_even,
    bool load_b,
    bool valid,
    const __nv_fp4x2_e2m1* rowA,
    const __nv_fp4x2_e2m1* vecB,
    const __nv_fp8_e4m3* rowS,
    const __nv_fp8_e4m3* vecS,
    SharedStorage& smem)
{
    for (int s = 0; s < STAGES; ++s) {
        int SF_offset = (global_k >> 3) + lane * 2;

        uint32_t smem_addr_A = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.A[buffer_idx][s][smem_offset_A]));
        uint32_t smem_addr_B = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.B[buffer_idx][s][smem_offset_B]));
        uint32_t smem_addr_SFA = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.SFA[buffer_idx][s][smem_sf_write_offset]));
        uint32_t smem_addr_SFB = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.SFB[buffer_idx][s][(lane / 2) * 4]));

        cp_async_ca<16>(smem_addr_A, rowA + global_k, valid);
        cp_async_ca<16>(smem_addr_B, vecB + global_k, valid && load_b);
        cp_async_ca<4>(smem_addr_SFA, rowS + SF_offset, valid && is_even);
        cp_async_ca<4>(smem_addr_SFB, vecS + SF_offset, valid && is_even && load_b);

        global_k += 256;
    }
}

__device__ __forceinline__ void load_fragments(
    uint64_t (&a_regs)[2],
    uint64_t (&b_regs)[2],
    uint16_t &sfa_regs,
    uint16_t &sfb_regs,
    int lane,
    int smem_pipe_read,
    int k_block,
    int smem_offset_A,
    int smem_offset_B,
    int smem_sf_offset,
    SharedStorage& smem) 
{
    uint32_t smem_addr_a = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.A[smem_pipe_read][k_block][smem_offset_A]));
    uint32_t smem_addr_b = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.B[smem_pipe_read][k_block][smem_offset_B]));
    
    asm volatile(
        "ld.shared.v2.u64 {%0, %1}, [%4];\n\t"
        "ld.shared.v2.u64 {%2, %3}, [%5];\n\t"
        : "=l"(a_regs[0]), "=l"(a_regs[1]),
          "=l"(b_regs[0]), "=l"(b_regs[1])
        : "r"(smem_addr_a), "r"(smem_addr_b)
    );

    uint32_t smem_addr_sfa = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.SFA[smem_pipe_read][k_block][smem_sf_offset]));
    uint32_t smem_addr_sfb = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.SFB[smem_pipe_read][k_block][lane * 2]));
    
    asm volatile(
        "ld.shared.u16 %0, [%2];\n\t"
        "ld.shared.u16 %1, [%3];\n\t"
        : "=h"(sfa_regs), "=h"(sfb_regs)
        : "r"(smem_addr_sfa), "r"(smem_addr_sfb)
    );
}



__device__ __forceinline__ __half 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);

    uint16_t out_half_bits;

    asm volatile(
        "{\n"
        ".reg .b8 byte0_0, byte0_1, byte0_2, byte0_3;\n"
        ".reg .b8 byte0_4, byte0_5, byte0_6, byte0_7;\n"
        ".reg .b8 byte0_8, byte0_9, byte0_10, byte0_11;\n"
        ".reg .b8 byte0_12, byte0_13, byte0_14, byte0_15;\n"
        ".reg .b8 byte1_0, byte1_1, byte1_2, byte1_3;\n"
        ".reg .b8 byte1_4, byte1_5, byte1_6, byte1_7;\n"
        ".reg .b8 byte1_8, byte1_9, byte1_10, byte1_11;\n"
        ".reg .b8 byte1_12, byte1_13, byte1_14, byte1_15;\n"

        ".reg .f16x2 accum_0, accum_1, accum_2, accum_3;\n"
        ".reg .f16x2 accum_4, accum_5, accum_6, accum_7;\n"
        ".reg .f16x2 accum_8, accum_9, accum_10, accum_11;\n"
        ".reg .f16x2 accum_12, accum_13, accum_14, accum_15;\n"

        ".reg .f16x2 sfa_f16x2;\n"
        ".reg .f16x2 sfb_f16x2;\n"
        ".reg .f16x2 sf_f16x2;\n"

        ".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_0_8, cvt_0_9, cvt_0_10, cvt_0_11;\n"
        ".reg .f16x2 cvt_0_12, cvt_0_13, cvt_0_14, cvt_0_15;\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 .f16x2 cvt_1_8, cvt_1_9, cvt_1_10, cvt_1_11;\n"
        ".reg .f16x2 cvt_1_12, cvt_1_13, cvt_1_14, cvt_1_15;\n"
        ".reg .f16 result_f16, lane0, lane1;\n"
        ".reg .f16x2 mul_f16x2_0, mul_f16x2_1;\n"

        "cvt.rn.f16x2.e4m3x2 sfa_f16x2, %9;\n"
        "cvt.rn.f16x2.e4m3x2 sfb_f16x2, %10;\n"

        "mov.b32 accum_0, 0;\n"
        "mov.b32 accum_1, 0;\n"
        "mov.b32 accum_2, 0;\n"
        "mov.b32 accum_3, 0;\n"
        "mov.b32 accum_4, 0;\n"
        "mov.b32 accum_5, 0;\n"
        "mov.b32 accum_6, 0;\n"
        "mov.b32 accum_7, 0;\n"
        "mov.b32 accum_8, 0;\n"
        "mov.b32 accum_9, 0;\n"
        "mov.b32 accum_10, 0;\n"
        "mov.b32 accum_11, 0;\n"
        "mov.b32 accum_12, 0;\n"
        "mov.b32 accum_13, 0;\n"
        "mov.b32 accum_14, 0;\n"
        "mov.b32 accum_15, 0;\n"

        "mul.rn.f16x2 sf_f16x2, sfa_f16x2, sfb_f16x2;\n"
        "mov.b32 {lane0, lane1}, sf_f16x2;\n"
        "mov.b32 mul_f16x2_0, {lane0, lane0};\n"
        "mov.b32 mul_f16x2_1, {lane1, lane1};\n"

        "mov.b32 {byte0_0, byte0_1, byte0_2, byte0_3}, %1;\n"
        "mov.b32 {byte0_4, byte0_5, byte0_6, byte0_7}, %2;\n"
        "mov.b32 {byte0_8, byte0_9, byte0_10, byte0_11}, %3;\n"
        "mov.b32 {byte0_12, byte0_13, byte0_14, byte0_15}, %4;\n"
        "mov.b32 {byte1_0, byte1_1, byte1_2, byte1_3}, %5;\n"
        "mov.b32 {byte1_4, byte1_5, byte1_6, byte1_7}, %6;\n"
        "mov.b32 {byte1_8, byte1_9, byte1_10, byte1_11}, %7;\n"
        "mov.b32 {byte1_12, byte1_13, byte1_14, byte1_15}, %8;\n"

        "cvt.rn.f16x2.e2m1x2 cvt_0_0,  byte0_0;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_1,  byte0_1;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_2,  byte0_2;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_3,  byte0_3;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_4,  byte0_4;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_5,  byte0_5;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_6,  byte0_6;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_7,  byte0_7;\n"

        "cvt.rn.f16x2.e2m1x2 cvt_0_8,  byte0_8;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_9,  byte0_9;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_10, byte0_10;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_11, byte0_11;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_12, byte0_12;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_13, byte0_13;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_14, byte0_14;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_0_15, byte0_15;\n"

        "cvt.rn.f16x2.e2m1x2 cvt_1_0,  byte1_0;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_1,  byte1_1;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_2,  byte1_2;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_3,  byte1_3;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_4,  byte1_4;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_5,  byte1_5;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_6,  byte1_6;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_7,  byte1_7;\n"

        "cvt.rn.f16x2.e2m1x2 cvt_1_8,  byte1_8;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_9,  byte1_9;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_10, byte1_10;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_11, byte1_11;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_12, byte1_12;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_13, byte1_13;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_14, byte1_14;\n"
        "cvt.rn.f16x2.e2m1x2 cvt_1_15, byte1_15;\n"

        "fma.rn.f16x2 accum_0,  cvt_0_0,  cvt_1_0,  accum_0;\n"
        "fma.rn.f16x2 accum_1,  cvt_0_1,  cvt_1_1,  accum_1;\n"
        "fma.rn.f16x2 accum_2,  cvt_0_2,  cvt_1_2,  accum_2;\n"
        "fma.rn.f16x2 accum_3,  cvt_0_3,  cvt_1_3,  accum_3;\n"
        "fma.rn.f16x2 accum_4,  cvt_0_4,  cvt_1_4,  accum_4;\n"
        "fma.rn.f16x2 accum_5,  cvt_0_5,  cvt_1_5,  accum_5;\n"
        "fma.rn.f16x2 accum_6,  cvt_0_6,  cvt_1_6,  accum_6;\n"
        "fma.rn.f16x2 accum_7,  cvt_0_7,  cvt_1_7,  accum_7;\n"

        "fma.rn.f16x2 accum_8,  cvt_0_8,  cvt_1_8,  accum_8;\n"
        "fma.rn.f16x2 accum_9,  cvt_0_9,  cvt_1_9,  accum_9;\n"
        "fma.rn.f16x2 accum_10, cvt_0_10, cvt_1_10, accum_10;\n"
        "fma.rn.f16x2 accum_11, cvt_0_11, cvt_1_11, accum_11;\n"
        "fma.rn.f16x2 accum_12, cvt_0_12, cvt_1_12, accum_12;\n"
        "fma.rn.f16x2 accum_13, cvt_0_13, cvt_1_13, accum_13;\n"
        "fma.rn.f16x2 accum_14, cvt_0_14, cvt_1_14, accum_14;\n"
        "fma.rn.f16x2 accum_15, cvt_0_15, cvt_1_15, accum_15;\n"

        "add.rn.f16x2 accum_0, accum_0, accum_1;\n"
        "add.rn.f16x2 accum_2, accum_2, accum_3;\n"
        "add.rn.f16x2 accum_4, accum_4, accum_5;\n"
        "add.rn.f16x2 accum_6, accum_6, accum_7;\n"
        "add.rn.f16x2 accum_8, accum_8, accum_9;\n"
        "add.rn.f16x2 accum_10, accum_10, accum_11;\n"
        "add.rn.f16x2 accum_12, accum_12, accum_13;\n"
        "add.rn.f16x2 accum_14, accum_14, accum_15;\n"

        "add.rn.f16x2 accum_0, accum_0, accum_2;\n"
        "add.rn.f16x2 accum_4, accum_4, accum_6;\n"
        "add.rn.f16x2 accum_8, accum_8, accum_10;\n"
        "add.rn.f16x2 accum_12, accum_12, accum_14;\n"

        "add.rn.f16x2 accum_0, accum_0, accum_4;\n"
        "add.rn.f16x2 accum_8, accum_8, accum_12;\n"

        "mul.rn.f16x2 accum_0, mul_f16x2_0, accum_0;\n"
        "mul.rn.f16x2 accum_8, mul_f16x2_1, accum_8;\n"

        "add.rn.f16x2 accum_0, accum_0, accum_8;\n"

        "mov.b32 {lane0, lane1}, accum_0;\n"
        "add.rn.f16 result_f16, lane0, lane1;\n"

        "mov.b16 %0, result_f16;\n"
        "}\n"
        : "=h"(out_half_bits)
        : "r"(a_regs_packed[0]), "r"(a_regs_packed[1]),
          "r"(a_regs_packed[2]), "r"(a_regs_packed[3]),
          "r"(b_regs_packed[0]), "r"(b_regs_packed[1]),
          "r"(b_regs_packed[2]), "r"(b_regs_packed[3]),
          "h"(sfa_regs), "h"(sfb_regs)
        : "memory"
    );

    union { uint16_t u; __half h; } conv;
    conv.u = out_half_bits;
    return conv.h;
}


__global__ void __launch_bounds__(BLOCK_SIZE, 8)
gemv_kernel_fast(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;

    extern __shared__ __align__(128) uint8_t shared_storage[];
    SharedStorage &smem = *reinterpret_cast<SharedStorage*>(shared_storage);

    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;

    rowA += lane * 16;
    vecB += lane * 16;

    const int k = params.k;

    float sum = 0.f;
    int global_k = 0;
    const int smem_offset_A = tid * 16;
    const int smem_offset_B = lane * 16;
    const int smem_sf_write_offset = (tid / 2) * 4;
    const int smem_sf_offset = tid * 2;
    const bool load_b = rib == 0;
    const bool is_even = (lane % 2 == 0);

    // prefetch to buffer 0
    cpasync_load(
        0,
        global_k,
        smem_offset_A,
        smem_offset_B,
        smem_sf_write_offset,
        lane,
        is_even,
        load_b,
        true,
        rowA,
        vecB,
        rowS,
        vecS,
        smem
    );
    asm volatile("cp.async.commit_group;\n" ::);
    asm volatile("cp.async.wait_all;\n" ::);
    __syncthreads();

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

    int smem_pipe_read = 0;
    int smem_pipe_write = 1;

    // prefetch buffer 0 stage 0 to registers 0
    load_fragments(
        a_regs[0],
        b_regs[0],
        sfa_regs[0],
        sfb_regs[0],
        lane,
        0, //smem_pipe_read
        0, //k_block
        smem_offset_A,
        smem_offset_B,
        smem_sf_offset,
        smem
    );
    asm volatile("cp.async.commit_group;\n" ::);

    int iters = params.k / (THREADS_PER_ROW * 16);

    int idx = 0;
    while (idx < iters) {
        int smem_pipe_read_curr = smem_pipe_read;

        for (int k_block = 0; k_block < STAGES; ++k_block)
        {
            if (k_block == STAGES-1)
            {
                asm volatile("cp.async.wait_all;\n" ::);
                __syncthreads();

                smem_pipe_read_curr = smem_pipe_read;
            }

            auto k_block_next = (k_block + 1) % STAGES;
            int frag_idx_next = (k_block + 1) & 1;

            load_fragments(
                a_regs[frag_idx_next],
                b_regs[frag_idx_next],
                sfa_regs[frag_idx_next],
                sfb_regs[frag_idx_next],
                lane,
                smem_pipe_read_curr, //smem_pipe_read
                k_block_next, //k_block
                smem_offset_A,
                smem_offset_B,
                smem_sf_offset,
                smem
            );

            if (k_block == 0)
            {
                bool valid = (global_k < k);
                cpasync_load(
                    smem_pipe_write,
                    global_k,
                    smem_offset_A,
                    smem_offset_B,
                    smem_sf_write_offset,
                    lane,
                    is_even,
                    load_b,
                    valid,
                    rowA,
                    vecB,
                    rowS,
                    vecS,
                    smem);
                asm volatile("cp.async.commit_group;\n" ::);

                smem_pipe_write = smem_pipe_read;
                smem_pipe_read = (smem_pipe_read + 1) & 1;
            }

            
            int frag_idx = k_block & 1;
            __half h = block_scaled_fma_16x2fp4(
                a_regs[frag_idx], 
                b_regs[frag_idx], 
                sfa_regs[frag_idx], 
                sfb_regs[frag_idx]);
            sum += __half2float(h);
            
        }
        idx += STAGES;
    }

    asm volatile("cp.async.wait_all;\n" ::);
    __syncthreads();


    unsigned mask = 0xffffffffu;
    sum += __shfl_down_sync(mask, sum, 8, 16);
    sum += __shfl_down_sync(mask, sum, 4, 16);
    sum += __shfl_down_sync(mask, sum, 2, 16);
    sum += __shfl_down_sync(mask, sum, 1, 16);

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

static inline void launch_kernel_fast(Gemv_params &params, cudaStream_t stream)
{
    const int grid_x = (params.m + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
    size_t smem_size = sizeof(SharedStorage);
    dim3 grid(grid_x, 1, params.b);
    dim3 block(BLOCK_SIZE, 1, 1);
    gemv_kernel_fast<<<grid, block, smem_size>>>(params);
}

__global__ void __launch_bounds__(BLOCK_SIZE, 8)
gemv_kernel_k1024(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;

    // each thread processes 16 2xFP4 elements per iteration
    const int iters = params.k / (THREADS_PER_ROW * 16);

    for (int idx = 0; idx < iters; ++idx) {
        int block_base = idx * THREADS_PER_ROW + lane;
        int elem_base  = block_base * 16;      // 16 __nv_fp4x2 per (idx,lane)

        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); // 2 fp8 packed in u16
        uint64_t vecS_addr = reinterpret_cast<uint64_t>(vecS_u16 + block_base);

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

        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)
        );

        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);

        uint16_t out_half_bits;

        asm volatile(
            "{\n"
            ".reg .b8 byte0_0, byte0_1, byte0_2, byte0_3;\n"
            ".reg .b8 byte0_4, byte0_5, byte0_6, byte0_7;\n"
            ".reg .b8 byte0_8, byte0_9, byte0_10, byte0_11;\n"
            ".reg .b8 byte0_12, byte0_13, byte0_14, byte0_15;\n"
            ".reg .b8 byte1_0, byte1_1, byte1_2, byte1_3;\n"
            ".reg .b8 byte1_4, byte1_5, byte1_6, byte1_7;\n"
            ".reg .b8 byte1_8, byte1_9, byte1_10, byte1_11;\n"
            ".reg .b8 byte1_12, byte1_13, byte1_14, byte1_15;\n"

            ".reg .f16x2 accum_0, accum_1, accum_2, accum_3;\n"
            ".reg .f16x2 accum_4, accum_5, accum_6, accum_7;\n"
            ".reg .f16x2 accum_8, accum_9, accum_10, accum_11;\n"
            ".reg .f16x2 accum_12, accum_13, accum_14, accum_15;\n"

            ".reg .f16x2 sfa_f16x2;\n"
            ".reg .f16x2 sfb_f16x2;\n"
            ".reg .f16x2 sf_f16x2;\n"

            ".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_0_8, cvt_0_9, cvt_0_10, cvt_0_11;\n"
            ".reg .f16x2 cvt_0_12, cvt_0_13, cvt_0_14, cvt_0_15;\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 .f16x2 cvt_1_8, cvt_1_9, cvt_1_10, cvt_1_11;\n"
            ".reg .f16x2 cvt_1_12, cvt_1_13, cvt_1_14, cvt_1_15;\n"
            ".reg .f16 result_f16, lane0, lane1;\n"
            ".reg .f16x2 mul_f16x2_0, mul_f16x2_1;\n"

            "cvt.rn.f16x2.e4m3x2 sfa_f16x2, %9;\n"
            "cvt.rn.f16x2.e4m3x2 sfb_f16x2, %10;\n"

            "mov.b32 accum_0, 0;\n"
            "mov.b32 accum_1, 0;\n"
            "mov.b32 accum_2, 0;\n"
            "mov.b32 accum_3, 0;\n"
            "mov.b32 accum_4, 0;\n"
            "mov.b32 accum_5, 0;\n"
            "mov.b32 accum_6, 0;\n"
            "mov.b32 accum_7, 0;\n"
            "mov.b32 accum_8, 0;\n"
            "mov.b32 accum_9, 0;\n"
            "mov.b32 accum_10, 0;\n"
            "mov.b32 accum_11, 0;\n"
            "mov.b32 accum_12, 0;\n"
            "mov.b32 accum_13, 0;\n"
            "mov.b32 accum_14, 0;\n"
            "mov.b32 accum_15, 0;\n"

            "mul.rn.f16x2 sf_f16x2, sfa_f16x2, sfb_f16x2;\n"
            "mov.b32 {lane0, lane1}, sf_f16x2;\n"
            "mov.b32 mul_f16x2_0, {lane0, lane0};\n"
            "mov.b32 mul_f16x2_1, {lane1, lane1};\n"

            "mov.b32 {byte0_0, byte0_1, byte0_2, byte0_3}, %1;\n"
            "mov.b32 {byte0_4, byte0_5, byte0_6, byte0_7}, %2;\n"
            "mov.b32 {byte0_8, byte0_9, byte0_10, byte0_11}, %3;\n"
            "mov.b32 {byte0_12, byte0_13, byte0_14, byte0_15}, %4;\n"
            "mov.b32 {byte1_0, byte1_1, byte1_2, byte1_3}, %5;\n"
            "mov.b32 {byte1_4, byte1_5, byte1_6, byte1_7}, %6;\n"
            "mov.b32 {byte1_8, byte1_9, byte1_10, byte1_11}, %7;\n"
            "mov.b32 {byte1_12, byte1_13, byte1_14, byte1_15}, %8;\n"

            "cvt.rn.f16x2.e2m1x2 cvt_0_0,  byte0_0;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_0_1,  byte0_1;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_0_2,  byte0_2;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_0_3,  byte0_3;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_0_4,  byte0_4;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_0_5,  byte0_5;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_0_6,  byte0_6;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_0_7,  byte0_7;\n"

            "cvt.rn.f16x2.e2m1x2 cvt_0_8,  byte0_8;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_0_9,  byte0_9;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_0_10, byte0_10;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_0_11, byte0_11;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_0_12, byte0_12;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_0_13, byte0_13;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_0_14, byte0_14;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_0_15, byte0_15;\n"

            "cvt.rn.f16x2.e2m1x2 cvt_1_0,  byte1_0;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_1_1,  byte1_1;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_1_2,  byte1_2;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_1_3,  byte1_3;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_1_4,  byte1_4;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_1_5,  byte1_5;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_1_6,  byte1_6;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_1_7,  byte1_7;\n"

            "cvt.rn.f16x2.e2m1x2 cvt_1_8,  byte1_8;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_1_9,  byte1_9;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_1_10, byte1_10;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_1_11, byte1_11;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_1_12, byte1_12;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_1_13, byte1_13;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_1_14, byte1_14;\n"
            "cvt.rn.f16x2.e2m1x2 cvt_1_15, byte1_15;\n"

            "fma.rn.f16x2 accum_0,  cvt_0_0,  cvt_1_0,  accum_0;\n"
            "fma.rn.f16x2 accum_1,  cvt_0_1,  cvt_1_1,  accum_1;\n"
            "fma.rn.f16x2 accum_2,  cvt_0_2,  cvt_1_2,  accum_2;\n"
            "fma.rn.f16x2 accum_3,  cvt_0_3,  cvt_1_3,  accum_3;\n"
            "fma.rn.f16x2 accum_4,  cvt_0_4,  cvt_1_4,  accum_4;\n"
            "fma.rn.f16x2 accum_5,  cvt_0_5,  cvt_1_5,  accum_5;\n"
            "fma.rn.f16x2 accum_6,  cvt_0_6,  cvt_1_6,  accum_6;\n"
            "fma.rn.f16x2 accum_7,  cvt_0_7,  cvt_1_7,  accum_7;\n"

            "fma.rn.f16x2 accum_8,  cvt_0_8,  cvt_1_8,  accum_8;\n"
            "fma.rn.f16x2 accum_9,  cvt_0_9,  cvt_1_9,  accum_9;\n"
            "fma.rn.f16x2 accum_10, cvt_0_10, cvt_1_10, accum_10;\n"
            "fma.rn.f16x2 accum_11, cvt_0_11, cvt_1_11, accum_11;\n"
            "fma.rn.f16x2 accum_12, cvt_0_12, cvt_1_12, accum_12;\n"
            "fma.rn.f16x2 accum_13, cvt_0_13, cvt_1_13, accum_13;\n"
            "fma.rn.f16x2 accum_14, cvt_0_14, cvt_1_14, accum_14;\n"
            "fma.rn.f16x2 accum_15, cvt_0_15, cvt_1_15, accum_15;\n"

            "add.rn.f16x2 accum_0, accum_0, accum_1;\n"
            "add.rn.f16x2 accum_2, accum_2, accum_3;\n"
            "add.rn.f16x2 accum_4, accum_4, accum_5;\n"
            "add.rn.f16x2 accum_6, accum_6, accum_7;\n"
            "add.rn.f16x2 accum_8, accum_8, accum_9;\n"
            "add.rn.f16x2 accum_10, accum_10, accum_11;\n"
            "add.rn.f16x2 accum_12, accum_12, accum_13;\n"
            "add.rn.f16x2 accum_14, accum_14, accum_15;\n"

            "add.rn.f16x2 accum_0, accum_0, accum_2;\n"
            "add.rn.f16x2 accum_4, accum_4, accum_6;\n"
            "add.rn.f16x2 accum_8, accum_8, accum_10;\n"
            "add.rn.f16x2 accum_12, accum_12, accum_14;\n"

            "add.rn.f16x2 accum_0, accum_0, accum_4;\n"
            "add.rn.f16x2 accum_8, accum_8, accum_12;\n"

            "mul.rn.f16x2 accum_0, mul_f16x2_0, accum_0;\n"
            "mul.rn.f16x2 accum_8, mul_f16x2_1, accum_8;\n"

            "add.rn.f16x2 accum_0, accum_0, accum_8;\n"

            "mov.b32 {lane0, lane1}, accum_0;\n"
            "add.rn.f16 result_f16, lane0, lane1;\n"

            "mov.b16 %0, result_f16;\n"
            "}\n"
            : "=h"(out_half_bits)
            : "r"(a_regs_packed[0]), "r"(a_regs_packed[1]),
              "r"(a_regs_packed[2]), "r"(a_regs_packed[3]),
              "r"(b_regs_packed[0]), "r"(b_regs_packed[1]),
              "r"(b_regs_packed[2]), "r"(b_regs_packed[3]),
              "h"(sfa_regs), "h"(sfb_regs)
            : "memory"
        );

        half h = *reinterpret_cast<half*>(&out_half_bits);
        sum += __half2float(h);
    }

    unsigned mask = 0xffffffffu;
    sum += __shfl_down_sync(mask, sum, 8, 16);
    sum += __shfl_down_sync(mask, sum, 4, 16);
    sum += __shfl_down_sync(mask, sum, 2, 16);
    sum += __shfl_down_sync(mask, sum, 1, 16);

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

static inline void launch_kernel_k1024(Gemv_params &params, cudaStream_t stream)
{
    const int grid_x = (params.m + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
    size_t smem_size = sizeof(SharedStorage);
    dim3 grid(grid_x, 1, params.b);
    dim3 block(BLOCK_SIZE, 1, 1);
    gemv_kernel_k1024<<<grid, block, 0, stream>>>(params);
}

// ============================================================================
// K=3584-specialized kernel 
// ============================================================================

__global__ void __launch_bounds__(BLOCK_SIZE, 8)
gemv_kernel_k3584(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;

    extern __shared__ __align__(128) uint8_t shared_storage[];
    SharedStorage &smem = *reinterpret_cast<SharedStorage*>(shared_storage);

    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;

    rowA += lane * 16;
    vecB += lane * 16;

    const int k = params.k;

    float sum = 0.f;
    int global_k = 0;
    const int smem_offset_A = tid * 16;
    const int smem_offset_B = lane * 16;
    const int smem_sf_write_offset = (tid / 2) * 4;
    const int smem_sf_offset = tid * 2;
    const bool load_b = rib == 0;
    const bool is_even = (lane % 2 == 0);

    // prefetch to buffer 0
    cpasync_load(
        0,
        global_k,
        smem_offset_A,
        smem_offset_B,
        smem_sf_write_offset,
        lane,
        is_even,
        load_b,
        true,
        rowA,
        vecB,
        rowS,
        vecS,
        smem
    );
    asm volatile("cp.async.commit_group;\n" ::);
    asm volatile("cp.async.wait_all;\n" ::);
    __syncthreads();

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

    int smem_pipe_read = 0;
    int smem_pipe_write = 1;

    // prefetch buffer 0 stage 0 to registers 0
    load_fragments(
        a_regs[0],
        b_regs[0],
        sfa_regs[0],
        sfb_regs[0],
        lane,
        0, //smem_pipe_read
        0, //k_block
        smem_offset_A,
        smem_offset_B,
        smem_sf_offset,
        smem
    );
    asm volatile("cp.async.commit_group;\n" ::);

    int total_stages = params.k / (THREADS_PER_ROW * 16);
    int full_iters = total_stages / STAGES;  // Number of complete STAGES-groups
    int tail_stages = total_stages % STAGES; // Remaining stages

    // Main loop: process complete groups of STAGES
    for (int idx = 0; idx < full_iters; ++idx) {
        int smem_pipe_read_curr = smem_pipe_read;

        for (int k_block = 0; k_block < STAGES; ++k_block)
        {
            if (k_block == STAGES-1)
            {
                asm volatile("cp.async.wait_all;\n" ::);
                __syncthreads();

                smem_pipe_read_curr = smem_pipe_read;
            }

            auto k_block_next = (k_block + 1) % STAGES;
            int frag_idx_next = (k_block + 1) & 1;

            load_fragments(
                a_regs[frag_idx_next],
                b_regs[frag_idx_next],
                sfa_regs[frag_idx_next],
                sfb_regs[frag_idx_next],
                lane,
                smem_pipe_read_curr,
                k_block_next,
                smem_offset_A,
                smem_offset_B,
                smem_sf_offset,
                smem
            );

            if (k_block == 0)
            {
                bool valid = (global_k < k);
                cpasync_load(
                    smem_pipe_write,
                    global_k,
                    smem_offset_A,
                    smem_offset_B,
                    smem_sf_write_offset,
                    lane,
                    is_even,
                    load_b,
                    valid,
                    rowA,
                    vecB,
                    rowS,
                    vecS,
                    smem);
                asm volatile("cp.async.commit_group;\n" ::);

                smem_pipe_write = smem_pipe_read;
                smem_pipe_read = (smem_pipe_read + 1) & 1;
            }

            int frag_idx = k_block & 1;
            __half h = block_scaled_fma_16x2fp4(
                a_regs[frag_idx], 
                b_regs[frag_idx], 
                sfa_regs[frag_idx], 
                sfb_regs[frag_idx]);
            sum += __half2float(h);
        }
    }

    // Tail loop: process remaining stages
    if (tail_stages > 0) {
        asm volatile("cp.async.wait_all;\n" ::);
        __syncthreads();

        for (int k_block = 0; k_block < tail_stages; ++k_block)
        {
            int frag_idx = k_block & 1;
            int frag_idx_next = (k_block + 1) & 1;

            if (k_block < tail_stages - 1) {
                load_fragments(
                    a_regs[frag_idx_next],
                    b_regs[frag_idx_next],
                    sfa_regs[frag_idx_next],
                    sfb_regs[frag_idx_next],
                    lane,
                    smem_pipe_read,
                    k_block + 1,
                    smem_offset_A,
                    smem_offset_B,
                    smem_sf_offset,
                    smem
                );
            }

            __half h = block_scaled_fma_16x2fp4(
                a_regs[frag_idx], 
                b_regs[frag_idx], 
                sfa_regs[frag_idx], 
                sfb_regs[frag_idx]);
            sum += __half2float(h);
        }
    }

    asm volatile("cp.async.wait_all;\n" ::);
    __syncthreads();


    unsigned mask = 0xffffffffu;
    sum += __shfl_down_sync(mask, sum, 8, 16);
    sum += __shfl_down_sync(mask, sum, 4, 16);
    sum += __shfl_down_sync(mask, sum, 2, 16);
    sum += __shfl_down_sync(mask, sum, 1, 16);

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

static inline void launch_kernel_k3584(Gemv_params &params, cudaStream_t stream)
{
    const int grid_x = (params.m + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
    size_t smem_size = sizeof(SharedStorage);
    dim3 grid(grid_x, 1, params.b);
    dim3 block(BLOCK_SIZE, 1, 1);
    gemv_kernel_k3584<<<grid, block, smem_size>>>(params);
}

// ============================================================================
// K=256-specialized kernel (your previous v1) -> gemv_kernel_k256 / launch_kernel_k256
// ============================================================================

__global__ void __launch_bounds__(BLOCK_SIZE, 8)
gemv_kernel_k256(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;

    float sum = 0.f;

    // Each thread does 1 16 group FP4 or (8 2xFP4)
    for (int idx = 0; idx < params.k / THREADS_PER_ROW / 8; ++idx) {
        int base = idx * 16;
        const int base_id = (idx * THREADS_PER_ROW + lane) * 8;

        __nv_fp8_storage_t sfa_storage = *reinterpret_cast<const __nv_fp8_storage_t*>(&rowS[base + lane]);
        __nv_fp8_storage_t sfb_storage = *reinterpret_cast<const __nv_fp8_storage_t*>(&vecS[base + lane]);
        __half sfa = __nv_cvt_fp8_to_halfraw(sfa_storage, __NV_E4M3);
        __half sfb = __nv_cvt_fp8_to_halfraw(sfb_storage, __NV_E4M3);
        __half scale = __hmul(sfa, sfb);

        __half2 acc = __float2half2_rn(0.0f);

        #pragma unroll
        for (int i = 0; i < 8; ++i) { // go over each individual 2xFP4
            const int id = base_id + i;
            __nv_fp4x2_storage_t a_storage = *reinterpret_cast<const __nv_fp4x2_storage_t*>(&rowA[id]);
            __nv_fp4x2_storage_t b_storage = *reinterpret_cast<const __nv_fp4x2_storage_t*>(&vecB[id]);
            __half2_raw a_raw = __nv_cvt_fp4x2_to_halfraw2(a_storage, __NV_E2M1);
            __half2_raw b_raw = __nv_cvt_fp4x2_to_halfraw2(b_storage, __NV_E2M1);

            const __half2 a_h2 = __half2(a_raw);
            const __half2 b_h2 = __half2(b_raw);
            
            acc = __hfma2(a_h2, b_h2, acc);
        }

        __half fin = __hadd(__low2half(acc), __high2half(acc));
        __half h = __hmul(fin, scale);

        sum += __half2float(h);
    }

    // Reduce within the 16-thread subgroup (one output row)
    unsigned mask = 0xffffffffu;
    sum += __shfl_down_sync(mask, sum, 8, 16);
    sum += __shfl_down_sync(mask, sum, 4, 16);
    sum += __shfl_down_sync(mask, sum, 2, 16);
    sum += __shfl_down_sync(mask, sum, 1, 16);

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

static inline void launch_kernel_k256(Gemv_params &params, cudaStream_t stream)
{
    const int grid_x = (params.m + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
    dim3 grid(grid_x, 1, params.b);
    dim3 block(BLOCK_SIZE, 1, 1);
    gemv_kernel_k256<<<grid, block, 0, stream>>>(params);
}

// ============================================================================
// Python-facing function: sets up params and dispatches based on K
// ============================================================================

torch::Tensor nvfp4_gemv_dispatch(torch::Tensor A,
                                  torch::Tensor B,
                                  torch::Tensor C,
                                  torch::Tensor SFA,
                                  torch::Tensor SFB)
{
    auto sizes = A.sizes();
    const int64_t M = sizes[0];
    const int64_t K = sizes[1];
    const int64_t L = sizes[2];

    Gemv_params params{};
    params.b = static_cast<int>(L);
    params.m = static_cast<int>(M);
    params.k = static_cast<int>(K);
    params.real_k = static_cast<int>(K * 2);

    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   = static_cast<uint64_t>(A.stride(2));
    params.b_batch_stride   = static_cast<uint64_t>(B.stride(2));
    params.sfa_batch_stride = static_cast<uint64_t>(SFA.stride(2));
    params.sfb_batch_stride = static_cast<uint64_t>(SFB.stride(2));
    params.o_batch_stride   = static_cast<uint64_t>(C.stride(2));

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

    auto stream = at::cuda::getCurrentCUDAStream().stream();

    // Tiny host-side dispatch: effectively zero overhead vs kernel time
    if (params.k < 512) {
        launch_kernel_k256(params, stream);
    } 
    else if (params.k == 1024) {
        launch_kernel_k1024(params, stream);
    }
    else if (params.k % 1024) {
        launch_kernel_k3584(params, stream);
    }
    else {
        launch_kernel_fast(params, stream);
    }

    return C;
}
"""

# ---- build the module ----
nvfp4_module = load_inline(
    name="nvfp4_gemv",
    cpp_sources=[gemv_cpp],
    cuda_sources=[gemv_cuda],
    functions=["nvfp4_gemv_dispatch"],  # single Python-visible entry point
    extra_cuda_cflags=[
        "-std=c++17",
        "-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.nvfp4_gemv_dispatch(data[0], data[1], data[6], data[2], data[3])
scrolls · 1114 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 74603.

⋯ 55 unchanged lines
static constexpr int ROWS_PER_BLOCK = 8;
static constexpr int THREADS_PER_ROW = 16;
- static constexpr int BLOCK_SIZE = ROWS_PER_BLOCK * THREADS_PER_ROW; // 128
+ static constexpr int BLOCK_SIZE = ROWS_PER_BLOCK * THREADS_PER_ROW; // 128
+ static constexpr int BUFFER = 2;
+ static constexpr int STAGES = 4;
+ template <int LoadBytes>
+ __device__ __forceinline__ void cp_async_ca(uint32_t smem_addr, void const *global_ptr, bool pred_guard=true) {
+ asm volatile(
+ "{\n"
+ " .reg .pred p;\n"
+ " setp.ne.b32 p, %0, 0;\n"
+ " @p cp.async.ca.shared.global.L2::128B [%1], [%2], %3;\n"
+ "}\n"
+ :
+ :"r"((int)pred_guard), "r"(smem_addr), "l"(global_ptr), "n"(LoadBytes)
+ );
+ }
+
+ struct SharedStorage {
+ alignas(128) __nv_fp4x2_e2m1 A[BUFFER][STAGES][128 * 16];
+ alignas(128) __nv_fp4x2_e2m1 B[BUFFER][STAGES][THREADS_PER_ROW * 16];
+ alignas(128) __nv_fp8_e4m3 SFA[BUFFER][STAGES][128 * 2];
+ alignas(128) __nv_fp8_e4m3 SFB[BUFFER][STAGES][THREADS_PER_ROW * 2];
+ };
+
+
+ __device__ __forceinline__ void cpasync_load(
+ int buffer_idx,
+ int &global_k, // should be in 2xfp4
+ int smem_offset_A,
+ int smem_offset_B,
+ int smem_sf_write_offset,
+ int lane,
+ bool is_even,
+ bool load_b,
+ bool valid,
+ const __nv_fp4x2_e2m1* rowA,
+ const __nv_fp4x2_e2m1* vecB,
+ const __nv_fp8_e4m3* rowS,
+ const __nv_fp8_e4m3* vecS,
+ SharedStorage& smem)
+ {
+ for (int s = 0; s < STAGES; ++s) {
+ int SF_offset = (global_k >> 3) + lane * 2;
+
+ uint32_t smem_addr_A = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.A[buffer_idx][s][smem_offset_A]));
+ uint32_t smem_addr_B = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.B[buffer_idx][s][smem_offset_B]));
+ uint32_t smem_addr_SFA = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.SFA[buffer_idx][s][smem_sf_write_offset]));
+ uint32_t smem_addr_SFB = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.SFB[buffer_idx][s][(lane / 2) * 4]));
+
+ cp_async_ca<16>(smem_addr_A, rowA + global_k, valid);
+ cp_async_ca<16>(smem_addr_B, vecB + global_k, valid && load_b);
+ cp_async_ca<4>(smem_addr_SFA, rowS + SF_offset, valid && is_even);
+ cp_async_ca<4>(smem_addr_SFB, vecS + SF_offset, valid && is_even && load_b);
+
+ global_k += 256;
+ }
+ }
+
+ __device__ __forceinline__ void load_fragments(
+ uint64_t (&a_regs)[2],
+ uint64_t (&b_regs)[2],
+ uint16_t &sfa_regs,
+ uint16_t &sfb_regs,
+ int lane,
+ int smem_pipe_read,
+ int k_block,
+ int smem_offset_A,
+ int smem_offset_B,
+ int smem_sf_offset,
+ SharedStorage& smem)
+ {
+ uint32_t smem_addr_a = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.A[smem_pipe_read][k_block][smem_offset_A]));
+ uint32_t smem_addr_b = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.B[smem_pipe_read][k_block][smem_offset_B]));
+
+ asm volatile(
+ "ld.shared.v2.u64 {%0, %1}, [%4];\n\t"
+ "ld.shared.v2.u64 {%2, %3}, [%5];\n\t"
+ : "=l"(a_regs[0]), "=l"(a_regs[1]),
+ "=l"(b_regs[0]), "=l"(b_regs[1])
+ : "r"(smem_addr_a), "r"(smem_addr_b)
+ );
+
+ uint32_t smem_addr_sfa = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.SFA[smem_pipe_read][k_block][smem_sf_offset]));
+ uint32_t smem_addr_sfb = static_cast<uint32_t>(__cvta_generic_to_shared(&smem.SFB[smem_pipe_read][k_block][lane * 2]));
+
+ asm volatile(
+ "ld.shared.u16 %0, [%2];\n\t"
+ "ld.shared.u16 %1, [%3];\n\t"
+ : "=h"(sfa_regs), "=h"(sfb_regs)
+ : "r"(smem_addr_sfa), "r"(smem_addr_sfb)
+ );
+ }
+
+
+
+ __device__ __forceinline__ __half 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);
+
+ uint16_t out_half_bits;
+
+ asm volatile(
+ "{\n"
+ ".reg .b8 byte0_0, byte0_1, byte0_2, byte0_3;\n"
+ ".reg .b8 byte0_4, byte0_5, byte0_6, byte0_7;\n"
+ ".reg .b8 byte0_8, byte0_9, byte0_10, byte0_11;\n"
+ ".reg .b8 byte0_12, byte0_13, byte0_14, byte0_15;\n"
+ ".reg .b8 byte1_0, byte1_1, byte1_2, byte1_3;\n"
+ ".reg .b8 byte1_4, byte1_5, byte1_6, byte1_7;\n"
+ ".reg .b8 byte1_8, byte1_9, byte1_10, byte1_11;\n"
+ ".reg .b8 byte1_12, byte1_13, byte1_14, byte1_15;\n"
+
+ ".reg .f16x2 accum_0, accum_1, accum_2, accum_3;\n"
+ ".reg .f16x2 accum_4, accum_5, accum_6, accum_7;\n"
+ ".reg .f16x2 accum_8, accum_9, accum_10, accum_11;\n"
+ ".reg .f16x2 accum_12, accum_13, accum_14, accum_15;\n"
+
+ ".reg .f16x2 sfa_f16x2;\n"
+ ".reg .f16x2 sfb_f16x2;\n"
+ ".reg .f16x2 sf_f16x2;\n"
+
+ ".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_0_8, cvt_0_9, cvt_0_10, cvt_0_11;\n"
+ ".reg .f16x2 cvt_0_12, cvt_0_13, cvt_0_14, cvt_0_15;\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 .f16x2 cvt_1_8, cvt_1_9, cvt_1_10, cvt_1_11;\n"
+ ".reg .f16x2 cvt_1_12, cvt_1_13, cvt_1_14, cvt_1_15;\n"
+ ".reg .f16 result_f16, lane0, lane1;\n"
+ ".reg .f16x2 mul_f16x2_0, mul_f16x2_1;\n"
+
+ "cvt.rn.f16x2.e4m3x2 sfa_f16x2, %9;\n"
+ "cvt.rn.f16x2.e4m3x2 sfb_f16x2, %10;\n"
+
+ "mov.b32 accum_0, 0;\n"
+ "mov.b32 accum_1, 0;\n"
+ "mov.b32 accum_2, 0;\n"
+ "mov.b32 accum_3, 0;\n"
+ "mov.b32 accum_4, 0;\n"
+ "mov.b32 accum_5, 0;\n"
+ "mov.b32 accum_6, 0;\n"
+ "mov.b32 accum_7, 0;\n"
+ "mov.b32 accum_8, 0;\n"
+ "mov.b32 accum_9, 0;\n"
+ "mov.b32 accum_10, 0;\n"
+ "mov.b32 accum_11, 0;\n"
+ "mov.b32 accum_12, 0;\n"
+ "mov.b32 accum_13, 0;\n"
+ "mov.b32 accum_14, 0;\n"
+ "mov.b32 accum_15, 0;\n"
+
+ "mul.rn.f16x2 sf_f16x2, sfa_f16x2, sfb_f16x2;\n"
+ "mov.b32 {lane0, lane1}, sf_f16x2;\n"
+ "mov.b32 mul_f16x2_0, {lane0, lane0};\n"
+ "mov.b32 mul_f16x2_1, {lane1, lane1};\n"
+
+ "mov.b32 {byte0_0, byte0_1, byte0_2, byte0_3}, %1;\n"
+ "mov.b32 {byte0_4, byte0_5, byte0_6, byte0_7}, %2;\n"
+ "mov.b32 {byte0_8, byte0_9, byte0_10, byte0_11}, %3;\n"
+ "mov.b32 {byte0_12, byte0_13, byte0_14, byte0_15}, %4;\n"
+ "mov.b32 {byte1_0, byte1_1, byte1_2, byte1_3}, %5;\n"
+ "mov.b32 {byte1_4, byte1_5, byte1_6, byte1_7}, %6;\n"
+ "mov.b32 {byte1_8, byte1_9, byte1_10, byte1_11}, %7;\n"
+ "mov.b32 {byte1_12, byte1_13, byte1_14, byte1_15}, %8;\n"
+
+ "cvt.rn.f16x2.e2m1x2 cvt_0_0, byte0_0;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_0_1, byte0_1;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_0_2, byte0_2;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_0_3, byte0_3;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_0_4, byte0_4;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_0_5, byte0_5;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_0_6, byte0_6;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_0_7, byte0_7;\n"
+
+ "cvt.rn.f16x2.e2m1x2 cvt_0_8, byte0_8;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_0_9, byte0_9;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_0_10, byte0_10;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_0_11, byte0_11;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_0_12, byte0_12;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_0_13, byte0_13;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_0_14, byte0_14;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_0_15, byte0_15;\n"
+
+ "cvt.rn.f16x2.e2m1x2 cvt_1_0, byte1_0;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_1_1, byte1_1;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_1_2, byte1_2;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_1_3, byte1_3;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_1_4, byte1_4;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_1_5, byte1_5;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_1_6, byte1_6;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_1_7, byte1_7;\n"
+
+ "cvt.rn.f16x2.e2m1x2 cvt_1_8, byte1_8;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_1_9, byte1_9;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_1_10, byte1_10;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_1_11, byte1_11;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_1_12, byte1_12;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_1_13, byte1_13;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_1_14, byte1_14;\n"
+ "cvt.rn.f16x2.e2m1x2 cvt_1_15, byte1_15;\n"
+
+ "fma.rn.f16x2 accum_0, cvt_0_0, cvt_1_0, accum_0;\n"
+ "fma.rn.f16x2 accum_1, cvt_0_1, cvt_1_1, accum_1;\n"
+ "fma.rn.f16x2 accum_2, cvt_0_2, cvt_1_2, accum_2;\n"
+ "fma.rn.f16x2 accum_3, cvt_0_3, cvt_1_3, accum_3;\n"
+ "fma.rn.f16x2 accum_4, cvt_0_4, cvt_1_4, accum_4;\n"
+ "fma.rn.f16x2 accum_5, cvt_0_5, cvt_1_5, accum_5;\n"
+ "fma.rn.f16x2 accum_6, cvt_0_6, cvt_1_6, accum_6;\n"
+ "fma.rn.f16x2 accum_7, cvt_0_7, cvt_1_7, accum_7;\n"
+
+ "fma.rn.f16x2 accum_8, cvt_0_8, cvt_1_8, accum_8;\n"
+ "fma.rn.f16x2 accum_9, cvt_0_9, cvt_1_9, accum_9;\n"
+ "fma.rn.f16x2 accum_10, cvt_0_10, cvt_1_10, accum_10;\n"
+ "fma.rn.f16x2 accum_11, cvt_0_11, cvt_1_11, accum_11;\n"
+ "fma.rn.f16x2 accum_12, cvt_0_12, cvt_1_12, accum_12;\n"
+ "fma.rn.f16x2 accum_13, cvt_0_13, cvt_1_13, accum_13;\n"
+ "fma.rn.f16x2 accum_14, cvt_0_14, cvt_1_14, accum_14;\n"
+ "fma.rn.f16x2 accum_15, cvt_0_15, cvt_1_15, accum_15;\n"
+
+ "add.rn.f16x2 accum_0, accum_0, accum_1;\n"
+ "add.rn.f16x2 accum_2, accum_2, accum_3;\n"
+ "add.rn.f16x2 accum_4, accum_4, accum_5;\n"
+ "add.rn.f16x2 accum_6, accum_6, accum_7;\n"
+ "add.rn.f16x2 accum_8, accum_8, accum_9;\n"
+ "add.rn.f16x2 accum_10, accum_10, accum_11;\n"
+ "add.rn.f16x2 accum_12, accum_12, accum_13;\n"
+ "add.rn.f16x2 accum_14, accum_14, accum_15;\n"
+
+ "add.rn.f16x2 accum_0, accum_0, accum_2;\n"
+ "add.rn.f16x2 accum_4, accum_4, accum_6;\n"
+ "add.rn.f16x2 accum_8, accum_8, accum_10;\n"
+ "add.rn.f16x2 accum_12, accum_12, accum_14;\n"
+
+ "add.rn.f16x2 accum_0, accum_0, accum_4;\n"
+ "add.rn.f16x2 accum_8, accum_8, accum_12;\n"
+
+ "mul.rn.f16x2 accum_0, mul_f16x2_0, accum_0;\n"
+ "mul.rn.f16x2 accum_8, mul_f16x2_1, accum_8;\n"
+
+ "add.rn.f16x2 accum_0, accum_0, accum_8;\n"
+
+ "mov.b32 {lane0, lane1}, accum_0;\n"
+ "add.rn.f16 result_f16, lane0, lane1;\n"
+
+ "mov.b16 %0, result_f16;\n"
+ "}\n"
+ : "=h"(out_half_bits)
+ : "r"(a_regs_packed[0]), "r"(a_regs_packed[1]),
+ "r"(a_regs_packed[2]), "r"(a_regs_packed[3]),
+ "r"(b_regs_packed[0]), "r"(b_regs_packed[1]),
+ "r"(b_regs_packed[2]), "r"(b_regs_packed[3]),
+ "h"(sfa_regs), "h"(sfb_regs)
+ : "memory"
+ );
+
+ union { uint16_t u; __half h; } conv;
+ conv.u = out_half_bits;
+ return conv.h;
+ }
+
+
__global__ void __launch_bounds__(BLOCK_SIZE, 8)
gemv_kernel_fast(const __grid_constant__ Gemv_params params)
{
⋯ 3 unchanged lines
const int batch = blockIdx.z;
const int row = blockIdx.x * ROWS_PER_BLOCK + rib;
+ extern __shared__ __align__(128) uint8_t shared_storage[];
+ SharedStorage &smem = *reinterpret_cast<SharedStorage*>(shared_storage);
+
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;
+
+ rowA += lane * 16;
+ vecB += lane * 16;
+
+ const int k = params.k;
+
+ float sum = 0.f;
+ int global_k = 0;
+ const int smem_offset_A = tid * 16;
+ const int smem_offset_B = lane * 16;
+ const int smem_sf_write_offset = (tid / 2) * 4;
+ const int smem_sf_offset = tid * 2;
+ const bool load_b = rib == 0;
+ const bool is_even = (lane % 2 == 0);
+
+ // prefetch to buffer 0
+ cpasync_load(
+ 0,
+ global_k,
+ smem_offset_A,
+ smem_offset_B,
+ smem_sf_write_offset,
+ lane,
+ is_even,
+ load_b,
+ true,
+ rowA,
+ vecB,
+ rowS,
+ vecS,
+ smem
+ );
+ asm volatile("cp.async.commit_group;\n" ::);
+ asm volatile("cp.async.wait_all;\n" ::);
+ __syncthreads();
+
+ uint64_t a_regs[2][2], b_regs[2][2];
+ uint16_t sfa_regs[2], sfb_regs[2];
+
+ int smem_pipe_read = 0;
+ int smem_pipe_write = 1;
+
+ // prefetch buffer 0 stage 0 to registers 0
+ load_fragments(
+ a_regs[0],
+ b_regs[0],
+ sfa_regs[0],
+ sfb_regs[0],
+ lane,
+ 0, //smem_pipe_read
+ 0, //k_block
+ smem_offset_A,
+ smem_offset_B,
+ smem_sf_offset,
+ smem
+ );
+ asm volatile("cp.async.commit_group;\n" ::);
+
+ int iters = params.k / (THREADS_PER_ROW * 16);
+
+ int idx = 0;
+ while (idx < iters) {
+ int smem_pipe_read_curr = smem_pipe_read;
+
+ for (int k_block = 0; k_block < STAGES; ++k_block)
+ {
+ if (k_block == STAGES-1)
+ {
+ asm volatile("cp.async.wait_all;\n" ::);
+ __syncthreads();
+
+ smem_pipe_read_curr = smem_pipe_read;
+ }
+
+ auto k_block_next = (k_block + 1) % STAGES;
+ int frag_idx_next = (k_block + 1) & 1;
+
+ load_fragments(
+ a_regs[frag_idx_next],
+ b_regs[frag_idx_next],
+ sfa_regs[frag_idx_next],
+ sfb_regs[frag_idx_next],
+ lane,
+ smem_pipe_read_curr, //smem_pipe_read
+ k_block_next, //k_block
+ smem_offset_A,
+ smem_offset_B,
+ smem_sf_offset,
+ smem
+ );
+
+ if (k_block == 0)
+ {
+ bool valid = (global_k < k);
+ cpasync_load(
+ smem_pipe_write,
+ global_k,
+ smem_offset_A,
+ smem_offset_B,
+ smem_sf_write_offset,
+ lane,
+ is_even,
+ load_b,
+ valid,
+ rowA,
+ vecB,
+ rowS,
+ vecS,
+ smem);
+ asm volatile("cp.async.commit_group;\n" ::);
+
+ smem_pipe_write = smem_pipe_read;
+ smem_pipe_read = (smem_pipe_read + 1) & 1;
+ }
+
+
+ int frag_idx = k_block & 1;
+ __half h = block_scaled_fma_16x2fp4(
+ a_regs[frag_idx],
+ b_regs[frag_idx],
+ sfa_regs[frag_idx],
+ sfb_regs[frag_idx]);
+ sum += __half2float(h);
+
+ }
+ idx += STAGES;
+ }
+
+ asm volatile("cp.async.wait_all;\n" ::);
+ __syncthreads();
+
+
+ unsigned mask = 0xffffffffu;
+ sum += __shfl_down_sync(mask, sum, 8, 16);
+ sum += __shfl_down_sync(mask, sum, 4, 16);
+ sum += __shfl_down_sync(mask, sum, 2, 16);
+ sum += __shfl_down_sync(mask, sum, 1, 16);
+
+ if (lane == 0) {
+ __half* out = (__half*)params.o_ptr + C_batch_base + row;
+ out[0] = __float2half(sum);
+ }
+ }
+
+ static inline void launch_kernel_fast(Gemv_params &params, cudaStream_t stream)
+ {
+ const int grid_x = (params.m + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
+ size_t smem_size = sizeof(SharedStorage);
+ dim3 grid(grid_x, 1, params.b);
+ dim3 block(BLOCK_SIZE, 1, 1);
+ gemv_kernel_fast<<<grid, block, smem_size>>>(params);
+ }
+
+ __global__ void __launch_bounds__(BLOCK_SIZE, 8)
+ gemv_kernel_k1024(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;
⋯ 209 unchanged lines
}
}
- static inline void launch_kernel_fast(Gemv_params &params, cudaStream_t stream)
+ static inline void launch_kernel_k1024(Gemv_params &params, cudaStream_t stream)
{
const int grid_x = (params.m + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
+ size_t smem_size = sizeof(SharedStorage);
dim3 grid(grid_x, 1, params.b);
dim3 block(BLOCK_SIZE, 1, 1);
- gemv_kernel_fast<<<grid, block, 0, stream>>>(params);
+ gemv_kernel_k1024<<<grid, block, 0, stream>>>(params);
}
// ============================================================================
+ // K=3584-specialized kernel
+ // ============================================================================
+
+ __global__ void __launch_bounds__(BLOCK_SIZE, 8)
+ gemv_kernel_k3584(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;
+
+ extern __shared__ __align__(128) uint8_t shared_storage[];
+ SharedStorage &smem = *reinterpret_cast<SharedStorage*>(shared_storage);
+
+ 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;
+
+ rowA += lane * 16;
+ vecB += lane * 16;
+
+ const int k = params.k;
+
+ float sum = 0.f;
+ int global_k = 0;
+ const int smem_offset_A = tid * 16;
+ const int smem_offset_B = lane * 16;
+ const int smem_sf_write_offset = (tid / 2) * 4;
+ const int smem_sf_offset = tid * 2;
+ const bool load_b = rib == 0;
+ const bool is_even = (lane % 2 == 0);
+
+ // prefetch to buffer 0
+ cpasync_load(
+ 0,
+ global_k,
+ smem_offset_A,
+ smem_offset_B,
+ smem_sf_write_offset,
+ lane,
+ is_even,
+ load_b,
+ true,
+ rowA,
+ vecB,
+ rowS,
+ vecS,
+ smem
+ );
+ asm volatile("cp.async.commit_group;\n" ::);
+ asm volatile("cp.async.wait_all;\n" ::);
+ __syncthreads();
+
+ uint64_t a_regs[2][2], b_regs[2][2];
+ uint16_t sfa_regs[2], sfb_regs[2];
+
+ int smem_pipe_read = 0;
+ int smem_pipe_write = 1;
+
+ // prefetch buffer 0 stage 0 to registers 0
+ load_fragments(
+ a_regs[0],
+ b_regs[0],
+ sfa_regs[0],
+ sfb_regs[0],
+ lane,
+ 0, //smem_pipe_read
+ 0, //k_block
+ smem_offset_A,
+ smem_offset_B,
+ smem_sf_offset,
+ smem
+ );
+ asm volatile("cp.async.commit_group;\n" ::);
+
+ int total_stages = params.k / (THREADS_PER_ROW * 16);
+ int full_iters = total_stages / STAGES; // Number of complete STAGES-groups
+ int tail_stages = total_stages % STAGES; // Remaining stages
+
+ // Main loop: process complete groups of STAGES
+ for (int idx = 0; idx < full_iters; ++idx) {
+ int smem_pipe_read_curr = smem_pipe_read;
+
+ for (int k_block = 0; k_block < STAGES; ++k_block)
+ {
+ if (k_block == STAGES-1)
+ {
+ asm volatile("cp.async.wait_all;\n" ::);
+ __syncthreads();
+
+ smem_pipe_read_curr = smem_pipe_read;
+ }
+
+ auto k_block_next = (k_block + 1) % STAGES;
+ int frag_idx_next = (k_block + 1) & 1;
+
+ load_fragments(
+ a_regs[frag_idx_next],
+ b_regs[frag_idx_next],
+ sfa_regs[frag_idx_next],
+ sfb_regs[frag_idx_next],
+ lane,
+ smem_pipe_read_curr,
+ k_block_next,
+ smem_offset_A,
+ smem_offset_B,
+ smem_sf_offset,
+ smem
+ );
+
+ if (k_block == 0)
+ {
+ bool valid = (global_k < k);
+ cpasync_load(
+ smem_pipe_write,
+ global_k,
+ smem_offset_A,
+ smem_offset_B,
+ smem_sf_write_offset,
+ lane,
+ is_even,
+ load_b,
+ valid,
+ rowA,
+ vecB,
+ rowS,
+ vecS,
+ smem);
+ asm volatile("cp.async.commit_group;\n" ::);
+
+ smem_pipe_write = smem_pipe_read;
+ smem_pipe_read = (smem_pipe_read + 1) & 1;
+ }
+
+ int frag_idx = k_block & 1;
+ __half h = block_scaled_fma_16x2fp4(
+ a_regs[frag_idx],
+ b_regs[frag_idx],
+ sfa_regs[frag_idx],
+ sfb_regs[frag_idx]);
+ sum += __half2float(h);
+ }
+ }
+
+ // Tail loop: process remaining stages
+ if (tail_stages > 0) {
+ asm volatile("cp.async.wait_all;\n" ::);
+ __syncthreads();
+
+ for (int k_block = 0; k_block < tail_stages; ++k_block)
+ {
+ int frag_idx = k_block & 1;
+ int frag_idx_next = (k_block + 1) & 1;
+
+ if (k_block < tail_stages - 1) {
+ load_fragments(
+ a_regs[frag_idx_next],
+ b_regs[frag_idx_next],
+ sfa_regs[frag_idx_next],
+ sfb_regs[frag_idx_next],
+ lane,
+ smem_pipe_read,
+ k_block + 1,
+ smem_offset_A,
+ smem_offset_B,
+ smem_sf_offset,
+ smem
+ );
+ }
+
+ __half h = block_scaled_fma_16x2fp4(
+ a_regs[frag_idx],
+ b_regs[frag_idx],
+ sfa_regs[frag_idx],
+ sfb_regs[frag_idx]);
+ sum += __half2float(h);
+ }
+ }
+
+ asm volatile("cp.async.wait_all;\n" ::);
+ __syncthreads();
+
+
+ unsigned mask = 0xffffffffu;
+ sum += __shfl_down_sync(mask, sum, 8, 16);
+ sum += __shfl_down_sync(mask, sum, 4, 16);
+ sum += __shfl_down_sync(mask, sum, 2, 16);
+ sum += __shfl_down_sync(mask, sum, 1, 16);
+
+ if (lane == 0) {
+ __half* out = (__half*)params.o_ptr + C_batch_base + row;
+ out[0] = __float2half(sum);
+ }
+ }
+
+ static inline void launch_kernel_k3584(Gemv_params &params, cudaStream_t stream)
+ {
+ const int grid_x = (params.m + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
+ size_t smem_size = sizeof(SharedStorage);
+ dim3 grid(grid_x, 1, params.b);
+ dim3 block(BLOCK_SIZE, 1, 1);
+ gemv_kernel_k3584<<<grid, block, smem_size>>>(params);
+ }
+
+ // ============================================================================
// K=256-specialized kernel (your previous v1) -> gemv_kernel_k256 / launch_kernel_k256
// ============================================================================
⋯ 116 unchanged lines
auto stream = at::cuda::getCurrentCUDAStream().stream();
// Tiny host-side dispatch: effectively zero overhead vs kernel time
- if (params.k == 128) {
+ if (params.k < 512) {
launch_kernel_k256(params, stream);
- } else {
+ }
+ else if (params.k == 1024) {
+ launch_kernel_k1024(params, stream);
+ }
+ else if (params.k % 1024) {
+ launch_kernel_k3584(params, stream);
+ }
+ else {
launch_kernel_fast(params, stream);
}
- cudaError_t err = cudaGetLastError();
- // Optional: check err and throw if needed
-
return C;
}
"""
scrolls · 720 diff lines total

Best evidence level for this revision: reported

JSON