Skip to content
KernelIndex
Search⌘K

submission 101911

tomaszki · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:46bf320fef802607cdf98470370bb18911f3067dee2bbeae415740a37214b845
license declaredunknown
license concludedunknown
authorstomaszki
imported2026-08-15

Techniques

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

fp4PyTorch reference implementation of NVFP4 block-scaled GEMV.
fp8__nv_fp8x2_storage_t a_pair =
shared-memoryextern __shared__ unsigned char shared_storage[];
vector-width = half2__half2 out_pair[4]) // 4 half2 → 8 results

Kernel source

submission.py331 lines
#!POPCORN leaderboard nvfp4_gemv

import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

# Kernel configuration parameters
sf_vec_size = 16


# Helper function for ceiling division
def ceil_div(a, b):
    return (a + b - 1) // b


# Helper function to convert scale factor tensor to blocked format
def to_blocked(input_matrix):
    rows, cols = input_matrix.shape

    # Please ensure rows and cols are multiples of 128 and 4 respectively
    n_row_blocks = ceil_div(rows, 128)
    n_col_blocks = ceil_div(cols, 4)

    padded = input_matrix
    blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
    rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)

    return rearranged.flatten()

def naive_pytorch(data: input_t) -> output_t:
    """
    PyTorch reference implementation of NVFP4 block-scaled GEMV.
    """
    a_ref, b_ref, sfa_ref_cpu, sfb_ref_cpu, _, _, c_ref = data

    # Get dimensions from MxNxL layout
    _, _, l = c_ref.shape

    # Call torch._scaled_mm to compute the GEMV result
    for l_idx in range(l):
        # Convert the scale factor tensor to blocked format
        scale_a = to_blocked(sfa_ref_cpu[:, :, l_idx])
        scale_b = to_blocked(sfb_ref_cpu[:, :, l_idx])
        # (m, k) @ (n, k).T -> (m, n)
        res = torch._scaled_mm(
            a_ref[:, :, l_idx],
            b_ref[:, :, l_idx].transpose(0, 1),
            scale_a.cuda(),
            scale_b.cuda(),
            bias=None,
            out_dtype=torch.float16,
        )
        c_ref[:, 0, l_idx] = res[:, 0]
    return c_ref

# CUDA SOURCE CODE

cuda_source = """
#include <cuda_fp4.h>
#include <cuda_fp8.h>
#include <cuda_fp16.h>


#define FULL_MASK 0xffffffff
#define ROWS_PER_BLOCK 32

__device__ void mul_fp4x8_to_half2(
    int a_packed,
    int b_packed,
    __half2  out_pair[4])   // 4 half2 → 8 results
{
    #pragma unroll
    for (int pair = 0; pair < 4; ++pair) {
        unsigned shift = 8 * pair;

        __nv_fp4x2_storage_t a_pair =
            static_cast<__nv_fp4x2_storage_t>((a_packed >> shift) & 0xFFu);
        __nv_fp4x2_storage_t b_pair =
            static_cast<__nv_fp4x2_storage_t>((b_packed >> shift) & 0xFFu);

        __half2_raw a_raw = __nv_cvt_fp4x2_to_halfraw2(a_pair, __NV_E2M1);
        __half2_raw b_raw = __nv_cvt_fp4x2_to_halfraw2(b_pair, __NV_E2M1);

        // __half2 has a constructor from __half2_raw in recent CUDA versions. :contentReference[oaicite:5]{index=5}
        __half2 a_h2(a_raw);
        __half2 b_h2(b_raw);

        out_pair[3 - pair] = __hmul2(a_h2, b_h2);
    }
}

__device__ void mul_fp8x8_to_half2(
    int2       a_packed,
    int2       b_packed,
    __half2    out_pair[4])   // 4 half2 → 8 results
{
    #pragma unroll
    for (int pair = 0; pair < 4; ++pair) {
        // Select which 32-bit word (x or y) and which 16-bit half inside it.
        int word_a = (pair < 2) ? a_packed.x : a_packed.y;
        int word_b = (pair < 2) ? b_packed.x : b_packed.y;

        unsigned shift = (pair & 1) * 16u;   // 0 or 16 bits

        __nv_fp8x2_storage_t a_pair =
            static_cast<__nv_fp8x2_storage_t>((static_cast<unsigned>(word_a) >> shift) & 0xFFFFu);
        __nv_fp8x2_storage_t b_pair =
            static_cast<__nv_fp8x2_storage_t>((static_cast<unsigned>(word_b) >> shift) & 0xFFFFu);

        // Convert fp8x2(e4m3) → half2_raw
        __half2_raw a_raw = __nv_cvt_fp8x2_to_halfraw2(a_pair, __NV_E4M3);
        __half2_raw b_raw = __nv_cvt_fp8x2_to_halfraw2(b_pair, __NV_E4M3);

        // __half2 has a constructor from __half2_raw in recent CUDA versions.
        __half2 a_h2(a_raw);
        __half2 b_h2(b_raw);

        out_pair[3 - pair] = __hmul2(a_h2, b_h2);
    }
}


__global__ void gemv_kernel(
    const __nv_fp4x2_storage_t* __restrict__ a,
    const __nv_fp4x2_storage_t* __restrict__ b,
    const __nv_fp8_e4m3* __restrict__ sfa,
    const __nv_fp8_e4m3* __restrict__ sfb,
    __half* __restrict__ c,
    int M,
    int K
) {
    extern __shared__ unsigned char shared_storage[];
    auto* b_shared = reinterpret_cast<__nv_fp4x2_storage_t*>(shared_storage);
    auto* sfb_shared = reinterpret_cast<__nv_fp8_e4m3*>(b_shared + (K / 2));
    __shared__ __half c_shared[32];

    b += blockIdx.y * (K / 2) * 128;
    sfb += blockIdx.y * (K / 16) * 128;

    for (int i = threadIdx.y * 32 + threadIdx.x; i < K / 8; i += blockDim.y * blockDim.x) {
        reinterpret_cast<int*>(b_shared)[i] = reinterpret_cast<const int*>(b)[i];
    }
    for (int i = threadIdx.y * 32 + threadIdx.x; i < K / 64; i += blockDim.y * blockDim.x) {
        reinterpret_cast<int*>(sfb_shared)[i] = reinterpret_cast<const int*>(sfb)[i];
    }
    __syncthreads();

    // Each warp computes one result and saves it to shared memory
    __half2 result_0 = __float2half2_rn(0.0f);
    __half2 result_1 = __float2half2_rn(0.0f);
    __half2 result_2 = __float2half2_rn(0.0f);
    __half2 result_3 = __float2half2_rn(0.0f);
    int offset = blockIdx.y * (K * M / 2) + (blockIdx.x * 32 + threadIdx.y) * (K / 2);
    a += offset;
    sfa += offset / 8;
    uchar4 a_packed;
    uchar4 b_packed;
    __half2 a_h2_0, a_h2_1, a_h2_2, a_h2_3, b_h2_0, b_h2_1, b_h2_2, b_h2_3;
    __half2 prod_h2_0, prod_h2_1, prod_h2_2, prod_h2_3;
    for (int i = threadIdx.x; i < K / 32; i += 32) {
        int4 a_packed_i4 = reinterpret_cast<const int4*>(a)[i];
        int4 b_packed_i4 = reinterpret_cast<int4*>(b_shared)[i];
        const uchar4* a_packed_u4 = reinterpret_cast<const uchar4*>(&a_packed_i4);
        const uchar4* b_packed_u4 = reinterpret_cast<const uchar4*>(&b_packed_i4);
        __nv_fp8x2_storage_t sfa_fp8x2 = reinterpret_cast<const __nv_fp8x2_storage_t*>(sfa)[i];
        __nv_fp8x2_storage_t sfb_fp8x2 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared)[i];
        __half2 sfa_h2 = __half2(__nv_cvt_fp8x2_to_halfraw2(sfa_fp8x2, __NV_E4M3));
        __half2 sfb_h2 = __half2(__nv_cvt_fp8x2_to_halfraw2(sfb_fp8x2, __NV_E4M3));
        __half2 sf_h2 = __hmul2(sfa_h2, sfb_h2);
        __half2 sf_low_h2 = __low2half2(sf_h2);
        __half2 sf_high_h2 = __high2half2(sf_h2);

        a_packed = a_packed_u4[0];
        b_packed = b_packed_u4[0];
        a_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.x), __NV_E2M1);
        a_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.y), __NV_E2M1);
        a_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.z), __NV_E2M1);
        a_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.w), __NV_E2M1);
        b_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.x), __NV_E2M1);
        b_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.y), __NV_E2M1);
        b_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.z), __NV_E2M1);
        b_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.w), __NV_E2M1);
        prod_h2_0 = __hmul2(a_h2_0, b_h2_0);
        prod_h2_1 = __hmul2(a_h2_1, b_h2_1);
        prod_h2_2 = __hmul2(a_h2_2, b_h2_2);
        prod_h2_3 = __hmul2(a_h2_3, b_h2_3);
        result_0 = __hadd2(result_0, __hmul2(prod_h2_0, sf_low_h2));
        result_1 = __hadd2(result_1, __hmul2(prod_h2_1, sf_low_h2));
        result_2 = __hadd2(result_2, __hmul2(prod_h2_2, sf_low_h2));
        result_3 = __hadd2(result_3, __hmul2(prod_h2_3, sf_low_h2));

        a_packed = a_packed_u4[1];
        b_packed = b_packed_u4[1];
        a_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.x), __NV_E2M1);
        a_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.y), __NV_E2M1);
        a_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.z), __NV_E2M1);
        a_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.w), __NV_E2M1);
        b_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.x), __NV_E2M1);
        b_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.y), __NV_E2M1);
        b_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.z), __NV_E2M1);
        b_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.w), __NV_E2M1);
        prod_h2_0 = __hmul2(a_h2_0, b_h2_0);
        prod_h2_1 = __hmul2(a_h2_1, b_h2_1);
        prod_h2_2 = __hmul2(a_h2_2, b_h2_2);
        prod_h2_3 = __hmul2(a_h2_3, b_h2_3);
        result_0 = __hadd2(result_0, __hmul2(prod_h2_0, sf_low_h2));
        result_1 = __hadd2(result_1, __hmul2(prod_h2_1, sf_low_h2));
        result_2 = __hadd2(result_2, __hmul2(prod_h2_2, sf_low_h2));
        result_3 = __hadd2(result_3, __hmul2(prod_h2_3, sf_low_h2));

        a_packed = a_packed_u4[2];
        b_packed = b_packed_u4[2];
        a_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.x), __NV_E2M1);
        a_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.y), __NV_E2M1);
        a_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.z), __NV_E2M1);
        a_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.w), __NV_E2M1);
        b_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.x), __NV_E2M1);
        b_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.y), __NV_E2M1);
        b_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.z), __NV_E2M1);
        b_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.w), __NV_E2M1);
        prod_h2_0 = __hmul2(a_h2_0, b_h2_0);
        prod_h2_1 = __hmul2(a_h2_1, b_h2_1);
        prod_h2_2 = __hmul2(a_h2_2, b_h2_2);
        prod_h2_3 = __hmul2(a_h2_3, b_h2_3);
        result_0 = __hadd2(result_0, __hmul2(prod_h2_0, sf_high_h2));
        result_1 = __hadd2(result_1, __hmul2(prod_h2_1, sf_high_h2));
        result_2 = __hadd2(result_2, __hmul2(prod_h2_2, sf_high_h2));
        result_3 = __hadd2(result_3, __hmul2(prod_h2_3, sf_high_h2));

        a_packed = a_packed_u4[3];
        b_packed = b_packed_u4[3];
        a_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.x), __NV_E2M1);
        a_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.y), __NV_E2M1);
        a_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.z), __NV_E2M1);
        a_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.w), __NV_E2M1);
        b_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.x), __NV_E2M1);
        b_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.y), __NV_E2M1);
        b_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.z), __NV_E2M1);
        b_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.w), __NV_E2M1);
        prod_h2_0 = __hmul2(a_h2_0, b_h2_0);
        prod_h2_1 = __hmul2(a_h2_1, b_h2_1);
        prod_h2_2 = __hmul2(a_h2_2, b_h2_2);
        prod_h2_3 = __hmul2(a_h2_3, b_h2_3);
        result_0 = __hadd2(result_0, __hmul2(prod_h2_0, sf_high_h2));
        result_1 = __hadd2(result_1, __hmul2(prod_h2_1, sf_high_h2));
        result_2 = __hadd2(result_2, __hmul2(prod_h2_2, sf_high_h2));
        result_3 = __hadd2(result_3, __hmul2(prod_h2_3, sf_high_h2));
    }


    // Reduce the result and store it in shared memory
    result_0 = __hadd2(result_0, result_1);
    result_2 = __hadd2(result_2, result_3);
    result_0 = __hadd2(result_0, result_2);
    float final_result_f = __half22float2(result_0).x + __half22float2(result_0).y;
    for (int offset = 16; offset > 0; offset /= 2) {
        final_result_f += __shfl_down_sync(FULL_MASK, final_result_f, offset);
    }
    if (threadIdx.x == 0) {
        c_shared[threadIdx.y] = __float2half_rn(final_result_f);
    }
    __syncthreads();
    
    // Write the result to global memory
    if (threadIdx.y == 0) {
        int c_offset = blockIdx.y * M + blockIdx.x * 32 + threadIdx.x;
        c[c_offset] = c_shared[threadIdx.x];
    }
}



torch::Tensor gemv_cuda(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c) {
    const int64_t M = a.size(0);
    const int64_t K = a.size(1) * 2;
    const int64_t L = a.size(2);

    dim3 block_dim(32, 32, 1);
    dim3 grid_dim(M / 32, L, 1);
    const auto* a_ptr = reinterpret_cast<const __nv_fp4x2_storage_t*>(a.data_ptr());
    const auto* b_ptr = reinterpret_cast<const __nv_fp4x2_storage_t*>(b.data_ptr());
    const auto* sfa_ptr = reinterpret_cast<const __nv_fp8_e4m3*>(sfa.data_ptr());
    const auto* sfb_ptr = reinterpret_cast<const __nv_fp8_e4m3*>(sfb.data_ptr());
    auto* c_ptr = reinterpret_cast<__half*>(c.data_ptr<c10::Half>());

    size_t shared_mem_bytes =
        (static_cast<size_t>(K) / 2) * sizeof(__nv_fp4x2_storage_t) +
        (static_cast<size_t>(K) / 16) * sizeof(__nv_fp8_e4m3);

    gemv_kernel<<<grid_dim, block_dim, shared_mem_bytes>>>(
        a_ptr,
        b_ptr,
        sfa_ptr,
        sfb_ptr,
        c_ptr,
        static_cast<int>(M),
        static_cast<int>(K)
    );
    return c;
}
"""


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

torch::Tensor gemv_cuda(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c);
"""

gemv_module = load_inline(
    name='gemv_cuda',
    cpp_sources=cpp_source,
    cuda_sources=cuda_source,
    functions=['gemv_cuda'],
    verbose=True,
    extra_cuda_cflags=['-arch=sm_100a', '-std=c++17', '-O3'],
)




def custom_kernel(
    data: input_t,
) -> output_t:
    """
    PyTorch reference implementation of NVFP4 block-scaled GEMV.
    """

    a, b, sfa, sfb, _, _, c = data

    return gemv_module.gemv_cuda(a, b, sfa, sfb, c)
scrolls · 331 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 100226.

⋯ 152 unchanged lines
int offset = blockIdx.y * (K * M / 2) + (blockIdx.x * 32 + threadIdx.y) * (K / 2);
a += offset;
sfa += offset / 8;
- for (int i = threadIdx.x; i < K / 8; i += 32) {
- uchar4 a_packed = reinterpret_cast<const uchar4*>(a)[i];
- uchar4 b_packed = reinterpret_cast<const uchar4*>(b_shared)[i];
- __half2 a_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.x), __NV_E2M1);
- __half2 a_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.y), __NV_E2M1);
- __half2 a_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.z), __NV_E2M1);
- __half2 a_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.w), __NV_E2M1);
- __half2 b_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.x), __NV_E2M1);
- __half2 b_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.y), __NV_E2M1);
- __half2 b_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.z), __NV_E2M1);
- __half2 b_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.w), __NV_E2M1);
- __half2 prod_h2_0 = __hmul2(a_h2_0, b_h2_0);
- __half2 prod_h2_1 = __hmul2(a_h2_1, b_h2_1);
- __half2 prod_h2_2 = __hmul2(a_h2_2, b_h2_2);
- __half2 prod_h2_3 = __hmul2(a_h2_3, b_h2_3);
+ uchar4 a_packed;
+ uchar4 b_packed;
+ __half2 a_h2_0, a_h2_1, a_h2_2, a_h2_3, b_h2_0, b_h2_1, b_h2_2, b_h2_3;
+ __half2 prod_h2_0, prod_h2_1, prod_h2_2, prod_h2_3;
+ for (int i = threadIdx.x; i < K / 32; i += 32) {
+ int4 a_packed_i4 = reinterpret_cast<const int4*>(a)[i];
+ int4 b_packed_i4 = reinterpret_cast<int4*>(b_shared)[i];
+ const uchar4* a_packed_u4 = reinterpret_cast<const uchar4*>(&a_packed_i4);
+ const uchar4* b_packed_u4 = reinterpret_cast<const uchar4*>(&b_packed_i4);
+ __nv_fp8x2_storage_t sfa_fp8x2 = reinterpret_cast<const __nv_fp8x2_storage_t*>(sfa)[i];
+ __nv_fp8x2_storage_t sfb_fp8x2 = reinterpret_cast<__nv_fp8x2_storage_t*>(sfb_shared)[i];
+ __half2 sfa_h2 = __half2(__nv_cvt_fp8x2_to_halfraw2(sfa_fp8x2, __NV_E4M3));
+ __half2 sfb_h2 = __half2(__nv_cvt_fp8x2_to_halfraw2(sfb_fp8x2, __NV_E4M3));
+ __half2 sf_h2 = __hmul2(sfa_h2, sfb_h2);
+ __half2 sf_low_h2 = __low2half2(sf_h2);
+ __half2 sf_high_h2 = __high2half2(sf_h2);
- __half sfa_h = __nv_cvt_fp8_to_halfraw(sfa[i / 2].__x, __NV_E4M3);
- __half sfb_h = __nv_cvt_fp8_to_halfraw(sfb_shared[i / 2].__x, __NV_E4M3);
- __half2 sf_h2 = __half2half2(__hmul(sfa_h, sfb_h));
- result_0 = __hadd2(result_0, __hmul2(prod_h2_0, sf_h2));
- result_1 = __hadd2(result_1, __hmul2(prod_h2_1, sf_h2));
- result_2 = __hadd2(result_2, __hmul2(prod_h2_2, sf_h2));
- result_3 = __hadd2(result_3, __hmul2(prod_h2_3, sf_h2));
+ a_packed = a_packed_u4[0];
+ b_packed = b_packed_u4[0];
+ a_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.x), __NV_E2M1);
+ a_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.y), __NV_E2M1);
+ a_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.z), __NV_E2M1);
+ a_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.w), __NV_E2M1);
+ b_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.x), __NV_E2M1);
+ b_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.y), __NV_E2M1);
+ b_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.z), __NV_E2M1);
+ b_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.w), __NV_E2M1);
+ prod_h2_0 = __hmul2(a_h2_0, b_h2_0);
+ prod_h2_1 = __hmul2(a_h2_1, b_h2_1);
+ prod_h2_2 = __hmul2(a_h2_2, b_h2_2);
+ prod_h2_3 = __hmul2(a_h2_3, b_h2_3);
+ result_0 = __hadd2(result_0, __hmul2(prod_h2_0, sf_low_h2));
+ result_1 = __hadd2(result_1, __hmul2(prod_h2_1, sf_low_h2));
+ result_2 = __hadd2(result_2, __hmul2(prod_h2_2, sf_low_h2));
+ result_3 = __hadd2(result_3, __hmul2(prod_h2_3, sf_low_h2));
+
+ a_packed = a_packed_u4[1];
+ b_packed = b_packed_u4[1];
+ a_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.x), __NV_E2M1);
+ a_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.y), __NV_E2M1);
+ a_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.z), __NV_E2M1);
+ a_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.w), __NV_E2M1);
+ b_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.x), __NV_E2M1);
+ b_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.y), __NV_E2M1);
+ b_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.z), __NV_E2M1);
+ b_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.w), __NV_E2M1);
+ prod_h2_0 = __hmul2(a_h2_0, b_h2_0);
+ prod_h2_1 = __hmul2(a_h2_1, b_h2_1);
+ prod_h2_2 = __hmul2(a_h2_2, b_h2_2);
+ prod_h2_3 = __hmul2(a_h2_3, b_h2_3);
+ result_0 = __hadd2(result_0, __hmul2(prod_h2_0, sf_low_h2));
+ result_1 = __hadd2(result_1, __hmul2(prod_h2_1, sf_low_h2));
+ result_2 = __hadd2(result_2, __hmul2(prod_h2_2, sf_low_h2));
+ result_3 = __hadd2(result_3, __hmul2(prod_h2_3, sf_low_h2));
+
+ a_packed = a_packed_u4[2];
+ b_packed = b_packed_u4[2];
+ a_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.x), __NV_E2M1);
+ a_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.y), __NV_E2M1);
+ a_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.z), __NV_E2M1);
+ a_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.w), __NV_E2M1);
+ b_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.x), __NV_E2M1);
+ b_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.y), __NV_E2M1);
+ b_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.z), __NV_E2M1);
+ b_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.w), __NV_E2M1);
+ prod_h2_0 = __hmul2(a_h2_0, b_h2_0);
+ prod_h2_1 = __hmul2(a_h2_1, b_h2_1);
+ prod_h2_2 = __hmul2(a_h2_2, b_h2_2);
+ prod_h2_3 = __hmul2(a_h2_3, b_h2_3);
+ result_0 = __hadd2(result_0, __hmul2(prod_h2_0, sf_high_h2));
+ result_1 = __hadd2(result_1, __hmul2(prod_h2_1, sf_high_h2));
+ result_2 = __hadd2(result_2, __hmul2(prod_h2_2, sf_high_h2));
+ result_3 = __hadd2(result_3, __hmul2(prod_h2_3, sf_high_h2));
+
+ a_packed = a_packed_u4[3];
+ b_packed = b_packed_u4[3];
+ a_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.x), __NV_E2M1);
+ a_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.y), __NV_E2M1);
+ a_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.z), __NV_E2M1);
+ a_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(a_packed.w), __NV_E2M1);
+ b_h2_0 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.x), __NV_E2M1);
+ b_h2_1 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.y), __NV_E2M1);
+ b_h2_2 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.z), __NV_E2M1);
+ b_h2_3 = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<__nv_fp4x2_storage_t>(b_packed.w), __NV_E2M1);
+ prod_h2_0 = __hmul2(a_h2_0, b_h2_0);
+ prod_h2_1 = __hmul2(a_h2_1, b_h2_1);
+ prod_h2_2 = __hmul2(a_h2_2, b_h2_2);
+ prod_h2_3 = __hmul2(a_h2_3, b_h2_3);
+ result_0 = __hadd2(result_0, __hmul2(prod_h2_0, sf_high_h2));
+ result_1 = __hadd2(result_1, __hmul2(prod_h2_1, sf_high_h2));
+ result_2 = __hadd2(result_2, __hmul2(prod_h2_2, sf_high_h2));
+ result_3 = __hadd2(result_3, __hmul2(prod_h2_3, sf_high_h2));
}
⋯ 62 unchanged lines
cuda_sources=cuda_source,
functions=['gemv_cuda'],
verbose=True,
- extra_cuda_cflags=['-arch=sm_100a'],
+ extra_cuda_cflags=['-arch=sm_100a', '-std=c++17', '-O3'],
)
scrolls · 130 diff lines total

Best evidence level for this revision: reported

JSON